Give execute_train the launch it is running (#3028)

This commit is contained in:
fzyzcjy
2026-09-26 20:21:41 +08:00
committed by GitHub
parent 95cf1849ef
commit 811e6c1c20
15 changed files with 300 additions and 70 deletions
+1
View File
@@ -193,6 +193,7 @@ def execute(args: ScriptArgs):
U.execute_train(
train_args=train_args,
config=args,
num_gpus_per_node=args.num_gpus_per_node,
megatron_model_type=args.megatron_model_type,
megatron_path=args.megatron_path,
@@ -231,6 +231,7 @@ def execute(args: ScriptArgs) -> None:
U.get_default_wandb_args(__file__, run_id=args.run_name),
]
),
config=args,
train_script="train_async.py",
num_gpus_per_node=args.num_gpus_per_node,
megatron_model_type="nemotron-3-nano-30b-a3b",
@@ -1,4 +1,5 @@
from miles.utils.external_utils.command_utils.base_backend import (
CommandUtilConfig,
ExecuteTrainConfig,
default_config,
resolve_extra_env_vars,
@@ -20,6 +21,7 @@ from miles.utils.external_utils.command_utils.common import (
from miles.utils.typer_utils import dataclass_cli
__all__ = [
"CommandUtilConfig",
"ExecuteTrainConfig",
"GENERATION_HARDWARE",
"NUM_GPUS_OF_HARDWARE",
@@ -30,28 +30,12 @@ from miles.utils.workers.types import ClusterBackend, DeployComponent, HotRestar
logger = logging.getLogger(__name__)
# This class can be extended by concrete scripts
@dataclass
class ExecuteTrainConfig:
cuda_core_dump: bool = False
num_nodes: int = field(default_factory=lambda: int(os.environ.get("SLURM_JOB_NUM_NODES", "1")))
extra_env_vars: str = ""
output_dir: str = "/root/shared_data"
class CommandUtilConfig:
cluster_backend: ClusterBackend = ClusterBackend.RAY
deploy_component: DeployComponent = DeployComponent.ALL
deploy_instance_id: str | None = None
hot_restart: str = ""
run_id: str = field(default_factory=create_run_id)
run_uuid: str | None = None
namespace: str = ""
helm_values: tuple[str, ...] = ()
skip_upgrade_check: bool = False
ci_run: bool = False
external_mooncake: bool = False
@property
def parsed_hot_restart(self) -> list[HotRestartComponent]:
return parse_hot_restart(self.hot_restart)
def create_backend(self) -> BaseCommandBackend:
match self.cluster_backend:
@@ -65,6 +49,26 @@ class ExecuteTrainConfig:
return RayCommandBackend(self)
# This class can be extended by concrete scripts
@dataclass
class ExecuteTrainConfig(CommandUtilConfig):
cuda_core_dump: bool = False
num_nodes: int = field(default_factory=lambda: int(os.environ.get("SLURM_JOB_NUM_NODES", "1")))
extra_env_vars: str = ""
output_dir: str = "/root/shared_data"
deploy_component: DeployComponent = DeployComponent.ALL
deploy_instance_id: str | None = None
hot_restart: str = ""
run_id: str = field(default_factory=create_run_id)
run_uuid: str | None = None
skip_upgrade_check: bool = False
external_mooncake: bool = False
@property
def parsed_hot_restart(self) -> list[HotRestartComponent]:
return parse_hot_restart(self.hot_restart)
def default_config(config_class: type = ExecuteTrainConfig) -> ExecuteTrainConfig:
return dataclass_from_env(config_class)
@@ -92,9 +96,10 @@ _PREPARE_CMD_ROLES = frozenset({TRAINER_ROLE})
class BaseCommandBackend(ABC):
def __init__(self, config: ExecuteTrainConfig) -> None:
def __init__(self, config: CommandUtilConfig) -> None:
from miles.utils.logging_utils import configure_logger_raw
assert isinstance(config, CommandUtilConfig), "config must be a CommandUtilConfig"
configure_logger_raw("launcher")
self.config = config
@@ -110,17 +115,30 @@ class BaseCommandBackend(ABC):
job_lifetime: Literal["independent", "launcher"] = "independent",
prepare_cmd: dict[str, str] | None = None,
extra_manifests: list[str] | None = None,
config: ExecuteTrainConfig | None = None,
) -> None:
if config is None:
assert isinstance(
self.config, ExecuteTrainConfig
), "execute_train requires an ExecuteTrainConfig, either as config or as the backend's config"
config = self.config
assert config.cluster_backend is self.config.cluster_backend, (
f"This backend was built to talk to {self.config.cluster_backend.value}, but the launch it is handed "
f"describes a run on {config.cluster_backend.value}, so everything this launch installs would be named "
f"for one cluster and installed onto the other; build the backend from the config of the launch"
)
assert job_lifetime in ("independent", "launcher")
extra_env_vars = extra_env_vars if extra_env_vars is not None else {}
# nothing reads this variable any more, so an old export would be ignored without a word
if "MILES_ROUTER_EXTERNAL_HOST" in {**os.environ, **resolve_extra_env_vars(extra_env_vars, self.config)}:
if "MILES_ROUTER_EXTERNAL_HOST" in {**os.environ, **resolve_extra_env_vars(extra_env_vars, config)}:
raise ValueError(
"MILES_ROUTER_EXTERNAL_HOST is no longer read. Pass --session-server-external-host for one host that "
"reaches every session server, or set MILES_NODE_EXTERNAL_IP on each node to its own address."
)
assert not (
self.config.parsed_hot_restart and self.config.cluster_backend is not ClusterBackend.KUBERNETES
config.parsed_hot_restart and config.cluster_backend is not ClusterBackend.KUBERNETES
), "--hot-restart is only supported on the kubernetes backend"
prepare_cmd = prepare_cmd if prepare_cmd is not None else {}
@@ -134,19 +152,17 @@ class BaseCommandBackend(ABC):
train_argv = shlex.split(train_args)
train_backend_fsdp = ArgvManipulator.get_effective(train_argv, "--train-backend") == "fsdp"
assert train_backend_fsdp == (megatron_model_type is None)
_assert_train_args_name_no_other_backend(train_argv, cluster_backend=self.config.cluster_backend.value)
_assert_train_args_name_no_other_deploy_component(
train_argv, deploy_component=self.config.deploy_component.value
)
train_args = f"{train_args} {_DEPLOY_COMPONENT_FLAG} {self.config.deploy_component.value}"
if self.config.deploy_instance_id is not None:
train_args = f"{train_args} --deploy-instance-id {self.config.deploy_instance_id}"
_assert_train_args_name_no_other_backend(train_argv, cluster_backend=config.cluster_backend.value)
_assert_train_args_name_no_other_deploy_component(train_argv, deploy_component=config.deploy_component.value)
train_args = f"{train_args} {_DEPLOY_COMPONENT_FLAG} {config.deploy_component.value}"
if config.deploy_instance_id is not None:
train_args = f"{train_args} --deploy-instance-id {config.deploy_instance_id}"
if config.run_uuid is not None:
_assert_train_args_name_no_other_run_uuid(train_argv, run_uuid=config.run_uuid)
train_args = f"{train_args} {_RUN_UUID_FLAG} {config.run_uuid}"
self._execute_train_inner(
ExecuteTrainRequest(
request=ExecuteTrainRequest(
train_args=train_args,
num_gpus_per_node=num_gpus_per_node,
megatron_model_type=megatron_model_type,
@@ -158,7 +174,8 @@ class BaseCommandBackend(ABC):
job_lifetime=job_lifetime,
prepare_cmd=prepare_cmd,
extra_manifests=extra_manifests if extra_manifests is not None else [],
)
),
config=config,
)
def convert_checkpoint(
@@ -251,11 +268,11 @@ class BaseCommandBackend(ABC):
f"--output-bf16-hf-path {path_dst} "
)
def api_server_host(self) -> str:
def api_server_host(self, config: ExecuteTrainConfig) -> str:
return "localhost"
@abstractmethod
def _execute_train_inner(self, request: ExecuteTrainRequest) -> None: ...
def _execute_train_inner(self, *, request: ExecuteTrainRequest, config: ExecuteTrainConfig) -> None: ...
def exec_command_cpu(self, cmd: str, capture_output: bool = False) -> str | None:
return run_shell_command(cmd, capture_output=capture_output)
@@ -1,7 +1,11 @@
from __future__ import annotations
from miles.utils.external_utils.command_utils.base_backend import BaseCommandBackend, ExecuteTrainRequest
from miles.utils.external_utils.command_utils.base_backend import (
BaseCommandBackend,
ExecuteTrainConfig,
ExecuteTrainRequest,
)
from miles.utils.external_utils.command_utils.common import chart_dir, repo_base_dir
from miles.utils.external_utils.command_utils.helm_backend import command_job
from miles.utils.external_utils.command_utils.helm_backend.launcher import entrypoint
@@ -9,8 +13,8 @@ from miles.utils.external_utils.command_utils.helm_backend.naming import Release
class KubernetesCommandBackend(BaseCommandBackend):
def _execute_train_inner(self, request: ExecuteTrainRequest) -> None:
entrypoint.execute_train(request=request, config=self.config)
def _execute_train_inner(self, *, request: ExecuteTrainRequest, config: ExecuteTrainConfig) -> None:
entrypoint.execute_train(request=request, config=config)
def exec_command_gpu(
self, cmd: str, capture_output: bool = False, num_gpus_per_node: int | None = None
@@ -26,7 +30,7 @@ class KubernetesCommandBackend(BaseCommandBackend):
num_nodes: int | None = None,
num_gpus_per_node: int | None = None,
) -> list[str | None]:
assert self.config.namespace, "Set ExecuteTrainConfig.namespace to run a command somewhere"
assert self.config.namespace, "Set CommandUtilConfig.namespace to run a command somewhere"
return command_job.run_on_nodes(
command_job.CommandJobContext(
namespace=self.config.namespace,
@@ -40,21 +44,21 @@ class KubernetesCommandBackend(BaseCommandBackend):
step="command",
)
def api_server_host(self) -> str:
assert self.config.run_id and self.config.namespace, (
def api_server_host(self, config: ExecuteTrainConfig) -> str:
assert config.run_id and config.namespace, (
"The api server of a kubernetes run answers on the orchestrator's pod, which is named after the "
"release; set ExecuteTrainConfig.run_id and .namespace before asking where that pod is"
"release; set the launch config's run_id and namespace before asking where that pod is"
)
assert not self.config.deploy_component.is_split(), (
assert not config.deploy_component.is_split(), (
f"The api server, and the mini ft controller polling it, answer for the cells of their own deployment, "
f"so a split run is refused one (--api-server-port 0) and nothing listens on the "
f"{self.config.deploy_component.value} deployment for this host to name"
f"{config.deploy_component.value} deployment for this host to name"
)
return RunNames.orchestrator_host(
release=ReleaseName(
run_id=self.config.run_id,
deploy_component=self.config.deploy_component,
deploy_instance_id=self.config.deploy_instance_id,
run_id=config.run_id,
deploy_component=config.deploy_component,
deploy_instance_id=config.deploy_instance_id,
).serialize(),
namespace=self.config.namespace,
namespace=config.namespace,
)
@@ -1,7 +1,11 @@
import os
import shlex
from miles.utils.external_utils.command_utils.base_backend import BaseCommandBackend, ExecuteTrainRequest
from miles.utils.external_utils.command_utils.base_backend import (
BaseCommandBackend,
ExecuteTrainConfig,
ExecuteTrainRequest,
)
from miles.utils.external_utils.command_utils.common import (
MOONCAKE_BACKEND_NAME,
OBJECT_STORE_BACKEND_FLAG,
@@ -21,7 +25,7 @@ from miles.utils.external_utils.ray_job import run_ray_job
class RayCommandBackend(BaseCommandBackend):
def _execute_train_inner(self, request: ExecuteTrainRequest) -> None:
def _execute_train_inner(self, *, request: ExecuteTrainRequest, config: ExecuteTrainConfig) -> None:
assert not request.extra_manifests, (
"extra_manifests are objects a helm release installs beside the run, and a ray launch installs no "
"release; launch onto kubernetes, or start what they describe yourself"
@@ -29,7 +33,7 @@ class RayCommandBackend(BaseCommandBackend):
external_ray = get_bool_env_var("MILES_SCRIPT_EXTERNAL_RAY")
master_addr = os.environ.get("MASTER_ADDR", "127.0.0.1")
mooncake_master_port = (
None if self.config.external_mooncake else self._resolve_mooncake_master_port(request.train_args)
None if config.external_mooncake else self._resolve_mooncake_master_port(request.train_args)
)
self._clean_up_previous_run(external_ray=external_ray)
@@ -50,7 +54,7 @@ class RayCommandBackend(BaseCommandBackend):
if (f := request.before_ray_job_submit) is not None:
f()
runtime_env_vars = train_env_vars(request, self._ray_env_vars(master_addr=master_addr), config=self.config)
runtime_env_vars = train_env_vars(request, self._ray_env_vars(master_addr=master_addr), config=config)
runtime_env_vars["PYTHONPATH"] = _pythonpath_with_sources(
request.megatron_path, runtime_env_vars.get("PYTHONPATH")
)
+1
View File
@@ -525,6 +525,7 @@ def _train(args: ScriptArgs) -> None:
backend = args.create_backend()
backend.execute_train(
train_args=train_args,
config=args,
num_gpus_per_node=args.num_gpus_per_node,
megatron_model_type=args.megatron_model_type,
extra_env_vars=extra_env_vars,
@@ -115,7 +115,7 @@ def run_ci(
+ "--mini-ft-controller-enable "
)
base_url = f"http://{config.create_backend().api_server_host()}:{API_SERVER_PORT}"
base_url = f"http://{config.create_backend().api_server_host(config)}:{API_SERVER_PORT}"
injector = spawn_fault_injector(
base_url=base_url,
seed=seed,
@@ -148,7 +148,7 @@ def run_realistic_gsm8k(
train_args += extra_train_args
run = Gsm8kRun(
base_url=f"http://{U.api_server_host()}:{API_SERVER_PORT}",
base_url=f"http://{U.api_server_host(config)}:{API_SERVER_PORT}",
config=config,
dump_dir=dump_dir,
train_args=train_args,
@@ -80,7 +80,7 @@ def _build_args(mode: FTTestMode, dump_dir: str, enable_dumper: bool = True) ->
def _inject_rollout_faults(
mode: FTTestMode, dump_dir: str, config: command_utils.ExecuteTrainConfig
) -> Iterator[None]:
base_url: str = f"http://{config.create_backend().api_server_host()}:{API_SERVER_PORT}"
base_url: str = f"http://{config.create_backend().api_server_host(config)}:{API_SERVER_PORT}"
print(f"Injecting into {ROLLOUT_CELL_TYPE} cells only, mean interval {CRASH_INTERVAL_SECONDS:.1f}s, seed {SEED}")
shutil.rmtree(dump_dir, ignore_errors=True)
@@ -21,9 +21,9 @@ class _RecordingBackend:
self.config = config
self._seen = seen
def api_server_host(self) -> str:
self._seen.asked_for_host.append(self.config)
return f"orchestrator-of-{self.config.run_id}"
def api_server_host(self, config: command_utils.ExecuteTrainConfig) -> str:
self._seen.asked_for_host.append(config)
return f"orchestrator-of-{config.run_id}"
def execute_train(self, **kwargs: object) -> None:
self._seen.trained.append(self.config)
@@ -9,11 +9,7 @@ from miles.utils.external_utils.command_utils.ray_backend.backend import RayComm
def test_workplace_launch_uses_the_configured_backend(monkeypatch: pytest.MonkeyPatch) -> None:
"""The workplace recipe submits through its backend instead of removed module exports."""
calls: list[dict[str, object]] = []
def capture(self: RayCommandBackend, **kwargs: object) -> None:
assert "config" not in kwargs
calls.append({**kwargs, "config": self.config})
monkeypatch.setattr(RayCommandBackend, "execute_train", capture)
monkeypatch.setattr(RayCommandBackend, "execute_train", lambda self, **kwargs: calls.append(kwargs))
module = runpy.run_path(
str(
Path(__file__).resolve().parents[3]
+2 -2
View File
@@ -91,8 +91,8 @@ def _called_name(func: ast.expr) -> str:
def _capture_backend(monkeypatch) -> list[ClusterBackend]:
chosen: list[ClusterBackend] = []
def _execute_train_inner(self, request) -> None:
chosen.append(self.config.cluster_backend)
def _execute_train_inner(self, *, request, config) -> None:
chosen.append(config.cluster_backend)
for backend_cls in (RayCommandBackend, KubernetesCommandBackend):
monkeypatch.setattr(backend_cls, "_execute_train_inner", _execute_train_inner)
@@ -40,13 +40,13 @@ RUN_ID = "260101-000000-000"
RUN_ID_SHORT_ENOUGH_FOR_FULL_OBJECT_NAMES = "260101-00"
SPLIT_RUN_UUID = "0123456789abcdef"
def _release(deploy_component: DeployComponent = DeployComponent.ALL) -> str:
return ReleaseName(run_id=RUN_ID, deploy_component=deploy_component, deploy_instance_id=None).serialize()
SPLIT_RUN_UUID = "0123456789abcdef"
def _config(run_id: str = RUN_ID, deploy_component: DeployComponent = DeployComponent.ALL) -> ExecuteTrainConfig:
return ExecuteTrainConfig(
cluster_backend=ClusterBackend.KUBERNETES,
@@ -319,17 +319,19 @@ class TestExecuteTrainTellsThePodsWhichPartOfTheRunTheyAre:
class TestApiServerHost:
def test_a_whole_run_answers_on_its_own_orchestrator(self):
"""The api server runs beside the orchestration script, which is a pod of the run's only release."""
host = KubernetesCommandBackend(_config()).api_server_host()
config = _config()
host = KubernetesCommandBackend(config).api_server_host(config)
assert host == f"{_release()}-orchestrator.{NAMESPACE}.svc"
@pytest.mark.parametrize("component", [DeployComponent.PRIMARY, DeployComponent.TRAINER])
def test_no_deployment_of_a_split_run_has_an_api_server_to_name(self, component):
"""A split run is refused an api server, so any host answered here would only ever time out."""
backend = KubernetesCommandBackend(_config(deploy_component=component))
config = _config(deploy_component=component)
backend = KubernetesCommandBackend(config)
with pytest.raises(AssertionError, match="--api-server-port 0"):
backend.api_server_host()
backend.api_server_host(config)
def _values_of_release(train_argv: list[str], *, run_id: str, release: str) -> dict[str, Any]:
@@ -8,10 +8,17 @@ from typing import Literal
import pytest
import typer
from miles.utils.external_utils.command_utils import base_backend
from miles.utils.external_utils.command_utils.base_backend import ExecuteTrainConfig, default_config, resolve_hardware
from miles.utils.external_utils.command_utils import CommandUtilConfig, base_backend
from miles.utils.external_utils.command_utils.base_backend import (
ExecuteTrainConfig,
ExecuteTrainRequest,
default_config,
resolve_extra_env_vars,
resolve_hardware,
)
from miles.utils.external_utils.command_utils.ray_backend.backend import RayCommandBackend
from miles.utils.typer_utils import SCRIPT_ENV_VAR_PREFIX, dataclass_cli
from miles.utils.workers.types import ClusterBackend
from miles.utils.workers.types import ClusterBackend, DeployComponent
@pytest.fixture(autouse=True)
@@ -47,6 +54,51 @@ class TestResolveHardware:
assert tuple(root.handlers) == expected_handlers
assert root.level == logging.WARNING
def test_supported_explicit_value_bypasses_detection_while_auto_uses_it(self, monkeypatch):
"""Explicit hardware bypasses detection, while auto resolves to a supported detected profile."""
detected: list[None] = []
def detect_hardware() -> str:
detected.append(None)
return "H100"
monkeypatch.setattr(base_backend, "detect_hardware", detect_hardware)
assert resolve_hardware(_HardwareConfig(hardware="H100")) == "H100"
assert detected == []
assert resolve_hardware(_HardwareConfig(hardware="auto")) == "H100"
assert detected == [None]
@pytest.mark.parametrize(
("configured_hardware", "detected_hardware"),
[("unsupported", "H100"), ("auto", "unsupported")],
)
def test_unsupported_explicit_or_detected_value_is_rejected(
self, configured_hardware: str, detected_hardware: str, monkeypatch
):
"""Neither explicit nor detected hardware may escape the config's supported profile literal."""
monkeypatch.setattr(base_backend, "detect_hardware", lambda: detected_hardware)
with pytest.raises(AssertionError, match="has no verified profile"):
resolve_hardware(_HardwareConfig(hardware=configured_hardware))
class TestResolveExtraEnvVars:
def test_config_extra_env_vars_override_the_callers_values(self):
"""Parsed config variables override duplicates while preserving caller-only variables."""
config = ExecuteTrainConfig(extra_env_vars="SHARED=from_config CONFIG_ONLY=kept")
resolved = resolve_extra_env_vars(
extra_env_vars={"SHARED": "from_caller", "CALLER_ONLY": "kept"},
config=config,
)
assert resolved == {
"SHARED": "from_config",
"CALLER_ONLY": "kept",
"CONFIG_ONLY": "kept",
}
class TestAScriptReadsItsLauncherConfigFromTheEnvironment:
def test_an_unset_environment_leaves_every_default_alone(self, monkeypatch):
@@ -131,3 +183,153 @@ class TestAHotRestartIsRefusedOutsideKubernetes:
config = ExecuteTrainConfig(cluster_backend=ClusterBackend.RAY)
assert config.parsed_hot_restart == []
class TestCommandUtilConfig:
def test_backend_config_only_contains_connection_fields(self):
"""Launch-only fields must remain on ExecuteTrainConfig rather than every command backend."""
assert [field.name for field in dataclasses.fields(CommandUtilConfig)] == [
"cluster_backend",
"namespace",
"helm_values",
"ci_run",
]
def test_backend_rejects_an_unrelated_config_type(self):
"""A backend must not retain an object that lacks its cluster connection fields."""
with pytest.raises(AssertionError, match="CommandUtilConfig"):
RayCommandBackend(object())
class TestExecuteTrainConfigSelection:
def test_an_explicit_config_for_another_backend_is_refused_before_launch(
self, monkeypatch: pytest.MonkeyPatch
) -> None:
"""A launch for another cluster must be refused before the backend performs any action."""
commands: list[str] = []
monkeypatch.setattr(base_backend, "run_shell_command", lambda command, **kwargs: commands.append(command))
backend = ExecuteTrainConfig(cluster_backend=ClusterBackend.RAY).create_backend()
launch_config = ExecuteTrainConfig(cluster_backend=ClusterBackend.KUBERNETES)
with pytest.raises(AssertionError, match="built to talk to ray"):
backend.execute_train(
train_args="--train-backend fsdp",
num_gpus_per_node=8,
megatron_model_type=None,
config=launch_config,
)
assert commands == []
def test_explicit_config_overrides_the_backend_default(self, monkeypatch):
"""A caller can launch a different deployment through an existing cluster backend."""
recorded: list[tuple[ExecuteTrainRequest, ExecuteTrainConfig]] = []
monkeypatch.setattr(
RayCommandBackend,
"_execute_train_inner",
lambda self, *, request, config: recorded.append((request, config)),
)
backend_config = ExecuteTrainConfig()
launch_config = ExecuteTrainConfig(deploy_component=DeployComponent.TRAINER)
backend_config.create_backend().execute_train(
train_args="--train-backend fsdp",
num_gpus_per_node=8,
megatron_model_type=None,
config=launch_config,
)
assert recorded[0][1] is launch_config
def test_omitted_config_uses_the_backend_execute_train_config(self, monkeypatch):
"""Existing launchers can keep constructing a backend and calling execute_train without config."""
recorded: list[tuple[ExecuteTrainRequest, ExecuteTrainConfig]] = []
monkeypatch.setattr(
RayCommandBackend,
"_execute_train_inner",
lambda self, *, request, config: recorded.append((request, config)),
)
config = ExecuteTrainConfig(deploy_component=DeployComponent.TRAINER)
config.create_backend().execute_train(
train_args="--train-backend fsdp", num_gpus_per_node=8, megatron_model_type=None
)
assert recorded[0][1] is config
def test_omitted_config_refuses_a_connection_only_backend(self):
"""A backend without launch fields cannot guess the execute_train configuration."""
backend = CommandUtilConfig().create_backend()
with pytest.raises(AssertionError, match="ExecuteTrainConfig"):
backend.execute_train(train_args="--train-backend fsdp", num_gpus_per_node=8, megatron_model_type=None)
def _launched_train_argv(monkeypatch, *, train_args: str, config: ExecuteTrainConfig) -> list[str]:
recorded: list[ExecuteTrainRequest] = []
monkeypatch.setattr(
RayCommandBackend, "_execute_train_inner", lambda self, *, request, config: recorded.append(request)
)
config.create_backend().execute_train(train_args=train_args, num_gpus_per_node=8, megatron_model_type=None)
return recorded[0].train_args.split()
class TestTheRunUuidALaunchDrives:
def test_the_configured_run_uuid_reaches_the_pods(self, monkeypatch):
"""Only --deploy-component and --deploy-instance-id were appended, so an unsplit run minted a second uuid."""
argv = _launched_train_argv(
monkeypatch,
train_args="--train-backend fsdp",
config=ExecuteTrainConfig(run_uuid="0123456789abcdef"),
)
assert argv[argv.index("--run-uuid") + 1] == "0123456789abcdef"
def test_a_launch_that_names_no_run_leaves_the_arguments_alone(self, monkeypatch):
"""Every existing ray launch names none, and an empty flag would be worse than no flag."""
argv = _launched_train_argv(monkeypatch, train_args="--train-backend fsdp", config=ExecuteTrainConfig())
assert "--run-uuid" not in argv
def test_train_arguments_that_already_name_this_run_are_accepted(self, monkeypatch):
"""The helm launcher sets the flag itself, and a launch agreeing with it is not a conflict."""
argv = _launched_train_argv(
monkeypatch,
train_args="--train-backend fsdp --run-uuid 0123456789abcdef",
config=ExecuteTrainConfig(run_uuid="0123456789abcdef"),
)
assert argv.count("--run-uuid") == 2
def test_refuses_train_arguments_that_name_another_run(self, monkeypatch):
"""The uuid is what joins the parts of a split run, so two of them are two runs."""
with pytest.raises(AssertionError, match="--run-uuid"):
_launched_train_argv(
monkeypatch,
train_args="--train-backend fsdp --run-uuid fedcba9876543210",
config=ExecuteTrainConfig(run_uuid="0123456789abcdef"),
)
def test_the_component_and_instance_flags_are_still_appended_beside_it(self, monkeypatch):
"""The run uuid joins a split run, and these two are what tell its halves apart."""
argv = _launched_train_argv(
monkeypatch,
train_args="--train-backend fsdp",
config=ExecuteTrainConfig(
run_uuid="0123456789abcdef", deploy_component=DeployComponent.TRAINER, deploy_instance_id="actor"
),
)
assert argv[argv.index("--deploy-component") + 1] == "trainer"
assert argv[argv.index("--deploy-instance-id") + 1] == "actor"
assert argv[argv.index("--run-uuid") + 1] == "0123456789abcdef"
class TestApiServerHost:
def test_the_ray_api_server_is_reached_on_localhost(self) -> None:
"""The default backend must keep fault-tolerance clients on the local API server."""
config = ExecuteTrainConfig(cluster_backend=ClusterBackend.RAY)
assert config.create_backend().api_server_host(config) == "localhost"