Codeseys/composer-replication-framework
0
1"""composer_diloco.py — DiLoCo outer-loop wrapper for Composer Replication Framework.2 3Wraps `torchft.local_sgd.DiLoCo` with the framework's conventions:4- Sign convention is documented LOUDLY here once and tested via Spike 008.5- The wrapper exposes the same constructor shape as torchft's DiLoCo so a6 future swap-in of the upstream class is a one-line change.7- Vanilla DiLoCo (Douillard et al. 2023, arXiv:2311.08105) =8 `fragment_sync_delay=0`, single fragment. Streaming DiLoCo (Douillard et9 al., arXiv:2501.18512 "Streaming DiLoCo with overlapping communication";10 the separate Eager-Updates work is Kale et al., arXiv:2502.12996 — citation11 corrected per deepread finding V7) = non-zero delay, multiple fragments.12 Spike 008 uses vanilla; Streaming is configured by the same API.13 14Reference: `docs/adrs/ADR-003-diloco-impl.md`.15 16Sign convention (READ THIS BEFORE TOUCHING):17 DiLoCo defines pseudo-gradient as18 19 pseudograd = θ_initial - θ_local20 21 (per torchft's `_save_grads()` at line 324 of `torchft/local_sgd.py`).22 This is the **negative** of the local update direction (the local23 update *moved* params from θ_initial toward θ_local).24 25 Standard SGD subtracts gradients: `p.data ← p.data - lr * grad`.26 So when the outer optimizer runs after `restore_parameters()` puts27 p.data back to θ_initial:28 29 p.data ← θ_initial - lr * (θ_initial - θ_local)30 = θ_initial + lr * (θ_local - θ_initial)31 32 For `lr=1, momentum=0` this lands exactly at θ_local. For `lr<1` it33 interpolates between θ_initial and θ_local (the standard DiLoCo outer34 step). Adding Nesterov momentum accumulates the local-update direction35 across outer rounds.36 37 No negation in our outer optimizer wrapper. The test38 `test_diloco_pseudogradient_sign_convention` in39 `spikes/008-streaming-diloco/tests/test_diloco_smoke.py` pins this40 arithmetic and reports both the expected and the wrong-sign value on41 failure for fast diagnosis if torchft ever flips its convention.42"""43from __future__ import annotations44 45from typing import Any46 47import torch48 49# Import lazily — torchft is an optional dep at framework level.50_TORCHFT_AVAILABLE = False51DiLoCo: Any = None52Manager: Any = None53_DummyWork: Any = None54try:55 from torchft.local_sgd import DiLoCo as _DiLoCo # type: ignore[import]56 from torchft.manager import Manager as _Manager # type: ignore[import]57 from torchft.work import _DummyWork as __DummyWork # type: ignore[import]58 59 _TORCHFT_AVAILABLE = True60 DiLoCo = _DiLoCo61 Manager = _Manager62 _DummyWork = __DummyWork63except ImportError: # pragma: no cover — only hits in lighter-weight CI envs64 pass65 66 67def make_diloco_outer_loop(68 manager: Any,69 model_fragments: list[torch.nn.Module],70 inner_optimizer: torch.optim.Optimizer,71 *,72 outer_lr: float = 0.7,73 outer_momentum: float = 0.9,74 nesterov: bool = True,75 sync_every: int = 100,76 fragment_sync_delay: int = 0,77 fragment_update_alpha: float = 0.0,78) -> Any:79 """Construct a DiLoCo wrapper around `model_fragments` with default DiLoCo hyperparams.80 81 Default hyperparams (DiLoCo paper §3.2):82 outer_lr = 0.7, outer_momentum = 0.9, Nesterov83 84 Args:85 manager: torchft.Manager (or test mock with `.allreduce`, `.should_commit`,86 `.current_step`, `.start_quorum`)87 model_fragments: list of nn.Modules. For vanilla DiLoCo, pass [whole_model].88 For Streaming DiLoCo with N fragments, pass [frag_0, frag_1, ..., frag_N-1].89 inner_optimizer: any torch.optim.Optimizer. Steps every batch.90 outer_lr / outer_momentum / nesterov: outer SGD hyperparams.91 Override defaults only if you know why.92 sync_every: number of inner steps per outer round.93 fragment_sync_delay: 0 = vanilla DiLoCo (sync at outer round).94 >0 = Streaming DiLoCo with overlapped sync. Requires CUDA streams.95 fragment_update_alpha: 0 = full replacement of fragment params on sync.96 >0 = exponential mixing weight. Streaming DiLoCo only.97 98 Returns:99 A torchft.local_sgd.DiLoCo instance configured for the framework's100 conventions. Use as a context manager:101 with make_diloco_outer_loop(...) as outer:102 for step in range(N):103 inner_optimizer.zero_grad()104 loss = compute_loss(...)105 loss.backward()106 inner_optimizer.step() # outer sync fires automatically107 """108 if not _TORCHFT_AVAILABLE:109 raise RuntimeError(110 "torchft is not installed. `pip install torchft-nightly` to use DiLoCo."111 )112 113 outer_optimizer = torch.optim.SGD(114 [p for frag in model_fragments for p in frag.parameters()],115 lr=outer_lr,116 momentum=outer_momentum,117 nesterov=nesterov,118 )119 120 return DiLoCo(121 manager=manager,122 model_fragments=model_fragments,123 inner_optimizer=inner_optimizer,124 outer_optimizer=outer_optimizer,125 sync_every=sync_every,126 fragment_sync_delay=fragment_sync_delay,127 fragment_update_alpha=fragment_update_alpha,128 )129 130 131__all__ = [132 "make_diloco_outer_loop",133 "DiLoCo",134 "Manager",135 "_DummyWork",136 "_TORCHFT_AVAILABLE",137]138 