Normalizations
August 2026 – Vladislav KruglikovNormalization layers control activation scale. They make training less sensitive to initialization, learning rate, and changes in activation magnitude across layers.
Imagine each layer receives a vector of activations. During training, earlier layers keep changing, so the next layer sees inputs whose mean and variance drift around. That makes optimization harder: the layer is trying to learn while its input coordinate system keeps moving.
What normalization preserves
Normalization does not make every activation the same. It shifts and rescales activations into a standard coordinate frame.
For example, if an activation is a vector , LayerNorm first subtracts the mean of its coordinates, then divides by their standard deviation. The absolute location and overall size change, but the relative pattern remains: which feature is higher, which is lower, and by how much relative to the others.
This means vectors such as , , and become almost the same after normalization because they have the same internal pattern. But and remain different because the high and low features are swapped.
So normalization changes the question from «where is this activation in absolute space?» to «what is the pattern of highs and lows inside this activation?» Learnable scale and shift parameters can then put the normalized activation into whatever scale and offset are useful for the model.
Why learn a scale and shift if the next weights could adapt to standardized inputs? The problem is that fixed normalization forces the next operation to always receive zero-mean, unit-scale inputs, even if a different scale or mean would work better. Maybe a standard deviation of is better than , or maybe a positive mean is useful before a mean-sensitive activation such as sigmoid. For example, inputs with mean make sigmoid output almost , while inputs with mean make it output almost . If the next operation already has a bias, the normalization shift can be redundant because two consecutive shifts collapse into one. But when normalization is followed by an activation or a bias-free operation, the learned shift and scale give the model back control over the distribution it wants to use.
Batch norm
Batch normalization normalizes activations using statistics computed across the batch. In image models this works well because examples usually have the same shape. A batch of images has a regular tensor layout, so each channel statistic is estimated from many comparable positions across many examples.
For a feature or channel, BatchNorm computes the batch mean and variance, then applies a learned scale and shift:
where is the number of values used for the feature statistic.
This is a poor fit for language models. Text sequences have variable lengths, and token positions are not equally populated. Imagine a batch with one million examples where most sequences have length , but one sequence has length . The first token position has about one million examples contributing to its batch statistic. Positions through have only the long sequence contributing. Those later positions get a terrible estimate because their statistic is effectively computed from one example.
In distributed data parallel training, ordinary BatchNorm is local. Each worker computes its own mean and variance from its local mini-batch, so normalization needs no communication. This is fast, but small per-worker batches give noisy statistics, and the result can change when the same global batch is split across more workers.
SyncBatchNorm combines per-channel statistics across a group of workers, so they all normalize with statistics from the effective global batch. It communicates only summary statistics such as per-channel sums, squared sums, and counts, not full activations. This gives more stable statistics for small local batches but adds synchronization in forward and backward.
During evaluation, both normally use the running mean and variance collected during training. SyncBatchNorm changes how those running statistics are estimated. It does not need cross-worker synchronization during ordinary evaluation.
When per-worker batches are very small, common alternatives avoid batch statistics. GroupNorm normalizes channels within groups independently for each sample, with no cross-worker communication, and is common in detection and segmentation. LayerNorm and RMSNorm normalize the embedding dimension and are typical in Transformer and ViT models. InstanceNorm normalizes each sample and channel independently and is sometimes used for image generation and style transfer.
Layer norm
Layer normalization normalizes within each individual example, usually across the hidden dimension. It does not need other examples in the batch, so it works with variable sequence lengths and batch sizes.
For a hidden vector of size , LayerNorm computes statistics across its features:
where and are learned scale and shift vectors.
RMS norm
RMSNorm is a simpler variant of LayerNorm. Instead of subtracting the mean and dividing by the standard deviation, it normalizes by the root mean square of the hidden activations. It controls scale but does not recenter the activations.
For a hidden vector of size , RMSNorm is
where is a learned scale vector and prevents division by zero.
For example, let .
- mean
- standard deviation
- RMS
LayerNorm outputs approximately
while RMSNorm outputs approximately
RMSNorm deliberately preserves the vector's mean information while controlling its overall scale. This can be useful in Transformer residual streams, where removing the mean is not always necessary and may discard information.
This makes RMSNorm cheaper. Mean subtraction has low arithmetic intensity: it moves memory but does little math, so it is often memory-bound and does not use tensor cores well. Removing it can improve hardware utilization, especially in transformer blocks where small memory-bound operations can become visible overhead.
QK norm
Attention divides the query-key dot product by , where is the head dimension. If the products have zero mean and unit variance, their variances add:
so the dot product has standard deviation . Dividing by gives a logit with standard deviation approximately :
QK normalization additionally applies a per-token normalization, usually RMSNorm, to the query and key vectors before their dot product. It keeps their magnitudes bounded, preventing logits from growing too large and softmax from becoming overly sharp. Because it is computed independently for each token, it needs no batch statistics or cross-GPU communication.
The factor assumes that query and key components maintain roughly fixed, unit variance. Learned projections can violate that assumption, producing different vector norms across tokens or heads. Since
scaling by alone does not bound those norms. QK normalization controls them explicitly, making the dot product depend more on the angle between the vectors and keeping attention logits stable.
Before or after activation
Normalization can be placed before or after an activation such as ReLU, sigmoid, or GELU. These two choices mean different things.
If normalization is before the activation,
then the activation receives inputs with controlled scale and mean. This is useful for activations that are sensitive to input location. For example, sigmoid saturates when its input is very positive or very negative, so normalization before sigmoid can keep it in a range where gradients are still useful. The downside is that the activation can change the distribution again, so the output passed to the next layer is no longer guaranteed to be normalized.
If normalization is after the activation,
then the activation first changes the distribution, and normalization cleans up the result. For ReLU, this means the activation first clips negative values to zero, then normalization recenters and rescales the nonnegative output distribution.
There is no universal placement rule. Normalizing before the activation controls what the activation sees. Normalizing after the activation controls what the next layer sees.
Pre-norm and post-norm
In transformer blocks, normalization can be placed before or after the main sublayer. Pre-norm means the attention or MLP sees normalized input:
Post-norm means the normalization is applied after the residual update:
Pre-norm is usually easier to train in deep transformers because the residual path stays close to an identity path. Gradients can move backward through the residual connection without always passing through the normalization operation.
Post-norm can make the block output scale cleaner, but it is often less stable at depth. For modern LLMs, pre-norm or variants of it are the common choice.
QK normalization
QK normalization normalizes the query and key vectors before computing their dot product. Instead of using
attention uses normalized queries and keys:
Here, is a temperature that controls how sharp the attention distribution is.
The usual factor removes the growth caused by the head dimension. If the terms are roughly independent and each has variance , then
so dividing by leaves variance approximately . This removes the dependence on , but not on the scale . If query or key norms grow, can still become large. QK normalization controls those norms before the dot product.
This keeps attention logits from becoming too large and can make training more stable. Very large logit differences make the softmax nearly one-hot: one token receives almost all the attention and the others receive almost none. A sharp distribution is not inherently bad when that token is the right one, but it is brittle, cannot combine information from several tokens, and gives the softmax very small gradients when it saturates. If the selected token is wrong, the model also has little gradient signal to redirect its attention.
The trade-off is that the raw magnitudes of queries and keys no longer affect the dot product directly; a learnable scale or temperature can restore some of that control.