Team Ai
Modelpublic

Codeseys/composer-replication-framework

sourceHugging Facemitupdated 4mo agoView on Hugging Face
0likes
test_serverless_local.py401 linesDownload Raw Back to tests
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