THE CODE READER / PYTHON
Download raw .py ↓train_copy.py
Your snippet, in context. Explore the file or visualize its recorded example.
1"""Train a tiny transformer on copying integer sequences; no downloads or GPU.
2
3Keep this file next to torch_core.py. Run: python train_copy.py --steps 800
4The train/evaluation sequences are disjoint; evaluation uses greedy generation.
5"""
6from __future__ import annotations
7
8from typing import TypedDict
9from torch import Tensor
10import argparse
11import itertools
12import json
13import time
14import torch
15from torch.nn import functional as F
16from torch_core import TinyTransformer, generate
17
18class GenerationExample(TypedDict):
19 source: list[int]
20 expected: list[int]
21 generated: list[int]
22
23
24class TrainingResult(TypedDict):
25 seed: int
26 steps: int
27 train_sequences: int
28 held_out_sequences: int
29 initial_train_loss: float
30 held_out_loss: float
31 greedy_exact_match: float
32 seconds: float
33 numpy_torch_note: str
34 examples: list[GenerationExample]
35
36
37# region data
38def dataset(seed: int = 7) -> tuple[list[list[int]], list[list[int]]]:
39 sequences = [list(items) for length in [1, 2, 3]
40 for items in itertools.product(range(4, 12), repeat=length)]
41 generator = torch.Generator().manual_seed(seed)
42 order = torch.randperm(len(sequences), generator=generator).tolist()
43 cut = int(0.8 * len(order))
44 return ([sequences[i] for i in order[:cut]],
45 [sequences[i] for i in order[cut:]])
46
47
48def batch(sequences: list[list[int]]) -> tuple[Tensor, Tensor, Tensor]:
49 source = torch.zeros(len(sequences), 3, dtype=torch.long)
50 decoder_input = torch.zeros(len(sequences), 4, dtype=torch.long)
51 target = torch.zeros(len(sequences), 4, dtype=torch.long)
52 for i, tokens in enumerate(sequences):
53 source[i, :len(tokens)] = torch.tensor(tokens)
54 decoder_input[i, :len(tokens) + 1] = torch.tensor([1] + tokens)
55 target[i, :len(tokens) + 1] = torch.tensor(tokens + [2])
56 return source, decoder_input, target
57# endregion data
58
59# region loss
60def training_loss(model: TinyTransformer, source: Tensor, decoder_input: Tensor, target: Tensor) -> Tensor:
61 logits = model(source, decoder_input) # (B, T, vocabulary)
62 return F.cross_entropy(
63 logits.reshape(-1, logits.shape[-1]),
64 target.reshape(-1),
65 ignore_index=0, # Padding is not a prediction target.
66 ) # Pass raw logits; do not softmax first.
67# endregion loss
68
69# region train
70def train(steps: int = 800, seed: int = 7) -> tuple[TinyTransformer, TrainingResult]:
71 torch.set_num_threads(1)
72 torch.manual_seed(seed)
73 train_sequences, held_out = dataset(seed)
74 model = TinyTransformer()
75 optimizer = torch.optim.Adam(model.parameters(), lr=1e-3)
76 generator = torch.Generator().manual_seed(seed + 1)
77 initial_source, initial_input, initial_target = batch(train_sequences[:32])
78 initial_loss = training_loss(model, initial_source, initial_input, initial_target).item()
79 start = time.perf_counter()
80 for _ in range(steps):
81 indices = torch.randint(len(train_sequences), (32,), generator=generator)
82 source, decoder_input, target = batch([train_sequences[i] for i in indices.tolist()])
83 optimizer.zero_grad(set_to_none=True)
84 loss = training_loss(model, source, decoder_input, target)
85 loss.backward()
86 torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
87 optimizer.step()
88
89 source, decoder_input, target = batch(held_out)
90 model.eval()
91 with torch.no_grad():
92 held_out_loss = training_loss(model, source, decoder_input, target).item()
93 generated = generate(model, source)
94 generated = F.pad(generated, (0, target.shape[1] - generated.shape[1]))
95 exact_match = (generated == target).all(dim=1).float().mean().item()
96 result: TrainingResult = {
97 "seed": seed, "steps": steps, "train_sequences": len(train_sequences),
98 "held_out_sequences": len(held_out), "initial_train_loss": round(initial_loss, 6),
99 "held_out_loss": round(held_out_loss, 6), "greedy_exact_match": round(exact_match, 6),
100 "seconds": round(time.perf_counter() - start, 2),
101 "numpy_torch_note": "NumPy demonstrates forward passes; PyTorch trains the model.",
102 "examples": [
103 {"source": source[i].tolist(), "expected": target[i].tolist(),
104 "generated": generated[i].tolist()} for i in range(3)
105 ],
106 }
107 return model, result
108# endregion train
109
110if __name__ == "__main__":
111 parser = argparse.ArgumentParser()
112 parser.add_argument("--steps", type=int, default=800)
113 parser.add_argument("--seed", type=int, default=7)
114 args = parser.parse_args()
115 if args.steps < 1:
116 parser.error("--steps must be positive")
117 _, result = train(args.steps, args.seed)
118 print(json.dumps(result, indent=2))