Team Ai
Modelpublic

Codeseys/composer-replication-framework

sourceHugging Facemitupdated 4mo agoView on Hugging Face
0likes
allreduce.py327 linesDownload Raw Back to serverless
1"""ObjectStoreAllReduce — fsspec-backed pseudo-gradient exchange for DiLoCo.2 3DiLoCo's outer-loop sync writes the local pseudo-gradient (= θ_initial − θ_local)4to a shared location once per H ≈ 500-1000 inner steps, then averages across5all replicas before the outer SGD step. With cross-job NCCL unavailable on6most serverless backends, we use object storage as the rendezvous medium.7 8Communication pattern per outer round:91. Each replica writes its pseudo-gradient: PUT(rendezvous/round_N/rank_R.pt)102. Each replica reads all peer pseudo-gradients: GET × N113. Average locally → applied as `Manager.allreduce()` would have.12 13Backend support via fsspec: s3://, gs://, az://, hf://, file://.14The same code path works across all of them.15 16License compatibility: this module re-implements the contract of17`torchft.Manager.allreduce` through duck-typing — no torchft code is18copied. torchft itself is BSD-3.19"""20from __future__ import annotations21 22import io23import os24import time25from typing import Any26 27import torch28 29 30class ObjectStoreAllReduce:31    """fsspec-backed pseudo-gradient rendezvous.32 33    Each call to `allreduce(tensor, name)` blocks until all peers have34    written their version of `tensor` to the rendezvous location, then35    returns the average.36 37    Args:38        uri: fsspec URI like "s3://bucket/path/" or "file:///tmp/diloco/" or39            a plain path "/tmp/diloco/run42/" (treated as file://).40        rank: this replica's rank (0-indexed)41        world_size: total number of replicas42        round_id: optional, used to namespace successive sync rounds.43            If None, a monotonically increasing counter is used internally.44        timeout_s: per-allreduce timeout in seconds.45        poll_interval_s: how often to check for peer files.46    """47 48    def __init__(49        self,50        uri: str,51        rank: int,52        world_size: int,53        *,54        round_id: int | None = None,55        timeout_s: float = 1800.0,56        poll_interval_s: float = 1.0,57    ) -> None:58        if not (0 <= rank < world_size):59            raise ValueError(f"rank {rank} not in [0, {world_size})")60        self.uri = uri.rstrip("/") + "/"61        self.rank = rank62        self.world_size = world_size63        self.timeout_s = timeout_s64        self.poll_interval_s = poll_interval_s65        self._round_counter = 0 if round_id is None else round_id66 67        # Lazy fsspec init; deferred so that local-only smoke tests don't68        # require fsspec install in the dev environment.69        self._fs = None70        self._is_local = self.uri.startswith("file://") or self.uri.startswith("/")71        if self._is_local:72            local_path = self.uri.removeprefix("file://")73            os.makedirs(local_path, exist_ok=True)74            self._local_root = local_path75        else:76            self._init_fsspec()77 78    def _init_fsspec(self) -> None:79        try:80            import fsspec  # noqa: F40181        except ImportError as e:82            raise RuntimeError(83                "Non-local rendezvous requires fsspec; install with "84                "`pip install -e .[serverless]`. Got: " + repr(e)85            )86        import fsspec87        protocol = self.uri.split("://", 1)[0] if "://" in self.uri else "file"88        self._fs = fsspec.filesystem(protocol)89 90    @property91    def round_id(self) -> int:92        return self._round_counter93 94    def _round_dir(self, round_id: int) -> str:95        return f"round_{round_id:06d}"96 97    def _path_for(self, round_id: int, rank: int) -> str:98        return f"{self._round_dir(round_id)}/rank_{rank:04d}.pt"99 100    def _put(self, relpath: str, payload: bytes) -> None:101        if self._is_local:102            full = os.path.join(self._local_root, relpath)103            os.makedirs(os.path.dirname(full), exist_ok=True)104            tmp = full + ".tmp"105            with open(tmp, "wb") as f:106                f.write(payload)107            os.replace(tmp, full)  # atomic on POSIX108        else:109            full = self.uri + relpath110            assert self._fs is not None111            with self._fs.open(full, "wb") as f:112                f.write(payload)113 114    def _get(self, relpath: str) -> bytes:115        if self._is_local:116            full = os.path.join(self._local_root, relpath)117            with open(full, "rb") as f:118                return f.read()119        full = self.uri + relpath120        assert self._fs is not None121        with self._fs.open(full, "rb") as f:122            return f.read()123 124    def _exists(self, relpath: str) -> bool:125        if self._is_local:126            return os.path.exists(os.path.join(self._local_root, relpath))127        full = self.uri + relpath128        assert self._fs is not None129        return self._fs.exists(full)130 131    def allreduce(self, tensor: torch.Tensor, *, name: str | None = None) -> torch.Tensor:132        """Average `tensor` across all replicas via the object store.133 134        Args:135            tensor: the tensor to average. Modified in-place AND returned.136            name: ignored — provided for API compat with torchft.Manager.137 138        Returns:139            The averaged tensor (modifies in-place; returns the same object).140        """141        round_id = self._round_counter142        self._round_counter += 1143 144        # Serialize my tensor145        buf = io.BytesIO()146        torch.save({"rank": self.rank, "tensor": tensor.detach().cpu()}, buf)147        my_path = self._path_for(round_id, self.rank)148        self._put(my_path, buf.getvalue())149 150        # Wait for all peers151        deadline = time.time() + self.timeout_s152        peer_tensors: list[torch.Tensor] = []153        for peer_rank in range(self.world_size):154            peer_path = self._path_for(round_id, peer_rank)155            while not self._exists(peer_path):156                if time.time() > deadline:157                    raise TimeoutError(158                        f"ObjectStoreAllReduce: timed out waiting for "159                        f"rank {peer_rank} at {self.uri}{peer_path} "160                        f"(world_size={self.world_size}, round={round_id})"161                    )162                time.sleep(self.poll_interval_s)163            payload = self._get(peer_path)164            peer_data = torch.load(io.BytesIO(payload), weights_only=False)165            peer_tensors.append(peer_data["tensor"].to(tensor.device, dtype=tensor.dtype))166 167        # Compute average168        stacked = torch.stack(peer_tensors, dim=0)169        avg = stacked.mean(dim=0)170        tensor.copy_(avg)171        return tensor172 173 174# ---------------------------------------------------------------------175# MockManager — torchft.Manager-shaped object that uses ObjectStoreAllReduce176# ---------------------------------------------------------------------177 178 179class _ImmediateWork:180    """Work-shaped wrapper for an already-completed allreduce.181 182    `torchft.Manager.allreduce` returns a `torch.distributed.Work` (or183    `torchft.work._DummyWork`) which DiLoCo calls `.wait()` on inside184    `_StreamingDiLoCoFragment.perform_sync`. Our `ObjectStoreAllReduce`185    is synchronous — by the time it returns, the average is already in186    the tensor — so `.wait()` is a no-op.187 188    We deliberately don't subclass `torch.distributed._Work` to keep this189    module importable in environments without a full torch distributed190    build; DiLoCo only does `work.wait()`, nothing more.191    """192 193    __slots__ = ("_tensor",)194 195    def __init__(self, tensor: torch.Tensor) -> None:196        self._tensor = tensor197 198    def wait(self, *_args: Any, **_kwargs: Any) -> bool:199        return True200 201    def get_future(self) -> Any:202        # Torch >=2.x sometimes calls Work.get_future(); provide a satisfied203        # future so callers don't crash. We only need to be defensive here;204        # DiLoCo itself doesn't call this.205        try:206            import torch.futures as _f207 208            fut = _f.Future()209            fut.set_result(self._tensor)210            return fut211        except Exception:  # pragma: no cover — defensive only212            return None213 214 215class MockManager:216    """Drop-in replacement for `torchft.Manager` that delegates allreduce217    to `ObjectStoreAllReduce`.218 219    The torchft `DiLoCo` class accepts a `Manager` and calls its `.allreduce`220    method on the pseudo-gradient. By passing this mock instead, we route221    that call through the object store, leaving the rest of the DiLoCo222    machinery (sign convention, post-hook sequencing, etc.) untouched.223 224    Reference: `make_diloco_outer_loop` in225    `composer_replication/diloco/__init__.py` accepts an optional226    `manager=` kwarg; pass a `MockManager` to enable serverless DiLoCo.227 228    torchft.Manager surface audited from229    ``torchft/local_sgd.py`` (DiLoCo + _StreamingDiLoCoFragment paths) and230    ``torchft/manager.py``. Methods/attributes DiLoCo touches:231 232    * ``allreduce(tensor, should_quantize=...) -> Work`` — must return an233      object with ``.wait()`` (DiLoCo calls ``work.wait()`` in234      ``perform_sync``).235    * ``should_commit() -> bool`` — gates the outer-optimizer step.236    * ``start_quorum()`` — called once per outer round, before237      ``prepare_sync``.238    * ``current_step() -> int`` — used to pick the streaming-DiLoCo239      fragment for this round (``step % len(fragments)``).240    * ``disallow_state_dict_read()`` / ``allow_state_dict_read()`` —241      called every inner step from the optimizer pre/post hooks.242    * ``register_state_dict_fn(key, load_fn, save_fn)`` — called once243      per fragment from ``DiLoCo.__init__``.244    * ``_use_async_quorum`` (attribute) — DiLoCo's constructor refuses245      to start if this is truthy. Must exist and be False.246    * ``num_participants`` / ``rank`` — read by upstream callers.247    """248 249    def __init__(self, store: ObjectStoreAllReduce) -> None:250        self._store = store251        # torchft Manager attributes that DiLoCo consults at construction time252        # or in user code paths.253        self.num_participants = store.world_size254        self.rank = store.rank255        # DiLoCo.__init__ raises if this is truthy (line 622 of256        # torchft/local_sgd.py). Object-store sync is synchronous → False.257        self._use_async_quorum: bool = False258        # Mirror the upstream Manager's monotonic step counter. DiLoCo reads259        # this via current_step() to decide which fragment to sync each round.260        # Bumped from start_quorum() so it advances exactly once per outer round.261        self._step: int = 0262        # State-dict-fn registry: torchft uses this for fault-tolerant263        # checkpoint restore. We're single-shot serverless — record but never264        # invoke. Tests can introspect this dict to confirm registration.265        self._state_dict_fns: dict[str, tuple[Any, Any]] = {}266 267    # ---- Core collective ------------------------------------------------268    def allreduce(self, tensor: torch.Tensor, **_kwargs: Any) -> _ImmediateWork:269        # DiLoCo expects a Work-like return value (it stores it in a list270        # then calls .wait() later). Object-store all-reduce is synchronous,271        # so the tensor is already averaged when we hand back the wrapper.272        averaged = self._store.allreduce(tensor)273        return _ImmediateWork(averaged)274 275    # ---- Quorum / commit lifecycle -------------------------------------276    def should_commit(self) -> bool:277        # No fault-tolerance failover in serverless mode: every quorum278        # always commits. Replica failure is handled by the orchestration279        # layer (HF Jobs / Modal restart), not by DiLoCo skipping a round.280        return True281 282    def start_quorum(self) -> None:283        # The upstream Manager bumps its step counter inside the quorum284        # bookkeeping. Do the same so current_step() advances per round285        # and DiLoCo's fragment-rotation math matches across replicas.286        self._step += 1287 288    def wait_quorum(self) -> int:289        return self.num_participants290 291    # ---- Step counter ---------------------------------------------------292    def current_step(self) -> int:293        return self._step294 295    # ---- State-dict read gating ----------------------------------------296    # torchft uses these to make checkpoint restore thread-safe. In a297    # single-process serverless mock there's no concurrent reader, so they298    # are no-ops — but they MUST exist (DiLoCo's pre/post optimizer hooks299    # call them on every inner step).300    def allow_state_dict_read(self) -> None:301        pass302 303    def disallow_state_dict_read(self) -> None:304        pass305 306    # ---- Checkpoint hook registry --------------------------------------307    def register_state_dict_fn(308        self,309        key: str,310        load_fn: Any,311        save_fn: Any,312    ) -> None:313        # DiLoCo registers one (load, save) pair per fragment so torchft can314        # checkpoint the outer-optimizer state and original-parameter backup.315        # In serverless mode we capture the registration so tests can verify316        # it happened, but never invoke it — there's no HA failover.317        self._state_dict_fns[key] = (load_fn, save_fn)318 319    # ---- Convenience ----------------------------------------------------320    def is_leader(self) -> bool:321        # Not strictly required by DiLoCo but referenced in some torchft322        # integrations / our own code that may swap MockManager in.323        return self.rank == 0324 325 326__all__ = ["MockManager", "ObjectStoreAllReduce", "_ImmediateWork"]327