Codeseys/composer-replication-framework
0
1"""Replica entrypoint — what each serverless replica runs.2 3This is the script invoked by `LocalProcessExecutor`, `ModalExecutor`,4`HFJobsExecutor`, etc. It learns its rank from the `REPLICA_RANK` env5var, sets up `ObjectStoreAllReduce` against the shared rendezvous URI,6wraps it in a `MockManager`, and hands it off to the user's training7function.8 9Usage from an executor:10 11 >>> executor.launch_replicas(12 ... n_replicas=4,13 ... entrypoint="composer_replication.diloco.serverless.replica_entrypoint",14 ... entrypoint_args={15 ... "rendezvous_uri": "/tmp/run42/",16 ... "world_size": 4,17 ... "trainer_module": "my_project.trainer",18 ... "trainer_fn": "train",19 ... "trainer_kwargs": {"model_name": "Qwen/Qwen2.5-0.5B"},20 ... },21 ... )22 23The entrypoint expects:24- `REPLICA_RANK` env var set to the rank (0..world_size-1)25- `rendezvous_uri`: fsspec URI for object-store rendezvous26- `world_size`: total replicas27- `trainer_module`, `trainer_fn`: importable path to the user's train fn28- `trainer_kwargs`: dict passed to the user's train fn, plus an injected29 `manager` kwarg containing the `MockManager`30"""31from __future__ import annotations32 33import importlib34import os35from typing import Any36 37 38def main(39 rendezvous_uri: str,40 world_size: int,41 trainer_module: str,42 trainer_fn: str = "train",43 trainer_kwargs: dict[str, Any] | None = None,44) -> Any:45 """Entrypoint executed inside each replica.46 47 Args:48 rendezvous_uri: fsspec URI (or local path) for the rendezvous49 world_size: total replicas50 trainer_module: importable Python module containing the user's51 train function52 trainer_fn: name of the function to call (default "train")53 trainer_kwargs: kwargs passed to the train function54 55 Returns:56 Whatever the train function returns.57 """58 from composer_replication.diloco.serverless.allreduce import (59 MockManager,60 ObjectStoreAllReduce,61 )62 63 rank_str = os.environ.get("REPLICA_RANK")64 if rank_str is None:65 raise RuntimeError(66 "REPLICA_RANK env var not set. The serverless executor "67 "should set this for each replica."68 )69 rank = int(rank_str)70 71 if not (0 <= rank < world_size):72 raise ValueError(f"REPLICA_RANK={rank} not in [0, {world_size})")73 74 store = ObjectStoreAllReduce(75 uri=rendezvous_uri,76 rank=rank,77 world_size=world_size,78 )79 manager = MockManager(store)80 81 mod = importlib.import_module(trainer_module)82 fn = getattr(mod, trainer_fn)83 84 kwargs = dict(trainer_kwargs or {})85 kwargs["manager"] = manager # injected86 kwargs["rank"] = rank87 kwargs["world_size"] = world_size88 return fn(**kwargs)89 90 91if __name__ == "__main__":92 import argparse93 import json94 95 # Dual input contract (both backends supported):96 # * argv — SageMakerExecutor / LocalProcessExecutor pass the run config as97 # `--rendezvous/--world-size/--trainer-module` ContainerArguments.98 # * env — EKSExecutor (and any backend that prefers a pure-env contract,99 # since k8s Indexed Jobs already inject REPLICA_RANK via the downward API)100 # pass the SAME values as RENDEZVOUS_URI / WORLD_SIZE / TRAINER_MODULE101 # env vars. The argv flags are therefore NOT `required=True`: when absent102 # we fall back to the env vars, and only error if NEITHER source supplies103 # a mandatory field. This is the R3 fix — previously the argparse block104 # hard-required argv, so an EKS pod (env-only) crashed at arg-parsing.105 parser = argparse.ArgumentParser()106 parser.add_argument("--rendezvous", default=None)107 parser.add_argument("--world-size", type=int, default=None)108 parser.add_argument("--trainer-module", default=None)109 parser.add_argument("--trainer-fn", default=None)110 parser.add_argument("--trainer-kwargs-json", default=None)111 args = parser.parse_args()112 113 def _resolve(arg_val, env_key, *, required, cast=lambda x: x):114 if arg_val is not None:115 return arg_val116 env_val = os.environ.get(env_key)117 if env_val is not None:118 return cast(env_val)119 if required:120 raise SystemExit(121 f"replica_entrypoint: missing '{env_key}' — supply it via the "122 f"argv flag or the {env_key} environment variable "123 f"(EKSExecutor uses env; SageMaker/Local use argv)."124 )125 return None126 127 rendezvous = _resolve(args.rendezvous, "RENDEZVOUS_URI", required=True)128 world_size = _resolve(args.world_size, "WORLD_SIZE", required=True, cast=int)129 trainer_module = _resolve(args.trainer_module, "TRAINER_MODULE", required=True)130 trainer_fn = _resolve(args.trainer_fn, "TRAINER_FN", required=False) or "train"131 kwargs_json = _resolve(132 args.trainer_kwargs_json, "TRAINER_KWARGS_JSON", required=False133 ) or "{}"134 135 main(136 rendezvous_uri=rendezvous,137 world_size=world_size,138 trainer_module=trainer_module,139 trainer_fn=trainer_fn,140 trainer_kwargs=json.loads(kwargs_json),141 )142 