Part of Language AI Handbook
Explains how batch size affects gradient noise and generalization, apply the linear scaling rule for learning rates, identify the critical batch size.
Choose your expertise level to adjust how many terms are explained. Beginners see more tooltips, experts see fewer to maintain reading flow. Hover over underlined terms for instant definitions.
Article links
Make inline references clickable
Large Batch Training
Training a neural network means processing data in batches: you take a chunk of examples, compute gradients, and update weights. The batch size you choose is one of the most consequential decisions in the entire training setup. It shapes how quickly you learn, how stable the training is, how well the model generalizes, and how efficiently you use your hardware.
For most of deep learning's history, small to medium batch sizes (32 to 256) were the default. Researchers discovered through painful experience that simply cranking up the batch size made models train faster per epoch but converge to worse solutions. The relationship between batch size and model quality seemed to be a fundamental constraint. Many practitioners accepted this as an immutable law of optimization, choosing their largest tolerable batch size and moving on.
That changed in the mid-2010s, driven partly by practical necessity. As datasets grew to billions of tokens and models reached hundreds of billions of parameters, training on small batches became prohibitively slow. Distributing training across thousands of GPUs required large aggregate batch sizes just to keep the hardware busy. A single forward pass on 4,096 GPUs processes one mini-batch per step; if your batch size is 256, most of your hardware sits idle waiting for gradient synchronization. Researchers had to figure out how to train effectively at scales that were previously unthinkable.
The result was a collection of techniques, with the most important insight being the linear scaling rule: when you multiply the batch size by a factor , multiply the learning rate by the same factor . This rule, backed by empirical and theoretical analysis, made it possible to train models on massive batches without degrading quality. Combined with the LAMB optimizer and careful attention to batch size limits, large batch training became standard practice.
In this chapter, you will understand why batch size affects learning in the way it does, how the linear scaling rule works and why it is justified, where the fundamental limits of batch scaling lie, and how the LAMB optimizer extends these ideas to the pre-training regimes of BERT and GPT.
Why Batch Size Matters
Before examining the mechanics of scaling, it helps to understand what gradient descent computes and why batch size changes its character.
The Gradient and Its Estimator
The true objective you want to minimize is the loss averaged over your entire dataset:
where:
- : the full-dataset loss as a function of model parameters
- : total number of training examples
- : the loss on the -th example
The true gradient of this loss with respect to the parameters is:
Computing this exactly would require a pass over all examples for each parameter update. For datasets with millions or billions of samples, that is computationally intractable.
Stochastic gradient descent avoids this by computing a noisy estimate using a mini-batch of size :
where:
- : the mini-batch gradient estimate
- : the batch size
- : the set of randomly sampled example indices
This estimate is unbiased: its expected value equals the true gradient. But it has variance that depends inversely on :
where is the per-sample gradient variance. Doubling the batch size halves the gradient variance. With a larger batch, your estimate of the true gradient direction becomes more accurate. This is simply the central limit theorem applied to gradient estimation: the sample mean of independent random vectors has variance times the variance of a single sample.
Mini-batch stochastic gradient descent computes parameter updates from a small random subset of the training data rather than the full dataset. Each such subset is called a mini-batch. The update provides an unbiased but noisy estimate of the true gradient, trading computational efficiency for stochasticity.
Noise and the Quality of Updates
The noise in the gradient estimate plays a dual role, and understanding both sides is essential for making good batch size decisions.
On the positive side, gradient noise acts as a form of implicit regularization. Noisy updates prevent the optimizer from settling into sharp minima, those loss landscape features with high curvature in many directions. Sharp minima tend to correspond to solutions that overfit to training data. The parameters that minimize the training loss in a sharp minimum are very specific: small perturbations to the weights increase the training loss significantly. When you deploy the model, the test distribution differs slightly from training, and this difference looks like a perturbation to the weights. Sharp minima therefore generalize poorly.
Flat minima, where the loss is low and relatively insensitive to small perturbations, generalize better precisely because they tolerate the implicit shift between train and test. Gradient noise helps find these flat regions because the stochastic updates effectively prevent the optimizer from getting trapped in narrow, sharp basins. This connection between batch size, noise, and generalization was studied in depth by Keskar et al. (2017), who found that models trained with large batches consistently converged to sharper minima than those trained with small batches under the same budget.
On the negative side, noise means you sometimes step in the wrong direction. If the gradient noise is large relative to the true signal, you spend many updates wandering randomly rather than making consistent progress toward the minimum. With batch size 1, each update might point nearly opposite to the true gradient if you happen to sample an outlier example. The optimizer makes progress on average (the estimate is unbiased) but wastes many steps on unhelpful updates.
This tradeoff gives batch size its characteristic effect: small batches take many individually noisy steps, making slower but ultimately more stable progress, while large batches take fewer but more accurate steps, converging faster per epoch but risking poor generalization.
Curvature and the Role of the Learning Rate
The appropriate learning rate depends on the curvature of the loss landscape. In regions of high curvature (sharp narrow valleys), a large learning rate causes divergent oscillations as each update overshoots the local minimum and bounces to the other side. In regions of low curvature (flat plateaus), a small learning rate makes progress unnecessarily slow.
When you increase the batch size, you reduce gradient noise. This changes the effective signal-to-noise ratio of each update. If you keep the learning rate fixed while increasing the batch size, the optimizer takes fewer but more accurate steps per epoch. The effective step size relative to the gradient signal has changed, and the training dynamics shift. Specifically, with a fixed learning rate, each more-accurate large-batch step covers less distance in gradient-normalized space than the noisy small-batch step would have.
This is the core insight behind learning rate scaling rules: batch size and learning rate are not independent settings. They are coupled through the statistics of the gradient estimate. Increasing batch size without adjusting learning rate is like making more accurate measurements and then using fewer of them without increasing your confidence accordingly.
Historical Context and Motivation
Understanding where the linear scaling rule came from helps contextualize why it works when it does and fails when it does not. For most of the 2010s, practitioners treated batch size as primarily a hardware constraint. You used the largest batch that fit in GPU memory, typically 32 to 256, and accepted the associated learning rate.
The breakthrough came when distributed training made truly large batches possible. A cluster of 16 GPUs, each processing a batch of 32, produces an effective batch of 512 per step. A cluster of 512 GPUs produces an effective batch of 16,384. At these scales, the question of how to set the learning rate became urgent, because the naive answer (use the same learning rate as the single-GPU setting) clearly did not work.
The Facebook AI Research team's 2017 paper "Accurate, Large Minibatch SGD: Training ImageNet in 1 Hour" provided the systematic answer. They trained ResNet-50 on ImageNet using 256 GPUs in one hour, matching the accuracy of a week-long single-GPU run. The key contributions were the linear scaling rule, the warmup strategy, and careful attention to batch normalization behavior at large scale. This work catalyzed a wave of research and engineering that ultimately enabled the training of GPT-3, PaLM, and similar models at scales previously considered impossible.
The Linear Scaling Rule
The most influential contribution to large batch training was the systematic analysis of how learning rate should scale with batch size. The rule itself is simple to state, but its derivation illuminates exactly why it works and where it breaks down.
Derivation from SGD Dynamics
Consider SGD with batch size and learning rate . The derivation below assumes the loss is locally linear over the region swept by consecutive steps, which is the key approximation. After one step, parameters change by:
Now consider scaling the batch size by a factor to batch size . The new gradient estimate:
has variance , which is times smaller than the variance of .
The key intuition is this: with batch size , you need fewer steps to process the same amount of data. If you want the model after processing samples to be in the same state regardless of batch size, you need the cumulative weight update to be the same. One large-batch step with batch size should replicate what small-batch steps with batch size would have achieved.
More formally, steps of SGD with learning rate and batch size produce:
One step with batch size and learning rate produces:
When the loss landscape is locally linear (a first-order approximation), these two updates are equal in expectation, since :
This is the formal justification for the linear scaling rule.
When multiplying the batch size by , multiply the learning rate by . Formally, if the optimal learning rate for batch size is , then the optimal learning rate for batch size is .
Intuition Behind the Derivation
The derivation is mathematically clean, but the intuition is worth dwelling on. The argument assumes that consecutive small-batch updates do not change the gradient significantly, so the gradient computed on the first mini-batch is approximately the same as the gradient on the second, third, and so on. This is the local linearity assumption: the loss function behaves like a flat plane over the small region swept by SGD steps.
When this assumption holds, the small-batch steps are equivalent to a single step in the average gradient direction, scaled by . A single large-batch step computes this average gradient more accurately (lower variance) and takes a proportionally larger step (with ), arriving at the same expected location.
This framing also shows what can go wrong. If the model parameters change substantially over steps (high curvature, large gradients, or large learning rate), then the gradient estimated at step 1 points in a different direction from the true gradient at step . The equivalence breaks down, and simply scaling the learning rate linearly overestimates how large a step you can safely take.
Why the Approximation Breaks Down
The linear scaling rule is exact only in the limit where the loss is locally linear (equivalently, where the gradient does not change significantly within a single step). In practice, the loss landscape has curvature, and the approximation breaks down at two extremes.
At small batch sizes, the gradient noise is so high that the local linearity assumption is violated. The gradient estimated on one mini-batch points in a very different direction than the gradient estimated on the next, so the simple equivalence argument does not hold.
At very large batch sizes, a different problem emerges. As you keep scaling the learning rate up proportionally to batch size, eventually the learning rate becomes so large that a single step overshoots the local minimum. The update magnitude exceeds the scale over which the linear approximation is valid. This is why there is an upper limit to how far the linear scaling rule can be pushed.
The practical range where the rule works reliably varies by architecture and dataset. For ImageNet training with ResNets, it holds up to batch sizes of around 8,192. For large language models trained on massive corpora, the effective range depends on model size, architecture, and stage of training. The rule is a powerful starting point, not an infallible prescription.
The Square Root Scaling Alternative
Before the linear scaling rule was fully established, researchers proposed an alternative: the square root scaling rule, which suggests using learning rate when scaling batch size by . This rule comes from a different analysis of noise dynamics.
The argument runs as follows. The variance of the gradient estimate scales as . The standard deviation (noise amplitude) scales as . If you want to keep the signal-to-noise ratio of your update constant as you increase batch size, you should scale the learning rate by , because the noise has decreased by .
In practice, the square root rule tends to be too conservative for moderate batch size increases. The linear rule matches empirical results better for scaling up to 10x or 100x the original batch size. For very large scaling factors (beyond 1,000x), neither rule alone is sufficient, and more sophisticated techniques are needed. The linear rule has become the dominant standard, but knowing the square root rule exists helps you understand why some older papers use different scaling conventions.
Warmup as a Practical Necessity
At the start of training, the network parameters are far from their final values. The loss landscape in this early phase is highly non-linear, with large curvature in many directions. The linear scaling rule's underlying assumption of local linearity is worst at the beginning.
This is why large batch training almost always uses a warmup phase. As we discussed in the Learning Rate Warmup chapter, you start with a small learning rate and gradually increase it to the target value over some number of steps. For large batch training, warmup is often essential for stable convergence.
The Facebook AI Research paper found that jumping directly to the linearly scaled learning rate caused instability in the first epoch. Using a warmup of 5 epochs resolved this. The intuition is that during warmup, the network moves from a random initialization to a region of parameter space where the local linearity assumption begins to hold, making the scaled learning rate safe. Put differently, warmup gives the optimizer time to find a basin before it starts taking large steps inside that basin.
The optimal warmup duration scales roughly with the batch scaling factor . If you are scaling batch size by 8x, a warmup of 5 epochs may suffice. If you are scaling by 1,000x (as in large LLM training), you may need to warm up for thousands of steps. The general guideline is to warm up for at least 1-3% of total training steps, but this varies significantly across tasks.
Worked Example: Training with Different Batch Sizes
Let us work through a concrete example to make the scaling rule tangible. Suppose you have determined through experimentation that batch size trains well with learning rate using SGD.
You now want to scale to a 64-GPU cluster where, to keep all GPUs busy, you need to use at least 16,384 examples per step.
The scaling factor is:
The linearly scaled learning rate is:
For warmup, a common practice is to linearly increase the learning rate from a small initial value to the target over the first few epochs, often 5 epochs for ImageNet-scale tasks.
At the end of training, the model has processed the same number of gradient updates relative to data seen. If the original run did 100 epochs of 1,000 steps with batch size 256, the scaled run does 100 epochs of roughly 16 steps with batch size 16,384. The same amount of data is processed overall, but in far fewer, more accurate steps.
In practice, this specific example is near the edge of where the rule holds. For batch sizes beyond roughly 8,192 on ImageNet, researchers have found that extra techniques (gradient noise injection, LARS/LAMB, more aggressive warmup) are needed to maintain quality. This example sits at 16,384, so you would likely need at least LAMB or careful warmup tuning to close the generalization gap.
Batch Size Limits
The linear scaling rule describes how to adjust the learning rate when you scale up batch size. But there is a separate question: is there a limit to how large a batch you should use, regardless of learning rate scaling?
The answer is yes, and this limit is sometimes called the critical batch size.
The Critical Batch Size
The critical batch size is the point at which further increasing the batch size no longer improves the speed of convergence (measured in number of examples processed). Below this size, larger batches speed up convergence. Above it, you get no benefit from more parallelism.
Intuitively, this happens because beyond a certain batch size, the gradient estimate is already accurate enough. The bottleneck for learning is no longer gradient noise but the optimization geometry itself: the curvature of the loss landscape, saddle points, and movement across a complex high-dimensional surface. Adding more samples to a batch that is already providing a low-variance gradient estimate is simply redundant.
Think of it this way: if the true gradient points northeast, a batch that gives you an estimate pointing northeast-by-north is good enough. Doubling the batch size to get an estimate pointing northeast-by-northeast gives you a more accurate estimate but does not let you take a bigger step. You are still limited by the curvature of the loss surface, which determines the maximum safe step size.
Formally, the critical batch size is related to the ratio of gradient noise to gradient signal. McCandlish et al. (2018) formalized this as:
where:
- : the critical batch size
- : the Hessian of the loss (capturing curvature)
- : the per-sample gradient covariance matrix
- : the true gradient vector
- : the matrix trace operator
This expression measures the ratio of gradient noise (numerator) to gradient signal (denominator) in a curvature-aware way. The critical batch size is larger when the gradient noise is high relative to the signal, and smaller when the signal is strong.
The numerator, , captures the total noise in the gradient weighted by curvature. Noise in directions of high curvature matters more because those directions are the ones where accurate gradient estimation most strongly affects convergence. The denominator, , is the curvature-weighted gradient signal, measuring how much the gradient points in useful (high-curvature) directions. The ratio gives a curvature-sensitive signal-to-noise measure, and the critical batch size is the point at which increasing batch size no longer improves this ratio relative to the cost of the extra computation.
In practice, you cannot easily compute this quantity directly because you cannot compute the full Hessian for large models. But you can estimate the critical batch size empirically: train with progressively larger batches (adjusting the learning rate proportionally) and measure how many examples you need to process to reach a target validation loss. The batch size at which this number stops decreasing is approximately the critical batch size.
Empirical Observations
Research has found that the critical batch size varies widely across tasks and architectures:
- For image classification on CIFAR-10 with small models, critical batch sizes can be as low as a few hundred.
- For ImageNet with ResNet-50, the critical batch size is roughly 1,000 to 2,000 (much lower than the 8,192 that works with careful tuning, suggesting the extra scaling is not free from a convergence-per-example perspective).
- For language model pre-training (GPT-style), the critical batch size grows as training progresses. Early in training, when gradients are large and noisy, large batches help. Later in training, when the model is near convergence, smaller batches would suffice.
- For transformer-based language models, Kaplan et al. (2020) found that the critical batch size scales with the number of tokens processed, not just model size, suggesting a dynamic rather than fixed optimal batch size.
The implication for practitioners is important: there is an efficiency frontier. At batch sizes below the critical value, you are leaving GPU parallelism on the table. At batch sizes above it, you are wasting compute on redundant gradient averaging. Optimal large-scale training aims to hit this frontier.
However, hardware constraints often force you above the critical batch size anyway. When you have 10,000 GPUs, you need an enormous per-step batch just to utilize all the hardware. In these cases, the extra batch scaling beyond the critical size is a tax paid for hardware utilization. Engineers accept some inefficiency (in examples processed per useful gradient step) in exchange for faster wall-clock time.
The Generalization Gap Revisited
Even at or below the critical batch size, large batch training can produce models that generalize slightly worse than small batch training with the same total compute budget. This is the "generalization gap" of large batch training, documented by Keskar et al. (2017).
The mechanism is the sharpness of minima. Large batch training, with its accurate gradient estimates, tends to converge to sharp minima where the loss is low but curvature is high. Small batch training, with its inherent noise, naturally explores the loss landscape more and finds flatter minima. Flat minima generalize better because small perturbations to the weights (representing test distribution shift) cause only small increases in the loss.
Several approaches address this gap:
- Ghost batch normalization: Run batch normalization with a smaller "ghost" batch size even when processing a larger batch, preserving the regularization effect of BN statistics computed on smaller groups. This decouples batch normalization from the training batch size.
- Linear warmup: Helps the optimizer find a good basin before scaling to the full learning rate, giving it time to identify flat regions.
- Sharpness-aware minimization (SAM): Explicitly penalizes sharp minima during training, regardless of batch size, by seeking parameters whose neighborhood has uniformly low loss.
- Data augmentation: Increases effective noise through input perturbations, partially compensating for the reduced gradient noise of large batches. More aggressive augmentation with large batches can close much of the generalization gap.
- Gradient noise injection: Deliberately adding noise to gradients to simulate the effect of smaller batches, while keeping the large batch for hardware efficiency.
In practice, for language model pre-training at the scales used in production (batch sizes of millions of tokens), the generalization gap is usually small and manageable with standard techniques. The model capacity and data volume dominate, and the batch size effect becomes a secondary concern.
The LAMB Optimizer
The linear scaling rule works well for SGD and its momentum variants, but modern large-scale models are trained with adaptive optimizers like Adam. Scaling Adam to large batch sizes introduces a different challenge, and the LAMB optimizer was developed specifically to address it.
The Problem with Adam at Large Batch Sizes
Adam maintains per-parameter learning rates, scaling each parameter's update by the square root of its historical gradient variance. This adaptive rescaling lets Adam work reliably across diverse architectures and tasks. It handles the varying gradient magnitudes across parameters automatically, requiring less hyperparameter tuning than SGD.
When you scale the batch size with Adam and apply the linear scaling rule to the overall learning rate, something subtle goes wrong. Adam's adaptive update for parameter is roughly:
where:
- : the first moment (exponential moving average of gradients)
- : the second moment (exponential moving average of squared gradients)
- : a small constant for numerical stability
The problem is that the ratio is essentially normalized: it tells you which direction to move but not how far. The global learning rate controls the step magnitude, but it applies uniformly to all parameters.
Different layers in a deep network have very different gradient magnitudes. Embedding layers may have gradients many orders of magnitude smaller than the final output layer. When you apply a single global learning rate, you are either learning too fast in the final layers or too slowly in the embeddings. This tension becomes acute at large learning rates, which is what large batch training requires.
With small batches and modest learning rates, this imbalance is tolerable because the absolute update magnitude stays small enough that no layer diverges. But when you multiply the learning rate by a factor of 64 or 256, a layer that was marginally stable becomes explosively unstable. The layer-agnostic nature of the global learning rate, which Adam handles gracefully at moderate scales, becomes a critical failure mode at large learning rates.
There is also a subtler issue: the second moment estimate in Adam accumulates gradient variance over time. At the start of training with a large batch, the gradient estimates are already low-variance (that is the point of the large batch). This means the second moment estimates are small, making the factor large, amplifying the effective per-parameter learning rate at exactly the moment when stability is most critical.
LARS: The Predecessor for Convolutional Networks
Before LAMB, there was LARS (Layer-wise Adaptive Rate Scaling), developed for large batch training of convolutional networks. LARS scales the learning rate for each layer based on the ratio of the layer's weight norm to its gradient norm:
where:
- : the effective learning rate for layer
- : the global learning rate
- : the L2 norm of the parameters in layer
- : the L2 norm of the gradients for layer
The intuition is that the update should be proportional to the current weight scale of the layer. A layer with large weights can tolerate larger absolute updates without destabilizing. Conversely, a layer with small gradients relative to its weights is already learning slowly; the trust ratio gives it permission to take a larger effective step.
LARS also provides a natural safety check: if the gradient norm becomes very large (indicating an unstable update direction), the trust ratio shrinks automatically, reducing the effective learning rate for that layer. This self-limiting behavior prevents the cascading divergence that can occur when one layer's unstable gradients propagate through the network.
LARS worked well for training ResNets on ImageNet with very large batch sizes (up to 32,768 or more) without loss of accuracy. However, it was designed for SGD-style updates and did not incorporate the adaptive learning rate logic of Adam. For transformer architectures with weight tying, mixed precision, and complex gradient flows, pure SGD is rarely competitive with Adam. A solution that combined LARS's layer-level safety with Adam's per-parameter adaptivity was needed.
LAMB: LARS Meets Adam
LAMB (Layer-wise Adaptive Moments optimizer for Batch training), introduced by You et al. (2019) for BERT pre-training, combines Adam's adaptive per-parameter scaling with LARS's layer-wise trust ratio.
The LAMB update rule is:
where:
- : the gradient at step
- : the first moment (bias-corrected: )
- : the second moment (bias-corrected: )
- : exponential decay rates (typically 0.9 and 0.999)
- : the Adam update direction (before layer scaling)
- : the weight decay coefficient
- : the trust ratio (the LARS-style layer scaling)
The trust ratio scales each layer's update. When the Adam update direction has small norm relative to the weight norm, the trust ratio is large, allowing the layer to take bigger steps. When the update is large relative to the weights, the trust ratio shrinks, giving a safety net against destabilizing steps. This self-regulating property is what makes LAMB stable at the large learning rates required by large batch training.
LAMB (Layer-wise Adaptive Moments for Batch training) is an optimizer that combines Adam's per-parameter adaptive scaling with a per-layer trust ratio that keeps update magnitudes proportional to layer weight norms. It enables stable training with very large batch sizes (32,768 to 131,072 for BERT-scale models) by preventing any layer from taking an oversized step.
Why LAMB Enables Larger Batch Sizes
The key insight behind LAMB's effectiveness is that the trust ratio decouples the update direction from the update magnitude. Adam computes an excellent direction for each parameter (using the first and second moment estimates). LAMB then scales the magnitude of the resulting layer-wise update to be appropriate for that layer's current weight scale.
This two-level structure mirrors the intuition behind the linear scaling rule but operates at finer granularity. The global learning rate is scaled up proportionally with batch size. But each layer's effective step is additionally modulated by its own trust ratio, preventing any single layer from taking catastrophically large steps when is large.
The trust ratio also adapts over training. Early in training, when weights are small (near initialization) and gradients are large (the model is far from convergence), trust ratios tend to be small, limiting the step size. As training progresses and the model approaches a good solution, weights grow and gradients shrink, and the trust ratio naturally allows larger steps. This dynamic adaptation is not an explicit design choice but an emergent property of the weight-norm-to-gradient-norm ratio.
In the original BERT paper and follow-up experiments, LAMB enabled training BERT-base in 76 minutes on 1,024 TPU chips using a batch size of 65,536, compared to roughly 3 days on 16 TPU chips with batch size 256. The key was that LAMB maintained model quality (BERT's downstream task performance on GLUE benchmarks) across this dramatic batch size increase. Without LAMB, Adam at the equivalent linearly scaled learning rate was unstable, causing training divergence or significant accuracy degradation.
Comparing LARS, LAMB, and Adam
Understanding when to use each optimizer guides practical decisions:
- Adam is the default for most tasks at moderate batch sizes (up to a few thousand). It handles heterogeneous gradient magnitudes across parameters naturally and is well-understood. Use Adam unless you have specific large-batch scaling requirements.
- LARS is appropriate for convolutional networks with SGD-style training at large batch sizes. It is simpler than LAMB and well-validated on image classification tasks. If your architecture does not benefit from Adam's adaptivity (ResNets trained with SGD and cosine decay, for instance), LARS is a solid choice.
- LAMB is the right tool for transformer models trained with Adam at very large batch sizes. The combination of Adam's per-parameter adaptivity and LARS's layer-wise safety net matches the needs of architectures like BERT and GPT where both are required for stable large-batch training.
The trust ratio in LAMB is computed per parameter tensor (one trust ratio per weight matrix or bias vector), not per individual parameter. This is intentional: computing a trust ratio per parameter would approach Adam's behavior (the second moment already provides per-parameter scaling). The layer-level granularity gives LAMB its distinct character.
Code Implementation
We will implement and compare several training scenarios: standard SGD at small batch size, SGD at large batch size without scaling, SGD at large batch size with linear scaling, and LAMB. For tractability, we use a simple convolutional model on CIFAR-10.
Setup and Imports
Data Loading
We use CIFAR-10 with standard preprocessing. To make the experiments run quickly, we use a subset of the training data.
import numpy as np
import torch
import torchvision
import torchvision.transforms as transforms
from torch.utils.data import Dataset, Subset
torch.manual_seed(42)
np.random.seed(42)
transform = transforms.Compose(
[
transforms.ToTensor(),
transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5)),
]
)
class HFImageDataset(Dataset):
"""Adapt a Hugging Face image split to the (tensor, label) pairs DataLoader expects.
We source CIFAR-10 from the Hub rather than through
`torchvision.datasets.CIFAR10(download=True)`, whose only mirror is the
original university host — a single server with no CDN that commonly
delivers the 170 MB archive at well under 100 KB/s.
"""
def __init__(self, split, transform):
self.split = split
self.transform = transform
def __len__(self):
return len(self.split)
def __getitem__(self, idx):
record = self.split[idx]
return self.transform(record["img"].convert("RGB")), record["label"]
try:
from datasets import load_dataset
cifar = load_dataset("uoft-cs/cifar10")
full_train = HFImageDataset(cifar["train"], transform)
full_test = HFImageDataset(cifar["test"], transform)
except Exception:
full_train = torchvision.datasets.FakeData(
size=10000, image_size=(3, 32, 32), num_classes=10, transform=transform
)
full_test = torchvision.datasets.FakeData(
size=2000, image_size=(3, 32, 32), num_classes=10, transform=transform
)
# Use a compact subset so the complete experiment remains practical to rerun
train_subset = Subset(full_train, indices=range(1000))
test_subset = Subset(full_test, indices=range(500))Model Definition
import torch.nn as nn
class SmallCNN(nn.Module):
"""A simple 3-layer CNN for CIFAR-10."""
def __init__(self):
super().__init__()
self.features = nn.Sequential(
nn.Conv2d(3, 32, kernel_size=3, padding=1),
nn.ReLU(),
nn.MaxPool2d(2),
nn.Conv2d(32, 64, kernel_size=3, padding=1),
nn.ReLU(),
nn.MaxPool2d(2),
)
self.classifier = nn.Sequential(
nn.Flatten(),
nn.Linear(64 * 8 * 8, 256),
nn.ReLU(),
nn.Linear(256, 10),
)
def forward(self, x):
return self.classifier(self.features(x))LAMB Optimizer Implementation
We implement a simplified LAMB to show the core mechanics. Production use would rely on the torch-optimizer or apex packages.
import torch
import torch.optim as optim
class LAMB(optim.Optimizer):
"""
Layer-wise Adaptive Moments for Batch training.
Simplified implementation following You et al. (2019).
"""
def __init__(
self, params, lr=1e-3, betas=(0.9, 0.999), eps=1e-6, weight_decay=0.01
):
defaults = dict(lr=lr, betas=betas, eps=eps, weight_decay=weight_decay)
super().__init__(params, defaults)
@torch.no_grad()
def step(self, closure=None):
loss = None
if closure is not None:
with torch.enable_grad():
loss = closure()
for group in self.param_groups:
beta1, beta2 = group["betas"]
for p in group["params"]:
if p.grad is None:
continue
grad = p.grad
state = self.state[p]
if len(state) == 0:
state["step"] = 0
state["exp_avg"] = torch.zeros_like(p)
state["exp_avg_sq"] = torch.zeros_like(p)
exp_avg, exp_avg_sq = state["exp_avg"], state["exp_avg_sq"]
state["step"] += 1
t = state["step"]
# Adam moment updates
exp_avg.mul_(beta1).add_(grad, alpha=1 - beta1)
exp_avg_sq.mul_(beta2).addcmul_(grad, grad, value=1 - beta2)
# Bias correction
bias_correction1 = 1 - beta1**t
bias_correction2 = 1 - beta2**t
m_hat = exp_avg / bias_correction1
v_hat = exp_avg_sq / bias_correction2
# Adam update direction
adam_update = m_hat / (v_hat.sqrt() + group["eps"])
# Weight decay in update direction
if group["weight_decay"] != 0:
adam_update.add_(p, alpha=group["weight_decay"])
# LARS trust ratio (layer-wise)
weight_norm = p.norm(p=2)
update_norm = adam_update.norm(p=2)
if weight_norm == 0 or update_norm == 0:
trust_ratio = 1.0
else:
trust_ratio = weight_norm / update_norm
# Apply scaled update
p.add_(adam_update, alpha=-group["lr"] * trust_ratio)
return lossTraining Loop
from torch.utils.data import DataLoader
def train_one_epoch(model, loader, optimizer, device):
"""Train for one epoch and return average loss."""
model.train()
total_loss = 0.0
criterion = nn.CrossEntropyLoss()
for xb, yb in loader:
xb, yb = xb.to(device), yb.to(device)
optimizer.zero_grad()
loss = criterion(model(xb), yb)
loss.backward()
optimizer.step()
total_loss += loss.item()
return total_loss / len(loader)
def evaluate(model, loader, device):
"""Evaluate model accuracy on a dataset."""
model.eval()
correct = 0
total = 0
with torch.no_grad():
for xb, yb in loader:
xb, yb = xb.to(device), yb.to(device)
preds = model(xb).argmax(dim=1)
correct += (preds == yb).sum().item()
total += len(yb)
return correct / total
def run_experiment(
batch_size,
lr,
optimizer_name="sgd",
epochs=20,
warmup_epochs=0,
device="cpu",
):
"""Run a single training experiment and return loss/accuracy curves."""
torch.manual_seed(42)
model = SmallCNN().to(device)
train_loader = DataLoader(
train_subset, batch_size=batch_size, shuffle=True, num_workers=0
)
test_loader = DataLoader(
test_subset, batch_size=256, shuffle=False, num_workers=0
)
if optimizer_name == "sgd":
optimizer = optim.SGD(
model.parameters(), lr=lr, momentum=0.9, weight_decay=1e-4
)
elif optimizer_name == "adam":
optimizer = optim.Adam(model.parameters(), lr=lr, weight_decay=1e-4)
elif optimizer_name == "lamb":
optimizer = LAMB(model.parameters(), lr=lr, weight_decay=0.01)
train_losses = []
test_accs = []
base_lr = lr / max(warmup_epochs, 1)
for epoch in range(epochs):
# Linear warmup
if warmup_epochs > 0 and epoch < warmup_epochs:
current_lr = base_lr + (lr - base_lr) * epoch / warmup_epochs
for pg in optimizer.param_groups:
pg["lr"] = current_lr
elif warmup_epochs > 0 and epoch == warmup_epochs:
for pg in optimizer.param_groups:
pg["lr"] = lr
epoch_loss = train_one_epoch(model, train_loader, optimizer, device)
acc = evaluate(model, test_loader, device)
train_losses.append(epoch_loss)
test_accs.append(acc)
return train_losses, test_accs
device = "cpu"
epochs = 20Now we run four experiments: small batch baseline, large batch without scaling, large batch with linear scaling, and large batch with LAMB.
# Experiment 1: Small batch baseline (B=64, lr=0.05)
base_batch = 64
base_lr_sgd = 0.05
loss_small, acc_small = run_experiment(
batch_size=base_batch,
lr=base_lr_sgd,
optimizer_name="sgd",
epochs=epochs,
warmup_epochs=0,
device=device,
)
# Experiment 2: Large batch without LR scaling (B=512, same lr=0.05)
large_batch = 512
k = large_batch // base_batch # scaling factor = 8
loss_large_no_scale, acc_large_no_scale = run_experiment(
batch_size=large_batch,
lr=base_lr_sgd,
optimizer_name="sgd",
epochs=epochs,
warmup_epochs=0,
device=device,
)
# Experiment 3: Large batch with linear scaling (B=512, lr=0.4=8*0.05)
scaled_lr = base_lr_sgd * k
loss_large_scaled, acc_large_scaled = run_experiment(
batch_size=large_batch,
lr=scaled_lr,
optimizer_name="sgd",
epochs=epochs,
warmup_epochs=3,
device=device,
)
# Experiment 4: Large batch with LAMB
lamb_lr = 0.01
loss_lamb, acc_lamb = run_experiment(
batch_size=large_batch,
lr=lamb_lr,
optimizer_name="lamb",
epochs=epochs,
warmup_epochs=3,
device=device,
)Base batch size: 64, base LR: 0.05 Large batch size: 512, scale factor k: 8 Linearly scaled LR: 0.400 Final test accuracy (small batch SGD): 0.380 Final test accuracy (large batch, unscaled LR): 0.306 Final test accuracy (large batch, scaled LR): 0.208 Final test accuracy (large batch, LAMB): 0.340
The compact experiment illustrates both the motivation for scaling and its limits. The small-batch baseline reaches the highest final accuracy. The unscaled large-batch run trails it, but applying the full 8x learning-rate increase is too aggressive here and performs worse. LAMB is the strongest of the large-batch configurations, although it still does not match the small-batch baseline. This is why the linear rule is a starting point rather than a guarantee: the scaled rate must remain within the stable range of the particular model, data, and training budget.
Training Dynamics Visualization


Simulating the Gradient Variance Reduction
To make the statistical intuition concrete, we can directly compute how gradient variance changes with batch size on a fixed parameter snapshot.
torch.manual_seed(42)
# Use a fixed model and fixed mini dataset for reproducibility
ref_model = SmallCNN()
criterion_ref = nn.CrossEntropyLoss()
# Pre-load a set of samples
sample_loader = DataLoader(
train_subset, batch_size=1, shuffle=True, num_workers=0
)
gradients = []
# Collect per-sample gradients for first conv layer bias
for i, (xb, yb) in enumerate(sample_loader):
if i >= 128:
break
ref_model.zero_grad()
loss_ref = criterion_ref(ref_model(xb), yb)
loss_ref.backward()
grad = ref_model.features[0].bias.grad.detach().clone()
gradients.append(grad)
gradients = torch.stack(gradients) # shape: (200, 32)
# Compute variance of gradient estimates at different batch sizes
batch_sizes_var = [1, 2, 4, 8, 16, 32, 64]
variances = []
for bs in batch_sizes_var:
n_batches = len(gradients) // bs
batch_means = [
gradients[i * bs : (i + 1) * bs].mean(0) for i in range(n_batches)
]
batch_means_t = torch.stack(batch_means)
var = batch_means_t.var(0).mean().item()
variances.append(var)Gradient estimate variance by batch size:
Batch Size | Variance | Var / Var(B=1)
----------------------------------------------
1 | 0.000237 | 1.0000
2 | 0.000119 | 0.5007
4 | 0.000061 | 0.2549
8 | 0.000037 | 0.1540
16 | 0.000018 | 0.0758
32 | 0.000009 | 0.0362
64 | 0.000006 | 0.0233The variance decreases inversely with batch size. When we move from batch size 1 to batch size 8, the variance drops by a factor of roughly 8. This is exactly the relationship derived from statistics. The linear scaling rule compensates for this: each individual update becomes more accurate but less frequent, so you can afford a proportionally larger step.

Estimating the Critical Batch Size
The critical batch size can be approximated empirically by measuring how many training examples you need to process to reach a target loss level, at different batch sizes with linearly scaled learning rates.
def compute_examples_to_threshold(losses, batch_size, threshold):
"""Compute cumulative examples processed to reach a loss threshold."""
dataset_size = len(train_subset)
cumulative_examples = 0
for epoch_loss in losses:
cumulative_examples += dataset_size
if epoch_loss <= threshold:
return cumulative_examples
return None
# Define loss threshold we want to reach
target_loss = 1.75
# Run experiments at several batch sizes with linearly scaled LR
exp_batch_sizes = [32, 64, 128, 256, 512, 1024]
base_lr_for_crit = 0.025 # smaller base to allow scaling
examples_to_thresh = {}
all_crit_losses = {}
for bs in exp_batch_sizes:
k_crit = max(bs // 32, 1)
lr_crit = min(base_lr_for_crit * k_crit, 0.4)
warmup_e = min(3, max(1, k_crit // 2))
losses_crit, _ = run_experiment(
batch_size=bs,
lr=lr_crit,
optimizer_name="sgd",
epochs=15,
warmup_epochs=warmup_e,
device=device,
)
all_crit_losses[bs] = losses_crit
n_examples = compute_examples_to_threshold(losses_crit, bs, target_loss)
examples_to_thresh[bs] = n_examplesTarget loss threshold: 1.75
Batch Size | Examples to Threshold
--------------------------------------
32 | 4,000
64 | 5,000
128 | 6,000
256 | Not reached
512 | Not reached
1024 | Not reached
This compact run does not show an initial efficiency gain from larger batches. Instead, examples required rise from batch size 32 through 128, and batches of 256 or more miss the target within the training budget. That outcome places the critical regime below 256 for this particular model, subset, learning-rate rule, and 15-epoch budget. In larger experiments the transition is often flatter, but the diagnostic principle is the same: once larger batches stop reducing the examples needed to reach a target, extra parallelism no longer improves statistical efficiency.
LAMB Trust Ratio Across Layers
A key property of LAMB is that the trust ratio varies significantly across layers. Early layers (especially embeddings and the first convolutional layers) tend to have different weight-to-gradient-norm ratios compared to later layers. We can visualize this directly on our trained model.
# Compute trust ratios for each layer of the LAMB-trained model
torch.manual_seed(42)
lamb_inspect_model = SmallCNN()
lamb_inspect_opt = LAMB(
lamb_inspect_model.parameters(), lr=lamb_lr, weight_decay=0.01
)
# Run a single forward/backward pass to compute gradients
inspect_loader = DataLoader(
train_subset, batch_size=large_batch, shuffle=True, num_workers=0
)
xb_inspect, yb_inspect = next(iter(inspect_loader))
criterion_inspect = nn.CrossEntropyLoss()
lamb_inspect_opt.zero_grad()
loss_inspect = criterion_inspect(lamb_inspect_model(xb_inspect), yb_inspect)
loss_inspect.backward()
# Collect trust ratios before the step (simulate first LAMB step)
layer_names = []
weight_norms = []
grad_norms = []
trust_ratios = []
for name, p in lamb_inspect_model.named_parameters():
if p.grad is not None and p.dim() > 1: # only weight tensors (not biases)
wn = p.norm(p=2).item()
gn = p.grad.norm(p=2).item()
if gn > 0:
tr = wn / gn
layer_names.append(name.replace(".weight", ""))
weight_norms.append(wn)
grad_norms.append(gn)
trust_ratios.append(tr)Layer | Weight Norm | Grad Norm | Trust Ratio -------------------------------------------------------------------------- features.0 | 3.3163 | 0.0150 | 221.2744 features.3 | 4.6290 | 0.0565 | 81.9249 classifier.1 | 9.2346 | 0.2364 | 39.0607 classifier.3 | 1.8339 | 0.0603 | 30.3973
The trust ratios vary considerably across layers, confirming that LAMB's layer-wise scaling is doing real work. Layers with large weight norms relative to their gradient norms receive larger effective learning rates, while layers where gradients are proportionally larger receive smaller effective learning rates.


Key Parameters
Understanding each hyperparameter and its role helps you make principled decisions when configuring large batch training.
-
batch_size: The number of training examples processed per gradient update. Should be chosen to balance hardware utilization and the critical batch size limit. The critical batch size varies by task; empirical measurement (as shown above) gives the best estimate.
-
learning_rate (): The global learning rate. Must be scaled linearly with batch size: when batch size is multiplied by . The base learning rate should be validated at a small batch size before scaling.
-
warmup_epochs: Number of epochs to linearly increase the learning rate from a small initial value to the target. Typically 5 epochs for ImageNet-scale tasks, often 1-3% of total training steps for LLM pre-training. A good rule of thumb is to warm up for at least as many steps as (the scaling factor).
-
weight_decay (): L2 regularization coefficient. Incorporated directly into the LAMB update as a decoupled weight decay term. Decoupled weight decay (AdamW-style, not the original Adam L2 approach) is important for LAMB to behave correctly.
-
beta1, beta2: Exponential decay rates for Adam's first and second moment estimates (LAMB). Standard values are 0.9 and 0.999. These generally do not need adjustment when scaling batch size.
-
trust_ratio_clip: (Optional) Maximum value of the LARS/LAMB trust ratio to prevent instability for layers near zero weight norm. A common default is to clip trust ratios to the range .
Practical Guidance: Putting It All Together
Moving from theory to a working large-batch training setup involves several steps that are easy to get wrong. Here is a practical checklist for configuring large batch training in a real project.
First, establish your baseline at a small batch size. Train your model at batch size 32 or 64 with the default learning rate for your optimizer. Verify that training converges stably and achieves the expected validation performance. This baseline is the reference you will compare everything against.
Second, choose your target batch size. This should be driven by your hardware: how many GPUs you have and how large a batch each can process. For multi-GPU training, the total batch is the per-GPU batch times the number of GPUs. Aim to hit the critical batch size approximately; going much higher wastes compute efficiency without improving convergence.
Third, apply the linear scaling rule. Multiply your baseline learning rate by . Add a warmup schedule starting from or a fixed small value, ramping to over the first 5% of training steps.
Fourth, choose your optimizer. Use Adam for small to medium batch sizes and standard transformer architectures. Use LAMB when you are training transformers with batch sizes above 8,192 tokens or when you observe training instability at the linearly scaled learning rate.
Fifth, monitor training stability carefully. Large batch training with a scaled learning rate can diverge quickly if something is misconfigured. Monitor gradient norms, weight norms, and per-layer update norms in the first few hundred steps. If any layer shows explosive gradient growth or the loss spikes, reduce the learning rate or extend the warmup.
Finally, validate generalization. Run your large-batch model on the held-out test set and compare to the small-batch baseline. If there is a meaningful accuracy gap, consider adding data augmentation, ghost batch normalization, or a longer training schedule to compensate.
Limitations and Practical Considerations
The linear scaling rule is an approximation, and its applicability depends on assumptions that are often only partially satisfied. Understanding where it breaks down guides practical decisions.
The most significant limitation is the relationship between batch size and generalization. While careful engineering can close most of the accuracy gap between large and small batch training, the gap rarely disappears entirely. Training with very large batches tends to converge to sharper minima. Sharp minima are more sensitive to distribution shift between training and test sets, leading to worse generalization at deployment. In practice, large batch pre-trained models often require more aggressive fine-tuning data augmentation to match the downstream performance of models trained with smaller batches.
A second practical constraint is compute efficiency. The critical batch size is typically much smaller than what modern distributed training hardware can support at full utilization. For example, a 1,024-GPU cluster might require a batch size of 32,768 per step just to keep all GPUs busy, while the critical batch size for the task might be 2,048. The extra 16x scaling is necessary for hardware efficiency but does not improve convergence speed per example. In these regimes, engineers often accept some efficiency loss in exchange for better model quality, choosing a batch size that balances hardware utilization with convergence efficiency.
LAMB's layer-wise trust ratio addresses the adaptive optimizer scaling problem, but it introduces its own sensitivity. The trust ratio computation depends on the L2 norms of weights and Adam update directions. Early in training, when weights are small and gradients are large, the trust ratio can produce very small effective learning rates for some layers. This is usually mitigated by trust ratio clipping and careful initialization, but it requires tuning. LAMB's original validation was primarily on BERT and ResNets. For other architectures with unusual weight scale distributions across layers (especially models with heavily normalized activations), the trust ratio may behave differently and require additional tuning.
Batch normalization also interacts with batch size in complex ways. Batch normalization statistics (mean and variance computed over the batch) change significantly as batch size changes. With very large batches, the statistics become very accurate, but they also incorporate more diversity in the batch, potentially changing the effective normalization behavior. Ghost batch normalization sidesteps this by computing BN statistics on smaller sub-batches, but it introduces the hyperparameter of ghost batch size. For language models using layer normalization rather than batch normalization, this issue is less acute.
The field has also moved toward gradient accumulation as an alternative to true large batch training. Rather than physically processing a large batch on one step, you accumulate gradients from multiple smaller batches before applying the update. This achieves the same effective batch size without requiring hardware that can fit the large batch in memory. Gradient accumulation is straightforward for plain SGD and Adam, but interacts subtly with batch normalization statistics and some regularization techniques. The next chapter on gradient accumulation explores this approach in depth.
Despite these limitations, large batch training has become indispensable for modern AI development. Training GPT-3 with batch sizes of millions of tokens, or BERT with batch sizes of hundreds of thousands, was only possible because of the combination of linear scaling, LAMB, warmup, and careful monitoring of training stability. The techniques in this chapter represent the engineering foundation that made the scaling era of language models possible.
Summary
Large batch training changes the statistical properties of gradient estimation. Larger batches provide more accurate gradient estimates (lower variance), enabling the use of larger learning rates that make each step more consequential.
The linear scaling rule states that when you increase the batch size by a factor , you should increase the learning rate by the same factor . This rule emerges from the requirement that the expected weight update after processing training examples should be the same regardless of batch size. The rule holds well in regimes where the loss landscape is locally smooth but breaks at extremes: very small batches have too much noise, and very large learning rates cause overshooting.
The critical batch size marks the practical upper bound on useful batch scaling. Beyond this size, additional parallelism does not reduce the total number of examples needed for convergence. It represents the point at which gradient noise is no longer the bottleneck for learning.
LAMB extends the linear scaling idea to adaptive optimizers. It combines Adam's per-parameter moment adaptation with a LARS-style layer-wise trust ratio, preventing any layer from taking an unstable step when the global learning rate is large. LAMB enabled BERT pre-training at batch sizes that reduced wall-clock training time from days to minutes on large TPU pods.
Warmup remains essential across all large batch techniques, bridging the early phase where the loss landscape is non-linear and the later phase where the linear scaling assumption holds.
The next chapter on weight decay examines complementary regularization techniques that interact with batch size decisions, completing the toolkit for large-scale training optimization.
Quiz
Ready to test your understanding? Take this quick quiz to reinforce what you've learned about large batch training, learning rate scaling, and the LAMB optimizer.
Large Batch Training Quiz
Reference
Citation details
Cite or share this article.
Continue with the full handbook
This chapter is part of Language AI Handbook. Use the handbook page to browse the complete table of contents and continue reading in sequence.
Explore Language AI HandbookStay up to date
Get articles, book updates, and news delivered to your inbox.
No spam, unsubscribe anytime.
Join the community
Sign in to remove popups, track your reading progress, and join the discussion.

Comments
No comments yet. Be the first to share your thoughts!