diff --git a/miles/ray/actor_group.py b/miles/ray/actor_group.py index 122e768649b..730eb494c20 100644 --- a/miles/ray/actor_group.py +++ b/miles/ray/actor_group.py @@ -78,7 +78,7 @@ def _allocate_gpus_for_actor(self, pg, num_gpus_per_actor, wandb_run_id: Optiona if self.args.use_routing_replay: env_vars["ENABLE_ROUTING_REPLAY"] = "1" - backend = os.environ.get("MILES_BACKEND", "megatron").lower() + backend = self.args.train_backend if backend == "megatron": from miles.backends.megatron_utils import MegatronTrainRayActor diff --git a/miles/utils/arguments.py b/miles/utils/arguments.py index ad7c7d9fe5d..c31f3e9c3b8 100644 --- a/miles/utils/arguments.py +++ b/miles/utils/arguments.py @@ -1,3 +1,4 @@ +import argparse import json import os from typing import Any, Dict @@ -88,6 +89,17 @@ def add_cluster_arguments(parser): return parser + def add_train_arguments(parser): + parser.add_argument( + "--train-backend", + type=str, + choices=["megatron", "fsdp"], + default="megatron", + help="The backend for training.", + ) + + return parser + # rollout def add_rollout_arguments(parser): parser.add_argument( @@ -928,6 +940,7 @@ def add_ci_arguments(parser): parser = add_custom_arguments(parser) parser = add_cluster_arguments(parser) + parser = add_train_arguments(parser) parser = add_rollout_arguments(parser) parser = add_data_arguments(parser) parser = add_eval_arguments(parser) @@ -966,7 +979,7 @@ def warning_for_unfinished_backend(backend: str): def parse_args(add_custom_arguments=None): add_miles_arguments = get_miles_extra_args_provider(add_custom_arguments) - backend = os.environ.get("MILES_BACKEND", "megatron").lower() + backend = parse_args_train_backend() if backend == "megatron": from miles.backends.megatron_utils import parse_args as megatron_parse_args from miles.backends.megatron_utils import set_default_megatron_args @@ -1013,6 +1026,16 @@ def parse_args(add_custom_arguments=None): return args +def parse_args_train_backend(): + if os.environ.get("MILES_BACKEND") is not None: + raise Exception("`MILES_BACKEND` is deprecated, please use --train-backend directly.") + + parser = argparse.ArgumentParser() + get_miles_extra_args_provider()(parser) + args_partial, _ = parser.parse_known_args() + return args_partial.train_backend + + def miles_validate_args(args): if args.kl_coef != 0 or args.use_kl_loss: if not os.path.exists(args.ref_load):