mirror of
https://github.com/radixark/miles.git
synced 2026-10-02 07:14:53 +08:00
Forbid a backend from overriding the exec methods (#3052)
This commit is contained in:
@@ -95,6 +95,9 @@ TRAINER_ROLE = "trainer"
|
||||
_PREPARE_CMD_ROLES = frozenset({TRAINER_ROLE})
|
||||
|
||||
|
||||
_DELEGATING_METHODS = ("execute_train", "exec_command_cpu", "exec_command_gpu", "exec_command_multi_node")
|
||||
|
||||
|
||||
class BaseCommandBackend(ABC):
|
||||
def __init__(self, config: CommandUtilConfig) -> None:
|
||||
from miles.utils.logging_utils import configure_logger_raw
|
||||
@@ -103,6 +106,15 @@ class BaseCommandBackend(ABC):
|
||||
configure_logger_raw("launcher")
|
||||
self.config = config
|
||||
|
||||
def __init_subclass__(cls, **kwargs: object) -> None:
|
||||
super().__init_subclass__(**kwargs)
|
||||
for name in _DELEGATING_METHODS:
|
||||
assert name not in vars(cls), (
|
||||
f"{cls.__name__} defines its own {name}, but every command any backend runs has to pass through "
|
||||
f"the one BaseCommandBackend defines, so that stubbing that single method out cannot be bypassed "
|
||||
f"by a backend nobody remembered to name; implement _{name}_inner instead"
|
||||
)
|
||||
|
||||
def execute_train(
|
||||
self,
|
||||
train_args: str,
|
||||
@@ -275,15 +287,34 @@ class BaseCommandBackend(ABC):
|
||||
def _execute_train_inner(self, *, request: ExecuteTrainRequest, config: ExecuteTrainConfig) -> None: ...
|
||||
|
||||
def exec_command_cpu(self, cmd: str, capture_output: bool = False) -> str | None:
|
||||
return self._exec_command_cpu_inner(cmd, capture_output=capture_output)
|
||||
|
||||
def exec_command_gpu(
|
||||
self, cmd: str, capture_output: bool = False, num_gpus_per_node: int | None = None
|
||||
) -> str | None:
|
||||
return self._exec_command_gpu_inner(cmd, capture_output=capture_output, num_gpus_per_node=num_gpus_per_node)
|
||||
|
||||
def exec_command_multi_node(
|
||||
self,
|
||||
cmd: str,
|
||||
capture_output: bool = False,
|
||||
num_nodes: int | None = None,
|
||||
num_gpus_per_node: int | None = None,
|
||||
) -> list[str | None]:
|
||||
return self._exec_command_multi_node_inner(
|
||||
cmd, capture_output=capture_output, num_nodes=num_nodes, num_gpus_per_node=num_gpus_per_node
|
||||
)
|
||||
|
||||
def _exec_command_cpu_inner(self, cmd: str, capture_output: bool = False) -> str | None:
|
||||
return run_shell_command(cmd, capture_output=capture_output)
|
||||
|
||||
@abstractmethod
|
||||
def exec_command_gpu(
|
||||
def _exec_command_gpu_inner(
|
||||
self, cmd: str, capture_output: bool = False, num_gpus_per_node: int | None = None
|
||||
) -> str | None: ...
|
||||
|
||||
@abstractmethod
|
||||
def exec_command_multi_node(
|
||||
def _exec_command_multi_node_inner(
|
||||
self,
|
||||
cmd: str,
|
||||
capture_output: bool = False,
|
||||
|
||||
@@ -16,14 +16,14 @@ class KubernetesCommandBackend(BaseCommandBackend):
|
||||
def _execute_train_inner(self, *, request: ExecuteTrainRequest, config: ExecuteTrainConfig) -> None:
|
||||
entrypoint.execute_train(request=request, config=config)
|
||||
|
||||
def exec_command_gpu(
|
||||
def _exec_command_gpu_inner(
|
||||
self, cmd: str, capture_output: bool = False, num_gpus_per_node: int | None = None
|
||||
) -> str | None:
|
||||
return self.exec_command_multi_node(
|
||||
cmd, capture_output=capture_output, num_nodes=1, num_gpus_per_node=num_gpus_per_node
|
||||
)[0]
|
||||
|
||||
def exec_command_multi_node(
|
||||
def _exec_command_multi_node_inner(
|
||||
self,
|
||||
cmd: str,
|
||||
capture_output: bool = False,
|
||||
|
||||
@@ -46,7 +46,7 @@ class RayCommandBackend(BaseCommandBackend):
|
||||
)
|
||||
|
||||
if mooncake_master_port is not None:
|
||||
start_mooncake_master(rpc_port=mooncake_master_port)
|
||||
start_mooncake_master(rpc_port=mooncake_master_port, run_command=self.exec_command_cpu)
|
||||
|
||||
for cmd in request.prepare_cmd.values():
|
||||
self.exec_command_multi_node(cmd)
|
||||
@@ -78,12 +78,12 @@ class RayCommandBackend(BaseCommandBackend):
|
||||
|
||||
return get_mooncake_master_port(train_argv)
|
||||
|
||||
def exec_command_gpu(
|
||||
def _exec_command_gpu_inner(
|
||||
self, cmd: str, capture_output: bool = False, num_gpus_per_node: int | None = None
|
||||
) -> str | None:
|
||||
return run_shell_command(cmd, capture_output=capture_output)
|
||||
|
||||
def exec_command_multi_node(
|
||||
def _exec_command_multi_node_inner(
|
||||
self,
|
||||
cmd: str,
|
||||
capture_output: bool = False,
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
import logging
|
||||
import shlex
|
||||
from collections.abc import Callable
|
||||
from pathlib import Path
|
||||
|
||||
import ray
|
||||
@@ -70,6 +71,7 @@ def start_mooncake_master(
|
||||
metrics_port: int = MOONCAKE_MASTER_METRICS_PORT,
|
||||
timeout: float = 30,
|
||||
log_path: str | Path = MOONCAKE_MASTER_LOG_PATH,
|
||||
run_command: Callable[[str], str | None] | None = None,
|
||||
) -> None:
|
||||
host = "127.0.0.1"
|
||||
if _is_tcp_server_ready(host, rpc_port):
|
||||
@@ -78,15 +80,17 @@ def start_mooncake_master(
|
||||
|
||||
log_path = Path(log_path)
|
||||
quoted_log_path = shlex.quote(str(log_path))
|
||||
run_shell_command(
|
||||
"pkill -x mooncake_master >/dev/null 2>&1 || true; "
|
||||
kill_master_cmd = f"pkill -f {shlex.quote(f'^mooncake_master --rpc_port {rpc_port} ')} >/dev/null 2>&1 || true"
|
||||
command_runner = run_command if run_command is not None else run_shell_command
|
||||
command_runner(
|
||||
f"{kill_master_cmd}; "
|
||||
f"(setsid mooncake_master --rpc_port {rpc_port} --metrics_port {metrics_port} "
|
||||
f"> {quoted_log_path} 2>&1 &)"
|
||||
)
|
||||
try:
|
||||
wait_for_server_ready(host, rpc_port, timeout=timeout)
|
||||
except RuntimeError as exc:
|
||||
run_shell_command("pkill -x mooncake_master >/dev/null 2>&1 || true")
|
||||
command_runner(kill_master_cmd)
|
||||
try:
|
||||
log_lines = log_path.read_text(errors="replace").splitlines()
|
||||
log_tail = "\n".join(log_lines[-100:]) or "<empty>"
|
||||
|
||||
@@ -35,7 +35,7 @@ def bare_environment(monkeypatch):
|
||||
def commands(monkeypatch):
|
||||
recorded = record_commands(monkeypatch)
|
||||
patch_helper(monkeypatch, "_check_has_nvlink", lambda self: False, backend_class=RayCommandBackend)
|
||||
for name in ("MILES_SCRIPT_EXTERNAL_RAY", "RAY_ADDRESS", "NCCL_NVLS_ENABLE", "WANDB_API_KEY"):
|
||||
for name in ("RAY_ADDRESS", "NCCL_NVLS_ENABLE", "WANDB_API_KEY"):
|
||||
monkeypatch.delenv(name, raising=False)
|
||||
monkeypatch.setenv("MILES_SCRIPT_ENABLE_RAY_SUBMIT", "1")
|
||||
monkeypatch.setenv("MASTER_ADDR", "127.0.0.1")
|
||||
|
||||
Reference in New Issue
Block a user