Team Ai
Modelpublic

Codeseys/composer-replication-framework

sourceHugging Facemitupdated 4mo agoView on Hugging Face
0likes
replica_entrypoint.py142 linesDownload Raw Back to serverless
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