Team Ai
Modelpublic

Codeseys/composer-replication-framework

sourceHugging Facemitupdated 4mo agoView on Hugging Face
0likes
test_serverless_diloco_integration.py353 linesDownload Raw Back to tests
1"""End-to-end MockManager × torchft.DiLoCo integration test.2 3Closes the Wave 13 cross-model adversarial-review gap (Suggestion 4):4the original MockManager was advertised as a drop-in for torchft.Manager5but only stubbed `.allreduce / .should_commit / .start_quorum`. DiLoCo's6real call surface (audited from `torchft/local_sgd.py` v2026-spring) also7includes `current_step()`, `disallow_state_dict_read()`,8`allow_state_dict_read()`, `register_state_dict_fn()`, and the9`_use_async_quorum` attribute — plus `allreduce()` must return a Work-like10object with `.wait()`, not a raw tensor.11 12This test runs ONE full DiLoCo outer round (sync_every inner steps + the13sync) against a tiny `nn.Linear(4, 4)` with `world_size=1` so the14object-store rendezvous is trivial. It verifies:15 161. Construction does not raise.172. Running through one full outer round does not raise AttributeError18   (which is what the old MockManager would have hit at `current_step()`).193. The model parameters change after the outer step fires (proving the20   outer SGD path actually executed end-to-end, not just that the21   inner-step hooks ran).224. The MockManager's step counter advanced exactly once (one outer round23   ⇒ one start_quorum bump).245. DiLoCo registered a state-dict fn per fragment.25"""26from __future__ import annotations27 28import pytest29import torch30 31torchft = pytest.importorskip(32    "torchft.local_sgd",33    reason="torchft must be installed to run the DiLoCo integration test",34)35 36from composer_replication.diloco import make_diloco_outer_loop37from composer_replication.diloco.serverless.allreduce import (38    MockManager,39    ObjectStoreAllReduce,40    _ImmediateWork,41)42 43 44def _make_store(tmp_path) -> ObjectStoreAllReduce:45    return ObjectStoreAllReduce(46        uri=str(tmp_path),47        rank=0,48        world_size=1,49        timeout_s=10.0,50        poll_interval_s=0.05,51    )52 53 54def test_mockmanager_has_full_diloco_call_surface(tmp_path):55    """Audited methods/attrs from torchft/local_sgd.py DiLoCo path must exist."""56    mgr = MockManager(_make_store(tmp_path))57    # Methods DiLoCo invokes58    for attr in (59        "allreduce",60        "should_commit",61        "start_quorum",62        "current_step",63        "disallow_state_dict_read",64        "allow_state_dict_read",65        "register_state_dict_fn",66        "wait_quorum",67        "is_leader",68    ):69        assert callable(getattr(mgr, attr)), f"MockManager missing method: {attr}"70    # Attributes DiLoCo reads at construction / runtime71    assert hasattr(mgr, "_use_async_quorum")72    assert mgr._use_async_quorum is False  # DiLoCo.__init__ rejects True73    assert hasattr(mgr, "num_participants")74    assert hasattr(mgr, "rank")75 76 77def test_mockmanager_allreduce_returns_workshaped(tmp_path):78    """DiLoCo stores the allreduce return in a list and calls `.wait()` later."""79    mgr = MockManager(_make_store(tmp_path))80    work = mgr.allreduce(torch.zeros(2, 2))81    # It must look like torch.distributed.Work / torchft._DummyWork82    assert hasattr(work, "wait"), "allreduce return must have .wait() (DiLoCo calls it)"83    assert callable(work.wait)84    # No-op .wait() must not raise on a synchronous mock.85    assert work.wait() is True86    # Defensive: get_future() should also work (some torch paths probe it).87    fut = work.get_future()88    assert fut is None or hasattr(fut, "wait")89    # Concrete type90    assert isinstance(work, _ImmediateWork)91 92 93def test_mockmanager_diloco_outer_round_completes(tmp_path):94    """Run one full inner+outer DiLoCo round and verify params change.95 96    With world_size=1 + MockManager → ObjectStoreAllReduce(file://), the97    rendezvous is single-process, so this test runs synchronously. We98    use `sync_every=4` and run exactly 4 inner-optimizer steps; at the99    4th step DiLoCo's post-hook fires `prepare_sync` then `perform_sync`,100    exercising the entire MockManager surface.101    """102    torch.manual_seed(0)103    model = torch.nn.Linear(4, 4, bias=False)104    initial_params = model.weight.detach().clone()105 106    inner_optim = torch.optim.SGD(model.parameters(), lr=0.1)107 108    store = _make_store(tmp_path)109    manager = MockManager(store)110 111    diloco = make_diloco_outer_loop(112        manager=manager,113        model_fragments=[model],114        inner_optimizer=inner_optim,115        outer_lr=0.7,116        outer_momentum=0.9,117        nesterov=True,118        sync_every=4,119        fragment_sync_delay=0,120        fragment_update_alpha=0.0,121    )122 123    # Sanity: DiLoCo registered a state-dict fn for our single fragment.124    assert len(manager._state_dict_fns) == 1, (125        f"expected 1 fragment registration, got {list(manager._state_dict_fns)}"126    )127 128    x = torch.randn(2, 4)129    target = torch.randn(2, 4)130 131    with diloco:132        for _ in range(4):  # exactly sync_every inner steps → one outer round133            inner_optim.zero_grad()134            loss = ((model(x) - target) ** 2).mean()135            loss.backward()136            # Must NOT raise AttributeError on current_step / state_dict_read /137            # register_state_dict_fn / etc. The original MockManager would have138            # crashed here on the very first step's _step_pre_hook calling139            # disallow_state_dict_read.140            inner_optim.step()141 142    # After exactly one outer round, the MockManager's step counter143    # should have advanced exactly once (start_quorum is called once).144    assert manager.current_step() == 1, (145        f"expected current_step()==1 after one outer round, got {manager.current_step()}"146    )147 148    # The outer SGD step actually fired ⇒ params differ from initial.149    final_params = model.weight.detach().clone()150    assert not torch.allclose(initial_params, final_params), (151        "model params unchanged after outer round — outer optimizer never ran"152    )153 154 155def _diloco_replica_one_outer_round(156    rendezvous_uri: str,157    world_size: int,158    sync_every: int,159) -> dict:160    """Top-level entry — must be importable for multiprocessing 'spawn'.161 162    Each replica:163      1. seeds torch with a SHARED seed for model init (DiLoCo's standard164         assumption: all replicas start with identical weights — DiLoCo165         only averages pseudo-gradients, not absolute weights, so divergent166         inits would never reconcile).167      2. builds nn.Linear(4, 4, bias=False) + SGD inner optimizer.168      3. trains on RANK-SPECIFIC data so each replica's inner-trained169         weights diverge during the inner loop (this is what gives the170         pseudo-gradient real cross-rank variance — without it, the171         averaging is observationally a no-op).172      4. runs `sync_every` inner steps inside `make_diloco_outer_loop` —173         this fires exactly one outer round.174      5. returns the final flattened weight vector and the pre-outer175         (purely-inner) weights.176 177    The test then asserts both ranks' final weights are identical178    (allclose), which proves the cross-replica allreduce of the179    pseudo-gradient ran end-to-end. The pre-outer weights MUST differ180    across ranks (proving rank-specific data drove divergence in the181    inner loop) — otherwise the convergence assertion is vacuous.182    """183    import os as _os184    import torch as _torch185    import torch.nn as _nn186 187    from composer_replication.diloco import make_diloco_outer_loop188    from composer_replication.diloco.serverless.allreduce import (189        MockManager,190        ObjectStoreAllReduce,191    )192 193    rank = int(_os.environ["REPLICA_RANK"])194 195    # SHARED init seed — both replicas start with identical weights, as196    # DiLoCo assumes. (DiLoCo averages pseudo-gradients, not weights, so197    # divergent inits would never reconcile and the convergence claim198    # would be incorrect.)199    _torch.manual_seed(0)200    model = _nn.Linear(4, 4, bias=False)201    initial = model.weight.detach().clone()202 203    inner_optim = _torch.optim.SGD(model.parameters(), lr=0.1)204 205    store = ObjectStoreAllReduce(206        rendezvous_uri,207        rank=rank,208        world_size=world_size,209        timeout_s=120.0,210        poll_interval_s=0.05,211    )212    manager = MockManager(store)213 214    diloco = make_diloco_outer_loop(215        manager=manager,216        model_fragments=[model],217        inner_optimizer=inner_optim,218        sync_every=sync_every,219    )220 221    # RANK-SPECIFIC data so the inner-trained weights diverge before the222    # outer sync — this is what makes "post-sync convergence" a real223    # property to verify rather than a tautology.224    _torch.manual_seed(100 + rank)225    x = _torch.randn(2, 4)226    target = _torch.randn(2, 4)227 228    with diloco:229        for _ in range(sync_every):230            inner_optim.zero_grad()231            loss = ((model(x) - target) ** 2).mean()232            loss.backward()233            inner_optim.step()234 235    final = model.weight.detach().clone()236    return {237        "rank": rank,238        "initial": initial.flatten().tolist(),239        "final": final.flatten().tolist(),240        "current_step": manager.current_step(),241    }242 243 244def test_mockmanager_diloco_multi_process_weights_converge(tmp_path):245    """Wave 14 (Suggestion 4): cross-replica weight convergence after one outer round.246 247    Spawns n_replicas=2 subprocesses with IDENTICAL initial weights248    (DiLoCo's standard assumption — it averages pseudo-gradients, not249    absolute weights) but RANK-SPECIFIC training data. After exactly250    one DiLoCo outer round, both replicas must end with IDENTICAL251    weights, because:252 253      pseudo_grad_i = init - inner_trained_i        # per-rank, differ254      avg_pseudo    = mean_i(pseudo_grad_i)         # same on all ranks255      final         = init - outer_lr * avg_pseudo  # same on all ranks256 257    This catches averaging-direction bugs that the world_size=1258    single-process test silently misses (a single-rank allreduce is a259    no-op and can hide bugs in the multi-rank averaging arithmetic, the260    file-staging round-id increment, or the weight redistribution after261    the outer SGD step).262    """263    import os as _os264    import tempfile as _tempfile265 266    from composer_replication.diloco.serverless import LocalProcessExecutor267 268    n_replicas = 2269    sync_every = 2270    with _tempfile.TemporaryDirectory() as td:271        rendezvous = _os.path.join(td, "diloco-multiproc-run")272        executor = LocalProcessExecutor()273        handles = executor.launch_replicas(274            n_replicas=n_replicas,275            entrypoint=f"{__name__}._diloco_replica_one_outer_round",276            entrypoint_args={277                "rendezvous_uri": rendezvous,278                "world_size": n_replicas,279                "sync_every": sync_every,280                "rank_env": "REPLICA_RANK",281            },282            timeout=180,283        )284        results = executor.collect(handles, timeout=180)285 286    # Diagnostic-friendly failure: surface per-rank error if any replica died.287    statuses = {r["rank"]: r["status"] for r in results}288    for rank in range(n_replicas):289        assert statuses[rank] == "succeeded", (290            f"rank {rank} failed: "291            f"{next(r for r in results if r['rank'] == rank).get('error')}"292        )293 294    payloads = sorted([r["result"] for r in results], key=lambda d: d["rank"])295    rank0, rank1 = payloads[0], payloads[1]296 297    # Sanity: each replica really did fire exactly one outer round.298    assert rank0["current_step"] == 1, rank0299    assert rank1["current_step"] == 1, rank1300 301    # Sanity: replicas STARTED with identical weights (DiLoCo assumption).302    assert rank0["initial"] == rank1["initial"], (303        "replicas started with different initial weights — DiLoCo only "304        "averages pseudo-gradients, not weights, so this would prevent "305        "convergence even with a perfectly correct allreduce"306    )307 308    # The actual property: after one full outer round both replicas must309    # have the SAME final weights. Tight tolerance because the only310    # arithmetic between them is SGD + a single allreduce-mean.311    final0 = torch.tensor(rank0["final"])312    final1 = torch.tensor(rank1["final"])313    if not torch.allclose(final0, final1, atol=1e-5, rtol=1e-5):314        max_abs_diff = (final0 - final1).abs().max().item()315        pytest.fail(316            "Multi-process DiLoCo did NOT converge to identical weights "317            "after one outer round.\n"318            f"  rank0 final = {final0.tolist()}\n"319            f"  rank1 final = {final1.tolist()}\n"320            f"  max|diff|   = {max_abs_diff}\n"321            "This indicates a real cross-replica-averaging bug "322            "(averaging direction, round-id desync, or weight redistribution)."323        )324 325 326def test_mockmanager_diloco_two_outer_rounds_step_counter(tmp_path):327    """Two outer rounds must bump current_step() to 2 (fragment rotation safety)."""328    torch.manual_seed(1)329    model = torch.nn.Linear(4, 4, bias=False)330    inner_optim = torch.optim.SGD(model.parameters(), lr=0.05)331 332    manager = MockManager(_make_store(tmp_path))333 334    diloco = make_diloco_outer_loop(335        manager=manager,336        model_fragments=[model],337        inner_optimizer=inner_optim,338        sync_every=2,339    )340 341    x = torch.randn(2, 4)342    target = torch.randn(2, 4)343 344    with diloco:345        for _ in range(4):  # 2 outer rounds at sync_every=2346            inner_optim.zero_grad()347            (((model(x) - target) ** 2).mean()).backward()348            inner_optim.step()349 350    assert manager.current_step() == 2, (351        f"expected current_step()==2 after two outer rounds, got {manager.current_step()}"352    )353