Core ML2017intermediate12 min read

Accurate, Large Minibatch SGD: Training ImageNet in 1 Hour

النزول التدريجي العشوائي بدفعات كبيرة ودقيقة: تدريب ImageNet في ساعة واحدة

Goyal, P. · Dollár, P. · Girshick, R. · Noordhuis, P. · Wesolowski, L. · Kyrola, A. · Tulloch, A. · Jia, Y. · He, K. — arXiv

The problem

Deep learning models get better with more data and bigger networks, but time grows too. using can help: split each minibatch across many GPUs. But this requires larger minibatches, and naively increasing causes optimization to diverge or to degrade. In 2017, training ResNet-50 on ImageNet took 29 hours on 8 GPUs, and nobody had shown that scaling to thousands of GPUs could maintain the same accuracy.

The contribution

Two simple but powerful techniques that allow minibatch sizes up to 8,192 images without accuracy loss: (1) a linear scaling rule — multiply the by k when the batch size is multiplied by k, and (2) gradual — start with a small learning rate and linearly ramp it up over the first 5 epochs. Combined with careful handling (keeping per-worker BN statistics at 32 samples), these techniques enabled training ResNet-50 on 256 GPUs in just 1 hour with the same top-1 accuracy as the 29-hour .

The impact

This paper became the standard recipe for large-batch distributed training. Its linear scaling rule and gradual warmup are now used in virtually every large-scale training run — from BERT to GPT to Vision Transformers. The paper proved that training time can be reduced linearly with more GPUs without sacrificing model quality, directly enabling the era of training on billions of examples. Its insights were later extended by LAMB and LARS optimizers to even larger batch sizes.

Imagine a team of painters working on a mural. With one painter, progress is slow but consistent. You could add 32 painters, each working on a section, and merge their work at the end of each round. But if you tell 256 painters to "go faster" without coordination, they'll paint over each other's edges and the mural turns to chaos.

This paper's insight is the coordination protocol: if you multiply the number of painters by 32×, you must also multiply the length of each brushstroke by 32× so that the combined effect per round stays the same. And during the first few rounds — while everyone is still finding their positions — you use shorter, gentler strokes before ramping up to full speed.

The problem: bigger batches break training

Training a deep network means looping over the dataset many times, updating the weights after each minibatch. With 8 GPUs and 32 images per GPU, the minibatch is 256 images and training ResNet-50 on ImageNet takes 29 hours.

The dream is simple: use 256 GPUs (32× more), split the minibatch equally, and finish in roughly 1 hour. Each GPU computes gradients on its local 32 images, then all GPUs average their gradients and take one shared step. This is synchronous data parallelism.

But here's the catch: to keep each GPU busy (high compute utilization), the total minibatch must grow to 8,192 images. Naively training with such a large batch causes the training loss to diverge or the final accuracy to drop by 1–2%. The community in 2017 did not know whether this was a fundamental generalization issue or merely an optimization problem that could be solved.

Open in Lab
See how data parallelism splits images across GPUs: each GPU computes local gradients, then all gradients are averaged via allreduce.
The demo wakes as you arrive…

The first key idea: the linear scaling rule

The core insight is deceptively simple. Imagine you normally take k small steps with learning rate η\eta, each on a different minibatch of nn images. Now instead, you want to take one big step on all knkn images combined. For the big step to have roughly the same effect as the k small ones, you need to set η^=kη\hat{\eta} = k\eta.

Why? After k small steps, the total weight update accumulates all k averages. The single big step averages over kn samples at once. If the gradients don't change much between consecutive steps — ∇l(x,wt)≈∇l(x,wt+j)\nabla l(x, w_t) \approx \nabla l(x, w_{t+j}) for small jj — then multiplying the learning rate by kk makes the big step match the accumulated small ones.

This is called the Linear Scaling Rule: when you multiply the batch size by kk, multiply the learning rate by kk.

η^=k⋅η\hat{\eta} = k \cdot \eta
Linear Scaling Rule — the simplest recipe for large-batch training — When the minibatch size grows by a factor of k (e.g., from 256 to 8192, so k=32), multiply the base learning rate η by the same factor k.

Think of it like adjusting your stride when walking vs. jogging. If you take 32 short steps, you cover a certain distance. To cover the same distance in 1 step, you need a stride 32× longer. The learning rate is the stride length of your optimization.

Open in Lab
Drag the batch size slider to see how the learning rate scales linearly, and compare the resulting training curves against the 256-batch baseline.
The demo wakes as you arrive…

The second key idea: gradual warmup

The linear scaling rule has an important assumption: gradients should be approximately constant across consecutive steps. Early in training, this assumption is badly violated. The network's weights change rapidly, so the gradient at step tt is very different from the gradient at step t+1t+1.

Starting with a learning rate of η^=3.2\hat{\eta} = 3.2 (32× the base rate) from the first iteration causes training to diverge immediately. The solution? Start small and ramp up.

The paper tried two approaches:

  • Constant warmup: use η=0.1\eta = 0.1 for 5 epochs, then jump to η^=3.2\hat{\eta} = 3.2. This actually made things worse — the sudden jump caused training error to spike and never recover.
  • Gradual warmup: start at η=0.1\eta = 0.1 and linearly increase it to η^=3.2\hat{\eta} = 3.2 over 5 epochs. This worked beautifully — the training curve matched the small-batch baseline almost perfectly.
Open in Lab
Compare three warmup strategies: none, constant, and gradual. Watch how each affects the training error curve relative to the small-batch baseline.
The demo wakes as you arrive…

Batch Normalization: a hidden dependency on batch size

There's a subtlety that the paper identifies as critical. Batch Normalization computes mean and variance statistics across the minibatch. This means the loss of each sample depends on all other samples in its minibatch — the loss function itself changes when you change the batch size.

The solution is elegant: keep per-worker BN statistics fixed at n=32n = 32 samples, regardless of how many GPUs you use. If you have 256 GPUs, the total batch is 8,192, but each GPU still computes BN statistics over its own 32 images. This way, the underlying loss function being optimized stays identical, and you avoid expensive cross-GPU communication for BN statistics.

The paper shows that computing BN over all 8,192 samples (which changes the effective loss) hurts accuracy. The lesson: treat BN batch size as a hyperparameter of the normalization, not of the distributed training.

Open in Lab
Toggle between local BN (per-GPU, 32 samples) and global BN (all 8192 samples) to see how the statistics and the effective loss function change.
The demo wakes as you arrive…

Practical pitfalls: details that break training

The paper identifies several implementation subtleties that are easy to get wrong in distributed training:

is not just scaling. The loss function includes an L2 regularization term λ2∥w∥2\frac{\lambda}{2}\|w\|^2. Scaling the cross-entropy loss (e.g., dividing by the total batch size) is not equivalent to scaling the learning rate, because the weight decay term must be handled separately. Getting this wrong silently produces a model that trains but has higher error.

correction. There are two common implementations of momentum . In one, the update tensor uu is independent of η\eta; in the other, v=ηuv = \eta u, which entangles the momentum buffer with the learning rate. When η\eta changes (as during warmup), the second form requires a correction factor ηt+1ηt\frac{\eta_{t+1}}{\eta_t} applied to the momentum buffer. Without this correction, the jump from warmup to full learning rate causes instability.

Gradient normalization. Each worker normalizes by its local batch size nn, but the allreduce sums (not averages) gradients. You must normalize by the total batch size knkn, not the local nn. Failing to do this is equivalent to using a different effective learning rate.

The recipe in code

Large-batch SGD with linear scaling and gradual warmuppython

Simplified to show the idea — not the real implementation.

import numpy as np

def get_lr(epoch, iteration, iters_per_epoch,
           base_lr=0.1, base_batch=256, actual_batch=8192,
           warmup_epochs=5):
    """Linear scaling + gradual warmup learning rate schedule."""
    k = actual_batch / base_batch          # scaling factor (e.g. 32)
    target_lr = base_lr * k                # linear scaling rule

    # Gradual warmup: linearly ramp from base_lr to target_lr
    warmup_iters = warmup_epochs * iters_per_epoch
    current_iter = epoch * iters_per_epoch + iteration
    if current_iter < warmup_iters:
        # Linear interpolation from base_lr to target_lr
        alpha = current_iter / warmup_iters
        return base_lr + alpha * (target_lr - base_lr)

    # After warmup: standard step decay at epochs 30, 60, 80
    if epoch < 30:
        return target_lr
    elif epoch < 60:
        return target_lr * 0.1
    elif epoch < 80:
        return target_lr * 0.01
    else:
        return target_lr * 0.001

# Key implementation detail: per-worker BN uses n=32 samples.
# The total batch is k*n = 8192, but BN stats are LOCAL.
# This keeps the loss function identical across batch sizes.

Results: 1 hour, same accuracy

The results tell a clear story. With the linear scaling rule and gradual warmup, ResNet-50 trained with a minibatch of 8,192 images achieved 23.74% top-1 error — compared to 23.60% for the 256-image baseline. That's only a 0.14% difference, well within the random variation of ±0.12%.

Even more telling: the training curves for large and small batches matched almost perfectly after the warmup phase. This suggests that the issue was purely about optimization in early training, not a fundamental generalization problem.

The paper tested batch sizes from 64 to 65,536. Accuracy held steady up to 8,192, then degraded at 16k and diverged at 64k. The sweet spot for ImageNet with ResNet-50 turned out to be around 8k — the point where the linear scaling rule's core assumption (stable gradients) starts to break even after warmup.

Open in Lab
Explore how validation error changes with batch size. Notice the stability up to 8k and the sharp degradation beyond.
The demo wakes as you arrive…

Does it generalize beyond ImageNet?

A critical test: do features learned with large batches transfer as well as those from small batches? The authors pre-trained ResNet-50 with batch sizes from 256 to 16k, then used each model to initialize Mask R-CNN for object detection on COCO.

The result: as long as ImageNet validation error was maintained (batch sizes up to 8k), the COCO detection AP was identical to the small-batch baseline — 35.8% box AP and 33.9% mask AP across all batch sizes. No generalization gap when transferring across datasets and across tasks.

The paper also showed the linear scaling rule works directly for training Mask R-CNN itself from 1 to 8 GPUs, confirming it generalizes beyond classification.

Systems: making 256 GPUs work together

The algorithm is only useful if communication doesn't become the bottleneck. The paper's system design achieves ~90% scaling efficiency from 8 to 256 GPUs using commodity 50 Gbit Ethernet — no specialized interconnects needed.

The key technique is overlapping communication with computation. During , as soon as gradients for one layer are computed, they are sent for allreduce while the next layer's gradients are still being calculated. This hides most of the communication latency behind useful work.

For the allreduce itself, the paper uses a recursive halving-doubling algorithm that outperforms the standard ring algorithm by 3× for typical gradient buffer sizes, operating in 2log⁡2(p)2\log_2(p) steps instead of 2(p−1)2(p-1) steps where pp is the number of servers.

Open in Lab
Watch how gradient computation and allreduce communication overlap during backpropagation, achieving near-linear speedup.
The demo wakes as you arrive…

The BN γ initialization trick

A small but effective trick: in each residual block, the final Batch Normalization layer's scaling parameter γ\gamma is initialized to 0 instead of the standard 1. This causes the residual block to initially behave as an identity mapping — the forward and backward signals flow through the skip connection only. As training progresses, γ\gamma grows from zero and the residual branch gradually contributes.

This improved accuracy for all batch sizes but was especially helpful for large batches, reducing the gap from 0.51% (with standard initialization) to 0.14%.

Why it changed everything

  1. 2014

    Krizhevsky: One weird trick

    First suggestion to scale learning rate linearly with batch size for CNNs, but reported 1% accuracy loss going from 128 to 1024 batch size.

  2. 2017

    This paper — Linear scaling + gradual warmup

    Showed batch sizes up to 8,192 with no accuracy loss. Trained ImageNet in 1 hour on 256 GPUs. Became the standard recipe for distributed training.

  3. 2017

    LARS — Layer-wise Adaptive Rate Scaling

    Extended large-batch training to 32k batch size by scaling each layer's learning rate individually based on its weight-to-gradient ratio.

  4. 2019

    LAMB — Layer-wise Adaptive Moments for Batch training

    Combined LARS-style per-layer scaling with Adam optimizer, enabling batch sizes up to 64k for BERT pre-training. Reduced BERT training time from 3 days to 76 minutes.

  5. 2020

    Large-batch training powers the scaling era

    GPT-3, DALL·E, and subsequent models all rely on large-batch techniques descended from this paper to train on hundreds of GPUs efficiently.

CitationGoyal, Dollár, Girshick, Noordhuis, Wesolowski, Kyrola, Tulloch, Jia, He. Accurate, Large Minibatch SGD: Training ImageNet in 1 Hour. arXiv, 2017.

Terms in this paper