mirror of
https://github.com/radixark/miles.git
synced 2026-10-02 07:14:53 +08:00
Give execute_train the launch it is running (#3028)
This commit is contained in:
@@ -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")
|
||||
)
|
||||
|
||||
@@ -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]
|
||||
|
||||
@@ -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)
|
||||
|
||||
+8
-6
@@ -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"
|
||||
|
||||
Reference in New Issue
Block a user