Team Ai
Modelpublic

Codeseys/composer-replication-framework

sourceHugging Facemitupdated 4mo agoView on Hugging Face
0likes
__init__.py138 linesDownload Raw Back to diloco
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