Refuse a multi-node command that a node pin cannot schedule (#3054)

This commit is contained in:
fzyzcjy
2026-09-26 20:35:29 +08:00
committed by GitHub
parent f4f49f1d17
commit 220bc2bb0d
2 changed files with 195 additions and 2 deletions
@@ -1,5 +1,6 @@
from __future__ import annotations
from pathlib import Path
from miles.utils.external_utils.command_utils.base_backend import (
BaseCommandBackend,
@@ -9,8 +10,12 @@ from miles.utils.external_utils.command_utils.base_backend import (
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
from miles.utils.external_utils.command_utils.helm_backend.launcher.values.helm_values_types import Scheduling
from miles.utils.external_utils.command_utils.helm_backend.launcher.values.misc import InfraInfo
from miles.utils.external_utils.command_utils.helm_backend.naming import ReleaseName, RunNames
_HOSTNAME_LABEL = "kubernetes.io/hostname"
class KubernetesCommandBackend(BaseCommandBackend):
def _execute_train_inner(self, *, request: ExecuteTrainRequest, config: ExecuteTrainConfig) -> None:
@@ -31,19 +36,44 @@ class KubernetesCommandBackend(BaseCommandBackend):
num_gpus_per_node: int | None = None,
) -> list[str | None]:
assert self.config.namespace, "Set CommandUtilConfig.namespace to run a command somewhere"
assert (
num_nodes is not None
), "kubernetes cannot infer num_nodes from a live cluster the way ray does; pass num_nodes explicitly"
chart = chart_dir(repo_base_dir=repo_base_dir)
self._assert_requested_nodes_schedulable(num_nodes=num_nodes, chart=chart)
return command_job.run_on_nodes(
command_job.CommandJobContext(
namespace=self.config.namespace,
chart_dir=chart_dir(repo_base_dir=repo_base_dir),
chart_dir=chart,
helm_values_files=tuple(self.config.helm_values),
gpus_per_node=num_gpus_per_node if num_gpus_per_node is not None else 1,
),
cmd,
capture_output=capture_output,
completions=num_nodes or 1,
completions=num_nodes,
step="command",
)
def _assert_requested_nodes_schedulable(self, *, num_nodes: int, chart: Path) -> None:
if num_nodes <= 1:
return
infra = InfraInfo.load(chart, list(self.config.helm_values))
scheduling = infra.scheduling
node_selector = (scheduling.node_selector if scheduling is not None else None) or {}
if (hostname := node_selector.get(_HOSTNAME_LABEL)) is not None:
raise AssertionError(
f"this command asks for {num_nodes} nodes, and infra.scheduling.nodeSelector pins every pod of this "
f"deployment to {hostname!r}, so every completion after the first would stay Pending for good while "
f"the first one holds its gpus until the job times out; ask for one node, or unpin the deployment"
)
hosts = _required_affinity_hostnames(scheduling)
assert hosts is None or len(hosts) >= num_nodes, (
f"this command asks for {num_nodes} nodes, but required node affinity permits at most "
f"{len(hosts)} hostnames: {sorted(hosts)}; ask for fewer nodes, or unpin the deployment"
)
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 "
@@ -62,3 +92,26 @@ class KubernetesCommandBackend(BaseCommandBackend):
).serialize(),
namespace=config.namespace,
)
def _required_affinity_hostnames(scheduling: Scheduling | None) -> set[str] | None:
if scheduling is None or not scheduling.affinity:
return None
affinity = scheduling.affinity.get("nodeAffinity") or {}
required = affinity.get("requiredDuringSchedulingIgnoredDuringExecution")
if required is None:
return None
allowed: set[str] = set()
for term in required.get("nodeSelectorTerms", []):
if not term:
continue
bounds = [
set(expression.get("values", []))
for expression in term.get("matchExpressions", [])
if expression.get("key") == _HOSTNAME_LABEL and expression.get("operator") == "In"
]
if not bounds:
return None
allowed.update(set.intersection(*bounds))
return allowed
@@ -0,0 +1,140 @@
import textwrap
from pathlib import Path
from typing import Any
import pytest
import yaml
from miles.utils.external_utils.command_utils.helm_backend import command_job
NAMESPACE = "rl"
def _backend(helm_values: tuple[str, ...]) -> Any:
pytest.importorskip("torch")
from miles.utils.external_utils.command_utils.base_backend import ExecuteTrainConfig
from miles.utils.external_utils.command_utils.helm_backend.backend import KubernetesCommandBackend
return KubernetesCommandBackend(ExecuteTrainConfig(namespace=NAMESPACE, helm_values=helm_values))
def _values_pinned_to(tmp_path: Path, hostname: str) -> tuple[str, ...]:
values = tmp_path / "infra-pinned.yaml"
values.write_text(
textwrap.dedent(
f"""
infra:
scheduling:
nodeSelector:
kubernetes.io/hostname: {hostname}
"""
)
)
return (str(values),)
def _record_completions(monkeypatch: pytest.MonkeyPatch) -> list[int]:
completions: list[int] = []
def fake_run_on_nodes(context: Any, cmd: str, **kwargs: Any) -> list[str | None]:
completions.append(kwargs["completions"])
return [None] * kwargs["completions"]
monkeypatch.setattr(command_job, "run_on_nodes", fake_run_on_nodes)
return completions
class TestTheNodesACommandAsksFor:
@pytest.mark.parametrize(
("terms", "num_nodes", "refused"),
[
([["gpu-1"]], 2, True),
([["gpu-1", "gpu-2"]], 3, True),
([["gpu-1"], ["gpu-2"]], 2, False),
([["gpu-1"], None], 2, False),
([["gpu-1"]], 1, False),
],
)
def test_required_affinity_host_bounds_respect_alternative_terms(
self,
monkeypatch: pytest.MonkeyPatch,
tmp_path: Path,
terms: list[list[str] | None],
num_nodes: int,
refused: bool,
) -> None:
"""Only a proven shortage across all affinity alternatives blocks submission."""
expressions = [
(
{"matchExpressions": [{"key": "kubernetes.io/hostname", "operator": "In", "values": hosts}]}
if hosts is not None
else {"matchExpressions": [{"key": "accelerator", "operator": "Exists"}]}
)
for hosts in terms
]
values = tmp_path / "affinity.yaml"
values.write_text(
yaml.safe_dump(
{
"infra": {
"scheduling": {
"affinity": {
"nodeAffinity": {
"requiredDuringSchedulingIgnoredDuringExecution": {
"nodeSelectorTerms": expressions
}
}
}
}
}
}
)
)
completions = _record_completions(monkeypatch)
backend = _backend((str(values),))
if refused:
with pytest.raises(AssertionError, match="required node affinity"):
backend.exec_command_multi_node("true", num_nodes=num_nodes, num_gpus_per_node=1)
assert completions == []
else:
backend.exec_command_multi_node("true", num_nodes=num_nodes, num_gpus_per_node=1)
assert completions == [num_nodes]
def test_refuses_a_multi_node_command_pinned_to_a_single_host(self, monkeypatch, tmp_path):
"""A two-node command under a single-host nodeSelector is refused before any Job reaches the cluster."""
completions = _record_completions(monkeypatch)
backend = _backend(_values_pinned_to(tmp_path, "gpu-1"))
with pytest.raises(AssertionError, match="gpu-1"):
backend.exec_command_multi_node("torchrun --nnodes={{nnodes}}", num_nodes=2, num_gpus_per_node=1)
assert completions == []
def test_runs_a_single_node_command_on_a_pinned_deployment(self, monkeypatch, tmp_path):
"""One node always fits the host it is pinned to, so the command is installed as before."""
completions = _record_completions(monkeypatch)
backend = _backend(_values_pinned_to(tmp_path, "gpu-1"))
backend.exec_command_multi_node("torchrun --nnodes={{nnodes}}", num_nodes=1, num_gpus_per_node=1)
assert completions == [1]
def test_refuses_a_command_that_never_says_how_many_nodes_it_wants(self, monkeypatch):
"""Ray reads that as every alive node; a Job would silently run a multi-node command on one node."""
completions = _record_completions(monkeypatch)
backend = _backend(())
with pytest.raises(AssertionError, match="pass num_nodes explicitly"):
backend.exec_command_multi_node("torchrun --nnodes={{nnodes}}", num_gpus_per_node=1)
assert completions == []
def test_runs_a_multi_node_command_when_no_host_is_pinned(self, monkeypatch):
"""Without a host pin the backend keeps letting the scheduler place every completion."""
completions = _record_completions(monkeypatch)
backend = _backend(())
backend.exec_command_multi_node("torchrun --nnodes={{nnodes}}", num_nodes=2, num_gpus_per_node=1)
assert completions == [2]