mirror of
https://github.com/rohitg00/ai-engineering-from-scratch.git
synced 2026-10-02 10:04:49 +08:00
feat(phase-19/46): gradient-accumulation
This commit is contained in:
@@ -0,0 +1,353 @@
|
||||
"""Gradient accumulation from scratch.
|
||||
|
||||
Effective batch size = micro batch size * accumulation steps. Accumulate
|
||||
gradients across several forward and backward passes, only step the
|
||||
optimizer after the last micro-batch. Tracks throughput against effective
|
||||
batch size so the curve is visible, not folklore.
|
||||
|
||||
Run: python3 code/main.py
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import math
|
||||
import time
|
||||
from dataclasses import dataclass, field, asdict
|
||||
from pathlib import Path
|
||||
from typing import Callable, Iterable, List
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
|
||||
|
||||
HERE = Path(__file__).parent
|
||||
OUT_DIR = HERE.parent / "outputs"
|
||||
LOG_PATH = OUT_DIR / "accum-curve.json"
|
||||
|
||||
|
||||
@dataclass
|
||||
class StepResult:
|
||||
step: int
|
||||
effective_batch: int
|
||||
micro_batch: int
|
||||
accum_steps: int
|
||||
loss: float
|
||||
grad_norm: float
|
||||
samples_per_sec: float
|
||||
wall_ms: float
|
||||
sync_calls: int
|
||||
|
||||
|
||||
@dataclass
|
||||
class CurvePoint:
|
||||
effective_batch: int
|
||||
accum_steps: int
|
||||
micro_batch: int
|
||||
avg_loss: float
|
||||
samples_per_sec: float
|
||||
median_step_ms: float
|
||||
sync_calls: int
|
||||
steps: int
|
||||
|
||||
|
||||
def seed_everything(seed: int) -> None:
|
||||
torch.manual_seed(seed)
|
||||
|
||||
|
||||
def synthetic_batch(batch_size: int, in_dim: int, out_dim: int, gen: torch.Generator) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
x = torch.randn(batch_size, in_dim, generator=gen)
|
||||
target = torch.randint(low=0, high=out_dim, size=(batch_size,), generator=gen)
|
||||
return x, target
|
||||
|
||||
|
||||
def make_model(in_dim: int, hidden: int, out_dim: int) -> nn.Module:
|
||||
return nn.Sequential(
|
||||
nn.Linear(in_dim, hidden),
|
||||
nn.GELU(),
|
||||
nn.Linear(hidden, hidden),
|
||||
nn.GELU(),
|
||||
nn.Linear(hidden, out_dim),
|
||||
)
|
||||
|
||||
|
||||
def global_grad_norm(model: nn.Module) -> float:
|
||||
total = 0.0
|
||||
for p in model.parameters():
|
||||
if p.grad is None:
|
||||
continue
|
||||
total += float(p.grad.detach().pow(2).sum().item())
|
||||
return math.sqrt(total)
|
||||
|
||||
|
||||
def zero_grads(model: nn.Module) -> None:
|
||||
for p in model.parameters():
|
||||
if p.grad is not None:
|
||||
p.grad.detach_()
|
||||
p.grad.zero_()
|
||||
|
||||
|
||||
def loss_scaled_for_accum(logits: torch.Tensor, target: torch.Tensor, accum_steps: int, loss_fn: Callable[[torch.Tensor, torch.Tensor], torch.Tensor]) -> torch.Tensor:
|
||||
raw = loss_fn(logits, target)
|
||||
return raw / accum_steps
|
||||
|
||||
|
||||
def train_one_optimizer_step(
|
||||
model: nn.Module,
|
||||
optimizer: torch.optim.Optimizer,
|
||||
micro_batches: List[tuple[torch.Tensor, torch.Tensor]],
|
||||
loss_fn: Callable[[torch.Tensor, torch.Tensor], torch.Tensor],
|
||||
*,
|
||||
no_sync_until_last: bool,
|
||||
sync_counter: List[int],
|
||||
) -> tuple[float, float]:
|
||||
"""Run accum_steps micro batches, accumulate grads, step once.
|
||||
|
||||
Returns (total_unscaled_loss, grad_norm).
|
||||
"""
|
||||
accum_steps = len(micro_batches)
|
||||
zero_grads(model)
|
||||
total = 0.0
|
||||
for i, (x, y) in enumerate(micro_batches):
|
||||
is_last = i == accum_steps - 1
|
||||
if no_sync_until_last and not is_last:
|
||||
with no_sync_context(model):
|
||||
logits = model(x)
|
||||
loss = loss_scaled_for_accum(logits, y, accum_steps, loss_fn)
|
||||
loss.backward()
|
||||
else:
|
||||
logits = model(x)
|
||||
loss = loss_scaled_for_accum(logits, y, accum_steps, loss_fn)
|
||||
loss.backward()
|
||||
sync_counter[0] += 1
|
||||
total += float(loss.detach().item()) * accum_steps
|
||||
grad_norm = global_grad_norm(model)
|
||||
optimizer.step()
|
||||
return total / accum_steps, grad_norm
|
||||
|
||||
|
||||
class _NoSyncCtx:
|
||||
def __init__(self, model: nn.Module):
|
||||
self.model = model
|
||||
|
||||
def __enter__(self):
|
||||
return self
|
||||
|
||||
def __exit__(self, exc_type, exc, tb):
|
||||
return False
|
||||
|
||||
|
||||
def no_sync_context(model: nn.Module):
|
||||
"""Stand-in for DDP no_sync.
|
||||
|
||||
In DDP this skips the all-reduce on the trailing backward. In this
|
||||
single-process demo there is no collective to skip, but we still
|
||||
surface the call site so the pattern reads the same on a real cluster.
|
||||
"""
|
||||
return _NoSyncCtx(model)
|
||||
|
||||
|
||||
def run_config(
|
||||
effective_batch: int,
|
||||
accum_steps: int,
|
||||
*,
|
||||
in_dim: int,
|
||||
hidden: int,
|
||||
out_dim: int,
|
||||
num_steps: int,
|
||||
lr: float,
|
||||
seed: int,
|
||||
) -> CurvePoint:
|
||||
assert effective_batch % accum_steps == 0, "effective_batch must divide by accum_steps"
|
||||
micro_batch = effective_batch // accum_steps
|
||||
seed_everything(seed)
|
||||
gen = torch.Generator()
|
||||
gen.manual_seed(seed)
|
||||
model = make_model(in_dim, hidden, out_dim)
|
||||
optimizer = torch.optim.SGD(model.parameters(), lr=lr)
|
||||
loss_fn = nn.CrossEntropyLoss()
|
||||
|
||||
losses: List[float] = []
|
||||
step_times_ms: List[float] = []
|
||||
sync_counter = [0]
|
||||
total_samples = 0
|
||||
wall_start = time.perf_counter()
|
||||
for step in range(num_steps):
|
||||
t0 = time.perf_counter()
|
||||
micro_batches = [synthetic_batch(micro_batch, in_dim, out_dim, gen) for _ in range(accum_steps)]
|
||||
avg_loss, _grad_norm = train_one_optimizer_step(
|
||||
model,
|
||||
optimizer,
|
||||
micro_batches,
|
||||
loss_fn,
|
||||
no_sync_until_last=True,
|
||||
sync_counter=sync_counter,
|
||||
)
|
||||
wall_ms = (time.perf_counter() - t0) * 1000.0
|
||||
losses.append(avg_loss)
|
||||
step_times_ms.append(wall_ms)
|
||||
total_samples += effective_batch
|
||||
total_wall = time.perf_counter() - wall_start
|
||||
sps = total_samples / max(total_wall, 1e-6)
|
||||
step_times_ms.sort()
|
||||
median_ms = step_times_ms[len(step_times_ms) // 2]
|
||||
avg_loss = sum(losses) / len(losses)
|
||||
return CurvePoint(
|
||||
effective_batch=effective_batch,
|
||||
accum_steps=accum_steps,
|
||||
micro_batch=micro_batch,
|
||||
avg_loss=avg_loss,
|
||||
samples_per_sec=sps,
|
||||
median_step_ms=median_ms,
|
||||
sync_calls=sync_counter[0],
|
||||
steps=num_steps,
|
||||
)
|
||||
|
||||
|
||||
def sweep_effective_batches(
|
||||
*,
|
||||
micro_batch: int,
|
||||
accum_grid: Iterable[int],
|
||||
in_dim: int = 64,
|
||||
hidden: int = 128,
|
||||
out_dim: int = 16,
|
||||
num_steps: int = 25,
|
||||
lr: float = 0.05,
|
||||
seed: int = 0,
|
||||
) -> List[CurvePoint]:
|
||||
points: List[CurvePoint] = []
|
||||
for accum in accum_grid:
|
||||
eff = micro_batch * accum
|
||||
pt = run_config(
|
||||
effective_batch=eff,
|
||||
accum_steps=accum,
|
||||
in_dim=in_dim,
|
||||
hidden=hidden,
|
||||
out_dim=out_dim,
|
||||
num_steps=num_steps,
|
||||
lr=lr,
|
||||
seed=seed,
|
||||
)
|
||||
points.append(pt)
|
||||
return points
|
||||
|
||||
|
||||
def equivalence_check(
|
||||
*,
|
||||
in_dim: int = 32,
|
||||
hidden: int = 48,
|
||||
out_dim: int = 8,
|
||||
big_batch: int = 16,
|
||||
accum_steps: int = 4,
|
||||
lr: float = 0.1,
|
||||
seed: int = 7,
|
||||
) -> dict:
|
||||
"""One full batch step vs accum_steps micro-batches must match.
|
||||
|
||||
Scaled loss is `raw / accum_steps`; the accumulated gradient equals the
|
||||
full batch gradient up to floating point noise.
|
||||
"""
|
||||
assert big_batch % accum_steps == 0
|
||||
micro = big_batch // accum_steps
|
||||
|
||||
seed_everything(seed)
|
||||
gen_a = torch.Generator(); gen_a.manual_seed(seed)
|
||||
x, y = synthetic_batch(big_batch, in_dim, out_dim, gen_a)
|
||||
|
||||
seed_everything(seed)
|
||||
model_full = make_model(in_dim, hidden, out_dim)
|
||||
opt_full = torch.optim.SGD(model_full.parameters(), lr=lr)
|
||||
loss_fn = nn.CrossEntropyLoss()
|
||||
zero_grads(model_full)
|
||||
out = model_full(x)
|
||||
loss_full = loss_fn(out, y)
|
||||
loss_full.backward()
|
||||
full_params_before = [p.detach().clone() for p in model_full.parameters()]
|
||||
full_grads = [p.grad.detach().clone() for p in model_full.parameters()]
|
||||
opt_full.step()
|
||||
full_params_after = [p.detach().clone() for p in model_full.parameters()]
|
||||
|
||||
seed_everything(seed)
|
||||
model_accum = make_model(in_dim, hidden, out_dim)
|
||||
opt_accum = torch.optim.SGD(model_accum.parameters(), lr=lr)
|
||||
zero_grads(model_accum)
|
||||
chunks_x = list(torch.split(x, micro, dim=0))
|
||||
chunks_y = list(torch.split(y, micro, dim=0))
|
||||
for cx, cy in zip(chunks_x, chunks_y):
|
||||
scaled = loss_fn(model_accum(cx), cy) / accum_steps
|
||||
scaled.backward()
|
||||
accum_grads = [p.grad.detach().clone() for p in model_accum.parameters()]
|
||||
accum_params_before = [p.detach().clone() for p in model_accum.parameters()]
|
||||
opt_accum.step()
|
||||
accum_params_after = [p.detach().clone() for p in model_accum.parameters()]
|
||||
|
||||
grad_diffs = [
|
||||
float((a - b).abs().max().item())
|
||||
for a, b in zip(full_grads, accum_grads)
|
||||
]
|
||||
param_diffs = [
|
||||
float((a - b).abs().max().item())
|
||||
for a, b in zip(full_params_after, accum_params_after)
|
||||
]
|
||||
return {
|
||||
"max_grad_diff": max(grad_diffs),
|
||||
"max_param_diff": max(param_diffs),
|
||||
"params_init_match": all(
|
||||
torch.equal(a, b) for a, b in zip(full_params_before, accum_params_before)
|
||||
),
|
||||
}
|
||||
|
||||
|
||||
def write_curve(points: List[CurvePoint], path: Path) -> None:
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
payload = {
|
||||
"schema": "accum-curve.v1",
|
||||
"points": [asdict(p) for p in points],
|
||||
}
|
||||
path.write_text(json.dumps(payload, indent=2) + "\n")
|
||||
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
p = argparse.ArgumentParser()
|
||||
p.add_argument("--micro-batch", type=int, default=4)
|
||||
p.add_argument("--accum-grid", type=str, default="1,2,4,8,16")
|
||||
p.add_argument("--num-steps", type=int, default=25)
|
||||
p.add_argument("--seed", type=int, default=0)
|
||||
p.add_argument("--lr", type=float, default=0.05)
|
||||
p.add_argument("--no-write", action="store_true")
|
||||
return p.parse_args()
|
||||
|
||||
|
||||
def main() -> int:
|
||||
args = parse_args()
|
||||
accum_grid = [int(s) for s in args.accum_grid.split(",") if s.strip()]
|
||||
print("equivalence check (full batch vs accumulated)")
|
||||
eq = equivalence_check()
|
||||
print(json.dumps(eq, indent=2))
|
||||
assert eq["max_grad_diff"] < 1e-4, f"gradients diverge: {eq['max_grad_diff']}"
|
||||
assert eq["max_param_diff"] < 1e-4, f"params diverge: {eq['max_param_diff']}"
|
||||
print("equivalence holds. running sweep...")
|
||||
|
||||
points = sweep_effective_batches(
|
||||
micro_batch=args.micro_batch,
|
||||
accum_grid=accum_grid,
|
||||
num_steps=args.num_steps,
|
||||
lr=args.lr,
|
||||
seed=args.seed,
|
||||
)
|
||||
header = f"{'eff_batch':>10} {'accum':>5} {'micro':>5} {'sps':>10} {'median_ms':>10} {'syncs':>6} {'loss':>8}"
|
||||
print(header)
|
||||
for p in points:
|
||||
print(
|
||||
f"{p.effective_batch:>10} {p.accum_steps:>5} {p.micro_batch:>5} "
|
||||
f"{p.samples_per_sec:>10.1f} {p.median_step_ms:>10.2f} {p.sync_calls:>6} {p.avg_loss:>8.4f}"
|
||||
)
|
||||
if not args.no_write:
|
||||
write_curve(points, LOG_PATH)
|
||||
print(f"wrote {LOG_PATH}")
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit(main())
|
||||
@@ -0,0 +1,104 @@
|
||||
"""Tests for gradient accumulation core paths."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import sys
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
|
||||
HERE = Path(__file__).parent
|
||||
sys.path.insert(0, str(HERE))
|
||||
|
||||
import main as accum
|
||||
|
||||
|
||||
class EquivalenceTests(unittest.TestCase):
|
||||
def test_full_batch_matches_accumulated(self):
|
||||
diffs = accum.equivalence_check()
|
||||
self.assertLess(diffs["max_grad_diff"], 1e-4)
|
||||
self.assertLess(diffs["max_param_diff"], 1e-4)
|
||||
self.assertTrue(diffs["params_init_match"])
|
||||
|
||||
def test_loss_scaled_for_accum_divides(self):
|
||||
import torch
|
||||
from torch import nn
|
||||
|
||||
logits = torch.tensor([[2.0, 0.5, -1.0]])
|
||||
target = torch.tensor([0])
|
||||
fn = nn.CrossEntropyLoss()
|
||||
raw = fn(logits, target)
|
||||
scaled = accum.loss_scaled_for_accum(logits, target, 4, fn)
|
||||
self.assertAlmostEqual(float(scaled.item()), float(raw.item()) / 4.0, places=6)
|
||||
|
||||
|
||||
class SweepTests(unittest.TestCase):
|
||||
def test_sweep_returns_one_point_per_accum(self):
|
||||
points = accum.sweep_effective_batches(
|
||||
micro_batch=2,
|
||||
accum_grid=[1, 2, 4],
|
||||
in_dim=16,
|
||||
hidden=24,
|
||||
out_dim=4,
|
||||
num_steps=5,
|
||||
lr=0.05,
|
||||
)
|
||||
self.assertEqual(len(points), 3)
|
||||
self.assertEqual([p.accum_steps for p in points], [1, 2, 4])
|
||||
self.assertEqual([p.effective_batch for p in points], [2, 4, 8])
|
||||
for p in points:
|
||||
self.assertGreater(p.samples_per_sec, 0.0)
|
||||
self.assertGreater(p.steps, 0)
|
||||
self.assertGreater(p.median_step_ms, 0.0)
|
||||
|
||||
def test_sync_calls_equal_step_count(self):
|
||||
points = accum.sweep_effective_batches(
|
||||
micro_batch=2,
|
||||
accum_grid=[1, 4],
|
||||
in_dim=16,
|
||||
hidden=24,
|
||||
out_dim=4,
|
||||
num_steps=7,
|
||||
lr=0.05,
|
||||
)
|
||||
for p in points:
|
||||
self.assertEqual(p.sync_calls, p.steps)
|
||||
|
||||
|
||||
class CurveOutputTests(unittest.TestCase):
|
||||
def test_write_curve_round_trip(self):
|
||||
import tempfile
|
||||
|
||||
points = accum.sweep_effective_batches(
|
||||
micro_batch=2,
|
||||
accum_grid=[1, 2],
|
||||
in_dim=8,
|
||||
hidden=12,
|
||||
out_dim=3,
|
||||
num_steps=3,
|
||||
lr=0.05,
|
||||
)
|
||||
with tempfile.TemporaryDirectory() as tmp:
|
||||
path = Path(tmp) / "curve.json"
|
||||
accum.write_curve(points, path)
|
||||
payload = json.loads(path.read_text())
|
||||
self.assertEqual(payload["schema"], "accum-curve.v1")
|
||||
self.assertEqual(len(payload["points"]), 2)
|
||||
self.assertIn("samples_per_sec", payload["points"][0])
|
||||
|
||||
def test_effective_batch_must_divide_accum(self):
|
||||
with self.assertRaises(AssertionError):
|
||||
accum.run_config(
|
||||
effective_batch=5,
|
||||
accum_steps=2,
|
||||
in_dim=8,
|
||||
hidden=12,
|
||||
out_dim=3,
|
||||
num_steps=2,
|
||||
lr=0.05,
|
||||
seed=0,
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,147 @@
|
||||
# Gradient Accumulation
|
||||
|
||||
> Train at an effective batch you cannot afford, one micro-batch at a time. Scale the loss, hold the optimizer step, and let the gradients pile up.
|
||||
|
||||
**Type:** Build
|
||||
**Languages:** Python
|
||||
**Prerequisites:** Phase 19 lessons 42 to 45
|
||||
**Time:** ~90 minutes
|
||||
|
||||
## Learning Objectives
|
||||
|
||||
- Derive the effective batch identity: `effective_batch = micro_batch * accum_steps`.
|
||||
- Implement loss-per-micro-batch scaling so the accumulated gradient matches a single full-batch backward.
|
||||
- Skip optimizer synchronization until the last micro-batch (sync-on-last-step).
|
||||
- Read a throughput against effective batch curve and explain the diminishing return.
|
||||
|
||||
## The Problem
|
||||
|
||||
You want to train at an effective batch of 512 because the loss curve is smoother and the optimizer step makes more sense at that scale. The accelerator on the desk holds 32 examples before it runs out of memory. Doubling the batch is not an option. Halving the model is not an option. The trick the field reached for in 2017 and never stopped using is to run 16 backward passes, let the gradients accumulate inside the parameter buffers, and only step the optimizer when the count reaches the target.
|
||||
|
||||
The risk is that the loss is no longer the same number it was at the bigger batch. The cross entropy of 16 mini-batches summed naively is 16 times the loss of one full batch. Without scaling, the gradient direction is correct but the magnitude is wrong, and the optimizer step is 16 times too big. The fix is one division. The fix is also easy to forget.
|
||||
|
||||
## The Concept
|
||||
|
||||
```mermaid
|
||||
flowchart LR
|
||||
start[start] --> zero[zero grads]
|
||||
zero --> mb1[micro batch 1: forward + scaled backward]
|
||||
mb1 --> mb2[micro batch 2: forward + scaled backward]
|
||||
mb2 --> dots[...]
|
||||
dots --> mbN[micro batch N: forward + scaled backward + sync]
|
||||
mbN --> step[optimizer step]
|
||||
step --> next[next effective step]
|
||||
```
|
||||
|
||||
The contract is short:
|
||||
|
||||
- Loss for each micro-batch is divided by `accum_steps` before `backward()`. PyTorch sums gradients into `param.grad` by default; the division pushes the running sum back into the right scale.
|
||||
- The optimizer step fires once per effective batch, after the last micro-batch's backward. Stepping mid-accumulation skews every parameter the rest of the run depends on.
|
||||
- The optimizer's state (momentum buffers, Adam moments) advances once per effective step, not once per micro-batch. The exponential moving averages would otherwise see the wrong frequency and burn through the schedule.
|
||||
- On a single device this is bookkeeping. On a multi-rank cluster the same pattern wraps the non-final micro-batches in a `no_sync` context that skips the gradient all-reduce; the last micro-batch reduces the full accumulated gradient in one pass instead of paying the network cost N times.
|
||||
|
||||
### The equivalence proof in code
|
||||
|
||||
```python
|
||||
loss = criterion(model(x_full), y_full)
|
||||
loss.backward()
|
||||
opt.step()
|
||||
```
|
||||
|
||||
is equivalent to
|
||||
|
||||
```python
|
||||
for x, y in chunks(x_full, y_full, n):
|
||||
scaled = criterion(model(x), y) / n
|
||||
scaled.backward()
|
||||
opt.step()
|
||||
```
|
||||
|
||||
up to floating point summation order. The accumulated gradient buffer at the end of the loop is the same tensor that a single full-batch backward would produce. The lesson code asserts this with a max-abs difference under 1e-4 in `equivalence_check`.
|
||||
|
||||
### Where the cost goes
|
||||
|
||||
Each micro-batch costs one forward and one backward. With accumulation you trade memory for time. The throughput curve in `outputs/accum-curve.json` shows what happens as the effective batch grows at fixed micro-batch:
|
||||
|
||||
```mermaid
|
||||
flowchart TD
|
||||
micro[fixed micro batch] --> small[small accum: low loss noise budget, high stepper churn]
|
||||
micro --> large[large accum: smooth loss, optimizer step rare]
|
||||
small --> sps1[samples per second saturates at hardware limit]
|
||||
large --> sps2[samples per second still hits hardware limit]
|
||||
sps1 --> note[total samples per optimizer step scales linearly with accum]
|
||||
sps2 --> note
|
||||
```
|
||||
|
||||
There is no free lunch. Doubling `accum_steps` doubles the wall time per optimizer step. What changes is the variance of the gradient estimate: at the same wall budget you have made fewer optimizer steps but each one was averaged over more samples. The literature treats large batch and small batch as different optimization problems; the lesson here is mechanical, not statistical.
|
||||
|
||||
## Build It
|
||||
|
||||
`code/main.py` is the runnable artifact. It does three things.
|
||||
|
||||
### Step 1: equivalence check
|
||||
|
||||
`equivalence_check()` builds two copies of the same network with the same seed. One sees a 16-sample batch in one forward pass. The other sees four 4-sample chunks with the loss divided by four. The function compares the gradient buffers before the optimizer step and the parameters after. The assertion is `max_abs_diff < 1e-4`.
|
||||
|
||||
### Step 2: sync-on-last-step pattern
|
||||
|
||||
`train_one_optimizer_step` walks micro-batches. For every micro-batch except the last it enters `no_sync_context(model)`. On a single process the context is a no-op; on DDP this is where the gradient all-reduce is skipped. The bookkeeping is the same regardless. A `sync_counter` records how many times we left the no_sync scope; for N micro-batches the count is one per effective step, not N.
|
||||
|
||||
### Step 3: the throughput curve
|
||||
|
||||
`sweep_effective_batches` runs the same model with a fixed micro-batch and a list of accumulation steps. For each setting it logs:
|
||||
|
||||
- `samples_per_sec`: total samples seen divided by wall time
|
||||
- `median_step_ms`: 50th percentile per effective step
|
||||
- `sync_calls`: collective points exercised
|
||||
- `avg_loss`: average across the sweep's optimizer steps
|
||||
|
||||
The output lands in `outputs/accum-curve.json` and is reusable from a notebook.
|
||||
|
||||
Run it:
|
||||
|
||||
```bash
|
||||
python3 code/main.py
|
||||
```
|
||||
|
||||
The script prints the equivalence diff, then the sweep table, then the JSON path. Exit code zero.
|
||||
|
||||
## Use It
|
||||
|
||||
In production training, gradient accumulation lives behind one knob. PyTorch's pattern is `accumulation_steps = effective_batch // (micro_batch * world_size)`. Frameworks that you are not allowed to use here wrap the same loop, but the steps are the same: scale the loss, skip sync on non-final micros, accumulate, step once.
|
||||
|
||||
Three patterns in the wild:
|
||||
|
||||
- The micro-batch size is chosen to saturate device memory. Anything smaller wastes accelerator cycles. Anything larger crashes.
|
||||
- The effective batch is chosen from a learning rate schedule. Large effective batches need scaled learning rates and warmup; this is the linear scaling rule talked about since 2017.
|
||||
- The accumulation count is the bridge between the two and the only knob you are free to tune at runtime without rewriting the data loader.
|
||||
|
||||
## Ship It
|
||||
|
||||
`outputs/skill-gradient-accumulation.md` captures the recipe so a peer can drop it into a new repo: scale loss by `accum_steps`, skip optimizer sync on non-final micros, step the optimizer once per effective batch, log throughput against effective batch as JSON so the trade is visible.
|
||||
|
||||
## Exercises
|
||||
|
||||
1. Re-run the sweep with `--num-steps 100` and plot samples per second against effective batch. Where does the curve flatten?
|
||||
2. Add a wrong scaling variant (no division) and show the parameter diff at step 1 against the reference.
|
||||
3. Swap SGD for AdamW and confirm the optimizer state advances once per effective step, not once per micro-batch.
|
||||
4. Introduce a real `DistributedDataParallel` wrapper and route the `no_sync_context` to its method. Confirm sync_calls drops by N-1 per effective batch.
|
||||
5. Modify the equivalence check to compare two different micro splits (2 by 8 vs 4 by 4) and explain any tolerance you need to relax.
|
||||
|
||||
## Key Terms
|
||||
|
||||
| Term | What people say | What it actually means |
|
||||
|------|-----------------|------------------------|
|
||||
| Micro batch | The batch you forward | The slice that fits in memory in a single forward pass |
|
||||
| Accum steps | Backward passes per step | Number of backwards summed before one optimizer step |
|
||||
| Effective batch | The batch | Micro batch times accum steps times data parallel world size |
|
||||
| Loss scaling | Divide by N | Per-micro-batch division so summed gradients match full batch |
|
||||
| Sync on last | Skip the rest | Only run the gradient collective on the last backward in the window |
|
||||
|
||||
## Further Reading
|
||||
|
||||
- PyTorch docs on `DistributedDataParallel.no_sync` for the production version of the sync-on-last-step trick.
|
||||
- Goyal et al., 2017, on linear scaling for large batch training, the canonical reason to care about effective batch.
|
||||
- PyTorch issue tracker on gradient accumulation interactions with mixed precision unscaling.
|
||||
- Phase 19 lessons 42 to 45 cover the model, data loader, optimizer, and trainer scaffolding this lesson assumes.
|
||||
- Phase 19 lesson 47 covers checkpoint and resume so a long accumulation run survives a wallclock cap.
|
||||
@@ -0,0 +1,35 @@
|
||||
{
|
||||
"schema": "accum-curve.v1",
|
||||
"points": [
|
||||
{
|
||||
"effective_batch": 4,
|
||||
"accum_steps": 1,
|
||||
"micro_batch": 4,
|
||||
"avg_loss": 2.7746317982673645,
|
||||
"samples_per_sec": 10589.599020053975,
|
||||
"median_step_ms": 0.36970898509025574,
|
||||
"sync_calls": 8,
|
||||
"steps": 8
|
||||
},
|
||||
{
|
||||
"effective_batch": 8,
|
||||
"accum_steps": 2,
|
||||
"micro_batch": 4,
|
||||
"avg_loss": 2.778739705681801,
|
||||
"samples_per_sec": 12129.73295382913,
|
||||
"median_step_ms": 0.6751659093424678,
|
||||
"sync_calls": 8,
|
||||
"steps": 8
|
||||
},
|
||||
{
|
||||
"effective_batch": 16,
|
||||
"accum_steps": 4,
|
||||
"micro_batch": 4,
|
||||
"avg_loss": 2.7735252678394318,
|
||||
"samples_per_sec": 12667.36269537429,
|
||||
"median_step_ms": 1.2691670563071966,
|
||||
"sync_calls": 8,
|
||||
"steps": 8
|
||||
}
|
||||
]
|
||||
}
|
||||
+33
@@ -0,0 +1,33 @@
|
||||
---
|
||||
name: gradient-accumulation
|
||||
description: Train at an effective batch larger than device memory by scaling micro-batch losses and stepping the optimizer once per window.
|
||||
version: 1.0.0
|
||||
phase: 19
|
||||
lesson: 46
|
||||
tags: [training, batch-size, distributed, scaling]
|
||||
---
|
||||
|
||||
## When to use
|
||||
|
||||
Effective batch is the lever that smooths the gradient and matches the learning rate schedule. When you cannot afford it in a single forward pass, this is the recipe.
|
||||
|
||||
## Recipe
|
||||
|
||||
1. Pick `micro_batch` as the largest size that fits in memory and saturates the accelerator.
|
||||
2. Pick `effective_batch` from the learning rate schedule.
|
||||
3. Set `accum_steps = effective_batch // (micro_batch * world_size)` and assert it divides evenly.
|
||||
4. Per micro batch: `loss = criterion(model(x), y) / accum_steps; loss.backward()`.
|
||||
5. On non-final micros, enter `model.no_sync()` to skip the gradient all-reduce in DDP.
|
||||
6. After the last micro batch, run `optimizer.step()` once. Zero gradients before the next window.
|
||||
7. The optimizer state advances once per effective batch; the learning rate schedule ticks once per effective batch.
|
||||
|
||||
## Logging
|
||||
|
||||
Emit a small JSON record per effective step with `samples_per_sec`, `median_step_ms`, `sync_calls`, `accum_steps`, `effective_batch`. Without this the cost trade is invisible.
|
||||
|
||||
## Failure modes
|
||||
|
||||
- Forgetting the `/ accum_steps` scaling: gradients explode by N.
|
||||
- Stepping mid-window: parameters drift.
|
||||
- Sync on every micro batch: network bound for no statistical gain.
|
||||
- Mixing this with mixed precision unscaling: scale the unscaled loss only.
|
||||
@@ -0,0 +1,78 @@
|
||||
{
|
||||
"lesson": "46-gradient-accumulation",
|
||||
"title": "Gradient Accumulation",
|
||||
"questions": [
|
||||
{
|
||||
"stage": "pre",
|
||||
"question": "What is the effective batch identity that gradient accumulation expresses?",
|
||||
"options": [
|
||||
"effective_batch = micro_batch / accum_steps",
|
||||
"effective_batch = micro_batch * accum_steps",
|
||||
"effective_batch = micro_batch + accum_steps",
|
||||
"effective_batch = world_size only"
|
||||
],
|
||||
"correct": 1,
|
||||
"explanation": "Micro batch is what fits in memory; accumulation steps are how many forward + backward passes pile gradients into the same buffer before one optimizer step."
|
||||
},
|
||||
{
|
||||
"stage": "pre",
|
||||
"question": "Why must the loss be divided by accum_steps before backward on each micro batch?",
|
||||
"options": [
|
||||
"It makes the loss curve prettier",
|
||||
"PyTorch sums gradients into param.grad by default, so without the division the accumulated gradient is N times too large and the optimizer step is N times too aggressive",
|
||||
"It avoids NaNs",
|
||||
"It is required for mixed precision only"
|
||||
],
|
||||
"correct": 1,
|
||||
"explanation": "The accumulation buffer is a sum. Per-micro-batch scaling by 1/N pushes that sum back into the same scale a single full-batch backward would produce."
|
||||
},
|
||||
{
|
||||
"stage": "check",
|
||||
"question": "When does the optimizer step run inside the accumulation loop?",
|
||||
"options": [
|
||||
"After every micro batch",
|
||||
"Only after the last micro batch in the accumulation window",
|
||||
"Twice per window for safety",
|
||||
"Never; the optimizer is replaced"
|
||||
],
|
||||
"correct": 1,
|
||||
"explanation": "Stepping mid-accumulation contaminates every parameter the rest of the run depends on. The optimizer step fires once per effective batch."
|
||||
},
|
||||
{
|
||||
"stage": "check",
|
||||
"question": "What does the no_sync context wrap on a real multi-GPU run?",
|
||||
"options": [
|
||||
"The forward pass",
|
||||
"The optimizer step",
|
||||
"The non-final micro batches, so the gradient all-reduce only fires after the last backward instead of N times",
|
||||
"The data loader"
|
||||
],
|
||||
"correct": 2,
|
||||
"explanation": "DDP's no_sync skips the gradient collective. Wrapping all but the last micro batch turns N collectives per effective step into one."
|
||||
},
|
||||
{
|
||||
"stage": "check",
|
||||
"question": "What does the equivalence_check function in main.py assert?",
|
||||
"options": [
|
||||
"That the loss converges",
|
||||
"That a single full-batch backward and an accum_steps chunked backward with loss / N produce the same gradient buffer and the same post-step parameters up to a small tolerance",
|
||||
"That throughput is constant",
|
||||
"That gradients are zero"
|
||||
],
|
||||
"correct": 1,
|
||||
"explanation": "The assertion bound is max-abs-diff under 1e-4. Without the loss scaling the diff blows up; with the scaling it sits at floating point noise."
|
||||
},
|
||||
{
|
||||
"stage": "post",
|
||||
"question": "Reading the throughput against effective batch curve, what stays roughly constant and what scales?",
|
||||
"options": [
|
||||
"Samples per second drops with larger effective batch",
|
||||
"Samples per second saturates near the hardware limit while wall time per optimizer step scales linearly with accum_steps; optimizer steps per second is what falls",
|
||||
"Both metrics scale linearly with accum",
|
||||
"Throughput goes up because the model gets faster"
|
||||
],
|
||||
"correct": 1,
|
||||
"explanation": "Each micro batch costs the same forward + backward. Accumulation buys statistical smoothing per optimizer step, not raw throughput; the curve is a useful reality check against folklore."
|
||||
}
|
||||
]
|
||||
}
|
||||
Reference in New Issue
Block a user