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.
The first key idea: the linear scaling rule
The core insight is deceptively simple. Imagine you normally take k small steps with learning rate , each on a different minibatch of images. Now instead, you want to take one big step on all images combined. For the big step to have roughly the same effect as the k small ones, you need to set .
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 — for small — then multiplying the learning rate by makes the big step match the accumulated small ones.
This is called the Linear Scaling Rule: when you multiply the batch size by , multiply the learning rate by .
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.
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 is very different from the gradient at step .
Starting with a learning rate of (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 for 5 epochs, then jump to . This actually made things worse — the sudden jump caused training error to spike and never recover.
- Gradual warmup: start at and linearly increase it to over 5 epochs. This worked beautifully — the training curve matched the small-batch baseline almost perfectly.
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 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.
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 . 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 is independent of ; in the other, , which entangles the momentum buffer with the learning rate. When changes (as during warmup), the second form requires a correction factor 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 , but the allreduce sums (not averages) gradients. You must normalize by the total batch size , not the local . Failing to do this is equivalent to using a different effective learning rate.
The recipe in code
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.
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 steps instead of steps where is the number of servers.
The BN γ initialization trick
A small but effective trick: in each residual block, the final Batch Normalization layer's scaling parameter 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, 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
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.
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.
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.
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.
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
- Stochastic Gradient Descent (SGD)الانحدار التدريجي العشوائي
- Mini-Batchالدفعة المصغرة
- Learning Rateمعدل التعلم
- Batch Normalizationتسوية الدفعات الحسابية
- Warmupالإحماء
- Data Parallelismتوازي البيانات
- Gradient Descentالانحدار التدريجي
- Momentumالزخم
- Convergenceالتقارب الحسابي
- Distributed Trainingالتدريب الموزَّع
- Weight Decayاضمحلال الأوزان