Merge pull request #207 from rohitg00/feat/phase-19-track-c2

feat(phase-19): track C train end-to-end lessons 46-49
This commit is contained in:
Rohit Ghumare
2026-05-27 19:35:27 +01:00
committed by GitHub
29 changed files with 3529 additions and 0 deletions
@@ -0,0 +1,353 @@
"""Gradient accumulation from scratch.
Effective batch size = micro batch size * accumulation steps. Accumulate
gradients across several forward and backward passes, only step the
optimizer after the last micro-batch. Tracks throughput against effective
batch size so the curve is visible, not folklore.
Run: python3 code/main.py
"""
from __future__ import annotations
import argparse
import json
import math
import time
from dataclasses import dataclass, field, asdict
from pathlib import Path
from typing import Callable, Iterable, List
import torch
from torch import nn
HERE = Path(__file__).parent
OUT_DIR = HERE.parent / "outputs"
LOG_PATH = OUT_DIR / "accum-curve.json"
@dataclass
class StepResult:
step: int
effective_batch: int
micro_batch: int
accum_steps: int
loss: float
grad_norm: float
samples_per_sec: float
wall_ms: float
sync_calls: int
@dataclass
class CurvePoint:
effective_batch: int
accum_steps: int
micro_batch: int
avg_loss: float
samples_per_sec: float
median_step_ms: float
sync_calls: int
steps: int
def seed_everything(seed: int) -> None:
torch.manual_seed(seed)
def synthetic_batch(batch_size: int, in_dim: int, out_dim: int, gen: torch.Generator) -> tuple[torch.Tensor, torch.Tensor]:
x = torch.randn(batch_size, in_dim, generator=gen)
target = torch.randint(low=0, high=out_dim, size=(batch_size,), generator=gen)
return x, target
def make_model(in_dim: int, hidden: int, out_dim: int) -> nn.Module:
return nn.Sequential(
nn.Linear(in_dim, hidden),
nn.GELU(),
nn.Linear(hidden, hidden),
nn.GELU(),
nn.Linear(hidden, out_dim),
)
def global_grad_norm(model: nn.Module) -> float:
total = 0.0
for p in model.parameters():
if p.grad is None:
continue
total += float(p.grad.detach().pow(2).sum().item())
return math.sqrt(total)
def zero_grads(model: nn.Module) -> None:
for p in model.parameters():
if p.grad is not None:
p.grad.detach_()
p.grad.zero_()
def loss_scaled_for_accum(logits: torch.Tensor, target: torch.Tensor, accum_steps: int, loss_fn: Callable[[torch.Tensor, torch.Tensor], torch.Tensor]) -> torch.Tensor:
raw = loss_fn(logits, target)
return raw / accum_steps
def train_one_optimizer_step(
model: nn.Module,
optimizer: torch.optim.Optimizer,
micro_batches: List[tuple[torch.Tensor, torch.Tensor]],
loss_fn: Callable[[torch.Tensor, torch.Tensor], torch.Tensor],
*,
no_sync_until_last: bool,
sync_counter: List[int],
) -> tuple[float, float]:
"""Run accum_steps micro batches, accumulate grads, step once.
Returns (total_unscaled_loss, grad_norm).
"""
accum_steps = len(micro_batches)
zero_grads(model)
total = 0.0
for i, (x, y) in enumerate(micro_batches):
is_last = i == accum_steps - 1
if no_sync_until_last and not is_last:
with no_sync_context(model):
logits = model(x)
loss = loss_scaled_for_accum(logits, y, accum_steps, loss_fn)
loss.backward()
else:
logits = model(x)
loss = loss_scaled_for_accum(logits, y, accum_steps, loss_fn)
loss.backward()
sync_counter[0] += 1
total += float(loss.detach().item()) * accum_steps
grad_norm = global_grad_norm(model)
optimizer.step()
return total / accum_steps, grad_norm
class _NoSyncCtx:
def __init__(self, model: nn.Module):
self.model = model
def __enter__(self):
return self
def __exit__(self, exc_type, exc, tb):
return False
def no_sync_context(model: nn.Module):
"""Stand-in for DDP no_sync.
In DDP this skips the all-reduce on the trailing backward. In this
single-process demo there is no collective to skip, but we still
surface the call site so the pattern reads the same on a real cluster.
"""
return _NoSyncCtx(model)
def run_config(
effective_batch: int,
accum_steps: int,
*,
in_dim: int,
hidden: int,
out_dim: int,
num_steps: int,
lr: float,
seed: int,
) -> CurvePoint:
assert effective_batch % accum_steps == 0, "effective_batch must divide by accum_steps"
micro_batch = effective_batch // accum_steps
seed_everything(seed)
gen = torch.Generator()
gen.manual_seed(seed)
model = make_model(in_dim, hidden, out_dim)
optimizer = torch.optim.SGD(model.parameters(), lr=lr)
loss_fn = nn.CrossEntropyLoss()
losses: List[float] = []
step_times_ms: List[float] = []
sync_counter = [0]
total_samples = 0
wall_start = time.perf_counter()
for step in range(num_steps):
t0 = time.perf_counter()
micro_batches = [synthetic_batch(micro_batch, in_dim, out_dim, gen) for _ in range(accum_steps)]
avg_loss, _grad_norm = train_one_optimizer_step(
model,
optimizer,
micro_batches,
loss_fn,
no_sync_until_last=True,
sync_counter=sync_counter,
)
wall_ms = (time.perf_counter() - t0) * 1000.0
losses.append(avg_loss)
step_times_ms.append(wall_ms)
total_samples += effective_batch
total_wall = time.perf_counter() - wall_start
sps = total_samples / max(total_wall, 1e-6)
step_times_ms.sort()
median_ms = step_times_ms[len(step_times_ms) // 2]
avg_loss = sum(losses) / len(losses)
return CurvePoint(
effective_batch=effective_batch,
accum_steps=accum_steps,
micro_batch=micro_batch,
avg_loss=avg_loss,
samples_per_sec=sps,
median_step_ms=median_ms,
sync_calls=sync_counter[0],
steps=num_steps,
)
def sweep_effective_batches(
*,
micro_batch: int,
accum_grid: Iterable[int],
in_dim: int = 64,
hidden: int = 128,
out_dim: int = 16,
num_steps: int = 25,
lr: float = 0.05,
seed: int = 0,
) -> List[CurvePoint]:
points: List[CurvePoint] = []
for accum in accum_grid:
eff = micro_batch * accum
pt = run_config(
effective_batch=eff,
accum_steps=accum,
in_dim=in_dim,
hidden=hidden,
out_dim=out_dim,
num_steps=num_steps,
lr=lr,
seed=seed,
)
points.append(pt)
return points
def equivalence_check(
*,
in_dim: int = 32,
hidden: int = 48,
out_dim: int = 8,
big_batch: int = 16,
accum_steps: int = 4,
lr: float = 0.1,
seed: int = 7,
) -> dict:
"""One full batch step vs accum_steps micro-batches must match.
Scaled loss is `raw / accum_steps`; the accumulated gradient equals the
full batch gradient up to floating point noise.
"""
assert big_batch % accum_steps == 0
micro = big_batch // accum_steps
seed_everything(seed)
gen_a = torch.Generator(); gen_a.manual_seed(seed)
x, y = synthetic_batch(big_batch, in_dim, out_dim, gen_a)
seed_everything(seed)
model_full = make_model(in_dim, hidden, out_dim)
opt_full = torch.optim.SGD(model_full.parameters(), lr=lr)
loss_fn = nn.CrossEntropyLoss()
zero_grads(model_full)
out = model_full(x)
loss_full = loss_fn(out, y)
loss_full.backward()
full_params_before = [p.detach().clone() for p in model_full.parameters()]
full_grads = [p.grad.detach().clone() for p in model_full.parameters()]
opt_full.step()
full_params_after = [p.detach().clone() for p in model_full.parameters()]
seed_everything(seed)
model_accum = make_model(in_dim, hidden, out_dim)
opt_accum = torch.optim.SGD(model_accum.parameters(), lr=lr)
zero_grads(model_accum)
chunks_x = list(torch.split(x, micro, dim=0))
chunks_y = list(torch.split(y, micro, dim=0))
for cx, cy in zip(chunks_x, chunks_y):
scaled = loss_fn(model_accum(cx), cy) / accum_steps
scaled.backward()
accum_grads = [p.grad.detach().clone() for p in model_accum.parameters()]
accum_params_before = [p.detach().clone() for p in model_accum.parameters()]
opt_accum.step()
accum_params_after = [p.detach().clone() for p in model_accum.parameters()]
grad_diffs = [
float((a - b).abs().max().item())
for a, b in zip(full_grads, accum_grads)
]
param_diffs = [
float((a - b).abs().max().item())
for a, b in zip(full_params_after, accum_params_after)
]
return {
"max_grad_diff": max(grad_diffs),
"max_param_diff": max(param_diffs),
"params_init_match": all(
torch.equal(a, b) for a, b in zip(full_params_before, accum_params_before)
),
}
def write_curve(points: List[CurvePoint], path: Path) -> None:
path.parent.mkdir(parents=True, exist_ok=True)
payload = {
"schema": "accum-curve.v1",
"points": [asdict(p) for p in points],
}
path.write_text(json.dumps(payload, indent=2) + "\n")
def parse_args() -> argparse.Namespace:
p = argparse.ArgumentParser()
p.add_argument("--micro-batch", type=int, default=4)
p.add_argument("--accum-grid", type=str, default="1,2,4,8,16")
p.add_argument("--num-steps", type=int, default=25)
p.add_argument("--seed", type=int, default=0)
p.add_argument("--lr", type=float, default=0.05)
p.add_argument("--no-write", action="store_true")
return p.parse_args()
def main() -> int:
args = parse_args()
accum_grid = [int(s) for s in args.accum_grid.split(",") if s.strip()]
print("equivalence check (full batch vs accumulated)")
eq = equivalence_check()
print(json.dumps(eq, indent=2))
assert eq["max_grad_diff"] < 1e-4, f"gradients diverge: {eq['max_grad_diff']}"
assert eq["max_param_diff"] < 1e-4, f"params diverge: {eq['max_param_diff']}"
print("equivalence holds. running sweep...")
points = sweep_effective_batches(
micro_batch=args.micro_batch,
accum_grid=accum_grid,
num_steps=args.num_steps,
lr=args.lr,
seed=args.seed,
)
header = f"{'eff_batch':>10} {'accum':>5} {'micro':>5} {'sps':>10} {'median_ms':>10} {'syncs':>6} {'loss':>8}"
print(header)
for p in points:
print(
f"{p.effective_batch:>10} {p.accum_steps:>5} {p.micro_batch:>5} "
f"{p.samples_per_sec:>10.1f} {p.median_step_ms:>10.2f} {p.sync_calls:>6} {p.avg_loss:>8.4f}"
)
if not args.no_write:
write_curve(points, LOG_PATH)
print(f"wrote {LOG_PATH}")
return 0
if __name__ == "__main__":
raise SystemExit(main())
@@ -0,0 +1,104 @@
"""Tests for gradient accumulation core paths."""
from __future__ import annotations
import json
import sys
import unittest
from pathlib import Path
HERE = Path(__file__).parent
sys.path.insert(0, str(HERE))
import main as accum
class EquivalenceTests(unittest.TestCase):
def test_full_batch_matches_accumulated(self):
diffs = accum.equivalence_check()
self.assertLess(diffs["max_grad_diff"], 1e-4)
self.assertLess(diffs["max_param_diff"], 1e-4)
self.assertTrue(diffs["params_init_match"])
def test_loss_scaled_for_accum_divides(self):
import torch
from torch import nn
logits = torch.tensor([[2.0, 0.5, -1.0]])
target = torch.tensor([0])
fn = nn.CrossEntropyLoss()
raw = fn(logits, target)
scaled = accum.loss_scaled_for_accum(logits, target, 4, fn)
self.assertAlmostEqual(float(scaled.item()), float(raw.item()) / 4.0, places=6)
class SweepTests(unittest.TestCase):
def test_sweep_returns_one_point_per_accum(self):
points = accum.sweep_effective_batches(
micro_batch=2,
accum_grid=[1, 2, 4],
in_dim=16,
hidden=24,
out_dim=4,
num_steps=5,
lr=0.05,
)
self.assertEqual(len(points), 3)
self.assertEqual([p.accum_steps for p in points], [1, 2, 4])
self.assertEqual([p.effective_batch for p in points], [2, 4, 8])
for p in points:
self.assertGreater(p.samples_per_sec, 0.0)
self.assertGreater(p.steps, 0)
self.assertGreater(p.median_step_ms, 0.0)
def test_sync_calls_equal_step_count(self):
points = accum.sweep_effective_batches(
micro_batch=2,
accum_grid=[1, 4],
in_dim=16,
hidden=24,
out_dim=4,
num_steps=7,
lr=0.05,
)
for p in points:
self.assertEqual(p.sync_calls, p.steps)
class CurveOutputTests(unittest.TestCase):
def test_write_curve_round_trip(self):
import tempfile
points = accum.sweep_effective_batches(
micro_batch=2,
accum_grid=[1, 2],
in_dim=8,
hidden=12,
out_dim=3,
num_steps=3,
lr=0.05,
)
with tempfile.TemporaryDirectory() as tmp:
path = Path(tmp) / "curve.json"
accum.write_curve(points, path)
payload = json.loads(path.read_text())
self.assertEqual(payload["schema"], "accum-curve.v1")
self.assertEqual(len(payload["points"]), 2)
self.assertIn("samples_per_sec", payload["points"][0])
def test_effective_batch_must_divide_accum(self):
with self.assertRaises(AssertionError):
accum.run_config(
effective_batch=5,
accum_steps=2,
in_dim=8,
hidden=12,
out_dim=3,
num_steps=2,
lr=0.05,
seed=0,
)
if __name__ == "__main__":
unittest.main()
@@ -0,0 +1,147 @@
# Gradient Accumulation
> Train at an effective batch you cannot afford, one micro-batch at a time. Scale the loss, hold the optimizer step, and let the gradients pile up.
**Type:** Build
**Languages:** Python
**Prerequisites:** Phase 19 lessons 42 to 45
**Time:** ~90 minutes
## Learning Objectives
- Derive the effective batch identity: `effective_batch = micro_batch * accum_steps`.
- Implement loss-per-micro-batch scaling so the accumulated gradient matches a single full-batch backward.
- Skip optimizer synchronization until the last micro-batch (sync-on-last-step).
- Read a throughput against effective batch curve and explain the diminishing return.
## The Problem
You want to train at an effective batch of 512 because the loss curve is smoother and the optimizer step makes more sense at that scale. The accelerator on the desk holds 32 examples before it runs out of memory. Doubling the batch is not an option. Halving the model is not an option. The trick the field reached for in 2017 and never stopped using is to run 16 backward passes, let the gradients accumulate inside the parameter buffers, and only step the optimizer when the count reaches the target.
The risk is that the loss is no longer the same number it was at the bigger batch. The cross entropy of 16 mini-batches summed naively is 16 times the loss of one full batch. Without scaling, the gradient direction is correct but the magnitude is wrong, and the optimizer step is 16 times too big. The fix is one division. The fix is also easy to forget.
## The Concept
```mermaid
flowchart LR
start[start] --> zero[zero grads]
zero --> mb1[micro batch 1: forward + scaled backward]
mb1 --> mb2[micro batch 2: forward + scaled backward]
mb2 --> dots[...]
dots --> mbN[micro batch N: forward + scaled backward + sync]
mbN --> step[optimizer step]
step --> next[next effective step]
```
The contract is short:
- Loss for each micro-batch is divided by `accum_steps` before `backward()`. PyTorch sums gradients into `param.grad` by default; the division pushes the running sum back into the right scale.
- The optimizer step fires once per effective batch, after the last micro-batch's backward. Stepping mid-accumulation skews every parameter the rest of the run depends on.
- The optimizer's state (momentum buffers, Adam moments) advances once per effective step, not once per micro-batch. The exponential moving averages would otherwise see the wrong frequency and burn through the schedule.
- On a single device this is bookkeeping. On a multi-rank cluster the same pattern wraps the non-final micro-batches in a `no_sync` context that skips the gradient all-reduce; the last micro-batch reduces the full accumulated gradient in one pass instead of paying the network cost N times.
### The equivalence proof in code
```python
loss = criterion(model(x_full), y_full)
loss.backward()
opt.step()
```
is equivalent to
```python
for x, y in chunks(x_full, y_full, n):
scaled = criterion(model(x), y) / n
scaled.backward()
opt.step()
```
up to floating point summation order. The accumulated gradient buffer at the end of the loop is the same tensor that a single full-batch backward would produce. The lesson code asserts this with a max-abs difference under 1e-4 in `equivalence_check`.
### Where the cost goes
Each micro-batch costs one forward and one backward. With accumulation you trade memory for time. The throughput curve in `outputs/accum-curve.json` shows what happens as the effective batch grows at fixed micro-batch:
```mermaid
flowchart TD
micro[fixed micro batch] --> small[small accum: low loss noise budget, high stepper churn]
micro --> large[large accum: smooth loss, optimizer step rare]
small --> sps1[samples per second saturates at hardware limit]
large --> sps2[samples per second still hits hardware limit]
sps1 --> note[total samples per optimizer step scales linearly with accum]
sps2 --> note
```
There is no free lunch. Doubling `accum_steps` doubles the wall time per optimizer step. What changes is the variance of the gradient estimate: at the same wall budget you have made fewer optimizer steps but each one was averaged over more samples. The literature treats large batch and small batch as different optimization problems; the lesson here is mechanical, not statistical.
## Build It
`code/main.py` is the runnable artifact. It does three things.
### Step 1: equivalence check
`equivalence_check()` builds two copies of the same network with the same seed. One sees a 16-sample batch in one forward pass. The other sees four 4-sample chunks with the loss divided by four. The function compares the gradient buffers before the optimizer step and the parameters after. The assertion is `max_abs_diff < 1e-4`.
### Step 2: sync-on-last-step pattern
`train_one_optimizer_step` walks micro-batches. For every micro-batch except the last it enters `no_sync_context(model)`. On a single process the context is a no-op; on DDP this is where the gradient all-reduce is skipped. The bookkeeping is the same regardless. A `sync_counter` records how many times we left the no_sync scope; for N micro-batches the count is one per effective step, not N.
### Step 3: the throughput curve
`sweep_effective_batches` runs the same model with a fixed micro-batch and a list of accumulation steps. For each setting it logs:
- `samples_per_sec`: total samples seen divided by wall time
- `median_step_ms`: 50th percentile per effective step
- `sync_calls`: collective points exercised
- `avg_loss`: average across the sweep's optimizer steps
The output lands in `outputs/accum-curve.json` and is reusable from a notebook.
Run it:
```bash
python3 code/main.py
```
The script prints the equivalence diff, then the sweep table, then the JSON path. Exit code zero.
## Use It
In production training, gradient accumulation lives behind one knob. PyTorch's pattern is `accumulation_steps = effective_batch // (micro_batch * world_size)`. Frameworks that you are not allowed to use here wrap the same loop, but the steps are the same: scale the loss, skip sync on non-final micros, accumulate, step once.
Three patterns in the wild:
- The micro-batch size is chosen to saturate device memory. Anything smaller wastes accelerator cycles. Anything larger crashes.
- The effective batch is chosen from a learning rate schedule. Large effective batches need scaled learning rates and warmup; this is the linear scaling rule talked about since 2017.
- The accumulation count is the bridge between the two and the only knob you are free to tune at runtime without rewriting the data loader.
## Ship It
`outputs/skill-gradient-accumulation.md` captures the recipe so a peer can drop it into a new repo: scale loss by `accum_steps`, skip optimizer sync on non-final micros, step the optimizer once per effective batch, log throughput against effective batch as JSON so the trade is visible.
## Exercises
1. Re-run the sweep with `--num-steps 100` and plot samples per second against effective batch. Where does the curve flatten?
2. Add a wrong scaling variant (no division) and show the parameter diff at step 1 against the reference.
3. Swap SGD for AdamW and confirm the optimizer state advances once per effective step, not once per micro-batch.
4. Introduce a real `DistributedDataParallel` wrapper and route the `no_sync_context` to its method. Confirm sync_calls drops by N-1 per effective batch.
5. Modify the equivalence check to compare two different micro splits (2 by 8 vs 4 by 4) and explain any tolerance you need to relax.
## Key Terms
| Term | What people say | What it actually means |
|------|-----------------|------------------------|
| Micro batch | The batch you forward | The slice that fits in memory in a single forward pass |
| Accum steps | Backward passes per step | Number of backwards summed before one optimizer step |
| Effective batch | The batch | Micro batch times accum steps times data parallel world size |
| Loss scaling | Divide by N | Per-micro-batch division so summed gradients match full batch |
| Sync on last | Skip the rest | Only run the gradient collective on the last backward in the window |
## Further Reading
- PyTorch docs on `DistributedDataParallel.no_sync` for the production version of the sync-on-last-step trick.
- Goyal et al., 2017, on linear scaling for large batch training, the canonical reason to care about effective batch.
- PyTorch issue tracker on gradient accumulation interactions with mixed precision unscaling.
- Phase 19 lessons 42 to 45 cover the model, data loader, optimizer, and trainer scaffolding this lesson assumes.
- Phase 19 lesson 47 covers checkpoint and resume so a long accumulation run survives a wallclock cap.
@@ -0,0 +1,35 @@
{
"schema": "accum-curve.v1",
"points": [
{
"effective_batch": 4,
"accum_steps": 1,
"micro_batch": 4,
"avg_loss": 2.7746317982673645,
"samples_per_sec": 10589.599020053975,
"median_step_ms": 0.36970898509025574,
"sync_calls": 8,
"steps": 8
},
{
"effective_batch": 8,
"accum_steps": 2,
"micro_batch": 4,
"avg_loss": 2.778739705681801,
"samples_per_sec": 12129.73295382913,
"median_step_ms": 0.6751659093424678,
"sync_calls": 8,
"steps": 8
},
{
"effective_batch": 16,
"accum_steps": 4,
"micro_batch": 4,
"avg_loss": 2.7735252678394318,
"samples_per_sec": 12667.36269537429,
"median_step_ms": 1.2691670563071966,
"sync_calls": 8,
"steps": 8
}
]
}
@@ -0,0 +1,33 @@
---
name: gradient-accumulation
description: Train at an effective batch larger than device memory by scaling micro-batch losses and stepping the optimizer once per window.
version: 1.0.0
phase: 19
lesson: 46
tags: [training, batch-size, distributed, scaling]
---
## When to use
Effective batch is the lever that smooths the gradient and matches the learning rate schedule. When you cannot afford it in a single forward pass, this is the recipe.
## Recipe
1. Pick `micro_batch` as the largest size that fits in memory and saturates the accelerator.
2. Pick `effective_batch` from the learning rate schedule.
3. Set `accum_steps = effective_batch // (micro_batch * world_size)` and assert it divides evenly.
4. Per micro batch: `loss = criterion(model(x), y) / accum_steps; loss.backward()`.
5. On non-final micros, enter `model.no_sync()` to skip the gradient all-reduce in DDP.
6. After the last micro batch, run `optimizer.step()` once. Zero gradients before the next window.
7. The optimizer state advances once per effective batch; the learning rate schedule ticks once per effective batch.
## Logging
Emit a small JSON record per effective step with `samples_per_sec`, `median_step_ms`, `sync_calls`, `accum_steps`, `effective_batch`. Without this the cost trade is invisible.
## Failure modes
- Forgetting the `/ accum_steps` scaling: gradients explode by N.
- Stepping mid-window: parameters drift.
- Sync on every micro batch: network bound for no statistical gain.
- Mixing this with mixed precision unscaling: scale the unscaled loss only.
@@ -0,0 +1,78 @@
{
"lesson": "46-gradient-accumulation",
"title": "Gradient Accumulation",
"questions": [
{
"stage": "pre",
"question": "What is the effective batch identity that gradient accumulation expresses?",
"options": [
"effective_batch = micro_batch / accum_steps",
"effective_batch = micro_batch * accum_steps",
"effective_batch = micro_batch + accum_steps",
"effective_batch = world_size only"
],
"correct": 1,
"explanation": "Micro batch is what fits in memory; accumulation steps are how many forward + backward passes pile gradients into the same buffer before one optimizer step."
},
{
"stage": "pre",
"question": "Why must the loss be divided by accum_steps before backward on each micro batch?",
"options": [
"It makes the loss curve prettier",
"PyTorch sums gradients into param.grad by default, so without the division the accumulated gradient is N times too large and the optimizer step is N times too aggressive",
"It avoids NaNs",
"It is required for mixed precision only"
],
"correct": 1,
"explanation": "The accumulation buffer is a sum. Per-micro-batch scaling by 1/N pushes that sum back into the same scale a single full-batch backward would produce."
},
{
"stage": "check",
"question": "When does the optimizer step run inside the accumulation loop?",
"options": [
"After every micro batch",
"Only after the last micro batch in the accumulation window",
"Twice per window for safety",
"Never; the optimizer is replaced"
],
"correct": 1,
"explanation": "Stepping mid-accumulation contaminates every parameter the rest of the run depends on. The optimizer step fires once per effective batch."
},
{
"stage": "check",
"question": "What does the no_sync context wrap on a real multi-GPU run?",
"options": [
"The forward pass",
"The optimizer step",
"The non-final micro batches, so the gradient all-reduce only fires after the last backward instead of N times",
"The data loader"
],
"correct": 2,
"explanation": "DDP's no_sync skips the gradient collective. Wrapping all but the last micro batch turns N collectives per effective step into one."
},
{
"stage": "check",
"question": "What does the equivalence_check function in main.py assert?",
"options": [
"That the loss converges",
"That a single full-batch backward and an accum_steps chunked backward with loss / N produce the same gradient buffer and the same post-step parameters up to a small tolerance",
"That throughput is constant",
"That gradients are zero"
],
"correct": 1,
"explanation": "The assertion bound is max-abs-diff under 1e-4. Without the loss scaling the diff blows up; with the scaling it sits at floating point noise."
},
{
"stage": "post",
"question": "Reading the throughput against effective batch curve, what stays roughly constant and what scales?",
"options": [
"Samples per second drops with larger effective batch",
"Samples per second saturates near the hardware limit while wall time per optimizer step scales linearly with accum_steps; optimizer steps per second is what falls",
"Both metrics scale linearly with accum",
"Throughput goes up because the model gets faster"
],
"correct": 1,
"explanation": "Each micro batch costs the same forward + backward. Accumulation buys statistical smoothing per optimizer step, not raw throughput; the curve is a useful reality check against folklore."
}
]
}
@@ -0,0 +1,498 @@
"""Checkpoint save and resume from scratch.
Full checkpoint dict: model state, optimizer state, scheduler state,
loss history, current step, RNG state (python random, numpy, torch CPU,
torch CUDA if present). Atomic save by writing to a temp file and then
renaming. Sharded save splits the model state by parameter group so a
single shard is small enough to load on demand. Resume continues mid
epoch with deterministic loss within tolerance.
Run: python3 code/main.py
"""
from __future__ import annotations
import argparse
import hashlib
import json
import os
import random
import tempfile
import time
from dataclasses import dataclass, field
from pathlib import Path
from typing import Any, Dict, List, Optional
import numpy as np
import torch
from torch import nn
HERE = Path(__file__).parent
OUT_DIR = HERE.parent / "outputs"
CHECKPOINT_SCHEMA = "ckpt.v1"
SHARD_SCHEMA = "ckpt-shard.v1"
@dataclass
class TrainState:
step: int
epoch: int
batch_in_epoch: int
losses: List[float] = field(default_factory=list)
def seed_everything(seed: int) -> None:
random.seed(seed)
np.random.seed(seed)
torch.manual_seed(seed)
if torch.cuda.is_available():
torch.cuda.manual_seed_all(seed)
def make_model(in_dim: int, hidden: int, out_dim: int) -> nn.Module:
return nn.Sequential(
nn.Linear(in_dim, hidden),
nn.GELU(),
nn.Linear(hidden, hidden),
nn.GELU(),
nn.Linear(hidden, out_dim),
)
def make_optimizer_and_scheduler(model: nn.Module, lr: float, total_steps: int):
opt = torch.optim.AdamW(model.parameters(), lr=lr)
sched = torch.optim.lr_scheduler.CosineAnnealingLR(opt, T_max=total_steps)
return opt, sched
def synthetic_loader(batch_size: int, num_batches: int, in_dim: int, out_dim: int, gen: torch.Generator):
for _ in range(num_batches):
x = torch.randn(batch_size, in_dim, generator=gen)
y = torch.randint(low=0, high=out_dim, size=(batch_size,), generator=gen)
yield x, y
def capture_rng_state() -> Dict[str, Any]:
state: Dict[str, Any] = {
"python": random.getstate(),
"numpy": np.random.get_state(),
"torch_cpu": torch.get_rng_state().tolist(),
}
if torch.cuda.is_available():
state["torch_cuda"] = [s.tolist() for s in torch.cuda.get_rng_state_all()]
return state
def restore_rng_state(state: Dict[str, Any]) -> None:
py = state.get("python")
if py is not None:
random.setstate(tuple_from_nested(py))
np_state = state.get("numpy")
if np_state is not None:
np.random.set_state(tuple_from_nested(np_state))
cpu = state.get("torch_cpu")
if cpu is not None:
torch.set_rng_state(torch.tensor(cpu, dtype=torch.uint8))
cuda = state.get("torch_cuda")
if cuda is not None and torch.cuda.is_available():
torch.cuda.set_rng_state_all([torch.tensor(s, dtype=torch.uint8) for s in cuda])
def tuple_from_nested(obj):
if isinstance(obj, list):
return tuple(tuple_from_nested(x) for x in obj)
return obj
def atomic_save(payload: Dict[str, Any], path: Path) -> Path:
path.parent.mkdir(parents=True, exist_ok=True)
tmp = tempfile.NamedTemporaryFile(
delete=False,
dir=str(path.parent),
prefix=path.name + ".",
suffix=".tmp",
)
tmp_path = Path(tmp.name)
tmp.close()
try:
torch.save(payload, tmp_path)
os.replace(tmp_path, path)
finally:
if tmp_path.exists():
try:
tmp_path.unlink()
except FileNotFoundError:
pass
return path
def atomic_write_json(payload: Dict[str, Any], path: Path) -> Path:
path.parent.mkdir(parents=True, exist_ok=True)
tmp = tempfile.NamedTemporaryFile(
mode="w",
delete=False,
dir=str(path.parent),
prefix=path.name + ".",
suffix=".tmp",
encoding="utf-8",
)
tmp_path = Path(tmp.name)
try:
json.dump(payload, tmp, indent=2)
tmp.write("\n")
tmp.close()
os.replace(tmp_path, path)
finally:
if tmp_path.exists():
try:
tmp_path.unlink()
except FileNotFoundError:
pass
return path
def file_sha256(path: Path) -> str:
h = hashlib.sha256()
with path.open("rb") as f:
for chunk in iter(lambda: f.read(1 << 16), b""):
h.update(chunk)
return h.hexdigest()
def save_checkpoint(
model: nn.Module,
optimizer: torch.optim.Optimizer,
scheduler: torch.optim.lr_scheduler._LRScheduler,
state: TrainState,
out_path: Path,
*,
schema: str = CHECKPOINT_SCHEMA,
extras: Optional[Dict[str, Any]] = None,
) -> Dict[str, Any]:
payload: Dict[str, Any] = {
"schema": schema,
"model": model.state_dict(),
"optimizer": optimizer.state_dict(),
"scheduler": scheduler.state_dict(),
"state": {
"step": state.step,
"epoch": state.epoch,
"batch_in_epoch": state.batch_in_epoch,
"losses": list(state.losses),
},
"rng": capture_rng_state(),
"wall_saved_at": time.time(),
}
if extras:
payload["extras"] = extras
atomic_save(payload, out_path)
return payload
def load_checkpoint(
path: Path,
model: nn.Module,
optimizer: torch.optim.Optimizer,
scheduler: torch.optim.lr_scheduler._LRScheduler,
) -> TrainState:
payload = torch.load(path, map_location="cpu", weights_only=False)
assert payload["schema"].startswith("ckpt"), f"unknown schema {payload['schema']}"
model.load_state_dict(payload["model"])
optimizer.load_state_dict(payload["optimizer"])
scheduler.load_state_dict(payload["scheduler"])
restore_rng_state(payload["rng"])
s = payload["state"]
return TrainState(
step=int(s["step"]),
epoch=int(s["epoch"]),
batch_in_epoch=int(s["batch_in_epoch"]),
losses=list(s["losses"]),
)
def shard_keys_by_prefix(state_dict: Dict[str, torch.Tensor], num_shards: int) -> Dict[int, List[str]]:
"""Round-robin allocate parameter keys across shards.
Production sharding usually goes by parameter group or by layer. The
round robin keeps the shards roughly the same size for the demo and
keeps the index easy to read.
"""
if num_shards < 1:
raise ValueError("num_shards must be >= 1")
keys = sorted(state_dict.keys())
shards: Dict[int, List[str]] = {i: [] for i in range(num_shards)}
for i, k in enumerate(keys):
shards[i % num_shards].append(k)
return shards
def save_sharded_checkpoint(
model: nn.Module,
optimizer: torch.optim.Optimizer,
scheduler: torch.optim.lr_scheduler._LRScheduler,
state: TrainState,
out_dir: Path,
*,
num_shards: int,
extras: Optional[Dict[str, Any]] = None,
) -> Dict[str, Any]:
out_dir.mkdir(parents=True, exist_ok=True)
model_sd = model.state_dict()
layout = shard_keys_by_prefix(model_sd, num_shards)
shard_files: List[Dict[str, Any]] = []
for shard_idx in range(num_shards):
keys = layout[shard_idx]
tensors = {k: model_sd[k] for k in keys}
shard_path = out_dir / f"model.shard-{shard_idx:03d}.pt"
atomic_save({"schema": SHARD_SCHEMA, "tensors": tensors, "keys": keys}, shard_path)
shard_files.append({
"shard": shard_idx,
"path": shard_path.name,
"num_params": len(keys),
"sha256": file_sha256(shard_path),
})
meta_path = out_dir / "meta.pt"
meta_payload = {
"schema": CHECKPOINT_SCHEMA + "-sharded",
"optimizer": optimizer.state_dict(),
"scheduler": scheduler.state_dict(),
"state": {
"step": state.step,
"epoch": state.epoch,
"batch_in_epoch": state.batch_in_epoch,
"losses": list(state.losses),
},
"rng": capture_rng_state(),
"wall_saved_at": time.time(),
"shards": shard_files,
"extras": extras or {},
}
atomic_save(meta_payload, meta_path)
index_payload = {
"schema": CHECKPOINT_SCHEMA + "-index",
"num_shards": num_shards,
"shards": shard_files,
"meta_sha256": file_sha256(meta_path),
"saved_at": meta_payload["wall_saved_at"],
"step": state.step,
}
atomic_write_json(index_payload, out_dir / "index.json")
return meta_payload
def load_sharded_checkpoint(
ckpt_dir: Path,
model: nn.Module,
optimizer: torch.optim.Optimizer,
scheduler: torch.optim.lr_scheduler._LRScheduler,
) -> TrainState:
index = json.loads((ckpt_dir / "index.json").read_text())
expected_sha = index["meta_sha256"]
meta_path = ckpt_dir / "meta.pt"
actual_sha = file_sha256(meta_path)
assert actual_sha == expected_sha, f"meta sha mismatch: {actual_sha} != {expected_sha}"
meta = torch.load(meta_path, map_location="cpu", weights_only=False)
merged: Dict[str, torch.Tensor] = {}
for shard in meta["shards"]:
shard_path = ckpt_dir / shard["path"]
actual = file_sha256(shard_path)
assert actual == shard["sha256"], f"shard sha mismatch: {shard['path']}"
body = torch.load(shard_path, map_location="cpu", weights_only=False)
assert body["schema"] == SHARD_SCHEMA
merged.update(body["tensors"])
model.load_state_dict(merged)
optimizer.load_state_dict(meta["optimizer"])
scheduler.load_state_dict(meta["scheduler"])
restore_rng_state(meta["rng"])
s = meta["state"]
return TrainState(
step=int(s["step"]),
epoch=int(s["epoch"]),
batch_in_epoch=int(s["batch_in_epoch"]),
losses=list(s["losses"]),
)
def step_one(
model: nn.Module,
optimizer: torch.optim.Optimizer,
scheduler: torch.optim.lr_scheduler._LRScheduler,
x: torch.Tensor,
y: torch.Tensor,
loss_fn,
) -> float:
optimizer.zero_grad()
loss = loss_fn(model(x), y)
loss.backward()
optimizer.step()
scheduler.step()
return float(loss.detach().item())
def train_until(
model: nn.Module,
optimizer: torch.optim.Optimizer,
scheduler: torch.optim.lr_scheduler._LRScheduler,
loss_fn,
state: TrainState,
*,
stop_step: int,
batches_per_epoch: int,
batch_size: int,
in_dim: int,
out_dim: int,
) -> TrainState:
while state.step < stop_step:
gen = torch.Generator()
gen.manual_seed(12345 + state.epoch)
for _ in range(state.batch_in_epoch):
torch.randn(batch_size, in_dim, generator=gen)
torch.randint(low=0, high=out_dim, size=(batch_size,), generator=gen)
while state.batch_in_epoch < batches_per_epoch and state.step < stop_step:
x = torch.randn(batch_size, in_dim, generator=gen)
y = torch.randint(low=0, high=out_dim, size=(batch_size,), generator=gen)
loss = step_one(model, optimizer, scheduler, x, y, loss_fn)
state.losses.append(loss)
state.step += 1
state.batch_in_epoch += 1
if state.batch_in_epoch >= batches_per_epoch:
state.epoch += 1
state.batch_in_epoch = 0
return state
def run_resume_demo(
*,
total_steps: int = 30,
interrupt_at: int = 12,
in_dim: int = 16,
hidden: int = 24,
out_dim: int = 4,
batch_size: int = 4,
batches_per_epoch: int = 5,
seed: int = 11,
ckpt_dir: Path,
sharded: bool = False,
num_shards: int = 3,
) -> Dict[str, Any]:
loss_fn = nn.CrossEntropyLoss()
seed_everything(seed)
m1 = make_model(in_dim, hidden, out_dim)
o1, s1 = make_optimizer_and_scheduler(m1, lr=0.01, total_steps=total_steps)
state_1 = TrainState(step=0, epoch=0, batch_in_epoch=0)
train_until(
m1, o1, s1, loss_fn, state_1,
stop_step=interrupt_at,
batches_per_epoch=batches_per_epoch,
batch_size=batch_size,
in_dim=in_dim,
out_dim=out_dim,
)
if sharded:
save_sharded_checkpoint(m1, o1, s1, state_1, ckpt_dir, num_shards=num_shards)
else:
save_checkpoint(m1, o1, s1, state_1, ckpt_dir / "ckpt.pt")
train_until(
m1, o1, s1, loss_fn, state_1,
stop_step=total_steps,
batches_per_epoch=batches_per_epoch,
batch_size=batch_size,
in_dim=in_dim,
out_dim=out_dim,
)
full_losses = list(state_1.losses)
seed_everything(seed)
m2 = make_model(in_dim, hidden, out_dim)
o2, s2 = make_optimizer_and_scheduler(m2, lr=0.01, total_steps=total_steps)
if sharded:
loaded = load_sharded_checkpoint(ckpt_dir, m2, o2, s2)
else:
loaded = load_checkpoint(ckpt_dir / "ckpt.pt", m2, o2, s2)
train_until(
m2, o2, s2, loss_fn, loaded,
stop_step=total_steps,
batches_per_epoch=batches_per_epoch,
batch_size=batch_size,
in_dim=in_dim,
out_dim=out_dim,
)
resumed_losses = list(loaded.losses)
suffix_full = full_losses[interrupt_at:]
suffix_resumed = resumed_losses[interrupt_at:]
if not suffix_full:
max_diff = 0.0
else:
max_diff = max(abs(a - b) for a, b in zip(suffix_full, suffix_resumed, strict=True))
return {
"interrupt_at": interrupt_at,
"total_steps": total_steps,
"max_loss_diff_after_resume": max_diff,
"full_losses": full_losses,
"resumed_losses": resumed_losses,
"sharded": sharded,
}
def parse_args() -> argparse.Namespace:
p = argparse.ArgumentParser()
p.add_argument("--total-steps", type=int, default=24)
p.add_argument("--interrupt-at", type=int, default=10)
p.add_argument("--sharded", action="store_true")
p.add_argument("--num-shards", type=int, default=3)
p.add_argument("--seed", type=int, default=11)
return p.parse_args()
def main() -> int:
args = parse_args()
with tempfile.TemporaryDirectory(prefix="ckpt-demo-") as scratch:
scratch_dir = Path(scratch)
print("running resume demo (single file checkpoint)")
single = run_resume_demo(
total_steps=args.total_steps,
interrupt_at=args.interrupt_at,
ckpt_dir=scratch_dir / "single",
sharded=False,
seed=args.seed,
)
print(json.dumps({k: v for k, v in single.items() if k not in ("full_losses", "resumed_losses")}, indent=2))
assert single["max_loss_diff_after_resume"] < 1e-4, "loss drifted after single-file resume"
print("running resume demo (sharded checkpoint)")
sharded = run_resume_demo(
total_steps=args.total_steps,
interrupt_at=args.interrupt_at,
ckpt_dir=scratch_dir / "sharded",
sharded=True,
num_shards=args.num_shards,
seed=args.seed,
)
print(json.dumps({k: v for k, v in sharded.items() if k not in ("full_losses", "resumed_losses")}, indent=2))
assert sharded["max_loss_diff_after_resume"] < 1e-4, "loss drifted after sharded resume"
summary = {
"schema": "resume-demo.v1",
"single": {
"max_loss_diff_after_resume": single["max_loss_diff_after_resume"],
"interrupt_at": single["interrupt_at"],
"total_steps": single["total_steps"],
},
"sharded": {
"max_loss_diff_after_resume": sharded["max_loss_diff_after_resume"],
"interrupt_at": sharded["interrupt_at"],
"total_steps": sharded["total_steps"],
"num_shards": args.num_shards,
},
}
atomic_write_json(summary, OUT_DIR / "resume-demo.json")
print(f"wrote {OUT_DIR / 'resume-demo.json'}")
return 0
if __name__ == "__main__":
raise SystemExit(main())
@@ -0,0 +1,126 @@
"""Tests for full checkpoint, atomic save, and sharded resume."""
from __future__ import annotations
import json
import sys
import tempfile
import unittest
from pathlib import Path
import torch
from torch import nn
HERE = Path(__file__).parent
sys.path.insert(0, str(HERE))
import main as ckpt
def _build_components(total_steps: int = 10, lr: float = 0.01):
model = ckpt.make_model(8, 12, 4)
opt, sched = ckpt.make_optimizer_and_scheduler(model, lr=lr, total_steps=total_steps)
return model, opt, sched
class AtomicSaveTests(unittest.TestCase):
def test_atomic_save_creates_no_partial_file(self):
with tempfile.TemporaryDirectory() as tmp:
target = Path(tmp) / "ckpt.pt"
payload = {"schema": "ckpt.v1", "value": torch.zeros(3)}
ckpt.atomic_save(payload, target)
self.assertTrue(target.exists())
siblings = [p.name for p in Path(tmp).iterdir() if p.name != target.name]
for name in siblings:
self.assertFalse(name.endswith(".tmp"), f"orphan tmp file left: {name}")
def test_atomic_write_json_round_trip(self):
with tempfile.TemporaryDirectory() as tmp:
target = Path(tmp) / "index.json"
ckpt.atomic_write_json({"schema": "ckpt.v1", "n": 7}, target)
payload = json.loads(target.read_text())
self.assertEqual(payload["n"], 7)
class CheckpointResumeTests(unittest.TestCase):
def test_single_file_round_trip_matches_state(self):
with tempfile.TemporaryDirectory() as tmp:
ckpt.seed_everything(0)
model, opt, sched = _build_components(total_steps=4)
state = ckpt.TrainState(step=3, epoch=1, batch_in_epoch=2, losses=[0.5, 0.4, 0.3])
target = Path(tmp) / "ckpt.pt"
ckpt.save_checkpoint(model, opt, sched, state, target)
ckpt.seed_everything(99)
model2, opt2, sched2 = _build_components(total_steps=4)
for (_k, v1), (_, v2) in zip(
model.state_dict().items(), model2.state_dict().items(), strict=True
):
self.assertFalse(torch.allclose(v1, v2))
restored = ckpt.load_checkpoint(target, model2, opt2, sched2)
self.assertEqual(restored.step, 3)
self.assertEqual(restored.epoch, 1)
self.assertEqual(restored.batch_in_epoch, 2)
self.assertEqual(restored.losses, [0.5, 0.4, 0.3])
for (k, v1), (_, v2) in zip(
model.state_dict().items(), model2.state_dict().items(), strict=True
):
self.assertTrue(torch.allclose(v1, v2), f"param diverged: {k}")
def test_mid_epoch_resume_continues_deterministically(self):
with tempfile.TemporaryDirectory() as tmp:
result = ckpt.run_resume_demo(
total_steps=14,
interrupt_at=5,
ckpt_dir=Path(tmp),
sharded=False,
seed=3,
)
self.assertLess(result["max_loss_diff_after_resume"], 1e-5)
class ShardedCheckpointTests(unittest.TestCase):
def test_sharded_round_trip(self):
with tempfile.TemporaryDirectory() as tmp:
result = ckpt.run_resume_demo(
total_steps=12,
interrupt_at=4,
ckpt_dir=Path(tmp),
sharded=True,
num_shards=3,
seed=5,
)
self.assertLess(result["max_loss_diff_after_resume"], 1e-5)
index = json.loads((Path(tmp) / "index.json").read_text())
self.assertEqual(index["num_shards"], 3)
self.assertEqual(len(index["shards"]), 3)
def test_sha_mismatch_is_detected(self):
with tempfile.TemporaryDirectory() as tmp:
ckpt.seed_everything(0)
model, opt, sched = _build_components(total_steps=4)
state = ckpt.TrainState(step=1, epoch=0, batch_in_epoch=1, losses=[0.9])
ckpt.save_sharded_checkpoint(model, opt, sched, state, Path(tmp), num_shards=2)
tampered = Path(tmp) / "model.shard-000.pt"
data = tampered.read_bytes()
tampered.write_bytes(data + b"\x00")
model2, opt2, sched2 = _build_components(total_steps=4)
with self.assertRaises(AssertionError):
ckpt.load_sharded_checkpoint(Path(tmp), model2, opt2, sched2)
class ShardLayoutTests(unittest.TestCase):
def test_shard_layout_is_round_robin_and_complete(self):
ckpt.seed_everything(0)
model = ckpt.make_model(4, 6, 3)
sd = model.state_dict()
layout = ckpt.shard_keys_by_prefix(sd, 3)
all_keys = sorted([k for v in layout.values() for k in v])
self.assertEqual(all_keys, sorted(sd.keys()))
sizes = [len(v) for v in layout.values()]
self.assertLessEqual(max(sizes) - min(sizes), 1)
if __name__ == "__main__":
unittest.main()
@@ -0,0 +1,153 @@
# Checkpoint Save and Resume
> Train interrupts kill runs; checkpoints let them continue. Save model, optimizer, scheduler, loss history, step counter, and RNG state, atomically, so a kill at any moment leaves a valid file on disk.
**Type:** Build
**Languages:** Python
**Prerequisites:** Phase 19 lessons 42 to 45
**Time:** ~90 minutes
## Learning Objectives
- Capture the full training state into a single payload that can be reloaded into a fresh process.
- Implement atomic save with write-to-temp then rename so a crash never leaves a half-written file.
- Restore the RNG state for Python, NumPy, and PyTorch so the post-resume loss matches the uninterrupted baseline.
- Build a sharded checkpoint layout for models that no longer fit in a single file, with hash-verified shards and a JSON index.
## The Problem
You set a training job for 18 hours. The wallclock cap is 4 hours. The cluster reboots at hour 11 because someone above your pay grade approved a kernel upgrade. Without checkpoints you start over. Without resume you also lose the optimizer state that took the first 11 hours to learn, so even if the model weights survived, the AdamW moments are gone and the next step lurches in a direction the training trajectory had already moved past.
The right artifact is a single file that holds everything needed to continue: model parameters, optimizer state, scheduler state, the loss history for plots, the current step and epoch and batch-in-epoch counters, and the RNG state for every source of randomness. Without the RNG state the resumed loss curve is a different curve. Same model, same data, different shuffle, different dropout mask, different number on the dashboard.
Atomic save is the other half of the contract. Writing into the final filename means a crash mid-write leaves a corrupt file; the resume reads garbage. Writing into a temporary file in the same directory and then renaming means a crash mid-write leaves the previous good file untouched. The rename is atomic on POSIX file systems.
## The Concept
```mermaid
flowchart TD
ckpt[checkpoint payload] --> m[model state_dict]
ckpt --> o[optimizer state_dict]
ckpt --> s[scheduler state_dict]
ckpt --> tr[train state: step, epoch, batch_in_epoch, losses]
ckpt --> rng[rng state: python, numpy, torch_cpu, torch_cuda]
ckpt --> meta[wall_saved_at, schema]
ckpt --> write[atomic write: tmp file then os.replace]
```
### The five state buckets
| Bucket | Why it matters |
|--------|----------------|
| Model | Weights and buffers; what the model is. |
| Optimizer | Momentum and adaptive moments; without these the next step is a different optimization problem. |
| Scheduler | Where the learning rate is on its curve; cosine schedules in particular care. |
| Train counters | Step, epoch, batch-in-epoch, plus the loss history that draws the dashboard. |
| RNG state | Determinism for dropout, data shuffling, and any sampling inside the model. |
### Atomic save
```mermaid
flowchart LR
payload[payload] --> tmpf[write to .ckpt.pt.XXXX.tmp]
tmpf --> rename[os.replace to ckpt.pt]
rename --> done[ckpt.pt is valid]
crash1[crash before rename] --> orig[ckpt.pt unchanged]
crash2[crash after rename] --> done
```
Two rules. First, the temporary file lives in the same directory as the target so the rename stays within the same file system; cross-device renames are not atomic. Second, the temporary name is unique per attempt so two writers do not stomp.
### Sharded checkpoints
When the model gets large the single-file payload becomes too big to load fast, too big to inspect, and too painful when a network share hiccups mid-read. The fix is to split the parameter state into shards and write a small index that ties them together.
```mermaid
flowchart LR
state[state_dict] --> split[split keys round robin into N shards]
split --> s0[model.shard-000.pt]
split --> s1[model.shard-001.pt]
split --> sN[model.shard-NNN.pt]
s0 --> idx[index.json]
s1 --> idx
sN --> idx
meta[meta.pt: optimizer + scheduler + train_state + rng] --> idx
```
The index records the shard count, the sha256 of each shard, and the sha256 of the meta file. The loader fails loudly when any hash mismatches. The shards can land on different physical disks; the meta is small and reads first.
### Resume continues mid epoch
A resume that snaps to the start of the next epoch wastes anywhere from minutes to a day. The fix is `(epoch, batch_in_epoch)` plus the RNG state. After load, the training loop fast-forwards the random number generator past the batches already consumed in the current epoch and continues from `batch_in_epoch`. The lesson code does this exactly; the assertion is that the loss trajectory after resume matches the uninterrupted baseline within 1e-4.
## Build It
`code/main.py` provides four primitives and a demo driver.
### Step 1: capture and restore RNG state
`capture_rng_state` returns a dict with Python's `random.getstate`, NumPy's `np.random.get_state`, and PyTorch CPU and CUDA RNG bytes. `restore_rng_state` reverses it. The CPU tensor is a uint8 byte buffer that PyTorch's RNG knows how to consume.
### Step 2: atomic save
`atomic_save` writes the payload to a temp file in the target directory, then `os.replace` swaps it into the final name. `atomic_write_json` does the same for the sharded index.
### Step 3: full checkpoint round trip
`save_checkpoint` packages the model, optimizer, scheduler, train state, and RNG into one dict. `load_checkpoint` reverses it and returns a `TrainState`. The schema field is the upgrade hook: future format changes bump the version string and the loader dispatches.
### Step 4: sharded variant
`save_sharded_checkpoint` round-robins the parameter keys across N shards, writes each shard with its own atomic save, writes a meta file with optimizer and scheduler and train state, and writes the JSON index with shard sha256s. `load_sharded_checkpoint` verifies every shard before merging.
### Step 5: resume demo
`run_resume_demo` trains a small model for `total_steps`, saves a checkpoint at `interrupt_at`, then continues. A second process restores the checkpoint and runs the remaining steps. The function returns the max absolute difference between the two loss trajectories after the interruption point. With RNG restored, the difference is zero or floating-point noise.
Run it:
```bash
python3 code/main.py
```
The single-file and sharded demos both assert max-diff under 1e-4. The summary lands in `outputs/resume-demo.json`.
## Use It
Production training stacks ship checkpointing as part of the trainer. The shape is the same: model + optimizer + scheduler + counters + RNG, written atomically, named by step so the latest is easy to find. Sharded layouts power large model loading with parallel reads; the index.json is what makes that work.
Three patterns to enforce:
- **Schema is a string in the payload.** Migrations branch on it. Without it you cannot evolve the format without breaking old runs.
- **Sha256 every shard.** A silently truncated download is the worst kind of bug; the loader fails fast or it fails late.
- **Keep checkpoint cadence honest.** Save every N steps and every wallclock-minute, whichever is shorter. Otherwise the long step that crashes wastes a full window of work.
## Ship It
`outputs/skill-checkpoint-save-resume.md` is the recipe for any new training script: payload shape, atomic write, RNG capture, sharded index. Drop the skill into a repo, wire `save_checkpoint` at the periodic save site, wire `load_checkpoint` at startup, and the run survives kills.
## Exercises
1. Replace round-robin sharding with sharding by parameter group (layers ending in `.weight` vs `.bias`). When is each layout preferable?
2. Extend the save loop to keep the last K checkpoints and prune older ones. What is the right K when the disk is small?
3. Add a `--ckpt-every-seconds` flag that triggers a save on a wallclock interval, not just step count.
4. Add a checksum verification path that runs at startup, scans every checkpoint in the directory, and reports which ones are corrupt.
5. Implement a `migrate_v1_to_v2` function that adds a new field to the payload and bumps the schema string. Make load tolerate both versions.
## Key Terms
| Term | What people say | What it actually means |
|------|-----------------|------------------------|
| Atomic save | "Write and pray" | Write to a temp file in the same directory, then os.replace into the target name |
| State dict | "The weights" | Model parameters and buffers, keyed by parameter name |
| Sharded checkpoint | "Big model file" | Multiple files, one per shard, plus a meta file and a JSON index with sha256s |
| RNG state | "Random seed" | Captured state for python random, numpy, torch CPU, torch CUDA; not just the seed |
| Mid-epoch resume | "Restart" | Fast-forward the RNG and continue from the next batch in the same epoch |
## Further Reading
- POSIX `rename` semantics for the atomicity claim that `os.replace` relies on.
- PyTorch documentation on `torch.save` and `torch.load`, including `map_location` for cross-device restores.
- Phase 19 lesson 46 covers the gradient accumulation that this lesson's checkpoint payload survives across.
- Phase 19 lesson 48 covers the distributed wrappers whose state dict format this scheme accommodates.
- The Linux kernel `fsync` documentation for the durability guarantee behind atomic rename.
@@ -0,0 +1,14 @@
{
"schema": "resume-demo.v1",
"single": {
"max_loss_diff_after_resume": 0.0,
"interrupt_at": 10,
"total_steps": 24
},
"sharded": {
"max_loss_diff_after_resume": 0.0,
"interrupt_at": 10,
"total_steps": 24,
"num_shards": 3
}
}
@@ -0,0 +1,52 @@
---
name: checkpoint-save-resume
description: Atomic, sharded checkpoints with full RNG capture so a killed run resumes mid-epoch with the same loss trajectory.
version: 1.0.0
phase: 19
lesson: 47
tags: [training, durability, resume, sharded-state]
---
## When to use
Any training run longer than the wallclock cap of the cluster, any run that must survive a node reboot, any model too large for a single payload.
## Payload shape
```python
{
"schema": "ckpt.v1",
"model": model.state_dict(),
"optimizer": opt.state_dict(),
"scheduler": sched.state_dict(),
"state": {"step": int, "epoch": int, "batch_in_epoch": int, "losses": [float, ...]},
"rng": {"python": ..., "numpy": ..., "torch_cpu": ..., "torch_cuda": ...},
"wall_saved_at": time.time(),
}
```
## Atomic save
1. Write the payload to a unique temp file in the same directory as the target.
2. `os.replace(tmp, target)` to swap atomically.
3. Never write directly to the target name.
## Sharded layout
- `model.shard-NNN.pt` per shard, round robin on keys or split by parameter group.
- `meta.pt` carries optimizer, scheduler, train state, RNG, and the shard manifest.
- `index.json` carries `sha256` for every shard and for `meta.pt`.
- Loader verifies every hash before merging.
## Mid-epoch resume
- Save `(epoch, batch_in_epoch)` next to `step`.
- Restore RNG state before the first batch of the resumed epoch.
- Fast-forward the generator past consumed batches.
## Failure modes
- Cross-device rename: not atomic, lose the previous file. Put temp in same directory.
- Forgetting RNG: resumed loss diverges from baseline. Run the demo's assertion.
- Forgetting optimizer state: next step lurches. Same diff blows up.
- Pruning the wrong checkpoint: keep last K plus best.
@@ -0,0 +1,78 @@
{
"lesson": "47-checkpoint-save-resume",
"title": "Checkpoint Save and Resume",
"questions": [
{
"stage": "pre",
"question": "Which set best describes a complete training checkpoint?",
"options": [
"Model weights only",
"Model state, optimizer state, scheduler state, step counters, loss history, and RNG state for python, numpy, torch CPU, torch CUDA",
"Just the loss history",
"Anything torch.save() can serialize"
],
"correct": 1,
"explanation": "Resume must walk the same trajectory. Without optimizer moments, scheduler position, and RNG state the resumed loss is a different curve."
},
{
"stage": "pre",
"question": "Why does the saver write to a temporary file and then rename instead of writing the target name directly?",
"options": [
"Faster I/O",
"POSIX rename within the same directory is atomic, so a crash mid-write leaves the previous good file in place rather than a half-written file at the target name",
"It dedupes the file",
"It bypasses fsync"
],
"correct": 1,
"explanation": "atomic_save calls os.replace from a temp file in the same directory. Cross-device renames lose atomicity, which is why the temp file must live in the target directory."
},
{
"stage": "check",
"question": "Why does a sharded checkpoint store a sha256 in index.json per shard?",
"options": [
"For SEO",
"So the loader can detect a truncated or tampered shard at load time instead of silently merging a bad state_dict",
"To compress the file",
"Because torch requires it"
],
"correct": 1,
"explanation": "load_sharded_checkpoint asserts each shard's actual hash matches the recorded one and the meta file's hash too. Silent corruption is the worst failure mode."
},
{
"stage": "check",
"question": "How does resume continue at the right offset inside the current epoch?",
"options": [
"It skips to the next epoch boundary",
"TrainState records (epoch, batch_in_epoch) and the RNG state is restored, so the loader fast-forwards the generator past the batches already consumed before continuing",
"It uses a special database",
"It always restarts from step 0"
],
"correct": 1,
"explanation": "Without (epoch, batch_in_epoch) you snap to the next epoch boundary and waste work. Without RNG you see different batches. Both are required."
},
{
"stage": "check",
"question": "What does the schema field on the payload buy you?",
"options": [
"Nothing",
"A migration hook: future format changes bump the schema string and the loader can dispatch on it instead of breaking old runs",
"A pretty filename",
"Compression"
],
"correct": 1,
"explanation": "The lesson uses ckpt.v1; the next version is ckpt.v2 with a migrate function. Without the field, format evolution requires breakage."
},
{
"stage": "post",
"question": "Reading the resume demo summary, what does max_loss_diff_after_resume near 0 prove?",
"options": [
"Nothing useful",
"That the post-resume loss trajectory matches the uninterrupted baseline within floating point noise, which means the RNG state and the optimizer state were both restored faithfully",
"That the model is good",
"That the disk is fast"
],
"correct": 1,
"explanation": "If you forget RNG you see a different curve and a non-trivial diff. If you forget optimizer state the diff is huge. The near-zero number is the contract."
}
]
}
@@ -0,0 +1,415 @@
"""Distributed data parallel from scratch on the gloo backend.
CUDA is not assumed. The demo simulates a multi-rank cluster by spawning
several worker processes with torch.multiprocessing and connecting them
through the gloo CPU backend. The same collective ops (all_reduce,
broadcast) you would use on a multi-GPU machine show up here; only the
device tag changes.
Three drills:
1. Show that a manual all_reduce of gradients across N ranks matches the
gradient a single process would compute on the concatenated input.
2. Wrap a model in a from-scratch DDP wrapper that broadcasts parameters
at construction and averages gradients in a post-backward hook.
3. Sketch FSDP parameter sharding by partitioning the parameter tensors
across ranks and gathering them for the forward pass.
Run: python3 code/main.py
"""
from __future__ import annotations
import argparse
import json
import os
import sys
import time
from dataclasses import dataclass
from pathlib import Path
from typing import Callable, Dict, List, Optional
import torch
import torch.distributed as dist
import torch.multiprocessing as mp
from torch import nn
HERE = Path(__file__).parent
OUT_DIR = HERE.parent / "outputs"
DEMO_PATH = OUT_DIR / "ddp-demo.json"
@dataclass
class RankResult:
rank: int
world_size: int
backend: str
final_loss: float
pre_param_sum: float
post_param_sum: float
grad_norm_after_all_reduce: float
fsdp_round_trip_ok: bool
def init_process_group(rank: int, world_size: int, backend: str, master_port: int) -> None:
os.environ["MASTER_ADDR"] = "127.0.0.1"
os.environ["MASTER_PORT"] = str(master_port)
loopback = "lo0" if sys.platform == "darwin" else "lo"
os.environ.setdefault("GLOO_SOCKET_IFNAME", loopback)
os.environ.setdefault("TP_SOCKET_IFNAME", loopback)
dist.init_process_group(backend=backend, rank=rank, world_size=world_size)
def shutdown_process_group() -> None:
if dist.is_initialized():
dist.destroy_process_group()
def make_model(in_dim: int, hidden: int, out_dim: int) -> nn.Module:
return nn.Sequential(
nn.Linear(in_dim, hidden),
nn.GELU(),
nn.Linear(hidden, out_dim),
)
def broadcast_module(module: nn.Module, src: int = 0) -> None:
for tensor in list(module.parameters()) + list(module.buffers()):
dist.broadcast(tensor.data, src=src)
def all_reduce_grads_(module: nn.Module, world_size: int) -> float:
"""Sum gradients across ranks, divide by world size, return l2 norm."""
total_sq = 0.0
for p in module.parameters():
if p.grad is None:
p.grad = torch.zeros_like(p.data)
dist.all_reduce(p.grad.data, op=dist.ReduceOp.SUM)
p.grad.data.div_(world_size)
total_sq += float(p.grad.data.pow(2).sum().item())
return total_sq ** 0.5
def shard_for_rank(x: torch.Tensor, rank: int, world_size: int) -> torch.Tensor:
total = x.shape[0]
per = total // world_size
remainder = total - per * world_size
start = rank * per + min(rank, remainder)
end = start + per + (1 if rank < remainder else 0)
return x[start:end]
class MinimalDDP(nn.Module):
"""Toy DistributedDataParallel.
On construction, broadcast every parameter from rank zero so all ranks
start from the same weights. On forward, run the wrapped module. After
backward, the trainer calls `sync_grads()` to all-reduce gradients.
A production DDP uses a post-backward gradient hook to overlap
communication with the backward pass and buckets parameters into
fixed-size chunks for efficient collective use. The shape of the
contract is the same; the bookkeeping above gets fancy.
"""
def __init__(self, module: nn.Module, world_size: int):
super().__init__()
self.module = module
self.world_size = world_size
if dist.is_initialized() and world_size > 1:
broadcast_module(self.module, src=0)
def forward(self, *args, **kwargs):
return self.module(*args, **kwargs)
def sync_grads(self) -> float:
if not dist.is_initialized() or self.world_size == 1:
return _grad_norm(self.module)
return all_reduce_grads_(self.module, self.world_size)
def _grad_norm(module: nn.Module) -> float:
total_sq = 0.0
for p in module.parameters():
if p.grad is None:
continue
total_sq += float(p.grad.data.pow(2).sum().item())
return total_sq ** 0.5
def fsdp_round_trip_sketch(module: nn.Module, world_size: int, rank: int) -> bool:
"""Sketch parameter sharding and gathering for the forward pass.
Each rank keeps a 1/world_size slice of every parameter. Before a
forward pass the full tensor is reconstructed with all_gather. After
the use, the full copy is dropped and only the slice remains. This
keeps the per-rank memory at 1/world_size of the model.
Gloo's all_gather requires equal output sizes per rank, so the flat
tensor is right-padded to a multiple of world_size before sharding
and the padding is dropped after the gather.
Returns True if the gathered tensor matches the original on every
rank.
"""
ok = True
for p in module.parameters():
full = p.data.detach().clone()
flat = full.flatten()
total = flat.numel()
per = (total + world_size - 1) // world_size
padded_total = per * world_size
pad = padded_total - total
if pad > 0:
padded = torch.cat([flat, torch.zeros(pad, dtype=flat.dtype)])
else:
padded = flat
my_slice = padded[rank * per:(rank + 1) * per].clone()
gathered = [torch.empty(per, dtype=flat.dtype) for _ in range(world_size)]
dist.all_gather(gathered, my_slice)
rebuilt_padded = torch.cat(gathered)
rebuilt = rebuilt_padded[:total].view_as(full)
if not torch.allclose(rebuilt, full):
ok = False
break
return ok
def manual_all_reduce_matches_single_process(
rank: int,
world_size: int,
in_dim: int,
out_dim: int,
batch_size: int,
) -> tuple[float, float]:
"""Each rank computes a gradient on its slice; all-reduce-mean recovers the
full-batch gradient up to numerical noise."""
torch.manual_seed(0)
full_x = torch.randn(batch_size * world_size, in_dim)
full_y = torch.randint(low=0, high=out_dim, size=(batch_size * world_size,))
my_x = shard_for_rank(full_x, rank, world_size)
my_y = shard_for_rank(full_y, rank, world_size)
torch.manual_seed(7)
model = make_model(in_dim, hidden=16, out_dim=out_dim)
broadcast_module(model, src=0)
loss_fn = nn.CrossEntropyLoss()
for p in model.parameters():
p.grad = None
out = model(my_x)
loss = loss_fn(out, my_y)
loss.backward()
norm_after = all_reduce_grads_(model, world_size)
if rank == 0:
torch.manual_seed(7)
ref_model = make_model(in_dim, hidden=16, out_dim=out_dim)
for p in ref_model.parameters():
p.grad = None
ref_loss = loss_fn(ref_model(full_x), full_y)
ref_loss.backward()
ref_norm = _grad_norm(ref_model)
diffs = []
for p, q in zip(model.parameters(), ref_model.parameters()):
diffs.append(float((p.grad.data - q.grad.data).abs().max().item()))
max_diff = max(diffs)
else:
ref_norm = 0.0
max_diff = 0.0
return norm_after, max_diff
def rank_main(
rank: int,
world_size: int,
backend: str,
master_port: int,
result_queue,
in_dim: int,
hidden: int,
out_dim: int,
batch_size: int,
num_steps: int,
lr: float,
seed: int,
) -> None:
try:
init_process_group(rank, world_size, backend, master_port)
torch.manual_seed(seed + rank)
grad_norm, max_diff = manual_all_reduce_matches_single_process(
rank, world_size, in_dim, out_dim, batch_size
)
torch.manual_seed(seed)
base = make_model(in_dim, hidden, out_dim)
ddp_model = MinimalDDP(base, world_size)
optimizer = torch.optim.SGD(ddp_model.parameters(), lr=lr)
loss_fn = nn.CrossEntropyLoss()
pre_param_sum = sum(float(p.data.sum().item()) for p in ddp_model.parameters())
torch.manual_seed(seed * 31 + rank)
local_loss = 0.0
for step in range(num_steps):
x = torch.randn(batch_size, in_dim)
y = torch.randint(low=0, high=out_dim, size=(batch_size,))
optimizer.zero_grad()
out = ddp_model(x)
loss = loss_fn(out, y)
loss.backward()
ddp_model.sync_grads()
optimizer.step()
local_loss = float(loss.detach().item())
fsdp_ok = fsdp_round_trip_sketch(ddp_model.module, world_size, rank)
post_param_sum = sum(float(p.data.sum().item()) for p in ddp_model.parameters())
result = RankResult(
rank=rank,
world_size=world_size,
backend=backend,
final_loss=local_loss,
pre_param_sum=pre_param_sum,
post_param_sum=post_param_sum,
grad_norm_after_all_reduce=grad_norm,
fsdp_round_trip_ok=fsdp_ok,
)
result_queue.put((rank, result.__dict__, max_diff))
except Exception as exc:
result_queue.put((rank, {"error": str(exc), "rank": rank}, -1.0))
finally:
shutdown_process_group()
def free_port() -> int:
import socket
s = socket.socket(socket.AF_INET, socket.SOCK_STREAM)
s.bind(("127.0.0.1", 0))
port = s.getsockname()[1]
s.close()
return port
def run_distributed_demo(
world_size: int = 2,
*,
backend: str = "gloo",
in_dim: int = 32,
hidden: int = 16,
out_dim: int = 4,
batch_size: int = 8,
num_steps: int = 6,
lr: float = 0.05,
seed: int = 0,
timeout: float = 60.0,
) -> Dict[str, object]:
ctx = mp.get_context("spawn")
result_queue = ctx.Queue()
port = free_port()
procs = []
for rank in range(world_size):
p = ctx.Process(
target=rank_main,
args=(
rank,
world_size,
backend,
port,
result_queue,
in_dim,
hidden,
out_dim,
batch_size,
num_steps,
lr,
seed,
),
)
p.start()
procs.append(p)
results: Dict[int, dict] = {}
max_diff = 0.0
deadline = time.time() + timeout
collected = 0
while collected < world_size and time.time() < deadline:
try:
rank, payload, diff = result_queue.get(timeout=1.0)
except Exception:
continue
results[rank] = payload
if diff > max_diff:
max_diff = diff
collected += 1
for p in procs:
p.join(timeout=max(1.0, deadline - time.time()))
if any(p.is_alive() for p in procs):
for p in procs:
if p.is_alive():
p.terminate()
raise RuntimeError("ranks did not finish in time")
if collected < world_size:
raise RuntimeError(f"only got {collected}/{world_size} results: {results}")
for rank, payload in results.items():
if "error" in payload:
raise RuntimeError(f"rank {rank} failed: {payload['error']}")
param_sums = {r: results[r]["post_param_sum"] for r in results}
spread = max(param_sums.values()) - min(param_sums.values())
losses = [results[r]["final_loss"] for r in results]
grad_norm = results[0]["grad_norm_after_all_reduce"]
return {
"world_size": world_size,
"backend": backend,
"param_sum_per_rank": param_sums,
"param_sum_spread": spread,
"losses": losses,
"grad_norm_after_all_reduce": grad_norm,
"manual_all_reduce_max_diff_vs_single_process": max_diff,
"fsdp_round_trip_all_ranks_ok": all(results[r]["fsdp_round_trip_ok"] for r in results),
}
def write_demo(payload: Dict[str, object], path: Path) -> None:
path.parent.mkdir(parents=True, exist_ok=True)
path.write_text(json.dumps({"schema": "ddp-demo.v1", **payload}, indent=2) + "\n")
def parse_args() -> argparse.Namespace:
p = argparse.ArgumentParser()
p.add_argument("--world-size", type=int, default=2)
p.add_argument("--backend", type=str, default="gloo")
p.add_argument("--num-steps", type=int, default=6)
p.add_argument("--batch-size", type=int, default=8)
p.add_argument("--seed", type=int, default=0)
p.add_argument("--no-write", action="store_true")
return p.parse_args()
def main() -> int:
args = parse_args()
if not dist.is_available():
print("torch.distributed not available; skipping the demo")
return 0
if args.backend == "gloo" and not dist.is_gloo_available():
print("gloo backend not compiled; cannot run on CPU. install a build with gloo support.")
return 1
print(f"running distributed demo: backend={args.backend}, world_size={args.world_size}")
result = run_distributed_demo(
world_size=args.world_size,
backend=args.backend,
num_steps=args.num_steps,
batch_size=args.batch_size,
seed=args.seed,
)
print(json.dumps(result, indent=2))
assert result["param_sum_spread"] < 1e-3, "parameters diverged across ranks"
assert result["fsdp_round_trip_all_ranks_ok"], "FSDP sketch round trip failed"
assert result["manual_all_reduce_max_diff_vs_single_process"] < 1e-4, "manual all-reduce mismatched single-process gradient"
if not args.no_write:
write_demo(result, DEMO_PATH)
print(f"wrote {DEMO_PATH}")
return 0
if __name__ == "__main__":
raise SystemExit(main())
@@ -0,0 +1,108 @@
"""Tests for the from-scratch DDP wrapper and FSDP sharding sketch.
The collective tests spawn worker processes through torch.multiprocessing
on the gloo backend; this works on CPU and does not require CUDA.
"""
from __future__ import annotations
import json
import os
import sys
import tempfile
import unittest
from pathlib import Path
import torch
HERE = Path(__file__).parent
sys.path.insert(0, str(HERE))
import main as ddp
class HelperTests(unittest.TestCase):
def test_shard_for_rank_partitions_evenly(self):
x = torch.arange(20)
all_slices = []
for rank in range(4):
sl = ddp.shard_for_rank(x, rank, 4)
all_slices.append(sl)
merged = torch.cat(all_slices)
self.assertTrue(torch.equal(merged, x))
sizes = [s.shape[0] for s in all_slices]
self.assertLessEqual(max(sizes) - min(sizes), 1)
def test_shard_for_rank_handles_remainder(self):
x = torch.arange(11)
sizes = [ddp.shard_for_rank(x, r, 3).shape[0] for r in range(3)]
self.assertEqual(sum(sizes), 11)
self.assertLessEqual(max(sizes) - min(sizes), 1)
class GradNormTests(unittest.TestCase):
def test_grad_norm_zero_when_no_grads(self):
model = ddp.make_model(4, 6, 3)
norm = ddp._grad_norm(model)
self.assertEqual(norm, 0.0)
def test_grad_norm_matches_manual_calc(self):
model = ddp.make_model(4, 6, 3)
x = torch.randn(2, 4)
y = torch.randint(low=0, high=3, size=(2,))
loss = torch.nn.CrossEntropyLoss()(model(x), y)
loss.backward()
norm = ddp._grad_norm(model)
expected = sum(float(p.grad.data.pow(2).sum().item()) for p in model.parameters()) ** 0.5
self.assertAlmostEqual(norm, expected, places=6)
class DistributedDemoTests(unittest.TestCase):
def setUp(self):
if not torch.distributed.is_available():
self.skipTest("torch.distributed not available")
if not torch.distributed.is_gloo_available():
self.skipTest("gloo backend not available")
def test_two_rank_param_sums_match(self):
result = ddp.run_distributed_demo(
world_size=2,
in_dim=16,
hidden=12,
out_dim=3,
batch_size=4,
num_steps=3,
seed=11,
)
self.assertEqual(result["world_size"], 2)
self.assertLess(result["param_sum_spread"], 1e-3)
self.assertTrue(result["fsdp_round_trip_all_ranks_ok"])
self.assertLess(result["manual_all_reduce_max_diff_vs_single_process"], 1e-3)
def test_three_rank_param_sums_match(self):
result = ddp.run_distributed_demo(
world_size=3,
in_dim=12,
hidden=10,
out_dim=3,
batch_size=4,
num_steps=2,
seed=5,
)
self.assertEqual(result["world_size"], 3)
self.assertLess(result["param_sum_spread"], 1e-3)
self.assertTrue(result["fsdp_round_trip_all_ranks_ok"])
class OutputTests(unittest.TestCase):
def test_write_demo_round_trip(self):
with tempfile.TemporaryDirectory() as tmp:
target = Path(tmp) / "demo.json"
ddp.write_demo({"world_size": 2}, target)
data = json.loads(target.read_text())
self.assertEqual(data["schema"], "ddp-demo.v1")
self.assertEqual(data["world_size"], 2)
if __name__ == "__main__":
unittest.main()
@@ -0,0 +1,160 @@
# Distributed Data Parallel and FSDP from Scratch
> Multi-rank training is two collectives and one rule. Broadcast the parameters at startup, average the gradients after backward, never let the ranks disagree about what step they are on.
**Type:** Build
**Languages:** Python
**Prerequisites:** Phase 19 lessons 42 to 45
**Time:** ~90 minutes
## Learning Objectives
- Bring up a process group across N ranks with the `gloo` backend, no special hardware.
- Implement a minimal DDP wrapper that broadcasts parameters at construction and all-reduces gradients after backward.
- Prove that the all-reduce of per-rank gradients matches a single-process gradient on the concatenated input.
- Sketch FSDP parameter sharding: each rank holds a slice, the full tensor is gathered for the forward pass and dropped after.
## The Problem
The model fits on one device. The dataset does not. The optimization budget says you want to see N times the examples per wallclock second. The first lever is data parallel: each rank runs the same model on a different slice of the batch, then averages gradients before the optimizer step. The second lever is FSDP: the model does not fit on one device either, so each rank holds a fraction of every parameter and reconstructs the full tensors layer by layer during the forward pass.
The pain is the bookkeeping. If parameters drift across ranks the run is silently corrupt. If you average gradients but not the loss the dashboard lies. If the collective backend cannot agree on a topology the run hangs forever. The fix is to write the collectives by hand once and never trust a wrapper you cannot reproduce.
This lesson runs on CPU. CUDA is not assumed. The `gloo` backend ships with every PyTorch build and accepts `torch.multiprocessing` workers; the same code switches to `nccl` on a multi-GPU node without changing structure.
## The Concept
```mermaid
flowchart TB
init[rank 0 process] --> seed[seed model on rank 0]
init --> spawn[spawn ranks 1..N-1]
spawn --> pg[init_process_group: backend, world_size, master_addr, master_port]
pg --> bcast[broadcast model parameters from rank 0]
bcast --> loop[training loop per rank]
loop --> shard[each rank: own slice of the batch]
shard --> fwd[forward + backward locally]
fwd --> ar[all_reduce gradients, divide by world_size]
ar --> step[optimizer.step on every rank with the same gradient]
step --> loop
```
### The two collectives that matter
| Collective | What it does | When |
|------------|--------------|------|
| `broadcast` | Copy a tensor from one rank to all others | Parameter init, scheduler state, any one-to-all sync |
| `all_reduce` | Sum (or mean, or max) a tensor across all ranks, every rank gets the result | Gradient averaging after backward |
| `all_gather` | Each rank contributes a tensor, every rank gets the concatenation | Logits collection, FSDP parameter unshard |
The DDP contract is `broadcast` at construction and `all_reduce` after backward. The FSDP sketch adds `all_gather` before each layer's forward pass.
### Gradient averaging matches single-process gradient
A model trained on a batch of B examples across N ranks must produce the same gradient as a single process training on a batch of N*B. The trick is that summing per-rank gradients and dividing by N gives the average loss gradient, which is what cross entropy with mean reduction would produce on the full batch. The lesson code asserts this with `max-abs-diff < 1e-3` between the manual all-reduce gradient and the reference single-process gradient.
### FSDP sketch
```mermaid
flowchart LR
param[full parameter] --> split[split into N equal flat shards]
split --> r0[rank 0 holds shard 0]
split --> r1[rank 1 holds shard 1]
split --> rN[rank N-1 holds shard N-1]
r0 --> gather[all_gather before forward]
r1 --> gather
rN --> gather
gather --> full[full tensor on every rank]
full --> fwd[forward through this layer]
fwd --> drop[drop full tensor, keep only the shard]
```
The memory win is exact: per-rank memory for parameters drops to 1/N. The cost is the gather, which is paid every forward pass. Production FSDP overlaps the gather with the previous layer's compute so the wallclock cost is much smaller than the naive accounting predicts. The lesson does the round-trip on every parameter and asserts the reconstruction is bit-equal to the original.
### CPU and the gloo backend
CUDA is the production target, but the same code paths exist on CPU. `gloo` is the CPU collective backend. It is slower than `nccl` on GPUs by orders of magnitude, but the API surface is identical. The lesson's process group is initialized with `backend="gloo"` and ranks are spawned with `torch.multiprocessing` rather than `torchrun`; both end up at the same `torch.distributed` calls. On a multi-GPU node, the only changes are `backend="nccl"`, device tensors, and `torchrun` to launch.
## Build It
`code/main.py` is the runnable artifact.
### Step 1: bring up the process group
```python
os.environ["MASTER_ADDR"] = "127.0.0.1"
os.environ["MASTER_PORT"] = str(port)
dist.init_process_group(backend="gloo", rank=rank, world_size=world_size)
```
`MASTER_ADDR` and `MASTER_PORT` are the rendezvous: every rank dials the same port on the same host. The lesson picks a free port via a bind-and-close trick to avoid collisions when several runs share a machine.
### Step 2: broadcast at construction
`MinimalDDP.__init__` walks every parameter and buffer and calls `dist.broadcast(tensor, src=0)`. Rank 0's values become the canonical init. Without this, each rank initializes with its own seed and the ranks diverge from step one.
### Step 3: all-reduce gradients after backward
```python
def all_reduce_grads_(module, world_size):
for p in module.parameters():
if p.grad is None:
p.grad = torch.zeros_like(p.data)
dist.all_reduce(p.grad.data, op=dist.ReduceOp.SUM)
p.grad.data.div_(world_size)
```
Every rank ends up with the same averaged gradient. The optimizer step is now a function of the same input on every rank, which is why the parameters stay in sync across the run.
### Step 4: prove the equivalence
`manual_all_reduce_matches_single_process` builds the same model on rank 0 and compares the post-all-reduce gradient against the gradient a single process would compute on the concatenated input. The max-abs-diff is around 1e-8.
### Step 5: FSDP round trip
`fsdp_round_trip_sketch` flattens each parameter, pads to a multiple of `world_size`, slices, all-gathers, and unpads. Every rank's reconstruction equals the original. This is the unshard step; the inverse (re-shard after the forward) is one slice off the gathered tensor.
Run it:
```bash
python3 code/main.py
```
Default world size is 2. Two CPU processes spawn, talk to each other through `gloo`, and exit zero. The output `outputs/ddp-demo.json` captures parameter sums per rank, the gradient norm after all-reduce, the FSDP round-trip result, and the manual-vs-reference gradient diff.
## Use It
Production training stacks call the same primitives. PyTorch's `DistributedDataParallel` adds: post-backward gradient hooks that overlap all-reduce with backward, bucketed all-reduce that combines several small gradients into one collective, and the `no_sync` context lesson 46 used.
PyTorch's FSDP adds: a flat parameter view per layer so each rank holds one contiguous buffer, overlap of the next layer's unshard with the current layer's compute, and optional CPU offload for the shards.
The shape stays the same: broadcast at startup, reduce after backward, shard parameters when they no longer fit.
## Ship It
`outputs/skill-distributed-fsdp-ddp.md` carries the recipe for a new training script: spin up the process group with `gloo` for CPU and `nccl` for GPU, wrap the model in a DDP shell that broadcasts at construction and reduces after backward, optionally shard parameters with the all_gather pattern from the FSDP sketch.
## Exercises
1. Run with `--world-size 4` and confirm the param spread stays under 1e-3 across the run.
2. Replace the manual averaging with `dist.all_reduce(op=dist.ReduceOp.AVG)` and time the difference.
3. Add a post-backward hook to the DDP wrapper so the all-reduce overlaps with the rest of the backward; measure the wallclock improvement.
4. Implement the FSDP re-shard step: after the forward pass, replace the full tensor with the local shard again. Confirm per-rank memory drops.
5. Switch the backend to `nccl` on a CUDA box. Note which environment variables change and which stay the same.
## Key Terms
| Term | What people say | What it actually means |
|------|-----------------|------------------------|
| Backend | "gloo or nccl" | The library that implements the collective ops; gloo is CPU, nccl is GPU |
| World size | "Total ranks" | Number of processes in the group; the group is the unit collectives operate on |
| Rank | "Worker id" | Process identifier within the group, zero indexed |
| All-reduce | "Sum the grads" | Sum a tensor across all ranks, every rank ends with the same result |
| Unshard | "Gather the params" | Reconstruct the full tensor from per-rank slices via all_gather |
## Further Reading
- PyTorch `torch.distributed` documentation for the collective semantics this lesson relies on.
- The `gloo` library's collective list, identical in shape to the CUDA-backed `nccl` primitives.
- Phase 19 lesson 46 for the gradient accumulation pattern that wraps the DDP all-reduce in `no_sync`.
- Phase 19 lesson 47 for the checkpoint layout that survives DDP and FSDP runs.
- PyTorch FSDP documentation for the production implementation of the parameter sharding sketched here.
@@ -0,0 +1,17 @@
{
"schema": "ddp-demo.v1",
"world_size": 2,
"backend": "gloo",
"param_sum_per_rank": {
"1": -1.2760528326034546,
"0": -1.2760528326034546
},
"param_sum_spread": 0.0,
"losses": [
1.6027354001998901,
1.378564715385437
],
"grad_norm_after_all_reduce": 0.5539114150168322,
"manual_all_reduce_max_diff_vs_single_process": 2.3283064365386963e-08,
"fsdp_round_trip_all_ranks_ok": true
}
@@ -0,0 +1,45 @@
---
name: distributed-fsdp-ddp
description: Bring up multi-rank training with a from-scratch DDP wrapper and an FSDP parameter sharding sketch on the gloo or nccl backend.
version: 1.0.0
phase: 19
lesson: 48
tags: [distributed, ddp, fsdp, collectives]
---
## When to use
The model fits on one device but you need more throughput (DDP). The model does not fit on one device (FSDP). Either case: a multi-rank training setup with the same code path.
## Bring up the process group
```python
os.environ["MASTER_ADDR"] = "127.0.0.1"
os.environ["MASTER_PORT"] = str(port)
dist.init_process_group(backend="gloo", rank=rank, world_size=world_size)
```
`gloo` is the CPU backend; `nccl` is the GPU backend. Both implement the same collective surface.
## Wrap the model
1. On rank 0, build the model from your seed.
2. Wrap it with the DDP shell.
3. The shell's `__init__` calls `dist.broadcast(p.data, src=0)` for every parameter and buffer.
4. After every `loss.backward()`, the trainer calls `sync_grads()`.
5. `sync_grads()` calls `dist.all_reduce(p.grad, op=SUM)` and `p.grad.div_(world_size)`.
6. Optimizer step on every rank with the same averaged gradient.
## Shard parameters (FSDP sketch)
1. Flatten each parameter, pad to a multiple of `world_size`.
2. Keep your shard locally; release the rest.
3. Before forward, `dist.all_gather(...)` to rebuild the full tensor on every rank.
4. After forward, drop the full tensor.
## Failure modes
- Skipping the broadcast: ranks start from different inits, diverge silently.
- Forgetting to divide after sum: gradients scaled by world_size, optimizer steps too big.
- Using cross-device rename for checkpoints: not atomic; same lesson 47 trap.
- Mixing CPU and CUDA tensors on the same collective: backend mismatch, run hangs.
@@ -0,0 +1,78 @@
{
"lesson": "48-distributed-fsdp-ddp",
"title": "Distributed Data Parallel and FSDP from Scratch",
"questions": [
{
"stage": "pre",
"question": "Which two collective ops are the core of a from-scratch DDP wrapper?",
"options": [
"send and recv",
"broadcast at construction so every rank starts from the same parameters, and all_reduce after backward so every rank ends with the averaged gradient",
"reduce_scatter only",
"barrier and barrier"
],
"correct": 1,
"explanation": "Broadcast at init keeps ranks in sync at step zero. All-reduce after backward keeps them in sync at every step that follows."
},
{
"stage": "pre",
"question": "Why does this lesson work on CPU without CUDA?",
"options": [
"It is a stub",
"torch.distributed ships a gloo backend that runs collectives over plain sockets, so torch.multiprocessing workers on CPU form a real process group; only the device tag changes on a GPU box",
"It uses simulated tensors",
"It does not use torch.distributed at all"
],
"correct": 1,
"explanation": "Gloo is the CPU collective backend. The same call sites work on nccl on GPU; the API surface is identical."
},
{
"stage": "check",
"question": "What does the manual_all_reduce_matches_single_process test prove?",
"options": [
"Nothing",
"That summing per-rank gradients and dividing by world_size recovers the gradient a single process would compute on the concatenated input, within floating point noise",
"That the optimizer converges",
"That CUDA is unnecessary"
],
"correct": 1,
"explanation": "The cross entropy mean reduction on the full batch equals the average of per-rank means. Sum-and-divide is the right operation; the test makes that visible."
},
{
"stage": "check",
"question": "What does the FSDP sketch in this lesson do per parameter on every forward pass?",
"options": [
"Nothing",
"Pads the flat parameter to a multiple of world_size, slices, all_gathers the slices on every rank, drops the pad, and verifies the reconstruction matches the original",
"Re-initializes the parameter",
"Computes a hash"
],
"correct": 1,
"explanation": "Each rank owns 1/world_size of every parameter. The gather rebuilds the full tensor for compute; production FSDP overlaps the gather with the previous layer's work."
},
{
"stage": "check",
"question": "Why is the gloo all_gather sized to world_size equal-size shards even when the parameter length is not divisible by world_size?",
"options": [
"Gloo's all_gather requires equal output tensor sizes per rank, so the flat parameter is right-padded to a multiple of world_size before slicing and the padding is dropped after the gather",
"It is a bug",
"It speeds up the network",
"Random choice"
],
"correct": 0,
"explanation": "Gloo (and nccl) all_gather expects same-size outputs. Padding is the cheapest fix; the sketch trims the result back to the original length."
},
{
"stage": "post",
"question": "Reading the demo output, param_sum_spread near zero across ranks tells you what?",
"options": [
"That the loss is low",
"That every rank ended the run at the same parameter values, which means the broadcast at construction and the all-reduce after every backward both worked; if either was broken the spread would grow over steps",
"That the network is fast",
"Nothing meaningful"
],
"correct": 1,
"explanation": "DDP's invariant is parameter parity across ranks. The spread is the audit; the demo asserts it stays under 1e-3."
}
]
}
@@ -0,0 +1,506 @@
"""Language model evaluation harness from scratch.
Task spec is a JSONL line per example with `prompt`, `targets`, and
`metric`. Five metrics ship: exact match for arithmetic, rouge-l F1 for
summary, executable check for code, accuracy for multiple choice, and
substring contains for generation. The runner batches examples by task,
runs them against a swappable model adapter, and emits a leaderboard
JSON with per-task and overall scores.
The model adapter is the seam. The default adapter is a deterministic
toy that pattern-matches the prompt; it has just enough behavior to make
the harness's scoring code exercise every metric. Swap the adapter for
an HTTP client, a local inference call, or a mock in tests.
Run: python3 code/main.py
"""
from __future__ import annotations
import argparse
import ast
import json
import operator
import re
import sys
import textwrap
import time
from collections import Counter
from dataclasses import dataclass, field, asdict
from pathlib import Path
from typing import Callable, Dict, Iterable, List, Optional, Protocol, Sequence
HERE = Path(__file__).parent
OUT_DIR = HERE.parent / "outputs"
TASKS_DIR = OUT_DIR / "tasks"
LEADERBOARD = OUT_DIR / "leaderboard.json"
@dataclass
class Example:
id: str
prompt: str
targets: List[str]
metric: str
extras: Dict[str, object] = field(default_factory=dict)
@dataclass
class TaskResult:
task: str
metric: str
score: float
correct: int
total: int
per_example: List[Dict[str, object]] = field(default_factory=list)
latency_ms: float = 0.0
@dataclass
class Leaderboard:
schema: str
timestamp: float
overall_score: float
tasks: List[TaskResult]
class ModelAdapter(Protocol):
def generate(self, prompts: Sequence[str]) -> List[str]:
...
@property
def name(self) -> str:
...
class ToyAdapter:
"""Deterministic adapter that pattern-matches each task.
The point is not to score well; the point is to give the harness a
fixed set of outputs to score against. Replace with a real client
when you ship the harness against a model.
"""
name = "toy.v1"
def generate(self, prompts: Sequence[str]) -> List[str]:
return [self._answer(p) for p in prompts]
def _answer(self, prompt: str) -> str:
text = prompt.strip()
if text.startswith("compute:"):
expr = text[len("compute:"):].strip()
try:
return str(safe_arith_eval(expr))
except Exception:
return ""
if text.startswith("summarize:"):
body = text[len("summarize:"):].strip()
sentences = re.split(r"(?<=[.!?])\s+", body)
return sentences[0] if sentences else body
if text.startswith("python:"):
body = text[len("python:"):].strip()
if "double" in body:
return "def f(x):\n return x * 2\n"
if "increment" in body:
return "def f(x):\n return x + 1\n"
if "square" in body:
return "def f(x):\n return x * x\n"
return "def f(x):\n return x\n"
if text.startswith("choose:"):
body = text[len("choose:"):].strip()
return body.split("|", 1)[0].strip()[:1].upper()
if text.startswith("write:"):
body = text[len("write:"):].strip()
return body
return text
_ARITH_OPS = {
ast.Add: operator.add,
ast.Sub: operator.sub,
ast.Mult: operator.mul,
ast.Div: operator.truediv,
ast.FloorDiv: operator.floordiv,
ast.Mod: operator.mod,
ast.UAdd: operator.pos,
ast.USub: operator.neg,
ast.Pow: operator.pow,
}
def safe_arith_eval(expr: str) -> float:
"""Evaluate a small arithmetic expression without exposing eval."""
tree = ast.parse(expr, mode="eval")
return _safe_eval(tree.body)
def _safe_eval(node: ast.AST) -> float:
if isinstance(node, ast.Constant) and isinstance(node.value, (int, float)):
return node.value
if isinstance(node, ast.BinOp) and type(node.op) in _ARITH_OPS:
return _ARITH_OPS[type(node.op)](_safe_eval(node.left), _safe_eval(node.right))
if isinstance(node, ast.UnaryOp) and type(node.op) in _ARITH_OPS:
return _ARITH_OPS[type(node.op)](_safe_eval(node.operand))
raise ValueError(f"unsafe node: {ast.dump(node)}")
def normalize(s: str) -> str:
return re.sub(r"\s+", " ", s.strip().lower())
def metric_exact_match(prediction: str, targets: List[str]) -> float:
norm_pred = normalize(prediction)
return 1.0 if any(normalize(t) == norm_pred for t in targets) else 0.0
def metric_substring_contains(prediction: str, targets: List[str]) -> float:
norm_pred = normalize(prediction)
return 1.0 if any(normalize(t) in norm_pred for t in targets) else 0.0
def metric_multiple_choice(prediction: str, targets: List[str]) -> float:
pred = prediction.strip()[:1].upper()
return 1.0 if pred in {t.strip()[:1].upper() for t in targets} else 0.0
def _tokens(s: str) -> List[str]:
return re.findall(r"[a-z0-9]+", s.lower())
def _lcs_length(a: List[str], b: List[str]) -> int:
if not a or not b:
return 0
prev = [0] * (len(b) + 1)
for ai in a:
cur = [0] * (len(b) + 1)
for j, bj in enumerate(b):
if ai == bj:
cur[j + 1] = prev[j] + 1
else:
cur[j + 1] = max(prev[j + 1], cur[j])
prev = cur
return prev[-1]
def metric_rouge_l(prediction: str, targets: List[str]) -> float:
pred = _tokens(prediction)
if not pred:
return 0.0
best = 0.0
for ref in targets:
ref_toks = _tokens(ref)
if not ref_toks:
continue
lcs = _lcs_length(pred, ref_toks)
if lcs == 0:
continue
prec = lcs / len(pred)
rec = lcs / len(ref_toks)
if prec + rec == 0:
continue
f1 = 2 * prec * rec / (prec + rec)
if f1 > best:
best = f1
return best
def metric_code_exec(prediction: str, targets: List[str], extras: Dict[str, object]) -> float:
"""Execute the prediction in a small namespace and compare against
expected outputs.
Targets is a list of stringified expected results; extras carries a
list of (input, output) pairs the function is checked against. The
code runs in a stripped builtins namespace so it cannot reach the
filesystem or network.
"""
pairs = extras.get("io_pairs") or []
if not isinstance(pairs, list) or not pairs:
return 0.0
safe_globals = {"__builtins__": {"range": range, "len": len, "min": min, "max": max, "abs": abs, "int": int, "float": float}}
local: Dict[str, object] = {}
try:
exec(prediction, safe_globals, local)
except Exception:
return 0.0
fn = local.get("f")
if not callable(fn):
return 0.0
correct = 0
for pair in pairs:
if not (isinstance(pair, list) and len(pair) == 2):
continue
x, expected = pair
try:
actual = fn(x)
except Exception:
continue
if actual == expected:
correct += 1
if not pairs:
return 0.0
return correct / len(pairs)
METRIC_FNS: Dict[str, Callable[..., float]] = {
"exact_match": lambda p, t, e: metric_exact_match(p, t),
"substring_contains": lambda p, t, e: metric_substring_contains(p, t),
"multiple_choice": lambda p, t, e: metric_multiple_choice(p, t),
"rouge_l": lambda p, t, e: metric_rouge_l(p, t),
"code_exec": lambda p, t, e: metric_code_exec(p, t, e),
}
def load_task_jsonl(path: Path) -> List[Example]:
examples: List[Example] = []
with path.open("r", encoding="utf-8") as f:
for line_num, raw in enumerate(f, start=1):
raw = raw.strip()
if not raw or raw.startswith("#"):
continue
try:
obj = json.loads(raw)
except json.JSONDecodeError as exc:
raise ValueError(f"bad json at {path}:{line_num}: {exc}") from exc
examples.append(Example(
id=str(obj.get("id", f"ex-{line_num}")),
prompt=obj["prompt"],
targets=list(obj["targets"]),
metric=obj["metric"],
extras=dict(obj.get("extras", {})),
))
return examples
def write_task_jsonl(examples: Iterable[Example], path: Path) -> None:
path.parent.mkdir(parents=True, exist_ok=True)
with path.open("w", encoding="utf-8") as f:
for ex in examples:
f.write(json.dumps({
"id": ex.id,
"prompt": ex.prompt,
"targets": ex.targets,
"metric": ex.metric,
**({"extras": ex.extras} if ex.extras else {}),
}) + "\n")
def run_task(
task_name: str,
examples: List[Example],
adapter: ModelAdapter,
*,
batch_size: int = 8,
) -> TaskResult:
if batch_size <= 0:
raise ValueError(f"batch_size must be > 0, got {batch_size}")
if not examples:
return TaskResult(task=task_name, metric="none", score=0.0, correct=0, total=0)
metric = examples[0].metric
assert all(ex.metric == metric for ex in examples), f"task {task_name} mixes metrics"
metric_fn = METRIC_FNS[metric]
per_example: List[Dict[str, object]] = []
correct_sum = 0.0
total = 0
start = time.perf_counter()
for i in range(0, len(examples), batch_size):
chunk = examples[i:i + batch_size]
prompts = [ex.prompt for ex in chunk]
outputs = adapter.generate(prompts)
if len(outputs) != len(chunk):
raise ValueError(
f"adapter returned {len(outputs)} outputs for {len(chunk)} prompts in task {task_name}"
)
for ex, out in zip(chunk, outputs, strict=True):
score = metric_fn(out, ex.targets, ex.extras)
correct_sum += score
total += 1
per_example.append({
"id": ex.id,
"prompt": ex.prompt,
"prediction": out,
"targets": ex.targets,
"score": score,
})
latency_ms = (time.perf_counter() - start) * 1000.0
score = correct_sum / total if total else 0.0
correct_int = int(round(correct_sum))
return TaskResult(
task=task_name,
metric=metric,
score=score,
correct=correct_int,
total=total,
per_example=per_example,
latency_ms=latency_ms,
)
def run_leaderboard(
tasks: Dict[str, List[Example]],
adapter: ModelAdapter,
*,
batch_size: int = 8,
) -> Leaderboard:
results: List[TaskResult] = []
for name in sorted(tasks):
result = run_task(name, tasks[name], adapter, batch_size=batch_size)
results.append(result)
if results:
overall = sum(r.score for r in results) / len(results)
else:
overall = 0.0
return Leaderboard(
schema="leaderboard.v1",
timestamp=time.time(),
overall_score=overall,
tasks=results,
)
def write_leaderboard(
board: Leaderboard,
path: Path,
*,
adapter_name: str,
include_per_example: bool = False,
) -> None:
path.parent.mkdir(parents=True, exist_ok=True)
payload = {
"schema": board.schema,
"timestamp": board.timestamp,
"overall_score": board.overall_score,
"adapter": adapter_name,
"tasks": [
{
"task": r.task,
"metric": r.metric,
"score": r.score,
"correct": r.correct,
"total": r.total,
"latency_ms": r.latency_ms,
**({"per_example": r.per_example} if include_per_example else {}),
}
for r in board.tasks
],
}
path.write_text(json.dumps(payload, indent=2) + "\n")
def build_arithmetic_task() -> List[Example]:
items = [("2 + 2", "4"), ("7 - 3", "4"), ("6 * 4", "24"), ("100 / 4", "25.0"), ("12 + 9", "21")]
return [Example(id=f"arith-{i:02d}", prompt=f"compute: {q}", targets=[a], metric="exact_match") for i, (q, a) in enumerate(items)]
def build_summary_task() -> List[Example]:
items = [
("Cats are mammals. Mammals are warm blooded.", "cats are mammals"),
("Python uses indentation. Indentation defines blocks.", "python uses indentation"),
("The river flows east. Boats pass slowly.", "the river flows east"),
("Storms approach the coast. Waves rise quickly.", "storms approach the coast"),
("Bread bakes at high heat. Crust forms last.", "bread bakes at high heat"),
]
return [Example(id=f"sum-{i:02d}", prompt=f"summarize: {p}", targets=[t], metric="rouge_l") for i, (p, t) in enumerate(items)]
def build_code_task() -> List[Example]:
items = [
("write a function f that doubles its input", "double", [[1, 2], [3, 6], [5, 10]]),
("write a function f that increments its input", "increment", [[1, 2], [5, 6], [10, 11]]),
("write a function f that squares its input", "square", [[2, 4], [3, 9], [4, 16]]),
("write a function f that doubles its input again", "double", [[7, 14], [9, 18]]),
("write a function f that increments its input again", "increment", [[0, 1], [2, 3]]),
]
return [
Example(
id=f"code-{i:02d}",
prompt=f"python: {prompt}",
targets=["ok"],
metric="code_exec",
extras={"io_pairs": pairs, "tag": tag},
)
for i, (prompt, tag, pairs) in enumerate(items)
]
def build_choice_task() -> List[Example]:
items = [
("A | mammal, B | reptile, C | bird", ["A"]),
("A | apple, B | car, C | tree", ["A"]),
("A | water, B | iron, C | wood", ["A"]),
("A | square, B | triangle, C | circle", ["A"]),
("A | bread, B | rock, C | leaf", ["A"]),
]
return [Example(id=f"mc-{i:02d}", prompt=f"choose: {q}", targets=t, metric="multiple_choice") for i, (q, t) in enumerate(items)]
def build_generation_task() -> List[Example]:
items = [
("hello world", ["hello"]),
("training language models", ["language"]),
("evaluation harness", ["evaluation"]),
("gradient accumulation step", ["gradient"]),
("distributed parameter sharding", ["distributed"]),
]
return [Example(id=f"gen-{i:02d}", prompt=f"write: {p}", targets=t, metric="substring_contains") for i, (p, t) in enumerate(items)]
def seed_fixture_tasks(target_dir: Path) -> Dict[str, Path]:
target_dir.mkdir(parents=True, exist_ok=True)
tasks = {
"arithmetic": build_arithmetic_task(),
"summary": build_summary_task(),
"code-exec": build_code_task(),
"multiple-choice": build_choice_task(),
"generation": build_generation_task(),
}
paths: Dict[str, Path] = {}
for name, examples in tasks.items():
path = target_dir / f"{name}.jsonl"
write_task_jsonl(examples, path)
paths[name] = path
return paths
def load_all_tasks(task_dir: Path) -> Dict[str, List[Example]]:
tasks: Dict[str, List[Example]] = {}
for path in sorted(task_dir.glob("*.jsonl")):
tasks[path.stem] = load_task_jsonl(path)
return tasks
def parse_args() -> argparse.Namespace:
p = argparse.ArgumentParser()
p.add_argument("--task-dir", type=Path, default=TASKS_DIR)
p.add_argument("--out", type=Path, default=LEADERBOARD)
p.add_argument("--batch-size", type=int, default=4)
p.add_argument("--include-per-example", action="store_true")
p.add_argument("--seed-fixtures", action="store_true")
return p.parse_args()
def main() -> int:
args = parse_args()
if args.seed_fixtures or not args.task_dir.exists() or not list(args.task_dir.glob("*.jsonl")):
print(f"seeding fixture tasks into {args.task_dir}")
seed_fixture_tasks(args.task_dir)
tasks = load_all_tasks(args.task_dir)
print(f"loaded {len(tasks)} tasks: {sorted(tasks)}")
adapter = ToyAdapter()
board = run_leaderboard(tasks, adapter, batch_size=args.batch_size)
write_leaderboard(
board,
args.out,
adapter_name=adapter.name,
include_per_example=args.include_per_example,
)
print(f"overall_score = {board.overall_score:.3f}")
for r in board.tasks:
print(f" {r.task:>16} metric={r.metric:>18} score={r.score:0.3f} ({r.correct}/{r.total}) latency_ms={r.latency_ms:.1f}")
print(f"wrote {args.out}")
return 0
if __name__ == "__main__":
raise SystemExit(main())
@@ -0,0 +1,129 @@
"""Tests for the eval harness: metric scoring, task loading, runner."""
from __future__ import annotations
import json
import sys
import tempfile
import unittest
from pathlib import Path
HERE = Path(__file__).parent
sys.path.insert(0, str(HERE))
import main as harness
class MetricTests(unittest.TestCase):
def test_exact_match_normalizes_case_and_whitespace(self):
self.assertEqual(harness.metric_exact_match(" Hello WORLD ", ["hello world"]), 1.0)
self.assertEqual(harness.metric_exact_match("hello", ["hi"]), 0.0)
def test_multiple_choice_uses_first_letter(self):
self.assertEqual(harness.metric_multiple_choice("a) the apple", ["A"]), 1.0)
self.assertEqual(harness.metric_multiple_choice("B", ["A"]), 0.0)
def test_rouge_l_returns_perfect_on_identical_strings(self):
score = harness.metric_rouge_l("the river flows east", ["the river flows east"])
self.assertAlmostEqual(score, 1.0, places=6)
def test_rouge_l_partial_overlap_below_one(self):
score = harness.metric_rouge_l("the river flows east", ["the river bends slowly"])
self.assertGreater(score, 0.0)
self.assertLess(score, 1.0)
def test_substring_contains_truthy_when_target_inside(self):
self.assertEqual(harness.metric_substring_contains("the gradient is good", ["gradient"]), 1.0)
self.assertEqual(harness.metric_substring_contains("no match here", ["gradient"]), 0.0)
def test_code_exec_runs_pairs_and_blocks_unsafe(self):
prediction = "def f(x):\n return x * 2\n"
pairs = {"io_pairs": [[1, 2], [3, 6]]}
self.assertEqual(harness.metric_code_exec(prediction, [], pairs), 1.0)
bad = "import os\ndef f(x):\n return os.getcwd()\n"
self.assertEqual(harness.metric_code_exec(bad, [], {"io_pairs": [[1, "anything"]]}), 0.0)
class SafeArithTests(unittest.TestCase):
def test_arith_eval_basic_ops(self):
self.assertEqual(harness.safe_arith_eval("1 + 2 * 3"), 7)
self.assertEqual(harness.safe_arith_eval("(8 - 2) / 3"), 2.0)
def test_arith_eval_rejects_calls(self):
with self.assertRaises(ValueError):
harness.safe_arith_eval("__import__('os').system('echo')")
class TaskIOTests(unittest.TestCase):
def test_round_trip_jsonl(self):
with tempfile.TemporaryDirectory() as tmp:
examples = harness.build_arithmetic_task()
path = Path(tmp) / "arith.jsonl"
harness.write_task_jsonl(examples, path)
loaded = harness.load_task_jsonl(path)
self.assertEqual(len(loaded), len(examples))
self.assertEqual(loaded[0].metric, "exact_match")
self.assertEqual(loaded[0].id, examples[0].id)
def test_load_skips_comments_and_blanks(self):
with tempfile.TemporaryDirectory() as tmp:
path = Path(tmp) / "comments.jsonl"
path.write_text("# comment\n\n{\"prompt\":\"compute: 1 + 1\",\"targets\":[\"2\"],\"metric\":\"exact_match\"}\n")
loaded = harness.load_task_jsonl(path)
self.assertEqual(len(loaded), 1)
class RunnerTests(unittest.TestCase):
def test_runner_full_score_with_toy_adapter(self):
with tempfile.TemporaryDirectory() as tmp:
paths = harness.seed_fixture_tasks(Path(tmp))
tasks = harness.load_all_tasks(Path(tmp))
self.assertEqual(len(tasks), 5)
adapter = harness.ToyAdapter()
board = harness.run_leaderboard(tasks, adapter, batch_size=3)
self.assertEqual(board.schema, "leaderboard.v1")
self.assertEqual(len(board.tasks), 5)
for r in board.tasks:
self.assertEqual(r.total, 5)
self.assertGreater(r.score, 0.5, msg=f"low score on {r.task}: {r.score}")
def test_runner_handles_mixed_failure(self):
class StubAdapter:
name = "stub"
def generate(self, prompts):
return ["" for _ in prompts]
with tempfile.TemporaryDirectory() as tmp:
paths = harness.seed_fixture_tasks(Path(tmp))
tasks = harness.load_all_tasks(Path(tmp))
board = harness.run_leaderboard(tasks, StubAdapter(), batch_size=2)
self.assertEqual(len(board.tasks), 5)
for r in board.tasks:
self.assertEqual(r.score, 0.0)
self.assertEqual(board.overall_score, 0.0)
class LeaderboardOutputTests(unittest.TestCase):
def test_leaderboard_json_contains_schema(self):
with tempfile.TemporaryDirectory() as tmp:
paths = harness.seed_fixture_tasks(Path(tmp) / "tasks")
tasks = harness.load_all_tasks(Path(tmp) / "tasks")
adapter = harness.ToyAdapter()
board = harness.run_leaderboard(tasks, adapter)
out = Path(tmp) / "leaderboard.json"
harness.write_leaderboard(board, out, adapter_name=adapter.name)
payload = json.loads(out.read_text())
self.assertEqual(payload["schema"], "leaderboard.v1")
self.assertEqual(payload["adapter"], adapter.name)
self.assertEqual(len(payload["tasks"]), 5)
for entry in payload["tasks"]:
self.assertNotIn("per_example", entry)
harness.write_leaderboard(board, out, adapter_name=adapter.name, include_per_example=True)
payload2 = json.loads(out.read_text())
for entry in payload2["tasks"]:
self.assertIn("per_example", entry)
if __name__ == "__main__":
unittest.main()
@@ -0,0 +1,194 @@
# Language Model Evaluation Harness
> A model that does well on a task you cannot define is a model that does well by accident. The harness is the task definition, the metric, the runner, and the leaderboard, in one short, swappable shape.
**Type:** Build
**Languages:** Python
**Prerequisites:** Phase 19 lessons 42 to 45
**Time:** ~90 minutes
## Learning Objectives
- Define a task as a JSONL file with `prompt`, `targets`, `metric`, and optional `extras` per example.
- Implement five metrics: exact match, rouge-l F1, executable check, multiple choice, and substring contains.
- Build a runner that batches examples per task and dispatches to a swappable model adapter.
- Emit a leaderboard JSON with per-task scores, latency, and an overall average that is reproducible.
## The Problem
A new language model lands every week. The marketing claim is that it does well. The honest question is: well at what? The honest answer is the leaderboard you wrote yourself, because the vendor's leaderboard is the one they tuned to.
Without a harness in your repo you compare two models by vibes. With a harness you compare them by score on a fixed task set with a fixed metric, on a JSON output you can diff. The harness is the contract between yesterday's run and today's run. Without it, regressions ship.
The trap is over-fitting the harness to a single model. The fix is the same trap in reverse: the harness is small enough to read in fifteen minutes, the tasks are small enough to ship in the repo, the metrics are written from scratch so a colleague can audit them, and the adapter is the only place model-specific code lives. Swap the adapter, the leaderboard moves; swap the tasks, the leaderboard moves. Nothing else should move.
## The Concept
```mermaid
flowchart TD
tasks[task JSONLs: prompt, targets, metric, extras] --> loader[load_all_tasks]
loader --> runner[run_leaderboard]
runner --> adapter[ModelAdapter.generate batch]
adapter --> metrics[METRIC_FNS dispatch by name]
metrics --> scores[per example score]
scores --> board[Leaderboard: per task + overall]
board --> out[leaderboard.json]
```
### Task spec
Every example is one JSONL line:
```json
{"id": "arith-00", "prompt": "compute: 2 + 2", "targets": ["4"], "metric": "exact_match"}
```
For metrics that need scoring helpers, `extras` carries the side payload:
```json
{
"id": "code-00",
"prompt": "python: write a function f that doubles its input",
"targets": ["ok"],
"metric": "code_exec",
"extras": {"io_pairs": [[1, 2], [3, 6]]}
}
```
A task is a `.jsonl` file under `outputs/tasks/`. The file name is the task name. All examples in a file share a metric.
### The five fixture tasks
| Task | Metric | What it tests |
|------|--------|---------------|
| arithmetic | exact_match | Token-level correctness on a deterministic answer |
| summary | rouge_l | Longest common subsequence F1 against a one-line reference summary |
| code-exec | code_exec | Executable test: the predicted function must satisfy a list of input-output pairs |
| multiple-choice | multiple_choice | First letter of the prediction must match an allowed letter |
| generation | substring_contains | Free-form text must contain at least one target substring |
### The metric contract
Every metric is a function from `(prediction, targets, extras) -> float in [0.0, 1.0]`. The harness averages the per-example scores to get a task score, then averages task scores to get the overall. The metric functions are tiny:
- `exact_match`: lowercase, collapse whitespace, equality.
- `substring_contains`: same normalization, substring test.
- `multiple_choice`: first character uppercased.
- `rouge_l`: LCS length divided by lengths of prediction and reference, F1 of precision and recall.
- `code_exec`: execute the prediction in a restricted namespace, call `f(x)` on every input-output pair, count matches.
The code_exec metric runs the prediction in a stripped builtins namespace. The lesson's test asserts that `import os` blows up because `os` is not in the namespace; you cannot reach the filesystem from a code prediction.
### The model adapter
```python
class ModelAdapter(Protocol):
def generate(self, prompts: Sequence[str]) -> List[str]: ...
@property
def name(self) -> str: ...
```
The adapter is the seam. The lesson ships `ToyAdapter`, a deterministic pattern matcher that returns the right answer for every prompt in the five fixture tasks. A real adapter calls the model and returns its output. The harness does not care which.
### The runner
`run_task` batches `batch_size` prompts at a time and dispatches to the metric function. `run_leaderboard` walks every task and averages. `write_leaderboard` emits JSON with a schema string so future format changes do not silently break dashboards.
```mermaid
flowchart LR
examples[N examples] --> batches[B-sized batches]
batches --> adapter[adapter.generate]
adapter --> per[per example score 0..1]
per --> avg[task score]
avg --> over[overall = mean of task scores]
```
## Build It
`code/main.py` is the runnable artifact.
### Step 1: seed fixture tasks
`seed_fixture_tasks(target_dir)` writes the five `.jsonl` files. The first run of `main.py` seeds them when the directory is empty.
### Step 2: load tasks
`load_all_tasks(task_dir)` reads every `.jsonl` and returns a dict from task name to a list of `Example` records. Comment lines starting with `#` and blank lines are skipped so contributors can annotate the files.
### Step 3: implement metrics
Each metric is a small function with a unit test. The lesson's test suite includes 13 cases covering normalization, partial overlap, code execution, and unsafe code rejection.
### Step 4: write the runner
`run_task` iterates batches and produces a `TaskResult` with score, correct count, total count, and latency. `run_leaderboard` walks all tasks and produces a `Leaderboard` with the overall average.
### Step 5: emit JSON
`write_leaderboard` serializes the board. The `--include-per-example` flag dumps the per-example records so you can diff predictions against the previous run when scores move.
Run it:
```bash
python3 code/main.py
```
The script seeds the fixtures on first run, scores them with the toy adapter (which gets every fixture right), and writes `outputs/leaderboard.json`. Overall score is 1.0 with the toy adapter; the stub adapter test in `test_main.py` shows the same harness produces 0.0 when the adapter cannot answer.
## Use It
To plug a real model in, write an adapter. The shape:
```python
class HttpAdapter:
name = "vendor.v1"
def __init__(self, endpoint, api_key):
self.endpoint = endpoint
self.api_key = api_key
def generate(self, prompts):
out = []
for prompt in prompts:
response = http_post(self.endpoint, prompt, self.api_key)
out.append(response["text"])
return out
```
Swap `ToyAdapter` for `HttpAdapter` at the top of `main()`. The harness, the tasks, the metrics, and the leaderboard stay the same.
Three patterns to enforce when shipping the harness in a real project:
- **Pin the task files.** The leaderboard.json carries hash-pinned task content or it carries the JSONLs alongside; otherwise the score moves when the task file does, and you cannot tell which.
- **Diff predictions, not just scores.** The `--include-per-example` flag lets you see what the model said the day the score dropped.
- **Cap the batch size.** Real adapters have rate limits. A small batch size keeps the harness compatible across vendors.
## Ship It
`outputs/skill-lm-eval-harness.md` carries the recipe: JSONL task spec, five metrics, swappable adapter, batched runner, leaderboard JSON with schema string. The task files in `outputs/tasks/` are the fixtures; copy them into a real project as starters.
## Exercises
1. Add a sixth task with a custom metric you write from scratch (BLEU-like overlap, BLEURT-like reference scoring, anything with a clear contract).
2. Extend `code_exec` to capture stdout and accept a list of expected stdouts as targets.
3. Add a leaderboard diff command: given two `leaderboard.json` files, print which tasks moved and by how much.
4. Cap latency per example. Wrap the adapter call in a timeout; surface a separate `timeouts` column in the leaderboard.
5. Pin task content with a sha256 in the leaderboard so a future reader can verify they scored the same tasks.
## Key Terms
| Term | What people say | What it actually means |
|------|-----------------|------------------------|
| Task spec | "The eval format" | JSONL file with prompt, targets, metric, optional extras per example |
| Metric | "How you score" | Function from (prediction, targets, extras) to a float in [0, 1] |
| Adapter | "The model client" | Object with a generate(prompts) -> list[str] method; the only model-specific code |
| Leaderboard | "The scoreboard" | JSON with per-task scores, total counts, latency, and an overall average |
| Code exec metric | "Run it and check" | Execute the prediction in a restricted namespace, compare against input-output pairs |
## Further Reading
- The original lm-evaluation-harness for the production reference, much larger but the same shape.
- HuggingFace's lighteval for an alternative implementation of the same contract.
- Phase 19 lesson 46 covers the gradient accumulation patterns used in the training stack the harness scores.
- Phase 19 lesson 47 covers the checkpoint format you score against; pin the checkpoint hash in the leaderboard.
- Phase 19 lesson 48 covers the distributed training stack that produced the model under test.
@@ -0,0 +1,48 @@
{
"schema": "leaderboard.v1",
"timestamp": 1779820461.3824089,
"overall_score": 1.0,
"adapter": "arithmetic",
"tasks": [
{
"task": "arithmetic",
"metric": "exact_match",
"score": 1.0,
"correct": 5,
"total": 5,
"latency_ms": 0.09666697587817907
},
{
"task": "code-exec",
"metric": "code_exec",
"score": 1.0,
"correct": 5,
"total": 5,
"latency_ms": 0.1051250146701932
},
{
"task": "generation",
"metric": "substring_contains",
"score": 1.0,
"correct": 5,
"total": 5,
"latency_ms": 0.02195802517235279
},
{
"task": "multiple-choice",
"metric": "multiple_choice",
"score": 1.0,
"correct": 5,
"total": 5,
"latency_ms": 0.011208001524209976
},
{
"task": "summary",
"metric": "rouge_l",
"score": 1.0,
"correct": 5,
"total": 5,
"latency_ms": 0.1154160127043724
}
]
}
@@ -0,0 +1,55 @@
---
name: lm-eval-harness
description: Minimal language model evaluation harness with JSONL task spec, five metrics, swappable adapter, and leaderboard JSON output.
version: 1.0.0
phase: 19
lesson: 49
tags: [evaluation, metrics, leaderboard, harness]
---
## When to use
Compare two models, two checkpoints, or two prompt templates against a fixed set of tasks. Anything that ships and that you need to monitor over time.
## Task spec
One JSONL line per example:
```json
{"id": "ex-001", "prompt": "...", "targets": ["..."], "metric": "exact_match", "extras": {}}
```
All examples in a file share a metric. The file name is the task name.
## Metrics
| Metric | Signature | Use for |
|--------|-----------|---------|
| exact_match | normalize lower + whitespace, equality | Arithmetic, factoid answers |
| substring_contains | target must appear in normalized prediction | Free-form generation with anchor words |
| multiple_choice | first letter match | A/B/C/D style questions |
| rouge_l | LCS F1 over tokenized text | Summary, paraphrase |
| code_exec | run prediction's `f` on io_pairs, count matches | Code generation |
All metrics return float in [0.0, 1.0]. Task score is the mean.
## Adapter
```python
class Adapter(Protocol):
name: str
def generate(self, prompts: list[str]) -> list[str]: ...
```
The adapter is the only model-specific code.
## Leaderboard JSON
Schema string, timestamp, per-task scores and latency, overall mean. Include per-example records when comparing runs so prediction-level regressions are visible.
## Failure modes
- Metric returns outside [0, 1]: overall score becomes uninterpretable.
- Mixed metrics in one task file: assertion fires; keep one metric per file.
- code_exec without restricted namespace: arbitrary code execution.
- No schema string: format evolution breaks downstream dashboards.
@@ -0,0 +1,5 @@
{"id": "arith-00", "prompt": "compute: 2 + 2", "targets": ["4"], "metric": "exact_match"}
{"id": "arith-01", "prompt": "compute: 7 - 3", "targets": ["4"], "metric": "exact_match"}
{"id": "arith-02", "prompt": "compute: 6 * 4", "targets": ["24"], "metric": "exact_match"}
{"id": "arith-03", "prompt": "compute: 100 / 4", "targets": ["25.0"], "metric": "exact_match"}
{"id": "arith-04", "prompt": "compute: 12 + 9", "targets": ["21"], "metric": "exact_match"}
@@ -0,0 +1,5 @@
{"id": "code-00", "prompt": "python: write a function f that doubles its input", "targets": ["ok"], "metric": "code_exec", "extras": {"io_pairs": [[1, 2], [3, 6], [5, 10]], "tag": "double"}}
{"id": "code-01", "prompt": "python: write a function f that increments its input", "targets": ["ok"], "metric": "code_exec", "extras": {"io_pairs": [[1, 2], [5, 6], [10, 11]], "tag": "increment"}}
{"id": "code-02", "prompt": "python: write a function f that squares its input", "targets": ["ok"], "metric": "code_exec", "extras": {"io_pairs": [[2, 4], [3, 9], [4, 16]], "tag": "square"}}
{"id": "code-03", "prompt": "python: write a function f that doubles its input again", "targets": ["ok"], "metric": "code_exec", "extras": {"io_pairs": [[7, 14], [9, 18]], "tag": "double"}}
{"id": "code-04", "prompt": "python: write a function f that increments its input again", "targets": ["ok"], "metric": "code_exec", "extras": {"io_pairs": [[0, 1], [2, 3]], "tag": "increment"}}
@@ -0,0 +1,5 @@
{"id": "gen-00", "prompt": "write: hello world", "targets": ["hello"], "metric": "substring_contains"}
{"id": "gen-01", "prompt": "write: training language models", "targets": ["language"], "metric": "substring_contains"}
{"id": "gen-02", "prompt": "write: evaluation harness", "targets": ["evaluation"], "metric": "substring_contains"}
{"id": "gen-03", "prompt": "write: gradient accumulation step", "targets": ["gradient"], "metric": "substring_contains"}
{"id": "gen-04", "prompt": "write: distributed parameter sharding", "targets": ["distributed"], "metric": "substring_contains"}
@@ -0,0 +1,5 @@
{"id": "mc-00", "prompt": "choose: A | mammal, B | reptile, C | bird", "targets": ["A"], "metric": "multiple_choice"}
{"id": "mc-01", "prompt": "choose: A | apple, B | car, C | tree", "targets": ["A"], "metric": "multiple_choice"}
{"id": "mc-02", "prompt": "choose: A | water, B | iron, C | wood", "targets": ["A"], "metric": "multiple_choice"}
{"id": "mc-03", "prompt": "choose: A | square, B | triangle, C | circle", "targets": ["A"], "metric": "multiple_choice"}
{"id": "mc-04", "prompt": "choose: A | bread, B | rock, C | leaf", "targets": ["A"], "metric": "multiple_choice"}
@@ -0,0 +1,5 @@
{"id": "sum-00", "prompt": "summarize: Cats are mammals. Mammals are warm blooded.", "targets": ["cats are mammals"], "metric": "rouge_l"}
{"id": "sum-01", "prompt": "summarize: Python uses indentation. Indentation defines blocks.", "targets": ["python uses indentation"], "metric": "rouge_l"}
{"id": "sum-02", "prompt": "summarize: The river flows east. Boats pass slowly.", "targets": ["the river flows east"], "metric": "rouge_l"}
{"id": "sum-03", "prompt": "summarize: Storms approach the coast. Waves rise quickly.", "targets": ["storms approach the coast"], "metric": "rouge_l"}
{"id": "sum-04", "prompt": "summarize: Bread bakes at high heat. Crust forms last.", "targets": ["bread bakes at high heat"], "metric": "rouge_l"}
@@ -0,0 +1,78 @@
{
"lesson": "49-lm-eval-harness",
"title": "Language Model Evaluation Harness",
"questions": [
{
"stage": "pre",
"question": "What four fields define a task example in the harness's JSONL format?",
"options": [
"name, score, payload, vendor",
"id, prompt, targets, metric, and an optional extras dict for side data the metric needs",
"input and output",
"weights and bias"
],
"correct": 1,
"explanation": "The JSONL line is the contract. The extras field lets the code_exec metric pass io_pairs without polluting the prompt."
},
{
"stage": "pre",
"question": "Why is the metric function signature (prediction, targets, extras) -> float?",
"options": [
"Random choice",
"It is the smallest signature that handles single-target string match, multi-reference rouge-l, and code_exec with side data, while keeping scores in a comparable [0.0, 1.0] range",
"It returns an integer",
"It needs the model object"
],
"correct": 1,
"explanation": "Floats in [0,1] mean per-task and overall scores are interpretable averages. The extras slot is how code_exec gets its io_pairs."
},
{
"stage": "check",
"question": "How does the code_exec metric resist malicious predictions?",
"options": [
"It does not",
"It runs the prediction with a stripped __builtins__ dict so the namespace exposes only a few safe names; import statements fail because the importer is not in scope",
"It uses sandbox containers",
"It rejects anything that contains def"
],
"correct": 1,
"explanation": "The lesson's safe namespace strips builtins down to a handful of names. The test asserts that import os returns score 0.0 instead of executing."
},
{
"stage": "check",
"question": "What does the model adapter abstraction buy you?",
"options": [
"Nothing",
"It is the only model-specific code in the harness; swap the adapter to point at a new vendor, and the tasks, metrics, runner, and leaderboard format stay the same",
"It makes the model faster",
"It is required by torch"
],
"correct": 1,
"explanation": "ToyAdapter in the lesson is a deterministic pattern matcher. An HttpAdapter for a real vendor has the same generate(prompts) -> list[str] surface."
},
{
"stage": "check",
"question": "Why does the leaderboard JSON carry a schema string like leaderboard.v1?",
"options": [
"For SEO",
"So future format changes bump the version and downstream dashboards can dispatch on it instead of breaking silently",
"It is random",
"It compresses the file"
],
"correct": 1,
"explanation": "Same trick as the checkpoint payload from lesson 47. The schema field is the migration hook."
},
{
"stage": "post",
"question": "Reading a leaderboard, overall_score is computed how, and what should you watch for when comparing two runs?",
"options": [
"Sum of correct counts",
"Mean of per-task scores (each in [0,1]); when comparing, also diff per-example predictions for tasks whose scores moved, because the score alone hides which examples regressed",
"Best task score",
"Random sample"
],
"correct": 1,
"explanation": "Mean of per-task means weights every task equally. Use --include-per-example to keep the prediction-level evidence next to the score so regressions are visible."
}
]
}