Optimization2019intermediate8 min read
Large Batch Optimization for Deep Learning: Training BERT in 76 Minutes
أمثَلة الدُّفعات الكبيرة للتعلُّم العميق: تدريب BERT في 76 دقيقة
You, Y. · Li, J. · Reddi, S. · Hseu, J. · Kumar, S. · Bhojanapalli, S. · Song, X. · Demmel, J. · Keutzer, K. · Hsieh, C.-J. — ICLR
The problem
large models like BERT on massive datasets takes days even with powerful hardware. The natural solution is to use larger batch sizes to parallelize computation, but naively increasing degrades model quality. , the best existing large-batch optimizer, works well for CNNs like ResNet but fails completely on -based models like BERT. There was no optimizer that could scale batch sizes for both CNNs and Transformers without sacrificing accuracy.
The contribution
LAMB (Layer-wise Adaptive Moments optimizer for Batch training): a new optimizer that combines 's per-parameter adaptive learning rates with layerwise trust ratios. The scales each layer's update by the ratio of the layer's weight norm to the update norm, preventing any single layer from taking a disproportionately large step. LAMB scales BERT training to batch size 32,868 with no accuracy loss, reducing training from 3 days to 76 minutes on a TPUv3 Pod. It also achieves state-of-the-art accuracy for ResNet-50 — the first adaptive solver to do so.
The impact
LAMB proved that large-batch training is not limited to CNNs — it works for Transformers too, if the optimizer respects each layer's scale. This insight enabled the large-scale distributed pre-training pipelines that power today's foundation models. The layerwise trust ratio concept influenced subsequent optimizers and became a key ingredient in scaling training to thousands of accelerators.
Imagine a highway with 100 lanes. With 8 cars, any speed limit works fine. But with 64,000 cars, a single speed limit causes chaos — sports cars go too slow, trucks go too fast, and everyone crashes.
LARS tried to fix this by giving each lane its own speed limit based on how heavy the cars are. This worked great on a straight highway (ResNet), but on a winding mountain road (BERT), the approach fell apart.
LAMB adds a second layer of intelligence: it not only adjusts speed per lane but also considers how fast each car is already going () and how bumpy the road is (gradient variance). The result: 64,000 cars arrive safely in 76 minutes instead of 3 days.
The problem: large batches break training
Deep learning training is embarrassingly parallel in one dimension: computing the gradient. Each example in a minibatch can be processed independently, so doubling the batch size halves the wall-clock time per iteration — in theory. But in practice, larger batches cause two problems.
First, larger batches reduce the noise in gradient estimates. This sounds helpful, but noise acts as implicit regularization that helps the model generalize. Remove it, and the model overfits or converges to sharp, fragile minima.
Second, the that works for a small batch is wrong for a large batch. The standard fix — linear scaling of learning rate with batch size — only works up to a point. Beyond that, training becomes unstable and the model diverges.
Why LARS fails on BERT
LARS (Layer-wise Adaptive Rate Scaling) was the first optimizer to use layerwise trust ratios for large-batch training. It scales each layer's learning rate by the ratio of the weight norm to the gradient norm. This worked brilliantly for ResNet on ImageNet — training that took hours could finish in minutes.
But LARS uses vanilla with momentum as its base optimizer. SGD with momentum treats all parameters in a layer identically, using the same learning rate for every weight. This is fine for convolutional layers where parameters have similar scales, but it is a poor match for Transformers. In BERT, the layer has dramatically different gradient statistics than the attention heads, and even within a single layer, some parameters are orders of magnitude more sensitive than others. LARS's per-layer scaling cannot compensate for this per-parameter heterogeneity.
The LAMB optimizer
LAMB builds on Adam — the standard optimizer for Transformers — and adds one critical ingredient: a layerwise trust ratio. The idea is elegant. First, compute the Adam update for each parameter as usual, including the exponential moving averages of the gradient (first moment) and the squared gradient (second moment). Then, for each layer, compute the ratio of the weight norm to the update norm. Multiply the layer's update by this ratio.
Think of it as a safety valve. If a layer's proposed update is huge relative to the layer's current weights, the trust ratio shrinks it. If the update is tiny, the trust ratio amplifies it. The effect is that every layer takes a step proportional to its own scale — no layer can dominate the update and destabilize training.
Simplified to show the idea — not the real implementation.
def lamb_update(params, grads, m, v, t, lr, beta1, beta2, eps, wd):
for i, (w, g) in enumerate(zip(params, grads)):
# Adam moments
m[i] = beta1 * m[i] + (1 - beta1) * g
v[i] = beta2 * v[i] + (1 - beta2) * g ** 2
# Bias correction
m_hat = m[i] / (1 - beta1 ** t)
v_hat = v[i] / (1 - beta2 ** t)
# Adam-style update + weight decay
update = m_hat / (v_hat.sqrt() + eps) + wd * w
# Trust ratio — LAMB's key addition
w_norm = w.norm()
u_norm = update.norm()
trust = w_norm / u_norm if w_norm > 0 and u_norm > 0 else 1.0
# Apply scaled update
w -= lr * trust * updateScaling BERT training
The original BERT was trained with a batch size of 256 for 1 million steps — about 3 days on 16 TPUv3 chips. The authors of LAMB scaled this progressively: batch size 512, then 4K, 8K, 16K, 32K, and finally 64K. At each scale, they verified that the model matched the original BERT's accuracy on downstream tasks.
The key finding: LAMB maintained BERT's accuracy at batch size 32,868, requiring only 8,599 iterations instead of 1 million. By pushing to 64K on a full TPUv3 Pod (1,024 chips), they reduced total training time to 76 minutes. The training recipe also used a gradual phase where the learning rate increases linearly for the first fraction of steps, followed by a polynomial decay.
Convergence guarantees
Beyond strong empirical results, the paper provides analysis for both LARS and LAMB in general nonconvex settings. The key insight is that can be understood as a form of preconditioned gradient descent, where the preconditioner is a block-diagonal matrix with each block scaled by the trust ratio.
The convergence rate depends on the average Lipschitz constant across layers rather than the maximum. This is significant: in deep networks, a few layers may have much larger Lipschitz constants than others. Standard optimizers are bottlenecked by the worst layer, while layerwise adaptive methods depend on the average — a much smaller quantity.
Practical training recipe
The paper provides several practical insights for large-batch training. The warmup phase is critical: the learning rate starts near zero and increases linearly for a fraction of total steps. Without warmup, large-batch training diverges immediately. After warmup, the learning rate follows polynomial decay.
is applied directly in the update (decoupled weight decay, as in ) rather than as L2 regularization. The hyperparameters β₁=0.9, β₂=0.999, and ε=10⁻⁶ work across both BERT and ResNet with minimal tuning. The only parameter that changes with batch size is the peak learning rate, which scales roughly linearly.
Large-batch training timeline
2017
Linear scaling rule
Goyal et al. showed ResNet-50 can be trained with batch size 8,192 by linearly scaling the learning rate, plus a warmup period. One hour instead of 29.
2017
LARS
You et al. introduced layerwise adaptive rate scaling. Trained ResNet on ImageNet in minutes by scaling batch size to 32K. Failed on attention models.
2019
LAMB
Combined Adam's per-parameter adaptivity with layerwise trust ratios. First optimizer to scale BERT to batch size 64K. Training time reduced from 3 days to 76 minutes.
2020
NVLAMB and scaled pre-training
NVIDIA adopted LAMB variants for pre-training Megatron-LM and subsequent large language models. Large-batch adaptive optimizers became standard in industry pipelines.
CitationYou, Li, Reddi, Hseu, Kumar, Bhojanapalli, Song, Demmel, Keutzer, Hsieh. Large Batch Optimization for Deep Learning: Training BERT in 76 Minutes. ICLR, 2020.
Terms in this paper
- LAMB Optimizerمُحسِّن LAMB
- LARSخوارزمية LARS
- Layerwise Adaptationالتكييف الطبقي
- Trust Ratioنسبة الثقة
- Large Batch Trainingالتدريب بالدفعات الكبيرة
- Learning Rateمعدل التعلم
- Warmupالإحماء
- Weight Decayاضمحلال الأوزان
- Batch Sizeحجم الدفعة الحسابية
- Adamخوارزمية آدام
- AdamWآدَم مع اضمحلال أوزان مفصول
- Momentumالزخم
- Convergenceالتقارب الحسابي
- Stochastic Gradient Descent (SGD)الانحدار التدريجي العشوائي
- Batch Normalizationتسوية الدفعات الحسابية