Forbid a backend from overriding the exec methods (#3052)

This commit is contained in:
fzyzcjy
2026-09-26 20:34:31 +08:00
committed by GitHub
parent 4daa1026cc
commit fb1d3861c8
5 changed files with 46 additions and 11 deletions
@@ -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>"
+1 -1
View File
@@ -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")