Codeseys/composer-replication-framework
0
1"""Verifies the serverless DiLoCo allreduce wraps correctly across local2multiprocessing replicas using `file://` rendezvous.3 4This is the core multi-process test for the serverless layer. It exercises5the real allreduce barrier (with concurrent processes), not just the6single-process API.7"""8from __future__ import annotations9 10import os11import sys12import tempfile13import time14 15import pytest16import torch17 18from composer_replication.diloco.serverless import (19 LocalProcessExecutor,20 ObjectStoreAllReduce,21 ReplicaHandle,22)23 24 25# ---------------------------------------------------------------------26# Single-process tests of ObjectStoreAllReduce primitives27# (don't need executor, just the file:// path + local manual orchestration)28# ---------------------------------------------------------------------29 30 31def test_object_store_allreduce_init_validates_rank():32 with tempfile.TemporaryDirectory() as td:33 with pytest.raises(ValueError, match="not in"):34 ObjectStoreAllReduce(td, rank=5, world_size=2)35 36 37def test_object_store_allreduce_local_paths_create_dir():38 """Local backend should mkdir on init."""39 with tempfile.TemporaryDirectory() as td:40 new_path = os.path.join(td, "subdir", "subsubdir")41 store = ObjectStoreAllReduce(new_path, rank=0, world_size=1)42 assert os.path.isdir(new_path)43 assert store.world_size == 144 45 46def test_object_store_allreduce_world_size_1_passthrough():47 """With world_size=1 it just averages the tensor with itself."""48 with tempfile.TemporaryDirectory() as td:49 store = ObjectStoreAllReduce(td, rank=0, world_size=1, timeout_s=10.0)50 t = torch.tensor([1.0, 2.0, 3.0])51 result = store.allreduce(t.clone())52 torch.testing.assert_close(result, t, atol=1e-6, rtol=1e-6)53 54 55def test_object_store_allreduce_round_id_increments():56 with tempfile.TemporaryDirectory() as td:57 store = ObjectStoreAllReduce(td, rank=0, world_size=1, timeout_s=10.0)58 t = torch.zeros(3)59 assert store.round_id == 060 store.allreduce(t.clone())61 assert store.round_id == 162 store.allreduce(t.clone())63 assert store.round_id == 264 65 66# ---------------------------------------------------------------------67# Multi-process tests (the real verification — local executor + spawn)68# ---------------------------------------------------------------------69 70 71def _replica_compute_and_sync(72 rendezvous_uri: str,73 world_size: int,74 rank_value: float,75) -> dict:76 """Top-level function — must be importable for multiprocessing 'spawn'.77 78 Each replica creates a tensor whose value is `rank_value * (rank+1)` and79 runs allreduce. The expected result is the mean of all replicas' tensors.80 """81 rank = int(os.environ["REPLICA_RANK"])82 store = ObjectStoreAllReduce(83 rendezvous_uri, rank=rank, world_size=world_size, timeout_s=120.0,84 )85 # tensor that depends on rank86 t = torch.full((4,), float(rank_value * (rank + 1)))87 pre = t.clone()88 averaged = store.allreduce(t)89 return {90 "rank": rank,91 "pre": pre.tolist(),92 "post": averaged.tolist(),93 "world_size": world_size,94 }95 96 97@pytest.mark.parametrize("n_replicas", [2, 3])98def test_local_executor_runs_allreduce_across_replicas(n_replicas):99 """End-to-end: 2-3 replica processes each call allreduce; result is the mean."""100 with tempfile.TemporaryDirectory() as td:101 rendezvous = os.path.join(td, "run")102 executor = LocalProcessExecutor()103 handles = executor.launch_replicas(104 n_replicas=n_replicas,105 entrypoint=f"{__name__}._replica_compute_and_sync",106 entrypoint_args={107 "rendezvous_uri": rendezvous,108 "world_size": n_replicas,109 "rank_value": 10.0,110 "rank_env": "REPLICA_RANK",111 },112 timeout=180,113 )114 assert len(handles) == n_replicas115 for i, h in enumerate(handles):116 assert h.rank == i117 assert h.backend_name == "local_process"118 119 results = executor.collect(handles, timeout=180)120 assert len(results) == n_replicas121 122 # Verify all succeeded123 for r in results:124 assert r["status"] == "succeeded", \125 f"rank {r['rank']} failed: {r.get('error')}"126 127 # Each replica created tensor full(rank_value * (rank+1)).128 # Expected mean = rank_value * (1+2+...+N) / N129 N = n_replicas130 expected_mean = 10.0 * (N * (N + 1) / 2) / N131 132 for r in results:133 post = r["result"]["post"]134 for v in post:135 assert abs(v - expected_mean) < 1e-4, \136 f"rank {r['rank']}: expected mean {expected_mean}, got {v}"137 138 139def _replica_two_round_sync(140 rendezvous_uri: str,141 world_size: int,142) -> dict:143 """Each replica does TWO consecutive allreduce calls; checks round_id increments."""144 rank = int(os.environ["REPLICA_RANK"])145 store = ObjectStoreAllReduce(146 rendezvous_uri, rank=rank, world_size=world_size, timeout_s=120.0,147 )148 t1 = torch.full((2,), float(rank))149 avg1 = store.allreduce(t1).clone()150 t2 = torch.full((2,), float(rank * 100))151 avg2 = store.allreduce(t2).clone()152 return {153 "rank": rank,154 "round_after_2_calls": store.round_id,155 "avg1": avg1.tolist(),156 "avg2": avg2.tolist(),157 }158 159 160def test_local_executor_handles_multiple_rounds():161 """Two consecutive rounds each give the right mean; round counter advances."""162 n_replicas = 3163 with tempfile.TemporaryDirectory() as td:164 rendezvous = os.path.join(td, "run-2round")165 executor = LocalProcessExecutor()166 handles = executor.launch_replicas(167 n_replicas=n_replicas,168 entrypoint=f"{__name__}._replica_two_round_sync",169 entrypoint_args={170 "rendezvous_uri": rendezvous,171 "world_size": n_replicas,172 },173 timeout=180,174 )175 results = executor.collect(handles, timeout=180)176 for r in results:177 assert r["status"] == "succeeded", r.get("error")178 assert r["result"]["round_after_2_calls"] == 2179 # mean of 0,1,2 = 1.0180 assert all(abs(v - 1.0) < 1e-4 for v in r["result"]["avg1"])181 # mean of 0,100,200 = 100.0182 assert all(abs(v - 100.0) < 1e-4 for v in r["result"]["avg2"])183 184 185# ---------------------------------------------------------------------186# Live-S3 smoke (F4 step 1): the file:// → s3:// transport gap.187#188# ObjectStoreAllReduce's S3 branches (_init_fsspec/_put/_exists/_get over189# s3fs) only have mock coverage; this exercises them against REAL S3 with190# concurrent OS processes, relying on S3's strong read-after-write191# consistency (the poll loop's _exists()→_get() assumption). Gated on192# AWS_SMOKE=1 so it never runs in ordinary CI / on machines without creds.193#194# Run it with:195# AWS_SMOKE=1 AWS_REGION=us-west-2 \196# DILOCO_S3_RENDEZVOUS=s3://<sagemaker-bucket>/diloco-rdv \197# pytest composer_replication/diloco/serverless/tests/test_serverless_local.py \198# -k s3_rendezvous -s199#200# Use a sagemaker-named bucket: stock AmazonSageMakerFullAccess only grants201# S3 on buckets whose name contains "sagemaker"/"aws-glue" — a custom-named202# bucket would 403 the first PUT and hang every peer until timeout_s (F4 §3).203# Verified PASS 2026-06-09 against204# s3://amazon-sagemaker-386931836011-us-west-2-7597bf4d9a3d/diloco-rdv/.205# ---------------------------------------------------------------------206 207 208def _s3_smoke_enabled() -> bool:209 return os.environ.get("AWS_SMOKE") == "1"210 211 212@pytest.mark.skipif(213 not _s3_smoke_enabled(),214 reason="live-S3 smoke; set AWS_SMOKE=1 (+ AWS creds, DILOCO_S3_RENDEZVOUS) to run",215)216@pytest.mark.parametrize("n_replicas", [2])217def test_s3_rendezvous_allreduce_across_replicas(n_replicas):218 """Real-S3 analogue of test_local_executor_runs_allreduce_across_replicas.219 220 Same property (N processes call allreduce, every rank ends with the221 cross-rank mean) but over an ``s3://`` rendezvous instead of a tmp dir,222 so it actually drives s3fs PUT/poll/GET and depends on S3 strong223 read-after-write consistency. This is the cheapest (≈$0, no GPU) closure224 of F4's documented "ObjectStoreAllReduce over s3:// never exercised225 against real S3" gap.226 """227 import uuid228 229 pytest.importorskip("s3fs", reason="s3fs required for the live-S3 smoke")230 import s3fs231 232 base = os.environ.get(233 "DILOCO_S3_RENDEZVOUS",234 "s3://amazon-sagemaker-386931836011-us-west-2-7597bf4d9a3d/diloco-rdv",235 ).rstrip("/")236 rendezvous = f"{base}/smoke-{uuid.uuid4().hex[:8]}/"237 238 executor = LocalProcessExecutor()239 handles = executor.launch_replicas(240 n_replicas=n_replicas,241 entrypoint=f"{__name__}._replica_compute_and_sync",242 entrypoint_args={243 "rendezvous_uri": rendezvous,244 "world_size": n_replicas,245 "rank_value": 10.0,246 "rank_env": "REPLICA_RANK",247 },248 timeout=300,249 )250 try:251 results = executor.collect(handles, timeout=300)252 253 for r in results:254 assert r["status"] == "succeeded", (255 f"rank {r['rank']} failed (S3 rendezvous {rendezvous}): "256 f"{r.get('error')}"257 )258 259 # Every rank must agree on the mean — only possible if each read the260 # SAME peer objects through S3 (proves the cross-process exchange).261 N = n_replicas262 expected_mean = 10.0 * (N * (N + 1) / 2) / N263 for r in results:264 for v in r["result"]["post"]:265 assert abs(v - expected_mean) < 1e-4, (266 f"rank {r['rank']}: expected S3-averaged mean {expected_mean}, "267 f"got {v}"268 )269 270 # Both ranks' pseudo-gradient objects must be present in S3.271 fs = s3fs.S3FileSystem()272 listing = fs.ls(rendezvous.replace("s3://", "") + "round_000000/")273 got = {os.path.basename(p) for p in listing}274 expected = {f"rank_{r:04d}.pt" for r in range(n_replicas)}275 assert expected <= got, f"missing rank objects in S3: {expected - got}"276 finally:277 # Best-effort cleanup so repeated smokes don't accrete prefixes.278 try:279 s3fs.S3FileSystem().rm(rendezvous.replace("s3://", ""), recursive=True)280 except Exception:281 pass282 283 284def _replica_that_raises(rendezvous_uri: str, world_size: int) -> dict:285 """Simulates a replica that crashes mid-run."""286 rank = int(os.environ["REPLICA_RANK"])287 if rank == 1:288 raise RuntimeError(f"Simulated crash on rank {rank}")289 return {"rank": rank, "ok": True}290 291 292def test_local_executor_reports_failed_replicas():293 """When a replica crashes, collect() reports it as failed without hanging294 (other ranks complete; the failed one should be reflected in the result).295 296 Note (Wave 18): timeouts bumped from 30s → 90s because this test was297 flaky in full-suite runs (passes individually but occasionally times298 out when other parallel multiprocessing tests contend for CPU).299 The 30s budget was tight for cold-start subprocess + import +300 rendezvous-file IO under contention; 90s gives comfortable headroom301 without changing the test's semantic intent (subprocess crashes302 surface as `failed` status, not hangs).303 """304 n_replicas = 2305 with tempfile.TemporaryDirectory() as td:306 rendezvous = os.path.join(td, "run-failure")307 executor = LocalProcessExecutor()308 handles = executor.launch_replicas(309 n_replicas=n_replicas,310 entrypoint=f"{__name__}._replica_that_raises",311 entrypoint_args={312 "rendezvous_uri": rendezvous,313 "world_size": n_replicas,314 },315 timeout=90,316 )317 results = executor.collect(handles, timeout=90)318 statuses = {r["rank"]: r["status"] for r in results}319 assert statuses[0] == "succeeded"320 assert statuses[1] == "failed"321 # Failure log should mention the simulated crash322 failure_log = next(r for r in results if r["rank"] == 1).get("error") or ""323 assert "Simulated crash" in failure_log324 325 326# ---------------------------------------------------------------------327# Sanity: MockManager is shape-compatible with torchft Manager surface328# ---------------------------------------------------------------------329 330 331def test_mock_manager_shape_compat():332 from composer_replication.diloco.serverless import MockManager333 with tempfile.TemporaryDirectory() as td:334 store = ObjectStoreAllReduce(td, rank=0, world_size=1, timeout_s=10.0)335 mgr = MockManager(store)336 # torchft.Manager surface (audited from torchft/local_sgd.py DiLoCo path)337 assert hasattr(mgr, "allreduce")338 assert hasattr(mgr, "should_commit")339 assert hasattr(mgr, "start_quorum")340 assert hasattr(mgr, "wait_quorum")341 assert hasattr(mgr, "current_step")342 assert hasattr(mgr, "disallow_state_dict_read")343 assert hasattr(mgr, "allow_state_dict_read")344 assert hasattr(mgr, "register_state_dict_fn")345 assert hasattr(mgr, "_use_async_quorum")346 assert mgr._use_async_quorum is False347 assert mgr.num_participants == 1348 assert mgr.rank == 0349 assert mgr.should_commit() is True350 # Single-replica allreduce: averaging is a passthrough, but the return351 # must be a Work-shaped object (DiLoCo calls .wait() on it). The352 # tensor itself is mutated in place by ObjectStoreAllReduce.353 t = torch.tensor([1.0, 2.0])354 buf = t.clone()355 work = mgr.allreduce(buf)356 assert hasattr(work, "wait") and callable(work.wait)357 assert work.wait() is True358 torch.testing.assert_close(buf, t, atol=1e-6, rtol=1e-6)359 360 361# ---------------------------------------------------------------------362# Public re-export surface (Wave 17a)363# ---------------------------------------------------------------------364 365 366def test_public_reexports_include_all_executors():367 """`from composer_replication.diloco.serverless import …` must368 surface every executor adapter the module's docstring claims, not369 just the LocalProcessExecutor.370 371 Wave 16's user-journey reviewer caught that ModalExecutor /372 HFJobsExecutor were defined in `modal.py` / `hf_jobs.py` but not373 re-exported from the package's `__init__.py`. Users who copied the374 docstring's `from composer_replication.diloco.serverless import375 ModalExecutor` line got an ImportError. Wave 17a added the missing376 re-exports; this test pins them.377 """378 import composer_replication.diloco.serverless as ss379 380 expected = {381 "LocalProcessExecutor",382 "ModalExecutor",383 "HFJobsExecutor",384 "MockManager",385 "ObjectStoreAllReduce",386 "ReplicaHandle",387 "ServerlessExecutor",388 }389 actual = set(ss.__all__)390 assert expected.issubset(actual), (391 f"Missing re-exports: {expected - actual}. "392 f"__all__ should include every executor adapter the package "393 f"docstring documents."394 )395 396 # Also verify each name is actually importable, not just listed.397 for name in expected:398 assert hasattr(ss, name), (399 f"{name} listed in __all__ but not present on package."400 )401 