Chapter 16
Normalization, Residuals & Precision
BatchNorm, LayerNorm, RMSNorm, residual streams, dropout, and bf16/fp8 arithmetic.
Deep networks are easier to train when activations stay at a predictable scale, gradients have short paths, and arithmetic does not overflow. Modern LLM blocks combine normalization, residual streams, dropout or other regularization, and mixed precision to make very large matrix stacks trainable. This chapter derives LayerNorm and RMSNorm, then connects them to residual design and low-precision training.
16.1 BatchNorm, LayerNorm, and RMSNorm
BatchNorm normalizes each feature across the minibatch during training [ioffe2015batch]: . It updates running means and variances, then uses those running statistics at inference time. That train/inference split is useful in convolutional nets but awkward for autoregressive language models, where batch composition, sequence lengths, and generation-time batch size change.
The axis choice is the whole difference. BatchNorm asks, "for this feature, what did the current batch look like?" LayerNorm asks, "for this example, what is the scale of its feature vector?" The second question has the same answer whether the example is trained alone, packed with other sequences, or decoded one token at a time. That is why LayerNorm-style methods fit sequence models so naturally.
LayerNorm instead normalizes across the last dimension of each example [ba2016layer]. For a row ,
RMSNorm removes the mean subtraction and normalizes by root mean square [zhang2019root]:
The learned gain and bias, or gain alone for RMSNorm, let the model choose the scale after normalization. The backward pass is a row-wise reduction. For LayerNorm, let . Then
where is the row standard deviation. RMSNorm uses the same idea but subtracts only the component introduced by the RMS denominator. The tests gradient-check both backward passes.
Two details are easy to miss in implementations. First, the parameter gradients reduce over every axis except the last one, because the same gain and bias are reused for all examples and positions. Second, the input gradient must have zero row sum for the mean-subtraction part of LayerNorm; adding the same constant to every coordinate before LayerNorm changes nothing, so the backward pass cannot send gradient in that direction.
def layer_norm_forward(x, gamma, beta, eps=1e-5):
mean = x.mean(axis=-1, keepdims=True)
centered = x - mean
variance = np.mean(centered * centered, axis=-1, keepdims=True)
inv_std = 1.0 / np.sqrt(variance + eps)
normalized = centered * inv_std
out = normalized * gamma + beta
cache = (normalized, inv_std, gamma)
return out, cache
def layer_norm_backward(dout, cache):
normalized, inv_std, gamma = cache
axes = tuple(range(dout.ndim - 1))
grad_gamma = np.sum(dout * normalized, axis=axes)
grad_beta = np.sum(dout, axis=axes)
grad_norm = dout * gamma
width = dout.shape[-1]
sum_grad = np.sum(grad_norm, axis=-1, keepdims=True)
sum_grad_norm = np.sum(grad_norm * normalized, axis=-1, keepdims=True)
dx = inv_std * (grad_norm - sum_grad / width
- normalized * sum_grad_norm / width)
return dx, grad_gamma, grad_beta
def rms_norm_forward(x, weight, eps=1e-8):
mean_square = np.mean(x * x, axis=-1, keepdims=True)
inv_rms = 1.0 / np.sqrt(mean_square + eps)
normalized = x * inv_rms
out = normalized * weight
cache = (x, normalized, inv_rms, weight)
return out, cache
def rms_norm_backward(dout, cache):
x, normalized, inv_rms, weight = cache
axes = tuple(range(dout.ndim - 1))
grad_weight = np.sum(dout * normalized, axis=axes)
grad_norm = dout * weight
width = dout.shape[-1]
dot = np.sum(grad_norm * x, axis=-1, keepdims=True)
dx = grad_norm * inv_rms - x * (inv_rms ** 3) * dot / width
return dx, grad_weight
16.2 Residual paths and where to normalize
A residual block adds a learned transformation to its input:
The gradient contains an identity path, , so a deep stack can pass signal backward even when is poorly conditioned. Residual connections made very deep vision networks trainable [he2015deep] and are now standard in Transformer blocks.
The residual stream is also a storage convention: each block writes an update into a shared state vector rather than replacing the whole representation. If an update is useful, later blocks can build on it; if it is not, the identity path lets the model learn a small branch and leave the stream mostly unchanged. That makes initialization and optimizer mistakes less catastrophic than in a pure stack with no skips.
Normalization can sit after the residual addition (post-norm) or before the sublayer (pre-norm). Post-norm writes , so the identity path passes through a normalization Jacobian. Pre-norm writes , leaving the residual stream itself as the direct path; this usually improves gradient flow in deep Transformers [xiong2020layer]. The trade-off is that the residual stream’s scale is less directly controlled, so implementations often add final normalization before the output head.
Dropout randomly zeros activations during training. Inverted dropout divides the kept values by the keep probability, so and inference can use the identity function. It is a regularizer, not a normalization layer: it changes the noise in the training computation but does not estimate feature statistics.
Use dropout only in training mode. At inference time randomness would make the same prompt produce different hidden states before sampling even begins, and scaling would no longer match the expected training computation. In many large-data LLM pretraining runs dropout rates are small or zero, but the mechanism remains important for smaller data, fine-tuning, and models outside language modeling.
def batch_norm_forward(x, gamma, beta, running_mean, running_var, training,
momentum=0.9, eps=1e-5):
if training:
mean = x.mean(axis=0)
var = x.var(axis=0)
running_mean *= momentum
running_mean += (1 - momentum) * mean
running_var *= momentum
running_var += (1 - momentum) * var
else:
mean = running_mean
var = running_var
normalized = (x - mean) / np.sqrt(var + eps)
return normalized * gamma + beta
def inverted_dropout(x, drop_probability, rng):
if not 0 <= drop_probability < 1:
raise ValueError("drop_probability must be in [0, 1)")
keep = rng.random(x.shape) >= drop_probability
return x * keep / (1 - drop_probability), keep
16.3 Mixed precision
Mixed precision keeps the fast path low precision while preserving sensitive state in float32. Float16 has a small exponent range: its largest finite value is 65,504, so large activations, losses, or gradients can overflow. bfloat16 keeps float32’s 8-bit exponent range but has fewer fraction bits, so it preserves range while rounding more coarsely; Appendix B shows the bit layout in Section B.6.
Range and precision are different failure modes. Overflow turns a finite value into infinity and usually destroys the step. Rounding error is quieter: the value stays finite, but small updates disappear because they do not change the stored low-precision number. bfloat16 mostly solves the first problem and makes the second more visible.
Loss scaling protects small fp16 gradients. Multiply the loss by a scale before backpropagation, compute scaled gradients, then divide the gradients by the same scale before the optimizer step. This moves tiny values into fp16’s representable range without changing the mathematical update, unless an overflow is detected and the step is skipped.
The other rule is to keep master weights and optimizer accumulators in float32. A low-precision copy can be used for matrix multiplies, but small updates should accumulate into the float32 master; otherwise rounding can erase them. Dot products and reductions are also commonly accumulated in float32, because summing many rounded terms compounds error [micikevicius2017].
This is why mixed precision is an engineering pattern, not just a dtype switch. Parameters, activations, gradients, reductions, optimizer moments, and communication buffers can each use a different format. The safe default for a scratch implementation is simple: do matmuls in the low-precision format being studied, but keep losses, reductions, master weights, and optimizer state in float32 unless a test proves the lower precision is safe.
def fp16_overflows_but_bfloat16_keeps_range():
large = np.array([1e5, 1e30], dtype=np.float32)
fp16 = large.astype(np.float16).astype(np.float32)
bf16 = round_to_bfloat16(large)
return fp16, bf16
def loss_scaled_gradient(gradient, scale):
scaled = (gradient * scale).astype(np.float16)
unscaled = scaled.astype(np.float32) / scale
return unscaled
|
In practice
|
Decoder-only LLMs usually use pre-norm residual blocks with LayerNorm or RMSNorm rather than BatchNorm, because generation should not depend on other examples in a batch. Many recent architectures use RMSNorm for a cheaper normalization path. Mixed precision is standard: fp16 needs loss scaling, whereas bf16 often avoids it because it shares float32’s exponent range. |
16.4 Teach it
The one-sentence version: normalization controls activation scale, residuals preserve a direct gradient path, and mixed precision keeps arithmetic fast without trusting low precision with state. Analogy: normalization is a leveler, residuals are skip roads around traffic, and float32 master weights are the ledger while fp16 or bf16 are the cash register. Board steps: (1) compute LayerNorm mean and variance per row; (2) compare RMSNorm’s RMS denominator; (3) draw a residual identity gradient; (4) show loss scaling and unscaling. Misconceptions: BatchNorm’s training statistics are not inference statistics; dropout scaling belongs at training time for inverted dropout; bf16 has range, not fp32 precision. Check for understanding: why does pre-norm leave a cleaner gradient path than post-norm?
16.5 Exercises
Explain why BatchNorm needs running statistics for inference, and why LayerNorm does not. Which axes are reduced by each method for an array of shape ?
Derive the LayerNorm input gradient in (16.3) from the centered input and row variance. Why do two row means appear in the formula?
Show that RMSNorm without is invariant to multiplying one row by a positive constant. Then explain how the residual equation creates an identity gradient path.
Implement inverted dropout and a tiny mixed-precision demo: show that fp16 overflows on large values where rounded bfloat16 remains finite, and show how loss scaling recovers a tiny fp16 gradient.
References
-
[micikevicius2017] P. Micikevicius et al. Mixed precision training. ICLR 2018. arXiv:1710.03740
-
[ba2016layer] J. L. Ba, J. R. Kiros, and G. E. Hinton. Layer Normalization. 2016. arXiv:1607.06450
-
[he2015deep] K. He et al. Deep Residual Learning for Image Recognition. 2015. arXiv:1512.03385
-
[ioffe2015batch] S. Ioffe and C. Szegedy. Batch Normalization: Accelerating Deep Network Training by Reducing Internal Covariate Shift. 2015. arXiv:1502.03167
-
[xiong2020layer] R. Xiong et al. On Layer Normalization in the Transformer Architecture. 2020. arXiv:2002.04745
-
[zhang2019root] B. Zhang and R. Sennrich. Root Mean Square Layer Normalization. 2019. arXiv:1910.07467