The training loop in Part 2 is correct and it's also a toy. It runs on the CPU in full precision, at a fixed learning rate, with no way to stop and resume. A real run — one that takes hours and costs money — wraps that same five-line core in a few layers of machinery. None of it is complicated. All of it is standard. This part assembles the loop you'll actually copy into projects.
Get on the GPU
Pick a device once, at the top, and move the model and every batch onto it. The model moves in place; tensors don't, so you reassign. non_blocking=True lets the host→device copy overlap with compute when the source is pinned memory (which your DataLoader is providing, from Part 4).
import torchdevice = "cuda" if torch.cuda.is_available() else "cpu"model = model.to(device)for images, labels in loader:images = images.to(device, non_blocking=True)labels = labels.to(device, non_blocking=True)...
Mixed precision: a nearly free 2×
Modern GPUs run half-precision matrix multiplies far faster than full precision, and use half the memory doing it. Automatic Mixed Precision runs the forward pass in bfloat16 or float16 where it's safe and keeps a float32 copy of the weights for the update. With float16 you also need a GradScaler, which multiplies the loss up before backward() so small gradients don't flush to zero, then unscales before the step.
import torchscaler = torch.amp.GradScaler("cuda")for images, labels in loader:images = images.to(device, non_blocking=True)labels = labels.to(device, non_blocking=True)opt.zero_grad()with torch.amp.autocast("cuda", dtype=torch.bfloat16):logits = model(images)loss = loss_fn(logits, labels)scaler.scale(loss).backward()scaler.step(opt)scaler.update()
Illustrative, ResNet-50-scale model on a recent GPU
AMP's advantage grows with batch size, because larger matmuls spend proportionally more time in the tensor cores that half precision unlocks. The memory saving also lets you fit the larger batch in the first place.
| batch | float32 | AMP (bf16) |
|---|---|---|
| bs 32 | 1 | 1.6 |
| bs 64 | 1 | 1.9 |
| bs 128 | 1 | 2.1 |
| bs 256 | 1 | 2.3 |
Learning-rate schedules
A fixed learning rate is a compromise: large enough to make early progress, small enough not to bounce around the minimum later. A schedule removes the compromise — start with a short linear warmup (so the first steps don't wreck freshly-initialised weights), then decay, usually along a cosine curve, toward zero.
from torch.optim.lr_scheduler import LinearLR, CosineAnnealingLR, SequentialLRwarmup = LinearLR(opt, start_factor=0.01, total_iters=500)cosine = CosineAnnealingLR(opt, T_max=total_steps - 500)scheduler = SequentialLR(opt, [warmup, cosine], milestones=[500])for step, (images, labels) in enumerate(loader):...scaler.step(opt)scaler.update()scheduler.step() # once per optimizer step
1,000-step run, 100-step warmup
Cosine + warmup (smooth) vs. step decay (drop 10× at the halfway mark). Warmup is the short ramp at the start; both end far below where they began. Hover to compare.
| step | cosine + warmup | step decay |
|---|---|---|
| 0 | 0.01 | 1 |
| 100 | 1 | 1 |
| 200 | 0.97 | 1 |
| 300 | 0.88 | 1 |
| 400 | 0.75 | 1 |
| 500 | 0.59 | 0.1 |
| 600 | 0.42 | 0.1 |
| 700 | 0.25 | 0.1 |
| 800 | 0.12 | 0.1 |
| 900 | 0.03 | 0.1 |
| 1000 | 0 | 0.1 |
Gradient clipping and accumulation
Clipping caps the global norm of the gradient before the step, which stops a single bad batch from throwing the weights off a cliff — standard practice for Transformers and RNNs. Accumulation runs several forward/backward passes before one optimizer step, so you get the training dynamics of a large batch on a GPU that can't hold one.
accum_steps = 4for step, (images, labels) in enumerate(loader):with torch.amp.autocast("cuda", dtype=torch.bfloat16):loss = loss_fn(model(images.to(device)), labels.to(device))loss = loss / accum_steps # average, don't sumscaler.scale(loss).backward()if (step + 1) % accum_steps == 0:scaler.unscale_(opt) # unscale before clippingtorch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)scaler.step(opt)scaler.update()opt.zero_grad()scheduler.step()
Checkpoints you can resume from
A checkpoint that only holds the model weights cannot resume a run — the optimizer's momentum buffers, the scheduler's step count, and the AMP scaler's state all matter. Save the lot in one dict.
import torchdef save_ckpt(path, epoch):torch.save({"epoch": epoch,"model": model.state_dict(),"opt": opt.state_dict(),"scheduler": scheduler.state_dict(),"scaler": scaler.state_dict(),}, path)def load_ckpt(path):ckpt = torch.load(path, map_location=device, weights_only=True)model.load_state_dict(ckpt["model"])opt.load_state_dict(ckpt["opt"])scheduler.load_state_dict(ckpt["scheduler"])scaler.load_state_dict(ckpt["scaler"])return ckpt["epoch"] + 1 # resume from the next epoch
A validation pass
@torch.no_grad()def validate(model, loader):model.eval()correct = total = 0for images, labels in loader:images, labels = images.to(device), labels.to(device)with torch.amp.autocast("cuda", dtype=torch.bfloat16):preds = model(images).argmax(dim=1)correct += (preds == labels).sum().item()total += labels.numel()model.train()return correct / total
The whole thing
import torchdevice = "cuda" if torch.cuda.is_available() else "cpu"model = model.to(device)opt = torch.optim.AdamW(model.parameters(), lr=3e-4, weight_decay=0.05)scaler = torch.amp.GradScaler("cuda")scheduler = build_scheduler(opt, total_steps)loss_fn = torch.nn.CrossEntropyLoss()start_epoch = load_ckpt("last.pt") if resume else 0best_acc = 0.0for epoch in range(start_epoch, epochs):model.train()for images, labels in train_loader:images = images.to(device, non_blocking=True)labels = labels.to(device, non_blocking=True)opt.zero_grad()with torch.amp.autocast("cuda", dtype=torch.bfloat16):loss = loss_fn(model(images), labels)scaler.scale(loss).backward()scaler.unscale_(opt)torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)scaler.step(opt)scaler.update()scheduler.step()acc = validate(model, val_loader)save_ckpt("last.pt", epoch)if acc > best_acc:best_acc = accsave_ckpt("best.pt", epoch)print(f"epoch {epoch} val_acc {acc:.4f} lr {scheduler.get_last_lr()[0]:.2e}")
- One device, chosen once; model moved in place, batches moved every step with
non_blocking=True. - AMP with
autocast+GradScaler— roughly 2× throughput and half the memory. - Warmup then cosine decay;
scheduler.step()once per optimizer step. - Clip the gradient norm; accumulate when the batch you want won't fit.
- Checkpoint model + optimizer + scheduler + scaler + epoch, or you can't truly resume.
@torch.no_grad()andmodel.eval()around validation; back tomodel.train()after.
That's the series. Five parts ago this was a single data structure; now it's a training run you could put on a cluster. Everything past here — torch.compile, FSDP and multi-node, quantisation, custom kernels — is a refinement of these same pieces, and none of it will surprise you once this loop is muscle memory.
References
- [1]Automatic Mixed Precision package — torch.amp · PyTorch documentation
- [2]Automatic Mixed Precision recipe · PyTorch tutorials
- [3]How to adjust learning rate — torch.optim.lr_scheduler · PyTorch documentation
- [4]SGDR: Stochastic Gradient Descent with Warm Restarts · Loshchilov & Hutter, ICLR 2017The cosine-annealing schedule.
- [5]Saving and loading a general checkpoint · PyTorch tutorials
- [6]Reproducibility · PyTorch documentation


