THE CODE READER / PYTHON

train_copy.py

Your snippet, in context. Explore the file or visualize its recorded example.

Download raw .py ↓
Complete file
PYTHON / LINE NUMBERS
  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))