Skip to content

Training loop and the autograd bigram

ModuleL0.5 · build · Python · Pass 2 · 4 to 6 h
You buildpython/tinyllm/train/loop.py: DataLoader, train_step, evaluate; and you take over python/tinyllm/lm/bigram.py from L0.0: BigramLogits (the table as a trainable Module) and BigramLM.sample on PCG32
Contractcourse/contracts/py/tinyllm/train/loop.pyi · course/contracts/py/tinyllm/lm/bigram.pyi
Testscourse/tests/L0.5/ (what they check: section 4); L0.0’s tests keep running against bigram.py as your regression suite
NeedsL0.1 no_grad · L0.2 F.embedding · L0.3 the losses the tests train with · L0.4 Module, Linear · M06.3 PCG32 · M10.2 SGD · M10.3 AdamW · M10.4 clip_grad_norm_ · M02.2 EMA (the smoothed loss) · rt.01 and M03.1 (BigramLM.logits still runs your C matmul) · reading: L0.0 (or --ref-deps)
Used byL10.0 your engine serves the table this module trains (inherited from L0.0 with bigram.py) · L0.6 runs L0.0’s suite, bigram.py included, as its regression · later: L2.2, L3.6, L6.1, and the capstone train with this loop · later: L5.5, L6.7
MilestoneMS-L0 (step 2: the autograd bigram reaches the count MLE; step 3: the digits MLP)
Optional depthKarpathy, “A Recipe for Training Neural Networks” (2019, free); Goodfellow, Bengio, Courville, Deep Learning, ch. 8
  • One training step is zero the gradients, compute the loss, refuse a non-finite loss, backpropagate, clip, update: in that order, every time (test_hand_example_train_step, test_zero_grad_every_step, test_nonfinite_loss_stops_before_the_update).
  • Data order comes from the seeded PCG32 through the spec’s Fisher-Yates shuffle, a fresh permutation each epoch, so two runs with one seed are bitwise identical (test_dataloader_shuffle_is_the_spec_permutation, test_same_seed_bitwise_identical).
  • Evaluation weights every batch by its rows, runs in eval mode under no_grad, and always hands the model back in the mode it found it (test_evaluate_weights_by_rows_and_restores_mode, test_evaluate_restores_mode_after_an_error).
  • The bigram is a convex problem whose minimum is the unsmoothed count model; your autograd reaches its NLL within 10−310^{-3} nats (test_autograd_bigram_reaches_the_count_mle).
  • sample now draws exactly as your Rust engine does, so a seed gives the same text in Python and Rust (test_sample_draws_like_the_engine).
Terminal window
ol start L0.5 # stubs loop.py; bigram.py is yours already: ol start prints its contract diff
ol tests L0.5 # read the test catalog first: rung R0
ol check L0.5 # also reruns L0.0's smoke tests against your bigram.py
ol check L0.5 --ref-deps # only if a dependency is not passing yet
ol diff L0.5 # after passing: your code against the reference

ol start never overwrites your bigram.py. It prints what the contract adds (BigramLogits, the PCG32 sampler): edit your Pass 1 file in place.

Your CLI gains two verbs, fixed by course/milestones/MS-L0.toml: train bigram --method autograd --data F --out DIR (full-batch training of BigramLogits on every byte pair of F, then the same model directory as Pass 1, from to_lm()) and train mlp --data digits.npz --hidden 64 --epochs 30 --ckpt DIR (a two-layer ReLU MLP on the UCI digits through DataLoader and train_step; standardize each pixel with the training set’s mean and standard deviation). --method counts stays the Pass 1 verb.


L0.1 to L0.4 give you every piece of a training run and no run. The digits MLP of MS-L0 needs batches drawn in a reproducible order, a step that does the five things of a step in the right order, and an evaluation that does not train. And the tracer model your engine has served since Pass 1 is still a table of counts: it cannot be improved, only recounted. This module writes the loop every later model trains with, then retrains the tracer’s own table with it. If the loop is right, gradient descent finds the count model again (the bigram problem has one minimum, and the counts are it), which makes the bigram the most precise end-to-end test of your autograd you will ever have: the answer is known to many digits.

SymbolMeaningType / shape
θ\thetaevery parameter of the modellist of Tensor
BBbatch size; nn rows in the datasetints
L(θ;batch)\mathcal{L}(\theta; \text{batch})the loss loss_fn returns, one elementTensor
η\etalearning ratefloat
ccgradient clipping threshold, global norm (M10.4)float
∥g∥=∑p∑igp,i2\lVert g \rVert = \sqrt{\sum_p \sum_i g_{p,i}^2}global gradient norm over all parametersfloat
W∈RV×VW \in \mathbb{R}^{V \times V}the bigram table, row ii = logits after token iifloat32[V, V]
CijC_{ij}, Ri=∑jCijR_i = \sum_j C_{ij}bigram counts and row totals (L0.0)
N=T−1N = T - 1predictions in a text of TT tokensint
uuone uniform draw in [0,1)[0, 1) from PCG32float

The loader. DataLoader(arrays, batch_size, shuffle, rng) holds named arrays that share their first dimension nn and yields dictionaries of rows. Without shuffle the order is 0,…,n−10, \dots, n - 1. With shuffle, when each epoch starts it builds the list 0,…,n−10, \dots, n - 1 and calls rng.shuffle on it: spec/pcg32.md’s Fisher-Yates (M06.3), going from the end and swapping ii with below(i + 1). Each epoch advances the generator, so epoch 2 has a new permutation and a seeded run repeats exactly. drop_last drops a final partial batch, so every step sees BB rows; len(loader) says how many batches an epoch has. Arrays of different lengths, or a shuffle without a generator, are errors: one pairs inputs with the wrong labels, the other hides an unseeded order.

The step. train_step(model, batch, loss_fn, opt, clip):

  1. opt.zero_grad(). Gradients accumulate across backward calls (L0.1), so without this step kk trains on the sum of the gradients of steps 1,…,k1, \dots, k.
  2. loss = loss_fn(model, batch): a one-element Tensor, or a pair (loss, {"acc": ...}) of the loss and extra metrics. More than one element is an error: backward() on a vector would need an upstream gradient nobody chose.
  3. If the loss is not finite, raise FloatingPointError before any update: one nan step turns every weight into nan and the run is lost; raising first lets the caller skip the batch.
  4. loss.backward().
  5. With clip, clip_grad_norm_(parameters, c) (M10.4) rescales all gradients together when ∥g∥>c\lVert g \rVert > c; report the norm before clipping, because that is the number that tells you training is unstable.
  6. opt.step(), then return {"loss": ..., "grad_norm": ...} plus the extra metrics.
  7. With ema (an M02.2 EMA), feed the loss to ema.update and report ema.value_debiased() as "loss_ema". One batch’s loss is noise; the debiased EMA is the curve a training log plots, and without the correction its first steps sit near zero. A non-finite loss raised in step 3, so it never reaches the average.

Evaluation. evaluate(model, loader, loss_fn) switches the model to eval mode (dropout off, L0.4), runs every batch under no_grad (no graph, no stored activations, L0.1), and averages the loss and metrics weighted by rows: a last batch of 2 rows counts as 2 rows, not as a full batch. It puts the model back in the mode it found it, in a finally, so an exception in a batch cannot leave a training run with dropout silently off.

The same model, trained. BigramLogits(V) is a Module with one parameter, weight of shape [V, V], starting at zeros (every next token equally likely, NLL ln⁡V\ln V). Its forward is a row lookup, F.embedding(weight, ids): the logits onehot(ids) @ W of L0.0 without the multiplications by zero, for ids of any shape (the trainer feeds [B, T] windows). Backward adds each position’s p−enextp - e_{\text{next}} into its row.

Why gradient descent finds the counts. The mean NLL of the table on a text is

L(W)=1N∑i,jCij (log⁡∑keWik−Wij),\mathcal{L}(W) = \frac{1}{N}\sum_{i,j} C_{ij}\,\big(\log\textstyle\sum_k e^{W_{ik}} - W_{ij}\big),

a sum over rows of convex functions (log-sum-exp is convex, −Wij-W_{ij} is linear). Setting the gradient of row ii to zero gives Ri softmax(Wi)=CiR_i\,\mathrm{softmax}(W_i) = C_i, so at the minimum softmax(Wi)j=Cij/Ri\mathrm{softmax}(W_i)_j = C_{ij}/R_i: the unsmoothed count model. Its NLL, −1N∑ijCijlog⁡(Cij/Ri)-\frac{1}{N}\sum_{ij} C_{ij}\log(C_{ij}/R_i), is the floor no bigram can beat on that text. Pass 1’s add-one model sits slightly above it; your trained table must get within 10−310^{-3} nats of the floor, never below it. Pairs that never occur push their logits toward −∞-\infty forever, so the loss approaches the floor without reaching it; AdamW (M10.3) with a large learning rate gets there in a few hundred full-batch steps.

Sampling like the engine. BigramLM.sample keeps its signature and switches its generator: one PCG32(seed) (stream 54) per call, one uniform() per token, and the Rust engine’s draw (L10.0): weights wj=ezj−max⁡zw_j = e^{z_j - \max z} of z=row/τz = \text{row}/\tau in float64, summed in id order, and the first jj whose running sum exceeds u⋅∑jwju \cdot \sum_j w_j. Same arithmetic, same order, same generator: generate --seed S and the engine’s completion with seed S give the same ids. to_lm() turns the trained module into a BigramLM holding a float32 copy of the table, the object your CLI saves as bigram.weight.

One training step. Linear(1, 1) with w=2w = 2, b=0b = 0; the batch x=[[1],[2]]x = [[1], [2]], y=[[3],[5]]y = [[3], [5]]; loss mse; SGD with η=0.1\eta = 0.1.

  1. Predictions wx+b=[2,4]wx + b = [2, 4]; residuals [2−3,4−5]=[−1,−1][2 - 3, 4 - 5] = [-1, -1]; loss ((−1)2+(−1)2)/2=1((-1)^2 + (-1)^2)/2 = 1.
  2. Gradients of 12∑rk2\frac{1}{2}\sum r_k^2: ∂/∂w=22∑rkxk=−1⋅1−1⋅2=−3\partial/\partial w = \frac{2}{2}\sum r_k x_k = -1 \cdot 1 - 1 \cdot 2 = -3, and ∂/∂b=∑rk=−2\partial/\partial b = \sum r_k = -2.
  3. Update: w=2−0.1⋅(−3)=2.3w = 2 - 0.1 \cdot (-3) = 2.3, b=0−0.1⋅(−2)=0.2b = 0 - 0.1 \cdot (-2) = 0.2.
  4. A second step must use only the new gradients: residuals [2.5−3,4.8−5]=[−0.5,−0.2][2.5 - 3, 4.8 - 5] = [-0.5, -0.2], so ∂/∂w=−0.5−0.4=−0.9\partial/\partial w = -0.5 - 0.4 = -0.9, ∂/∂b=−0.7\partial/\partial b = -0.7. A loop that forgot zero_grad would add the first step’s −3-3 and −2-2.
  5. With clip=1, the global norm is 32+22=13=3.606\sqrt{3^2 + 2^2} = \sqrt{13} = 3.606; it is reported, and both gradients are scaled by 1/3.6061/3.606 before the update.

The bigram, from zeros. BigramLogits(3) on the text abbacab (ids [0, 1, 1, 0, 2, 0, 1], N=6N = 6 predictions):

  1. Every row of zeros has softmax [1/3,1/3,1/3][1/3, 1/3, 1/3], so the NLL is ln⁡3=1.0986\ln 3 = 1.0986.
  2. Row a makes three predictions (the next tokens are b, c, b). Each contributes p−enextp - e_{\text{next}}, and the mean divides by 6: (3⋅[13,13,13]−[0,2,1])/6=[16,−16,0]\big(3 \cdot [\tfrac13, \tfrac13, \tfrac13] - [0, 2, 1]\big)/6 = [\tfrac16, -\tfrac16, 0]. Descent raises WabW_{ab} and lowers WaaW_{aa}; WacW_{ac} is already right on average.
  3. The floor: row a goes to bb twice and cc once (2/32/3, 1/31/3), row b to aa and bb once each, row c to aa always. The NLL is 16(2ln⁡32+ln⁡3+2ln⁡2+0)=3ln⁡36=0.5493\frac{1}{6}\big(2\ln\tfrac32 + \ln 3 + 2\ln 2 + 0\big) = \frac{3\ln 3}{6} = 0.5493, below the add-one model’s 0.83510.8351 of L0.0.

These are test_hand_example_train_step, test_zero_grad_every_step, test_clip_reports_the_norm_and_clips, and test_hand_example_bigram_gradient.

python/tinyllm/train/loop.py
class DataLoader:
def __init__(self, arrays: Mapping[str, NDArray], batch_size: int, shuffle: bool, rng, drop_last: bool = True)
def __iter__(self) -> Iterator[dict[str, NDArray]]; def __len__(self) -> int
def train_step(model, batch, loss_fn, opt, clip: Optional[float] = None,
ema: Optional[EMA] = None) -> dict[str, float] # EMA from tinyllm.num.ema (M02.2)
def evaluate(model, loader, loss_fn) -> dict[str, float] # mean loss, mean metrics, "n"
# python/tinyllm/lm/bigram.py (taken over from L0.0; BigramLM keeps its v0 API)
class BigramLogits(Module):
def __init__(self, vocab: int = 256); def forward(self, ids) -> Tensor; def to_lm(self) -> BigramLM
TestKINDChecksWhy it matters downstream
test_hand_example_train_stepunitsection 3: loss 1, gradients −3-3 and −2-2, then w=2.3w = 2.3, b=0.2b = 0.2you and the test agree on a step
test_zero_grad_every_stepunitthe second step uses −0.9-0.9 and −0.7-0.7 alonegradients accumulate (L0.1)
test_clip_reports_the_norm_and_clipsunitgrad_norm is 13\sqrt{13}, the update uses clipped gradientslong runs stay stable (M10.4)
test_nonfinite_loss_stops_before_the_updateboundarya nan loss raises and no weight movesone bad batch cannot poison the run
test_loss_fn_extras_and_shapeboundary(loss, metrics) pairs are reported; a vector loss is an erroraccuracy beside the loss
test_loss_ema_is_the_debiased_averageunitwith β=0.9\beta = 0.9 the losses 1 and 0.145 report loss_ema 1.0 and 0.55; a nan loss never reaches the EMAthe smoothed curve every later training log plots (M02.2)
test_dataloader_batches_in_orderunitin-order batches, drop_last, lenevaluation sees every row once
test_dataloader_shuffle_is_the_spec_permutationpropertythe order is the spec’s Fisher-Yates, new each epochthe Go and Rust ports shuffle the same way
test_dataloader_rejects_bad_inputboundaryunequal lengths, shuffle without a generatorsilent mispairing
test_evaluate_weights_by_rows_and_restores_modeunita 2-row batch counts 2 rows; eval mode, no_grad, mode restoredhonest validation numbers (L6.7)
test_evaluate_restores_mode_after_an_errorboundaryan exception still restores training modedropout cannot stay off
test_learns_xorlearninga 2-8-1 tanh MLP solves XOR and reaches the reference loss barevery piece of L0.1 to L0.5 on one path
test_learns_two_moonslearningminibatch AdamW with clipping reaches the bar and 95% accuracya real minibatch run
test_same_seed_bitwise_identicalpropertysame seed, identical weights and metrics; another seed differsresuming (L0.6) and regression tests
test_hand_example_bigram_gradientunitsection 3: NLL ln⁡3\ln 3, row a’s gradient [1/6,−1/6,0][1/6, -1/6, 0]the bigram trains the right rows
test_forward_is_a_row_gatherunitlogits are table rows for 1-D and [B, T] ids; bad ids failthe trainer feeds token windows
test_autograd_bigram_reaches_the_count_mlepropertywithin 10−310^{-3} nats of the floor, never belowMS-L0 step 2 in miniature
test_to_lm_serves_the_same_tableunitto_lm() gives the same logits through your C matmul, as a copythe engine serves what you trained
test_sample_draws_like_the_engineunitids equal the engine’s PCG32 draw at three temperaturesPython and Rust agree on a seed

The two learning tests compare against bars in course/fixtures/ref-thresholds.tsv: the reference’s mean plus three standard deviations over five seeds.

PitfallSymptomCaught by
1. no zero_grad at the start of the stepthe loss falls, then oscillates: every step adds all earlier gradientstest_zero_grad_every_step (mutant s01)
2. evaluation that averages batches, leaves dropout on, or restores the mode only on successvalidation loss biased by the last batch; noisy evaluations; training after a failed eval with dropout offtest_evaluate_weights_by_rows_and_restores_mode (mutants s10, s12, s14), test_evaluate_restores_mode_after_an_error (mutant s13)
3. shuffling by sorting uniforms, or one permutation for every epochruns are not reproducible across languages, or every epoch sees the same ordertest_dataloader_shuffle_is_the_spec_permutation (mutants s07, s08)
4. checking the loss for nan after the updatethe run is already lost when the error is raisedtest_nonfinite_loss_stops_before_the_update (mutant s05)
5. a new generator for every token, or a stream other than the engine’sPython and Rust disagree on the same seed; repeated draws of the same uutest_sample_draws_like_the_engine (mutants s20, s21)

| Forward | L2.2 | Registered call site uses this module. | | Forward | L3.6 | Registered call site uses this module. | | Forward | L5.5 | Registered call site uses this module. | | Forward | L6.1 | Registered call site uses this module. | | Forward | L6.7 | Registered call site uses this module. |

DirectionModuleHow it uses this
BackL0.1evaluate runs under no_grad; the loss is a Tensor
BackL0.2BigramLogits.forward is F.embedding
BackL0.3the tests (and your CLI) train on cross_entropy, mse, bce_with_logits
BackL0.4BigramLogits is a Module; the tests train Linear stacks
BackM06.3PCG32: the loader’s shuffle and the sampler’s draws
BackM10.2SGD steps the hand example
BackM10.3AdamW trains the MLPs and the bigram
BackM10.4clip_grad_norm_ in train_step
BackM02.2EMA.update and value_debiased give train_step’s loss_ema
BackPython referenceThis model has no native library boundary.
BackM03.1tl_matmul_f32 computes BigramLM.logits
ForwardL10.0your Rust engine serves the table to_lm() produces, unchanged (tl_arch = bigram)
ForwardL0.6takes over safetensors.py; its regression run of L0.0’s suite exercises your bigram.py too

Later passes (L2.2, L3.6, L6.1, the capstone) call train_step and evaluate as they are. If you skip this module, ol milestone MS-L0 cannot train either model.

Your pieceProduction equivalentWhat it addsWhere to look
DataLoadertorch.utils.data.DataLoaderworker processes, pinned memory, samplers, a generator argument for seeded shufflestorch/utils/data/dataloader.py
train_stepHF Trainer.training_step, Lightning LightningModulemixed precision, gradient accumulation, distributed reduction, callbackstransformers/trainer.py
non-finite guardtorch.cuda.amp.GradScalerskips steps with inf gradients and lowers the loss scale (L11.1)torch/amp/grad_scaler.py
PCG32 sampling shared with RustvLLM’s per-request seeded samplinga generator per request on GPU, the same draws across batch compositionsvLLM v1/sample/sampler.py