# Part 1: executable PyTorch snippets in lecture order.
# Source: src/sections1/*.html. Generated by check_part1_torch.py --export-text.
# Run with Python and PyTorch installed. No datasets are downloaded.
# This constructs fresh teaching models and illustrates individual SGD updates.
# It does not reproduce the already-trained worksheet values or its named axes.
# The companion checker separately loads toy1.json to verify those values.
# The name model receives only one update, so its sampled text can be implausible.

# vocabulary
import torch
from torch import nn
from torch.nn import functional as F
vocab = ["-"] + list("abcdefghijklmnopqrstuvwxyz")
stoi = {c: j for j, c in enumerate(vocab)}

# context
ctx = torch.tensor([[stoi[c] for c in "aab"]])
target = torch.tensor([stoi["i"]])
print(ctx.shape, target.shape)

# windows
w = 3
padded = [stoi["-"]] * w + [stoi[c] for c in "aabid-"]
X = torch.tensor([padded[j:j+w] for j in range(6)])
y = torch.tensor(padded[w:])

# one-hot
one_hot = F.one_hot(ctx, num_classes=27)
print(one_hot.shape)  # [1, 3, 27]

# single-character-lookup
embedding = nn.Embedding(27, 2)
char_id = torch.tensor([1])  # id("a") = 1
e_a = embedding(char_id)
print(e_a.shape)  # [1, 2]

# embedding
ctx = torch.tensor([[1, 1, 2]])  # one example: "a", "a", "b"
e = embedding(ctx)
print(e.shape)  # [1, 3, 2]

# word-analogy
man, woman, king, queen = torch.tensor(
    [[1., 1.], [3., 1.], [1., 3.], [3., 3.]])
candidate = king - man + woman
print(candidate)  # tensor([3., 3.])

# document-mean
doc_embedding = nn.Embedding(4, 2)
doc_ids = torch.tensor([[1, 2, 3]])  # bank lends money
doc_rows = doc_embedding(doc_ids)
e_doc = doc_rows.mean(dim=1)

# image-encoder
pixels = torch.tensor([[[[0., 1.], [1., 0.]]]])
image_encoder = nn.Conv2d(1, 4, kernel_size=2)
e_image = image_encoder(pixels).flatten(1)
print(e_image.shape)  # [1, 4]

# time-encoder
signal = torch.tensor([[[0., 1., 0., -1., 0., 1.]]])
time_encoder = nn.Conv1d(1, 3, kernel_size=3)
e_time = time_encoder(signal).mean(dim=-1)
print(e_time.shape)  # [1, 3]

# flatten
a0 = e.flatten(start_dim=1)
print(a0.shape)  # [1, 6]

# relu-rule
before_relu = torch.tensor([-2., 0., 3.])
after_relu = torch.relu(before_relu)
print(after_relu)  # tensor([0., 0., 3.])

# hidden
hidden = nn.Linear(6, 32)
a1 = torch.relu(hidden(a0))
print(a1.shape)  # [1, 32]

# weight-convention
W1 = hidden.weight.T
b1 = hidden.bias
check = torch.relu(a0 @ W1 + b1)
torch.testing.assert_close(a1, check)

# output
output = nn.Linear(32, 27)
z = output(a1)
print(z.shape)  # [1, 27]

# model
model_seq = nn.Sequential(
    embedding, nn.Flatten(1), hidden, nn.ReLU(), output
)
z_seq = model_seq(ctx)

# model-class
class NameMLP(nn.Module):
    def __init__(self, embedding, hidden, output):
        super().__init__()
        self.embedding = embedding
        self.hidden = hidden
        self.output = output
    def forward(self, ctx):
        a0 = self.embedding(ctx).flatten(1)
        a1 = torch.relu(self.hidden(a0))
        return self.output(a1)

model = NameMLP(embedding, hidden, output)
z = model(ctx)  # [1, 27], same logits as model_seq(ctx)

# batch-shapes
a0_batch = embedding(X[:4]).flatten(1)
a1_batch = torch.relu(hidden(a0_batch))
z_batch = output(a1_batch)
print(a0_batch.shape, a1_batch.shape, z_batch.shape)
# [4, 6], [4, 32], [4, 27]

# softmax
p = z.softmax(dim=-1)
print(p.shape, p.sum(dim=-1))  # [1, 27], one total

# stable-softmax
shifted = z - z.max(dim=-1, keepdim=True).values
weights = shifted.exp()
p = weights / weights.sum(dim=-1, keepdim=True)

# loss
loss = F.cross_entropy(z, target)
print(loss.item())

# named-parameters
for name, param in model.named_parameters():
    print(name, list(param.shape))

# optimizer
params = list(model.parameters())  # the five tensors above
optimizer = torch.optim.SGD(params, lr=0.1)
model.train()

# training
optimizer.zero_grad()
z = model(X)
loss = F.cross_entropy(z, y)
loss.backward()
optimizer.step()

# gradients
print(model.embedding.weight.grad.shape)  # [27, 2]
print(model.hidden.weight.grad.shape)     # [32, 6]

# sample-function
@torch.no_grad()
def sample_next(ctx, temperature=1.0):
    z = model(ctx)
    p = (z / temperature).softmax(dim=-1)
    return torch.multinomial(p, num_samples=1)

# generation-setup
model.eval()
ctx = torch.full((1, 3), stoi["-"], dtype=torch.long)
name = []
temperature = 1.0

# generation-loop
for _ in range(18):
    next_id = sample_next(ctx, temperature)
    if next_id.item() == stoi["-"]: break
    name.append(vocab[next_id.item()])
    ctx = torch.cat([ctx[:, 1:], next_id], dim=1)

# generation-choices
ctx = torch.tensor([[stoi["-"], stoi["s"], stoi["a"]]])
with torch.no_grad(): p = model(ctx).softmax(dim=-1)
greedy_id = p.argmax(dim=-1, keepdim=True)
sampled_id = torch.multinomial(p, num_samples=1)

# batch
batch_z = model(X)
batch_loss = F.cross_entropy(batch_z, y)

# tokenization
text = "deep learning is amazing"
char_tokens = list(text)
word_tokens = text.split()

# word-embedding
word_vocab = ["<unk>", "The", "cat", "sat", "on", "the"]
word_stoi = {t: j for j, t in enumerate(word_vocab)}
word_ctx = torch.tensor([[2, 3, 4]])  # cat sat on
word_embedding = nn.Embedding(len(word_vocab), 2)
word_e = word_embedding(word_ctx)

# word-output
word_hidden = nn.Linear(6, 32)
word_output = nn.Linear(32, len(word_vocab))
word_a1 = torch.relu(word_hidden(word_e.flatten(1)))
word_z = word_output(word_a1)

# longer-window
new_w, new_d = 5, 4
size_model = nn.Sequential(
    nn.Embedding(27, new_d), nn.Flatten(1),
    nn.Linear(new_w * new_d, 32), nn.ReLU(), nn.Linear(32, 27))
print(sum(p.numel() for p in size_model.parameters()))  # 1671

# forward-recap
e = embedding(ctx)
a0 = e.flatten(1)
a1 = torch.relu(hidden(a0))
z = output(a1)
p = z.softmax(dim=-1)
