Codeseys/composer-replication-framework
0
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 