The number that grows before the loss spikes
Loss spikes in transformer runs are preceded by growth of the largest attention logit. QK-norm, z-loss and QK-clip bound that quantity, so log it rather than waiting for the loss.
A loss spike in a transformer training run looks like an accident. The curve is descending, then in a few hundred steps it jumps, sometimes recovers, sometimes never does, and the post-mortem blames the learning rate, the data, the precision, or bad luck. It is rarely an accident. In the runs where the mechanism has been studied, the spike is the visible end of a process that started long before: the largest attention logit, the dot product of a query and a key divided by the square root of the head dimension, drifted upward until the softmax that consumes it saturated, and a saturated softmax passes back a gradient that is either nothing or enormous. The loss is the last thing to move. The logit moved first, and it was loggable the whole time.
The mechanism
Attention takes a query vector and a set of key vectors, computes their dot products, scales, and pushes the result through a softmax to get weights. The softmax is well behaved when the logits are of moderate size and badly behaved when one logit is much larger than the rest, because the weights then collapse onto that one key and the gradient with respect to every other key goes to zero, while the gradient with respect to the winning logit itself can become large and erratic. Nothing in the standard architecture bounds the size of the logits. The query and key projections are linear layers, and if their weights grow, or the residual stream feeding them grows, the logits grow with them, quadratically, since both factors of the dot product scale.
Wortsman and colleagues made this reproducible at small scale in Small-scale proxies for large-scale transformer training instabilities: the instability that large runs hit appears in small models trained at high learning rates, the growth of attention logits is one of its two documented sources, and the same mitigations that work at scale work in the small proxies, which is what makes the mechanism testable on a single machine. Their most useful finding for a practitioner is that the instability can be predicted before it emerges, from the scaling behaviour of activation and gradient norms, which is the same claim as this post's in more general form: the numbers that move first are not the loss.
How large is too large
The threshold at which a softmax saturates is a matter of arithmetic, and it depends on the context length. With n keys and a gap g between the largest logit and the rest, the weight on the largest key is about one over one plus n times e to the minus g, and the total weight left for every other key is about n times e to the minus g. Saturation, in the sense that the other keys together receive less than some small fraction ε, happens when g exceeds the natural log of n over ε. For a thousand keys and ε of one in a thousand, that is a gap of about 14; for a context of 128,000 and the same ε, about 19. A maximum logit in the twenties, with typical logits near zero, means the softmax has already stopped distributing gradient.
That derivation gives the logit runaway check its threshold. Log, every N steps and per layer, the maximum over the batch of the scaled query-key dot product. Compare it to the saturation gap for your context length, which is about 14 to 19 at the sizes anyone trains today, and treat a sustained approach toward that number as the warning it is. In a healthy run the maximum logit sits well below the threshold and stays there; in a run heading for a spike it climbs, steadily, over thousands of steps, with the loss curve giving no sign until the softmax has already saturated. The check costs nothing at training time, because the logits are computed anyway, and it moves the alarm from the loss, which fires after the damage, to the cause, which fires before.
Three fixes for one quantity
Once the quantity is named, the fixes that the field has converged on read as three ways of bounding it, and the diagram above places each. QK-norm applies a normalisation to the query and key vectors before the dot product, so that their magnitudes cannot grow and the logit is bounded by a learned scale; it was introduced for the largest vision transformers and it is now standard in many open language models. The z-loss, from the PaLM training recipe, adds a small penalty on the log of the softmax's normaliser at the output layer, which keeps the output logits from drifting to large values and, by the same mechanism, keeps the output softmax out of saturation. And QK-clip, described in the Kimi K2 technical report as the technique that lets its MuonClip optimiser train through 15.5 trillion tokens with no loss spike, watches the maximum attention logit directly and rescales the query and key projection weights whenever it exceeds a threshold, which is the runaway check with the response built in.
Why the loss is the wrong alarm
It is worth being explicit about why the loss curve fails as a monitor, because most training dashboards have it as the only one. The loss is an average over every token in the batch, and a softmax that has saturated in one head of one layer changes the prediction on a small fraction of tokens, which moves the average by less than the batch-to-batch noise. The run looks healthy. Meanwhile the gradient flowing back through that head has become unreliable, the optimiser state for the query and key projections accumulates a bad direction, and the weights grow along it, which makes the next batch's logits larger still. The loop closes over thousands of steps and the loss only reacts when enough heads have gone that the average moves, at which point the optimiser state is already poisoned and a rollback of a few thousand steps is the cheapest repair.
The maximum logit sees the same loop from the inside. It is not an average; it is the extreme value of exactly the quantity the softmax is sensitive to, per layer, and it rises monotonically through the loop rather than waiting for the loop's consequences to reach the average. That is the whole reason to log it. A monitor on the extreme catches a fault in one head; a monitor on the mean catches it once it has spread.
What to actually do with the number
The practice is short. Log the per-layer maximum attention logit at the same cadence as the loss, on the same dashboard, with the saturation gap for the run's context length drawn as a line. Alert on the logit crossing, not on the loss, and treat the alert as a reason to act before the spike rather than a curiosity: lower the learning rate, apply the clip, or, if the architecture allows a change mid-run, turn on the normalisation. Keep the logs for post-mortems, because a spike that arrives with no logit growth beforehand is a different failure, a data or a precision problem, and the absence of the warning is as informative as its presence.
The reason to prefer the log to the fix is that the fixes are not free. QK-norm changes the architecture and the checkpoints; the z-loss changes the objective; the clip changes the optimiser. Each is a decision a team should make with evidence, and the evidence is the number that grows before the loss spikes. A run that has never approached the saturation gap does not need any of them. A run that approaches it every few thousand steps needs one of them, and the log is how you find out which run you are in before the loss tells you the expensive way.
Get new posts by email
Occasional essays on engineering, AI, and building for the people technology leaves behind.
Subscribe with RSS