diff --git a/src/srtctl/backends/vllm.py b/src/srtctl/backends/vllm.py index 74f673b79..377606a2b 100644 --- a/src/srtctl/backends/vllm.py +++ b/src/srtctl/backends/vllm.py @@ -716,25 +716,31 @@ def build_worker_command( if frontend_type == "vllm": if mode != "agg": raise ValueError("frontend.type: vllm supports aggregate vLLM jobs only") - if is_multi_node: - raise ValueError("frontend.type: vllm currently supports single-node aggregate jobs only") config.pop("host", None) config.pop("port", None) config.pop("connector", None) config.setdefault("served-model-name", served_model_name) - cmd.extend( - [ - "vllm", - "serve", - model_arg, - "--host", - "0.0.0.0", - "--port", - str(runtime.frontend_port), - ] - ) + node_rank = endpoint_nodes.index(process.node) + cmd.extend(["vllm", "serve", model_arg]) + if node_rank == 0: + cmd.extend(["--host", "0.0.0.0", "--port", str(runtime.frontend_port)]) + if is_multi_node: + # vLLM-native multi-node serve (torchrun-style): the leader owns + # the OpenAI server; other node ranks run headless engine workers. + cmd.extend( + [ + "--master-addr", + leader_ip, + "--nnodes", + str(len(endpoint_nodes)), + "--node-rank", + str(node_rank), + ] + ) + if node_rank > 0: + cmd.append("--headless") if not self.set_cuda_visible_devices: device_ids = ",".join(str(i) for i in sorted(process.gpu_indices)) if device_ids: diff --git a/src/srtctl/core/schema.py b/src/srtctl/core/schema.py index 1263ddcf0..0ef7ae44a 100644 --- a/src/srtctl/core/schema.py +++ b/src/srtctl/core/schema.py @@ -1587,8 +1587,6 @@ def _validate_vllm_frontend(self): raise ValidationError("frontend.type: vllm supports aggregate jobs only, not disaggregated layouts") if self.resources.num_agg < 1: raise ValidationError("frontend.type: vllm requires resources.agg_workers >= 1") - if (self.resources.agg_nodes or 1) != 1: - raise ValidationError("frontend.type: vllm currently supports single-node aggregate jobs only") def _validate_het_jobs(self): """When ``resources.het_jobs`` is set to True, enforce supported shape. diff --git a/tests/test_configs.py b/tests/test_configs.py index 154b27060..36d9b0e2c 100644 --- a/tests/test_configs.py +++ b/tests/test_configs.py @@ -3122,3 +3122,131 @@ def mock_scontrol(cmd, **kwargs): assert extra_root.resolve() in runtime.container_mounts assert runtime.container_mounts[extra_root.resolve()] == Path("/extra") + + +class TestDirectVllmMultiNode: + """Multi-node aggregate support for the direct vllm frontend.""" + + def _make_config(self, **resource_overrides): + from srtctl.backends import VLLMProtocol, VLLMServerConfig + from srtctl.core.schema import FrontendConfig, ModelConfig, ResourceConfig, SrtConfig + + resources_kwargs = { + "gpu_type": "b200", + "gpus_per_node": 8, + "agg_nodes": 2, + "agg_workers": 1, + } + resources_kwargs.update(resource_overrides) + return SrtConfig( + name="t", + model=ModelConfig(path="/m", container="/c.sqsh", precision="fp4"), + resources=ResourceConfig(**resources_kwargs), + frontend=FrontendConfig(type="vllm", enable_multiple_frontends=False), + backend=VLLMProtocol( + vllm_config=VLLMServerConfig(aggregated={"tensor-parallel-size": 8}) + ), + ) + + def _make_processes(self, nodes): + from srtctl.core.topology import Process + + return [ + Process( + node=node, + gpu_indices=frozenset(range(8)), + sys_port=8081, + http_port=0, + endpoint_mode="agg", + endpoint_index=0, + node_rank=rank, + ) + for rank, node in enumerate(nodes) + ] + + def _build_command(self, process, endpoint_processes): + from pathlib import Path + from unittest.mock import MagicMock, patch + + from srtctl.backends import VLLMProtocol, VLLMServerConfig + + backend = VLLMProtocol( + vllm_config=VLLMServerConfig( + aggregated={"tensor-parallel-size": 8, "pipeline-parallel-size": 2} + ) + ) + runtime = MagicMock() + runtime.model_path = Path("/model") + runtime.is_hf_model = False + runtime.frontend_port = 9000 + + with patch("srtctl.core.slurm.get_hostname_ip", return_value="10.0.0.1"): + return backend.build_worker_command( + process=process, + endpoint_processes=endpoint_processes, + runtime=runtime, + frontend_type="vllm", + ) + + def test_schema_accepts_multi_node_aggregate(self): + """agg_nodes > 1 no longer trips the vllm-frontend load-time validation.""" + cfg = self._make_config() + assert cfg.resources.agg_nodes == 2 + assert cfg.frontend.type == "vllm" + + def test_schema_still_rejects_disaggregated(self): + import pytest + from marshmallow import ValidationError + + with pytest.raises(ValidationError, match="aggregate jobs only"): + self._make_config( + agg_nodes=None, + agg_workers=None, + prefill_nodes=1, + decode_nodes=1, + prefill_workers=1, + decode_workers=1, + ) + + def test_multi_node_leader_owns_port_and_coordination(self): + """Rank 0 keeps the OpenAI port and gets the torchrun-style flags.""" + leader, worker = self._make_processes(["node0", "node1"]) + + cmd = self._build_command(leader, [leader, worker]) + + assert cmd[:3] == ["vllm", "serve", "/model"] + assert cmd[cmd.index("--host") + 1] == "0.0.0.0" + assert cmd[cmd.index("--port") + 1] == "9000" + assert cmd[cmd.index("--master-addr") + 1] == "10.0.0.1" + assert cmd[cmd.index("--nnodes") + 1] == "2" + assert cmd[cmd.index("--node-rank") + 1] == "0" + assert "--headless" not in cmd + assert "dynamo.vllm" not in cmd + assert "--request-plane" not in cmd + + def test_multi_node_nonleader_runs_headless_without_port(self): + """Ranks > 0 are headless engine workers and must not bind the API port.""" + leader, worker = self._make_processes(["node0", "node1"]) + + cmd = self._build_command(worker, [leader, worker]) + + assert cmd[:3] == ["vllm", "serve", "/model"] + assert "--headless" in cmd + assert cmd[cmd.index("--node-rank") + 1] == "1" + assert cmd[cmd.index("--nnodes") + 1] == "2" + assert cmd[cmd.index("--master-addr") + 1] == "10.0.0.1" + assert "--host" not in cmd + assert "--port" not in cmd + + def test_single_node_command_has_no_multinode_flags(self): + """The original single-node command shape is unchanged.""" + (leader,) = self._make_processes(["node0"]) + + cmd = self._build_command(leader, [leader]) + + assert cmd[:3] == ["vllm", "serve", "/model"] + assert cmd[cmd.index("--port") + 1] == "9000" + assert "--nnodes" not in cmd + assert "--node-rank" not in cmd + assert "--master-addr" not in cmd + assert "--headless" not in cmd