The inner loop of pretraining is short enough to fit on a screen, which is why it is easy to get subtly wrong. Each piece, the optimizer, the learning-rate schedule, mixed precision, gradient accumulation, clipping, checkpointing, is well documented on its own. The failures happen at the seams: accumulation that averages over the wrong count, a schedule that restarts from zero after a resume, a checkpoint that saved the weights but not the data position, a loss spike that is handled by hand at three in the morning with a different fix each time. A long run will be interrupted and it will spike, so the loop must be designed for both.

This article treats the loop as one system and writes it out as pseudocode. It assumes the architecture and pipeline from Training a Small Language Model from Scratch and the seeded, immutable shards from Pretraining Data for Small Language Models. The underlying mathematics of each mechanism has deeper treatments elsewhere on this site, for example learning-rate schedules and gradient clipping and stability; here the concern is how they fit together.

Advertisement

The loop at a glance

One optimizer step, and the state that must survive a restartLoaderstep, seed to batchForward (bf16)fp32 norms, lossBackwardaccumulate K microClipglobal normAdamWfp32 masterSpike monitorloss vs rolling median, grad normRecovery policyskip step / rewind + skip windowspikeCheckpoint writerasync, atomic, retainedrewindCheckpoint contentsweights (fp32 master)optimizer m and vstep, schedule stateloader position, seedsRNG states, all ranksloss scaler if fp16skip list of bad windowsconfig, tokenizer, manifestMissing any one of these makes a resume diverge from the run it is supposed to continue.
One optimizer step and the state a checkpoint must hold. The spike monitor sits beside the step and can veto an update or trigger a rewind.

The loop below is the skeleton the rest of the article fills in. Its most important property is that nothing in it depends on history other than through explicit state: the batch is a function of the step and seed, the learning rate is a function of the step, and everything else lives in the checkpoint. That property is what makes resume exact and rewind possible.

def train(cfg, model, manifest, ckpt_dir):
    state = load_latest(ckpt_dir) or init_state(cfg, model)
    opt = build_adamw(model, cfg)                    # param groups below
    restore(model, opt, state)                       # weights, moments, RNG
    loader = DeterministicLoader(manifest, cfg.seed, skip=state.skip_list)
    monitor = SpikeMonitor(cfg.spike)

    while state.step < cfg.total_steps:
        lr = lr_at(state.step, cfg) * state.lr_scale # pure function of step
        set_lr(opt, lr)
        loss, n_tok = accumulate(model, loader.batch(state.step), cfg)
        gnorm = clip_global_norm(model.parameters(), cfg.clip)
        verdict = monitor.check(state.step, loss, gnorm)
        if verdict == "skip_step":
            opt.zero_grad(set_to_none=True)          # drop this update only
        elif verdict == "rewind":
            state = rewind(ckpt_dir, state, monitor.window())
            restore(model, opt, state)               # not just the optimizer
            loader = DeterministicLoader(manifest, cfg.seed, skip=state.skip_list)
            continue
        else:
            opt.step(); opt.zero_grad(set_to_none=True)
        state.step += 1
        log(state.step, loss=loss, gnorm=gnorm, lr=lr, tokens=n_tok)
        if state.step % cfg.ckpt_every == 0:
            save_async(ckpt_dir, state.snapshot(model, opt))

Optimizer: AdamW and parameter groups

AdamW remains the default for pretraining small models. The settings that recur across public recipes are β1 around 0.9, β2 around 0.95 rather than the 0.999 library default, weight decay around 0.1, and a small epsilon such as 1e-8. The lower β2 makes the second-moment estimate adapt faster when gradient statistics shift, which reduces the chance that a sudden large gradient meets a stale, too-small denominator and produces an oversized update, one of the mechanisms behind loss spikes.

Weight decay should apply to matrices, not to everything. Norm gains, biases if present, and often the embedding table are excluded, because decaying a norm gain toward zero fights the normalization itself and decaying embeddings penalizes rare tokens that receive few gradient updates. Getting the groups wrong does not crash anything, so assert it.

def build_adamw(model, cfg):
    decay, no_decay = [], []
    for name, p in model.named_parameters():
        if p.ndim < 2 or "norm" in name or name.endswith(".bias") \
           or ("embed" in name and cfg.no_decay_embeddings):
            no_decay.append(p)
        else:
            decay.append(p)
    assert len(decay) + len(no_decay) == len(list(model.parameters()))
    return AdamW([{"params": decay, "weight_decay": cfg.wd},       # ~0.1
                  {"params": no_decay, "weight_decay": 0.0}],
                 lr=cfg.peak_lr, betas=(0.9, 0.95), eps=1e-8)

Budget memory for it. With bf16 compute, fp32 master weights and fp32 Adam moments, training state costs roughly 16 bytes per parameter before activations: 4 for master weights, 8 for the two moments, 2 for bf16 weights and 2 or more for gradients. A 500M-parameter model therefore needs about 8 GB for state alone, which fits one accelerator comfortably; sharding optimizer state across data-parallel ranks is an option rather than a necessity at this scale.

Advertisement

Learning-rate schedule as a pure function

Two schedule families dominate. Warmup followed by cosine decay to a floor, often a tenth of the peak, is the long-standing default, but it bakes the total step count into every step: stop early or extend the run and the schedule is wrong. Warmup-stable-decay holds the peak rate for most of training and then decays over a final fraction, commonly somewhere between a tenth and a fifth of the run. It lets you branch a decay from any stable-phase checkpoint, which suits the cooldown experiments described in the first article and makes extending a run straightforward.

Whichever you use, write it as a pure function of the step and the config, never as an object that counts calls. A stateful scheduler that is not restored with the checkpoint, or that is stepped once per micro-batch instead of once per optimizer step, is a classic silent bug: after a resume the learning rate jumps back to warmup, or the run finishes its decay at a quarter of the way through.

def lr_at(step, c):
    if step < c.warmup:                              # linear warmup
        return c.peak_lr * (step + 1) / c.warmup
    if c.schedule == "cosine":
        t = (step - c.warmup) / max(1, c.total_steps - c.warmup)
        return c.min_lr + 0.5 * (c.peak_lr - c.min_lr) * (1 + cos(pi * min(t, 1.0)))
    decay_start = c.total_steps - c.decay_steps      # warmup-stable-decay
    if step < decay_start:
        return c.peak_lr
    t = (step - decay_start) / c.decay_steps
    return c.min_lr + (c.peak_lr - c.min_lr) * (1 - t)   # linear or 1-sqrt

Warmup is usually a few hundred to a few thousand steps; too short a warmup shows up as an early spike.

Mixed precision: where bf16 goes and where it does not

Run matrix multiplications in bf16 and keep the numerically sensitive pieces in fp32: the master copy of the weights and the optimizer moments, the softmax in attention, RMSNorm's variance computation, and the final logits and cross-entropy. bf16 has the same exponent range as fp32, so gradients do not underflow and no loss scaling is needed. If hardware forces fp16, a dynamic loss scaler is mandatory, its state belongs in the checkpoint, and steps where it detects overflow must be skipped, not applied.

A small auxiliary loss on the output logits helps stability at modest cost. The z-loss adds a penalty proportional to the squared log of the softmax normalizer, with a coefficient around 1e-4, which discourages logits from drifting to large magnitudes where bf16 rounding and softmax saturation cause trouble. Normalizing queries and keys before the attention product is a related architectural remedy for attention-logit growth, if your model uses it.

def loss_fn(model, x, y, c):
    with autocast(dtype=bf16):                       # matmuls in bf16
        logits = model(x)                            # norms and softmax upcast
    logits = logits.float()
    ce = cross_entropy(logits.view(-1, V), y.view(-1),
                       ignore_index=PAD, reduction="sum")
    log_z = logsumexp(logits, dim=-1)                # softmax normalizer
    z = c.z_loss * (log_z ** 2).sum()                # ~1e-4
    n_tok = (y != PAD).sum()
    return ce + z, n_tok

Accumulation and clipping, in the right order

When the global batch does not fit in memory, accumulate gradients across micro-batches before stepping. The subtle part is normalization. Averaging the mean loss of each micro-batch weights every micro-batch equally, even when they contain different numbers of real target tokens, which happens with padding, masking or variable-length final sequences. Summing token losses and dividing by the total token count across the whole accumulation window gives every token the same weight regardless of how it was grouped. With data parallelism, sum the token count across ranks, and scale by world size if the framework averages gradients.

Clip after all micro-batches have been accumulated and, for fp16, after unscaling; clip the global norm across all parameters, not per tensor, so the direction of the update is preserved. A threshold of 1.0 is common. Log the pre-clip norm every step: it is the earliest signal of trouble, and the spike monitor below reads it.

def accumulate(model, micro_batches, c):
    total_tok = all_reduce_sum(sum(count_targets(y) for _, y in micro_batches))
    total_loss = 0.0
    for i, (x, y) in enumerate(micro_batches):
        with no_grad_sync(model, enabled=i < len(micro_batches) - 1):
            loss_sum, _ = loss_fn(model, x, y, c)
            (loss_sum / total_tok).backward()        # token-weighted mean
        total_loss += loss_sum.item()
    return all_reduce_sum(total_loss) / total_tok, total_tok

Checkpoints that actually resume

A checkpoint is resumable only if restoring it reproduces the run. That requires the fp32 master weights, both optimizer moments, the optimizer step count used for bias correction, the training step, the loss scaler state if fp16, the data loader position, the random number generator states for every rank and device, and the skip list of bad data windows. It should also carry the config hash, tokenizer hash and shard manifest hash, and the loader should refuse to resume if any of them disagree with the current run.

Write checkpoints atomically: serialize to a temporary directory, fsync, write a small manifest with hashes of every file, then rename the directory into place. A reader treats a checkpoint without a valid manifest as absent. Copying state to host memory and persisting it on a background thread keeps saves from stalling training. Retain checkpoints on a schedule such as the last few, plus one every so many thousand steps, plus the last checkpoint before any decay phase, because rewinds and cooldown branches need older points.

def save_async(root, snap):                          # snap already on host
    def work():
        tmp = root / f".tmp-{snap.step}"
        write_tensors(tmp / "model.bin", snap.model)
        write_tensors(tmp / "optim.bin", snap.opt)
        write_json(tmp / "state.json", dict(
            step=snap.step, rng=snap.rng, skip_list=snap.skip_list,
            lr_scale=snap.lr_scale, rewinds=snap.rewinds, scaler=snap.scaler, cfg_hash=snap.cfg_hash,
            tok_hash=snap.tok_hash, manifest_hash=snap.manifest_hash))
        write_json(tmp / "MANIFEST", file_hashes(tmp))
        fsync_dir(tmp)
        rename(tmp, root / f"step-{snap.step:08d}")  # atomic publish
        apply_retention(root)
    background.submit(work)

Test resume on every new setup: run to step N, kill the process, resume, and compare the loss curve to an uninterrupted run. Nondeterministic kernels rule out bitwise equality, but the curves should overlap within noise. A visible jump at the resume point means some state is missing.

Loss-spike detection and recovery

Loss spikes are sudden jumps in training loss, sometimes recovering within a few hundred steps, sometimes leading to divergence. Proximate causes include a batch of unusual data meeting a stale second-moment estimate, and logits growing until softmax saturates. Prevention comes first: warmup, β2 around 0.95, z-loss, clipping, and a peak learning rate chosen by ablation rather than optimism. Recovery is for what gets through.

Detection compares the current loss with a robust rolling baseline such as the median over the last few hundred steps, scaled by the median absolute deviation, and watches the gradient norm, which often jumps a step or two before the loss does. The response should be policy, written in code, and graded. An isolated outlier step is skipped: gradients are discarded and the step counter advances. A sustained spike, one that persists for more than a few steps or reaches a hard threshold, triggers a rewind: reload the last checkpoint from before the spike began, add the data window that coincided with it to a persisted skip list, and continue. If the same region spikes again after the skip, the cause is probably the optimizer state or learning rate rather than the data, and the policy lowers the peak rate or escalates to a human.

class SpikeMonitor:
    def check(self, step, loss, gnorm):
        base, mad = self.hist.median(), self.hist.mad()
        z = (loss - base) / max(mad, 1e-3)
        if not isfinite(loss) or z > self.c.hard_z:
            return self.escalate(step)               # NaN or huge: rewind now
        if z > self.c.soft_z or gnorm > self.c.gnorm_mult * self.gn.median():
            self.consecutive += 1
            return "skip_step" if self.consecutive <= self.c.max_skips \
                   else self.escalate(step)
        self.consecutive = 0
        self.hist.push(loss); self.gn.push(gnorm)    # only clean steps
        return "ok"

def rewind(ckpt_dir, state, window):
    ck = latest_before(ckpt_dir, window.start - SAFETY_STEPS)
    new = load(ck)                                   # model, opt, rng, step
    new.skip_list = state.skip_list + [window]       # persisted in checkpoints
    new.rewinds, new.lr_scale = state.rewinds + 1, state.lr_scale
    if new.rewinds_in_region(window) >= 2:
        new.lr_scale *= 0.7                          # data was not the cause
    alert(f"rewind to {ck.step}, skipping {window}")
    return new

Two details matter. The skip list must live in the checkpoint and the loader, so that a later resume from any checkpoint does not replay the bad window. And skipped steps should still advance the step counter, so the learning-rate schedule and the data sequence stay aligned with the plan.

What to watch

Log per step: loss, pre-clip gradient norm, learning rate, tokens per second, and the fraction of steps clipped. Log per few hundred steps: validation loss per domain from the held-out splits, the maximum attention and output logit magnitudes, parameter and update norms per layer group, and the ratio of update norm to parameter norm, which should stay small and stable. Alert on non-finite values, sustained throughput drops, and validation loss rising on one domain while falling elsewhere, which usually means a data problem.

Failure modes

  • Stateful scheduler not restored. The learning rate returns to warmup after every resume.
  • Mean-of-means accumulation. Micro-batches with fewer tokens get extra weight and the effective loss is biased.
  • Decaying norm gains. Weight decay applied to every parameter quietly degrades the model.
  • Non-atomic checkpoints. A crash during save leaves a truncated file that looks like the latest checkpoint.
  • Unpersisted skip list. A later resume replays the data window that caused the spike.
  • Manual spike handling. Each incident gets a different ad-hoc fix, and the run cannot be reproduced.
  • Clipping as the only defense. A too-high learning rate is masked until the run diverges late.

Trade-offs

Frequent checkpoints shorten rewinds and lose less work on failure, at the cost of storage and some write bandwidth, which async saving mostly hides. Aggressive spike detection skips healthy but noisy steps; lax detection lets real spikes contaminate the optimizer state before a rewind. A lower peak learning rate reduces spikes but trains more slowly, which is why rewind machinery is worth building. Warmup-stable-decay adds flexibility for cooldowns and extensions; cosine is simpler when the budget is fixed and certain.

Key takeaway: Write the loop so that the batch and the learning rate are pure functions of the step, use AdamW with beta2 near 0.95 and no decay on norms, keep master weights, softmax, norms and the loss in fp32 with a small z-loss, normalize accumulated gradients by total tokens and clip the global norm after accumulation, save complete checkpoints atomically and asynchronously and prove resume works, and handle loss spikes with a coded policy that skips isolated steps, rewinds past sustained ones and persists the skipped data window.