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