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.

Open in Lab
Increase the batch size and watch how a single global learning rate causes training to diverge.
The demo wakes as you arrive…

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.

Open in Lab
Compare how LARS and LAMB handle heterogeneous layers. LARS applies one scale per layer; LAMB adapts per-parameter before applying the layerwise trust ratio.
The demo wakes as you arrive…

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.

mt=β1 mt−1+(1−β1) gtvt=β2 vt−1+(1−β2) gt2m^t=mt/(1−β1t),v^t=vt/(1−β2t)rt(i)=m^t(i)v^t(i)+ϵ+λ wt(i)wt+1(i)=wt(i)−η⋅ϕ ⁣(∥wt(i)∥∥rt(i)∥)⋅rt(i)\begin{aligned} m_t &= \beta_1 \, m_{t-1} + (1-\beta_1) \, g_t \\ v_t &= \beta_2 \, v_{t-1} + (1-\beta_2) \, g_t^2 \\ \hat{m}_t &= m_t / (1-\beta_1^t), \quad \hat{v}_t = v_t / (1-\beta_2^t) \\ r_t^{(i)} &= \frac{\hat{m}_t^{(i)}}{\sqrt{\hat{v}_t^{(i)}} + \epsilon} + \lambda \, w_t^{(i)} \\ w_{t+1}^{(i)} &= w_t^{(i)} - \eta \cdot \phi\!\left(\frac{\|w_t^{(i)}\|}{\|r_t^{(i)}\|}\right) \cdot r_t^{(i)} \end{aligned}
LAMB Update Rule — Lines 1–3 are identical to Adam: exponential moving averages of gradient and squared gradient, with bias correction. Line 4 adds weight decay directly to the update (like AdamW). Line 5 is the key addition: the trust ratio φ(‖w‖/‖r‖) scales the full update so that the step size is proportional to the layer's weight magnitude.
Open in Lab
Step through the LAMB algorithm to see how each component transforms the gradient into the final update.
The demo wakes as you arrive…
LAMB Optimizer — Core Looppython

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 * update

Scaling 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.

Open in Lab
See how increasing batch size reduces total training time while LAMB maintains model quality.
The demo wakes as you arrive…

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.

1T∑t=1TE∥∇f(xt)∥2  ≤  O ⁣(Lavg (f(x1)−f∗)T  +  Lavg σbT)\frac{1}{T}\sum_{t=1}^{T}\mathbb{E}\|\nabla f(x_t)\|^2 \;\leq\; O\!\left(\frac{L_{\mathrm{avg}}\,(f(x_1)-f^*)}{T} \;+\; \frac{L_{\mathrm{avg}}\,\sigma}{\sqrt{bT}}\right)
LAMB Convergence Bound — The convergence rate depends on the average Lipschitz constant L_avg rather than the maximum L_∞. This means LAMB benefits from layer heterogeneity — the very thing that makes standard optimizers struggle. b is the batch size and σ the gradient noise.

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

  1. 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.

  2. 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.

  3. 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.

  4. 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