Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion miles/ray/actor_group.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down
25 changes: 24 additions & 1 deletion miles/utils/arguments.py
Original file line number Diff line number Diff line change
@@ -1,3 +1,4 @@
import argparse
import json
import os
from typing import Any, Dict
Expand Down Expand Up @@ -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(
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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):
Expand Down
Loading