A pretraining run is the most expensive experiment most teams will ever perform, and it is usually observed through a handful of noisy line charts. When the loss jumps at 205 billion tokens, the questions that matter are concrete: was it the data, the learning rate, a bad node, or a code change in the last restart? Experiment tracking is the system that lets you answer those questions from records instead of memory.
This article is about what to record for LLM training and how to keep the records coherent across preemptions, rollbacks and many ranks. It is deliberately tool-agnostic: the same design works with Weights & Biases, MLflow, or a home-grown store. For running an MLflow server see MLflow for GPU teams; for the run manifest and determinism see training reproducibility.
Why LLM runs need more than a dashboard
Tracking for a classifier fine-tune is a solved problem: one run, a few hyperparameters, a validation score. LLM training breaks that model in four ways.
- Runs are long and discontinuous. A multi-week job is preempted, restarted on different nodes, sometimes with a fixed bug. One logical run becomes many process lifetimes.
- Steps are not comparable. Batch size often ramps up during training, and two runs with different global batch sizes have different meanings for step 10,000. Tokens seen is the honest x-axis.
- The interesting events are rare. Loss spikes, gradient-norm explosions and throughput drops last a few hundred steps out of hundreds of thousands. Coarse logging averages them away.
- Logging itself can slow training. Calling
.item()on a CUDA tensor forces the host to wait for the GPU. Doing that for ten scalars every step serialises work that would otherwise overlap.
What to record, and how often
Decide the metric list before the run, group it by purpose, and give each group a cadence. The table is a sensible default for a dense decoder model; adjust names, not structure.
| Group | Metrics | Cadence | Why |
|---|---|---|---|
| Progress | tokens_seen, step, epoch fraction per data source | every log step | the x-axis and data-mix audit |
| Optimisation | loss, grad_norm (pre-clip), clip ratio, lr, loss scale | every 1-10 steps | spikes and divergence show here first |
| Throughput | tokens/s, step time, MFU, data-loader wait | every 10-50 steps | slow nodes and input stalls |
| Model health | per-layer grad norm, weight norm, update/weight ratio, attention logit max | every 100-500 steps | localises a spike to a layer |
| System | GPU memory peak, ECC/Xid count, NCCL timeouts, host | every 30-60 s | separates hardware faults from maths |
| Evaluation | held-out loss per domain, benchmark scores | per checkpoint | the quantity you actually ship |
Two definitions deserve precision. MFU (model FLOPs utilisation) is achieved model FLOPs per second divided by the hardware peak. For a dense transformer a common approximation is 6 × parameters × tokens per second, plus an attention term at long context. Log the formula version with the run, because changing it silently moves every chart. Gradient norm should be the global norm before clipping; the value returned by torch.nn.utils.clip_grad_norm_ is exactly that, so log the return value rather than recomputing it.
Tokens seen is the x-axis
Log tokens seen as a first-class metric and make it the x-axis for every training chart. With batch-size warm-up, runs A and B can reach step 50,000 at very different token counts; plotting by step makes the smaller-batch run look worse when it is merely earlier.
Most trackers let you declare a custom step metric. In W&B that is define_metric; in MLflow you pass the token count as the step argument (an integer) and keep the optimiser step as an ordinary metric:
# Weights & Biases: chart train/* and eval/* against tokens, not the internal step
run = wandb.init(project="pretrain-7b", id=run_id, resume="allow", config=cfg)
wandb.define_metric("tokens_seen")
wandb.define_metric("train/*", step_metric="tokens_seen")
wandb.define_metric("eval/*", step_metric="tokens_seen")
# MLflow: reuse the run across restarts; token count is the step
with mlflow.start_run(run_id=existing_run_id):
mlflow.log_metrics({"train/loss": loss, "train/opt_step": step},
step=tokens_seen, synchronous=False)The synchronous=False flag in recent MLflow versions queues the call in a background thread. The wrapper below does the same thing for any backend, which matters more than which tracker you pick.
Logging from many ranks without slowing the GPU
Every data-parallel rank computes its own loss; the number worth logging is the mean across ranks. The cheap pattern keeps everything on the GPU until a log step, does one fused all_reduce, then converts to Python numbers once, on rank 0 only, and hands them to a bounded background queue so a slow tracking server can never stall the step.
import queue, threading, torch, torch.distributed as dist
class Tracker:
"""Rank-0 async logger. backend.log(dict, step) is W&B, MLflow or your own."""
def __init__(self, backend, maxsize=10_000):
self.q, self.backend = queue.Queue(maxsize), backend
self.dropped = 0
threading.Thread(target=self._drain, daemon=True).start()
def log(self, metrics, step):
try:
self.q.put_nowait((metrics, step))
except queue.Full: # never block the training loop
self.dropped += 1
def _drain(self):
while True:
m, s = self.q.get()
m["tracker/dropped"] = self.dropped
try:
self.backend.log(m, step=s)
except Exception: # network blips must not kill training
self.dropped += 1
LOG_EVERY = 10
acc = torch.zeros(3, device="cuda") # sum_loss, sum_tokens, n_micro
for step, batch in enumerate(loader, start=start_step):
loss = train_step(model, batch) # no .item() inside the step
ntok = batch.num_tokens.float() # 0-d CUDA tensor: non-pad tokens
acc += torch.stack([loss.detach() * ntok, ntok, torch.ones_like(ntok)])
gnorm = torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
tokens_seen += global_batch_tokens
if step % LOG_EVERY == 0:
stats = torch.cat([acc, gnorm.reshape(1).float()])
dist.all_reduce(stats) # one collective, 4 numbers
if dist.get_rank() == 0:
s = stats.tolist() # the only host sync
tracker.log({"train/loss": s[0] / s[1],
"train/grad_norm": s[3] / dist.get_world_size(),
"train/lr": sched.get_last_lr()[0],
"tokens_seen": tokens_seen}, step=tokens_seen)
acc.zero_()Two details are easy to get wrong. Weight the loss by tokens, not by micro-batch, or packed and padded batches bias the mean. And grad norm is already global after clipping on every rank, so dividing the summed value by world size merely recovers it; with sharded optimisers use the framework's own clip function (FSDP exposes one) so the norm covers every shard.
Run lineage: segments and child runs
Treat a training run as a tree. The logical run has a stable ID chosen before launch and written into the checkpoint, so a restarted process resumes the same tracker run instead of creating a stranger. Each process lifetime is a segment, recorded with its start token count, node list, git commit and container digest. A child run is created whenever you change the recipe on resume: rolling back after a spike and lowering the learning rate is a different experiment and must not be drawn on the same line.
def open_run(ckpt, cfg):
meta = ckpt.meta if ckpt else {}
run_id = meta.get("run_id") or new_run_id()
recipe = recipe_hash(cfg) # lr schedule, data mix, model shape
if meta and meta["recipe"] != recipe: # recipe changed on resume -> child
parent, run_id = run_id, new_run_id()
cfg = {**cfg, "parent_run": parent, "forked_from_ckpt": meta["ckpt_id"]}
run = backend.open(run_id, config=cfg) # resume if it exists
run.log_segment(start_tokens=meta.get("tokens_seen", 0),
git=git_sha(), image=container_digest(),
nodes=hostnames(), world_size=dist.get_world_size())
return run, run_id, recipeResuming means one more subtle thing: the tracker already holds points beyond the last checkpoint, logged by the process that died. When you resume from 200B tokens, the server still shows the doomed 200B-205B stretch. Backends also differ on what happens next: W&B ignores (with a warning) any log whose step is lower than the run's current step, so re-logging 200B-205B into the same run ID is silently dropped, while MLflow accepts out-of-order steps and draws both stretches on top of each other. That is another reason a rollback should become a child run rather than a same-ID resume. With W&B, call wandb.log(m) without step= and let the define_metric step metric supply the x-axis. Never let two segments' points interleave on the same x values unmarked.
Evaluations keyed by checkpoint
Evaluations run asynchronously on saved checkpoints, often hours later and on other hardware. Key every eval result by a checkpoint ID (a hash of the weight files plus the token count) and log it to the training run at that checkpoint's tokens_seen. That gives the answer to the only question leadership asks, "which checkpoint is best and why", without a spreadsheet. Store the eval harness version and prompt template hash with the score; a changed few-shot template can move a benchmark more than another 50B tokens.
Spike and divergence alerts
Alerting turns tracking from a history book into a smoke detector. A robust spike detector compares the current loss with an exponential moving average and its deviation, and a second rule watches the gradient norm, which usually moves first:
class SpikeDetector:
def __init__(self, alpha=0.01, z=6.0, warmup=200):
self.mu = self.var = None; self.n = 0
self.alpha, self.z, self.warmup = alpha, z, warmup
def update(self, x):
self.n += 1
if self.mu is None:
self.mu, self.var = x, 0.0
return False
d = x - self.mu
spike = self.n > self.warmup and d > self.z * (self.var ** 0.5 + 1e-8)
if not spike: # do not learn from the spike itself
self.mu += self.alpha * d
self.var = (1 - self.alpha) * (self.var + self.alpha * d * d)
return spike
loss_det, gn_det = SpikeDetector(), SpikeDetector(z=8.0)
# on rank 0 at a log step, fed the Python floats already pulled from `stats`
loss_spike = loss_det.update(s[0] / s[1]) # evaluate both: `or` would skip one
gn_spike = gn_det.update(s[3] / dist.get_world_size())
if loss_spike or gn_spike:
tracker.log({"alert/spike": 1, "alert/batch_ids": batch_ids_hash}, step=tokens_seen)
pager.notify(run_id, tokens_seen)Log the identifiers of the batches around each alert. When a spike recurs after rollback at the same data position, the data is the suspect; when it moves, look at optimiser state and hardware. Also alert on throughput: a single slow GPU drags every synchronous step, and a 15% tokens/s drop is usually a node problem, not a model problem.
Budgeting metric volume
Budget the metric volume before launch. Illustrative numbers for a 7B run of 1 trillion tokens at 4M tokens per step (about 250,000 steps):
| Group | Series | Every | Points |
|---|---|---|---|
| Optimisation + progress | 8 | 10 steps | 200,000 |
| Throughput | 5 | 50 steps | 25,000 |
| Per-layer health (32 layers × 4) | 128 | 250 steps | 128,000 |
| System (64 GPUs × 4) | 256 | 60 s for ~30 days | 11,000,000 |
The system row dominates by fifty times. That is normal and points to the right split: send per-GPU telemetry to the monitoring stack you already run (DCGM exporter into Prometheus, for example) and log only aggregates such as minimum tokens/s and maximum memory into the experiment tracker. Trackers are built for comparing runs, not for storing per-device time series at high resolution.
Failure modes
- Every rank logs. 64 ranks each opening a run produce 64 ghost runs and rate-limit the server. Gate all tracker calls on rank 0 and set the tracker to disabled elsewhere.
- Host syncs in the hot loop. A
loss.item()every step for a progress bar can cost several percent of throughput. Profile one step with and without logging. - New run on every restart. Without a persisted run ID, a preempted job leaves dozens of fragments that nobody can stitch together. Store the ID in the checkpoint.
- Untracked recipe changes. Someone edits the data mix during a restart and the chart draws one continuous line. Hash the recipe and fork on change.
- Blocking network calls. A tracker outage turns into a training outage. Use a bounded queue, count drops, and log the drop count as a metric.
- Mismatched x-axes. Comparing step-based charts across batch-size schedules produces confident, wrong conclusions.
Trade-offs
Hosted trackers give the best comparison UI and zero operations, at the price of sending metrics and config off-site and of per-seat or per-volume cost. Self-hosted MLflow keeps data in-house but makes you run a database and artifact store. A thin wrapper like the Tracker class keeps you portable: the training loop calls one interface, and switching backends is a configuration change. Finer logging catches shorter events but costs throughput and storage; the cadence table is the compromise most teams converge on, with per-layer stats switched to every step for a short window after an alert.
What to do next
- Write the metric list and cadence table for your next run before launch, including the MFU formula version.
- Make tokens_seen the x-axis for every training and eval chart.
- Generate the run ID before launch and persist it in every checkpoint; log a segment record on each start.
- Hash the recipe and fork a child run whenever it changes on resume.
- Replace per-step
.item()calls with oneall_reduceevery k steps on rank 0, behind an async bounded queue. - Add loss and grad-norm spike alerts that record nearby batch IDs.
- Key eval results by checkpoint ID and log them at that checkpoint's token count.
- Route per-GPU telemetry to monitoring, and only aggregates to the tracker. Then read training checkpointing to make the resume path solid.