From 69e9bd0983760e3f8065579a27b5560095ba7429 Mon Sep 17 00:00:00 2001 From: Biswa Panda Date: Mon, 3 Aug 2026 11:30:32 -0700 Subject: [PATCH 1/9] fix(nccl): preserve checkpoint layer boundaries --- src/prime_rl/inference/vllm/worker/nccl.py | 34 +++++++++++-------- .../unit/inference/test_nccl_weight_update.py | 15 ++++++++ 2 files changed, 35 insertions(+), 14 deletions(-) create mode 100644 tests/unit/inference/test_nccl_weight_update.py diff --git a/src/prime_rl/inference/vllm/worker/nccl.py b/src/prime_rl/inference/vllm/worker/nccl.py index 686bb856cf..d2f5f66cae 100644 --- a/src/prime_rl/inference/vllm/worker/nccl.py +++ b/src/prime_rl/inference/vllm/worker/nccl.py @@ -78,15 +78,22 @@ def __init__( self.communicator = PyNcclCommunicator(pg, device=device) @torch.no_grad() - def receive_state_dict(self): - """Receives the state dict of a model from the trainer master rank using NCCL communicator.""" + def receive_state_dicts( + self, + ) -> Generator[Generator[tuple[str, torch.Tensor], None, None], None, None]: + """Receive each trainer-broadcast state dict as a separate stream.""" logger.info("Receiving weights from trainer") num_state_dict_to_receive = receive_integer(self.communicator) logger.info(f"Receiving {num_state_dict_to_receive} layer state dicts") for layer_id in range(num_state_dict_to_receive): logger.info(f"Receiving state dict {layer_id + 1}/{num_state_dict_to_receive}") - for key, value in receive_state_dict(self.communicator): - yield key, value + yield receive_state_dict(self.communicator) + + @torch.no_grad() + def receive_state_dict(self): + """Receive trainer weights as one flat stream for kernel-format loading.""" + for state_dict in self.receive_state_dicts(): + yield from state_dict class NCCLWeightUpdateWorker(Worker): @@ -144,15 +151,14 @@ def update_weights_from_path(self, weight_dir: str) -> None: model = model_runner.model assert isinstance(model, Module) - state_iter = self.nccl_broadcast_receiver.receive_state_dict() if self.quantize_in_weight_transfer: - load_weights_kernel(model, state_iter) + load_weights_kernel(model, self.nccl_broadcast_receiver.receive_state_dict()) update_mla_absorbed_weights(model) - return - - load_weights_checkpoint_layerwise( - model, - state_iter, - self.model_runner.model_config, - self.vllm_config, - ) + else: + for state_iter in self.nccl_broadcast_receiver.receive_state_dicts(): + load_weights_checkpoint_layerwise( + model, + state_iter, + self.model_runner.model_config, + self.vllm_config, + ) diff --git a/tests/unit/inference/test_nccl_weight_update.py b/tests/unit/inference/test_nccl_weight_update.py new file mode 100644 index 0000000000..155fd98cf0 --- /dev/null +++ b/tests/unit/inference/test_nccl_weight_update.py @@ -0,0 +1,15 @@ +from prime_rl.inference.vllm.worker import nccl + + +def test_receiver_preserves_state_dict_boundaries(monkeypatch): + receiver = object.__new__(nccl.NCCLWeightBroadcastReceiver) + receiver.communicator = object() + streams = [iter([("layer.0", 0)]), iter([("layer.1", 1)])] + + monkeypatch.setattr(nccl, "receive_integer", lambda _communicator: len(streams)) + monkeypatch.setattr(nccl, "receive_state_dict", lambda _communicator: streams.pop(0)) + + assert [list(stream) for stream in receiver.receive_state_dicts()] == [ + [("layer.0", 0)], + [("layer.1", 1)], + ] From b12b30435e8858dbc6429b2ed2616d3a6b51f8e1 Mon Sep 17 00:00:00 2001 From: Biswa Panda Date: Mon, 3 Aug 2026 14:41:05 -0700 Subject: [PATCH 2/9] fix(nccl): bracket grouped layerwise reloads --- src/prime_rl/inference/vllm/worker/nccl.py | 15 ++++---- .../inference/vllm/worker/weight_transfer.py | 12 ++++++- .../unit/inference/test_nccl_weight_update.py | 35 ++++++++++++++++++- 3 files changed, 52 insertions(+), 10 deletions(-) diff --git a/src/prime_rl/inference/vllm/worker/nccl.py b/src/prime_rl/inference/vllm/worker/nccl.py index d2f5f66cae..952a05c142 100644 --- a/src/prime_rl/inference/vllm/worker/nccl.py +++ b/src/prime_rl/inference/vllm/worker/nccl.py @@ -8,7 +8,7 @@ from vllm.logger import init_logger from prime_rl.inference.vllm.worker.weight_transfer import ( - load_weights_checkpoint_layerwise, + load_weight_groups_checkpoint_layerwise, load_weights_kernel, update_mla_absorbed_weights, ) @@ -155,10 +155,9 @@ def update_weights_from_path(self, weight_dir: str) -> None: load_weights_kernel(model, self.nccl_broadcast_receiver.receive_state_dict()) update_mla_absorbed_weights(model) else: - for state_iter in self.nccl_broadcast_receiver.receive_state_dicts(): - load_weights_checkpoint_layerwise( - model, - state_iter, - self.model_runner.model_config, - self.vllm_config, - ) + load_weight_groups_checkpoint_layerwise( + model, + self.nccl_broadcast_receiver.receive_state_dicts(), + self.model_runner.model_config, + self.vllm_config, + ) diff --git a/src/prime_rl/inference/vllm/worker/weight_transfer.py b/src/prime_rl/inference/vllm/worker/weight_transfer.py index 545c62ef7f..f249bd29fd 100644 --- a/src/prime_rl/inference/vllm/worker/weight_transfer.py +++ b/src/prime_rl/inference/vllm/worker/weight_transfer.py @@ -15,12 +15,22 @@ def load_weights_checkpoint_layerwise( state_iter: Iterable[tuple[str, torch.Tensor]], model_config, vllm_config, +) -> None: + load_weight_groups_checkpoint_layerwise(model, (state_iter,), model_config, vllm_config) + + +def load_weight_groups_checkpoint_layerwise( + model: Module, + state_iters: Iterable[Iterable[tuple[str, torch.Tensor]]], + model_config, + vllm_config, ) -> None: logger.info("Reloading checkpoint-format weights with vLLM layerwise processing") device = next(model.parameters()).device with torch.device(device), set_current_vllm_config(vllm_config): initialize_layerwise_reload(model) - model.load_weights(state_iter) # type: ignore + for state_iter in state_iters: + model.load_weights(state_iter) # type: ignore finalize_layerwise_reload(model, model_config) diff --git a/tests/unit/inference/test_nccl_weight_update.py b/tests/unit/inference/test_nccl_weight_update.py index 155fd98cf0..b2b8a86466 100644 --- a/tests/unit/inference/test_nccl_weight_update.py +++ b/tests/unit/inference/test_nccl_weight_update.py @@ -1,4 +1,8 @@ -from prime_rl.inference.vllm.worker import nccl +from contextlib import nullcontext +from types import SimpleNamespace +from unittest.mock import Mock + +from prime_rl.inference.vllm.worker import nccl, weight_transfer def test_receiver_preserves_state_dict_boundaries(monkeypatch): @@ -13,3 +17,32 @@ def test_receiver_preserves_state_dict_boundaries(monkeypatch): [("layer.0", 0)], [("layer.1", 1)], ] + + +def test_layerwise_reload_brackets_all_state_dict_groups(monkeypatch): + events = [] + model = Mock() + model.parameters.return_value = iter([SimpleNamespace(device="cpu")]) + model.load_weights.side_effect = lambda state_iter: events.append(("load", list(state_iter))) + + monkeypatch.setattr(weight_transfer, "set_current_vllm_config", lambda _config: nullcontext()) + monkeypatch.setattr(weight_transfer, "initialize_layerwise_reload", lambda _model: events.append("initialize")) + monkeypatch.setattr( + weight_transfer, + "finalize_layerwise_reload", + lambda _model, _model_config: events.append("finalize"), + ) + + weight_transfer.load_weight_groups_checkpoint_layerwise( + model, + [iter([("layer.0", 0)]), iter([("layer.1", 1)])], + model_config=object(), + vllm_config=object(), + ) + + assert events == [ + "initialize", + ("load", [("layer.0", 0)]), + ("load", [("layer.1", 1)]), + "finalize", + ] From 9aec678fccd2285321d680b3af1848a00cc21b13 Mon Sep 17 00:00:00 2001 From: Biswa Panda Date: Mon, 3 Aug 2026 11:32:40 -0700 Subject: [PATCH 3/9] fix(inference): map external engine ranks --- src/prime_rl/inference/vllm/ranks.py | 39 +++++++ src/prime_rl/inference/vllm/server.py | 12 +- src/prime_rl/inference/vllm/worker/nccl.py | 27 +++-- src/prime_rl/inference/vllm/worker/nixl.py | 20 +++- src/prime_rl/utils/client.py | 108 +++++++++++++----- tests/unit/inference/test_nccl_rank.py | 90 +++++++++++++++ .../utils/test_external_engine_topology.py | 44 +++++++ 7 files changed, 299 insertions(+), 41 deletions(-) create mode 100644 src/prime_rl/inference/vllm/ranks.py create mode 100644 tests/unit/inference/test_nccl_rank.py create mode 100644 tests/unit/utils/test_external_engine_topology.py diff --git a/src/prime_rl/inference/vllm/ranks.py b/src/prime_rl/inference/vllm/ranks.py new file mode 100644 index 0000000000..8de43f4ace --- /dev/null +++ b/src/prime_rl/inference/vllm/ranks.py @@ -0,0 +1,39 @@ +def global_inference_rank( + *, + rank_offset: int, + data_parallel_index: int, + data_parallel_size: int, + worker_rank: int, + tensor_parallel_size: int, + pipeline_parallel_size: int, + inference_world_size: int, + prefill_context_parallel_size: int = 1, + engine_world_size: int | None = None, +) -> int: + """Map one vLLM worker to its rank in Prime's inference NCCL group.""" + model_parallel_size = tensor_parallel_size * pipeline_parallel_size * prefill_context_parallel_size + logical_data_parallel_size = data_parallel_size + if engine_world_size is not None: + if engine_world_size <= 0 or engine_world_size % model_parallel_size: + raise ValueError( + f"engine world size {engine_world_size} is not divisible by model parallel size {model_parallel_size}" + ) + logical_data_parallel_size = engine_world_size // model_parallel_size + # Dense vLLM EngineCore processes retain their global DP index but rewrite + # data_parallel_size to one. MoE EngineCore processes preserve the logical + # size, so keep validating that value rather than masking bad discovery. + if data_parallel_size != 1 and data_parallel_size != logical_data_parallel_size: + raise ValueError( + f"data parallel size {data_parallel_size} does not match engine-derived size " + f"{logical_data_parallel_size}" + ) + if not 0 <= data_parallel_index < logical_data_parallel_size: + raise ValueError( + f"data parallel index {data_parallel_index} is outside logical data parallel size " + f"{logical_data_parallel_size}" + ) + + rank = rank_offset + data_parallel_index * model_parallel_size + worker_rank % model_parallel_size + if not 0 <= rank < inference_world_size: + raise ValueError(f"calculated inference rank {rank} is outside inference world size {inference_world_size}") + return rank diff --git a/src/prime_rl/inference/vllm/server.py b/src/prime_rl/inference/vllm/server.py index 728574aba9..984ed47eed 100644 --- a/src/prime_rl/inference/vllm/server.py +++ b/src/prime_rl/inference/vllm/server.py @@ -138,11 +138,21 @@ async def init_broadcaster(request: Request): timeout = data.get("timeout") rank_offset = data.get("rank_offset") inference_world_size = data.get("inference_world_size") + engine_world_size = data.get("engine_world_size") quantize_in_weight_transfer = data.get("quantize_in_weight_transfer", False) session_id = data.get("session_id", "default") await engine_client(request).collective_rpc( "init_broadcaster", - args=(host, port, rank_offset, inference_world_size, timeout, quantize_in_weight_transfer, session_id), + args=( + host, + port, + rank_offset, + inference_world_size, + timeout, + quantize_in_weight_transfer, + session_id, + engine_world_size, + ), ) return {"status": "ok"} diff --git a/src/prime_rl/inference/vllm/worker/nccl.py b/src/prime_rl/inference/vllm/worker/nccl.py index 952a05c142..90b7fc6681 100644 --- a/src/prime_rl/inference/vllm/worker/nccl.py +++ b/src/prime_rl/inference/vllm/worker/nccl.py @@ -7,6 +7,7 @@ from vllm.distributed.utils import StatelessProcessGroup from vllm.logger import init_logger +from prime_rl.inference.vllm.ranks import global_inference_rank from prime_rl.inference.vllm.worker.weight_transfer import ( load_weight_groups_checkpoint_layerwise, load_weights_kernel, @@ -108,27 +109,37 @@ def init_broadcaster( timeout: int, quantize_in_weight_transfer: bool = False, session_id: str = "default", + engine_world_size: int | None = None, ) -> None: """Initialize the NCCL broadcast receiver. Args: rank_offset: Starting GPU offset for this server in the global inference group. inference_world_size: Total number of inference GPUs across all servers. + engine_world_size: Number of inference GPUs assigned to this server. """ del session_id self.quantize_in_weight_transfer = quantize_in_weight_transfer - # Use the worker's device index directly as the local rank. - # The previous dp_group-based computation broke in vLLM v1 multiprocess - # DP mode where each worker is a separate process with a singleton - # DP group (rank_in_group is always 0). - local_rank = self.device.index - global_rank_inference = rank_offset + local_rank + if engine_world_size is None: + global_rank_inference = rank_offset + self.device.index + else: + parallel_config = self.parallel_config + global_rank_inference = global_inference_rank( + rank_offset=rank_offset, + data_parallel_index=parallel_config.data_parallel_index, + data_parallel_size=parallel_config.data_parallel_size, + worker_rank=self.rank, + tensor_parallel_size=parallel_config.tensor_parallel_size, + pipeline_parallel_size=parallel_config.pipeline_parallel_size, + prefill_context_parallel_size=parallel_config.prefill_context_parallel_size, + inference_world_size=inference_world_size, + engine_world_size=engine_world_size, + ) logger.info( - f"Worker [local_rank={local_rank} rank_offset={rank_offset}] " + f"Worker [worker_rank={self.rank} rank_offset={rank_offset}] " f"-> [global_rank={global_rank_inference} inference_world_size={inference_world_size}]" ) - self.nccl_broadcast_receiver = NCCLWeightBroadcastReceiver( host=host, port=port, diff --git a/src/prime_rl/inference/vllm/worker/nixl.py b/src/prime_rl/inference/vllm/worker/nixl.py index d6990211a1..5959b3a6bc 100644 --- a/src/prime_rl/inference/vllm/worker/nixl.py +++ b/src/prime_rl/inference/vllm/worker/nixl.py @@ -18,6 +18,7 @@ from vllm.config import set_current_vllm_config from vllm.logger import init_logger +from prime_rl.inference.vllm.ranks import global_inference_rank from prime_rl.inference.vllm.worker.weight_transfer import update_mla_absorbed_weights from prime_rl.trainer.rl.broadcast.nixl.agent import MemDesc, NixlAgent, make_agent_name, set_ucx_env_defaults from prime_rl.trainer.rl.broadcast.nixl.cuda_malloc_memory import ( @@ -96,9 +97,24 @@ def init_broadcaster( timeout: int, quantize_in_weight_transfer: bool = False, session_id: str = "default", + engine_world_size: int | None = None, ) -> None: - del inference_world_size, quantize_in_weight_transfer - global_rank = rank_offset + self.device.index + del quantize_in_weight_transfer + if engine_world_size is None: + global_rank = rank_offset + self.device.index + else: + parallel_config = self.parallel_config + global_rank = global_inference_rank( + rank_offset=rank_offset, + data_parallel_index=parallel_config.data_parallel_index, + data_parallel_size=parallel_config.data_parallel_size, + worker_rank=self.rank, + tensor_parallel_size=parallel_config.tensor_parallel_size, + pipeline_parallel_size=parallel_config.pipeline_parallel_size, + prefill_context_parallel_size=parallel_config.prefill_context_parallel_size, + inference_world_size=inference_world_size, + engine_world_size=engine_world_size, + ) server_url = f"{host}:{port}" set_ucx_env_defaults() self.nixl_agent = NixlAgent(make_agent_name("inference", global_rank)) diff --git a/src/prime_rl/utils/client.py b/src/prime_rl/utils/client.py index 66aadb4f29..d660b497cb 100644 --- a/src/prime_rl/utils/client.py +++ b/src/prime_rl/utils/client.py @@ -395,6 +395,8 @@ async def init_nccl_broadcast( timeout: int, inference_world_size: int | None = None, quantize_in_weight_transfer: bool = False, + *, + engine_world_sizes: list[int] | None = None, ) -> None: """Initialize NCCL broadcast on all inference servers. @@ -404,32 +406,46 @@ async def init_nccl_broadcast( """ logger = get_logger() + has_explicit_engine_world_sizes = engine_world_sizes is not None if inference_world_size is None: - inference_world_size = len(admin_clients) + if engine_world_sizes is not None: + inference_world_size = sum(engine_world_sizes) + else: + inference_world_size = len(admin_clients) logger.warning( f"inference_world_size not provided, defaulting to {inference_world_size} (one GPU per admin client)" ) - gpus_per_server = inference_world_size // len(admin_clients) + if engine_world_sizes is None: + if inference_world_size % len(admin_clients) != 0: + raise ValueError("inference_world_size must be divisible by the number of admin clients") + engine_world_sizes = [inference_world_size // len(admin_clients)] * len(admin_clients) + if len(engine_world_sizes) != len(admin_clients): + raise ValueError("one engine world size is required for each admin client") + rank_offsets = _rank_offsets(engine_world_sizes, inference_world_size) logger.info( f"Initializing NCCL broadcast: {len(admin_clients)} servers, " - f"inference_world_size={inference_world_size}, gpus_per_server={gpus_per_server}" + f"inference_world_size={inference_world_size}, engine_world_sizes={engine_world_sizes}" ) - async def _init_nccl_broadcast(admin_client: AsyncClient, rank_offset: int) -> None: + async def _init_nccl_broadcast( + admin_client: AsyncClient, + rank_offset: int, + engine_world_size: int, + ) -> None: + payload = { + "host": host, + "port": port, + "rank_offset": rank_offset, + "inference_world_size": inference_world_size, + "timeout": timeout, + "quantize_in_weight_transfer": quantize_in_weight_transfer, + } + if has_explicit_engine_world_sizes: + payload["engine_world_size"] = engine_world_size try: - response = await admin_client.post( - "/init_broadcaster", - json={ - "host": host, - "port": port, - "rank_offset": rank_offset, - "inference_world_size": inference_world_size, - "timeout": timeout, - "quantize_in_weight_transfer": quantize_in_weight_transfer, - }, - ) + response = await admin_client.post("/init_broadcaster", json=payload) response.raise_for_status() except httpx.HTTPStatusError as e: if e.response.status_code == 404: @@ -438,8 +454,10 @@ async def _init_nccl_broadcast(admin_client: AsyncClient, rank_offset: int) -> N await asyncio.gather( *[ - _init_nccl_broadcast(admin_client, client_num * gpus_per_server) - for client_num, admin_client in enumerate(admin_clients) + _init_nccl_broadcast(admin_client, rank_offset, engine_world_size) + for admin_client, rank_offset, engine_world_size in zip( + admin_clients, rank_offsets, engine_world_sizes, strict=True + ) ] ) @@ -451,31 +469,61 @@ async def init_nixl_broadcast( timeout: int, inference_world_size: int, session_id: str, + *, + engine_world_sizes: list[int] | None = None, ) -> None: """Configure every vLLM worker for NIXL + ModelExpress pulls.""" - workers_per_server = inference_world_size // len(admin_clients) - - async def initialize(admin_client: AsyncClient, rank_offset: int) -> None: + has_explicit_engine_world_sizes = engine_world_sizes is not None + if engine_world_sizes is None: + if inference_world_size % len(admin_clients) != 0: + raise ValueError("inference_world_size must be divisible by the number of admin clients") + engine_world_sizes = [inference_world_size // len(admin_clients)] * len(admin_clients) + if len(engine_world_sizes) != len(admin_clients): + raise ValueError("one engine world size is required for each admin client") + rank_offsets = _rank_offsets(engine_world_sizes, inference_world_size) + + async def initialize(admin_client: AsyncClient, rank_offset: int, engine_world_size: int) -> None: + payload = { + "host": host, + "port": port, + "rank_offset": rank_offset, + "inference_world_size": inference_world_size, + "timeout": timeout, + "quantize_in_weight_transfer": False, + "session_id": session_id, + } + if has_explicit_engine_world_sizes: + payload["engine_world_size"] = engine_world_size await _admin_post( admin_client, "/init_broadcaster", timeout_s=max(ADMIN_TIMEOUT_S, timeout), - json={ - "host": host, - "port": port, - "rank_offset": rank_offset, - "inference_world_size": inference_world_size, - "timeout": timeout, - "quantize_in_weight_transfer": False, - "session_id": session_id, - }, + json=payload, ) await asyncio.gather( - *[initialize(admin_client, index * workers_per_server) for index, admin_client in enumerate(admin_clients)] + *[ + initialize(admin_client, rank_offset, engine_world_size) + for admin_client, rank_offset, engine_world_size in zip( + admin_clients, rank_offsets, engine_world_sizes, strict=True + ) + ] ) +def _rank_offsets(engine_world_sizes: list[int], inference_world_size: int) -> list[int]: + if not engine_world_sizes or any(isinstance(size, bool) or size <= 0 for size in engine_world_sizes): + raise ValueError("engine world sizes must be positive integers") + if sum(engine_world_sizes) != inference_world_size: + raise ValueError("engine world sizes do not match inference_world_size") + offsets: list[int] = [] + offset = 0 + for world_size in engine_world_sizes: + offsets.append(offset) + offset += world_size + return offsets + + async def prefill_logprobs(openai: AsyncOpenAI, model: str, token_ids: list[int]) -> list[float]: """Prefill-score ``token_ids`` under ``model`` via ``/inference/v1/generate`` + ``prompt_logprobs`` (the prime-rl server-side extension in diff --git a/tests/unit/inference/test_nccl_rank.py b/tests/unit/inference/test_nccl_rank.py new file mode 100644 index 0000000000..dc2955814d --- /dev/null +++ b/tests/unit/inference/test_nccl_rank.py @@ -0,0 +1,90 @@ +import pytest + +from prime_rl.inference.vllm.ranks import global_inference_rank + + +def test_dense_aggregate_uses_engine_size_when_vllm_rewrites_dp_size(): + actual = { + global_inference_rank( + rank_offset=0, + data_parallel_index=dp_index, + data_parallel_size=1, + worker_rank=tp_rank, + tensor_parallel_size=2, + pipeline_parallel_size=1, + inference_world_size=8, + engine_world_size=8, + ) + for dp_index in range(4) + for tp_rank in range(2) + } + + assert actual == set(range(8)) + + +def test_already_global_moe_worker_ranks_are_not_double_counted(): + actual = { + global_inference_rank( + rank_offset=0, + data_parallel_index=dp_index, + data_parallel_size=4, + worker_rank=dp_index * 2 + tp_rank, + tensor_parallel_size=2, + pipeline_parallel_size=1, + inference_world_size=8, + ) + for dp_index in range(4) + for tp_rank in range(2) + } + + assert actual == set(range(8)) + + +def test_rank_offset_composes_with_pipeline_parallel_rank(): + actual = { + global_inference_rank( + rank_offset=8, + data_parallel_index=dp_index, + data_parallel_size=2, + worker_rank=model_parallel_rank, + tensor_parallel_size=2, + pipeline_parallel_size=2, + inference_world_size=16, + ) + for dp_index in range(2) + for model_parallel_rank in range(4) + } + + assert actual == set(range(8, 16)) + + +def test_global_inference_rank_rejects_out_of_bounds_rank(): + with pytest.raises(ValueError, match="outside inference world size"): + global_inference_rank( + rank_offset=2, + data_parallel_index=3, + data_parallel_size=4, + worker_rank=0, + tensor_parallel_size=2, + pipeline_parallel_size=1, + inference_world_size=8, + ) + + +def test_prefill_context_parallel_ranks_are_unique(): + actual = { + global_inference_rank( + rank_offset=0, + data_parallel_index=0, + data_parallel_size=1, + worker_rank=worker_rank, + tensor_parallel_size=1, + pipeline_parallel_size=1, + prefill_context_parallel_size=2, + inference_world_size=2, + engine_world_size=2, + ) + for worker_rank in range(2) + } + + assert actual == {0, 1} diff --git a/tests/unit/utils/test_external_engine_topology.py b/tests/unit/utils/test_external_engine_topology.py new file mode 100644 index 0000000000..591bd8c78a --- /dev/null +++ b/tests/unit/utils/test_external_engine_topology.py @@ -0,0 +1,44 @@ +import asyncio +from unittest.mock import AsyncMock, MagicMock + +from prime_rl.utils.client import _rank_offsets, init_nccl_broadcast + + +def test_rank_offsets_support_heterogeneous_engines(): + assert _rank_offsets([1, 3, 2], inference_world_size=6) == [0, 1, 4] + + +def test_nccl_broadcast_forwards_explicit_engine_sizes(): + clients = [AsyncMock(), AsyncMock()] + for client in clients: + response = MagicMock() + response.raise_for_status = MagicMock() + client.post.return_value = response + + asyncio.run( + init_nccl_broadcast( + clients, + host="127.0.0.1", + port=29519, + timeout=1200, + inference_world_size=4, + engine_world_sizes=[1, 3], + ) + ) + + assert [call.kwargs["json"]["rank_offset"] for client in clients for call in client.post.await_args_list] == [0, 1] + assert [call.kwargs["json"]["engine_world_size"] for client in clients for call in client.post.await_args_list] == [ + 1, + 3, + ] + + +def test_nccl_broadcast_preserves_legacy_payload(): + client = AsyncMock() + response = MagicMock() + response.raise_for_status = MagicMock() + client.post.return_value = response + + asyncio.run(init_nccl_broadcast([client], "127.0.0.1", 29519, 1200, 1)) + + assert "engine_world_size" not in client.post.await_args.kwargs["json"] From 76907a2cfb0d003eed731d15c726ced9b9ac96ea Mon Sep 17 00:00:00 2001 From: Biswa Panda Date: Mon, 3 Aug 2026 11:29:26 -0700 Subject: [PATCH 4/9] refactor(vllm): delegate token serving to vllm 0.26 --- src/prime_rl/inference/vllm/routed_experts.py | 22 +- src/prime_rl/inference/vllm/server.py | 25 +- src/prime_rl/inference/vllm/serving_tokens.py | 347 ++---------------- tests/unit/inference/test_serving_tokens.py | 327 +++++------------ 4 files changed, 149 insertions(+), 572 deletions(-) diff --git a/src/prime_rl/inference/vllm/routed_experts.py b/src/prime_rl/inference/vllm/routed_experts.py index 5f73512fe5..a32d93ded5 100644 --- a/src/prime_rl/inference/vllm/routed_experts.py +++ b/src/prime_rl/inference/vllm/routed_experts.py @@ -1,11 +1,10 @@ from __future__ import annotations -from collections.abc import AsyncIterator +import io from typing import Any import numpy as np import pybase64 -from vllm.outputs import RequestOutput def serialize_routed_experts(routed_experts: Any, start: int = 0) -> dict[str, Any] | None: @@ -33,16 +32,9 @@ def serialize_routed_experts(routed_experts: Any, start: int = 0) -> dict[str, A } -class RoutedExpertsCapture: - def __init__(self, generator: AsyncIterator[RequestOutput], start: int = 0): - self._generator = generator - self._start = start - self.routed_experts: dict[int, dict[str, Any]] = {} - - async def __aiter__(self): - async for request_output in self._generator: - for output in request_output.outputs: - encoded = serialize_routed_experts(getattr(output, "routed_experts", None), start=self._start) - if encoded is not None: - self.routed_experts[output.index] = encoded - yield request_output +def compact_vllm_routed_experts(encoded: str | None, start: int = 0) -> dict[str, Any] | None: + """Convert vLLM's base64 ``.npy`` payload to Prime's compact payload.""" + if encoded is None: + return None + array = np.load(io.BytesIO(pybase64.b64decode(encoded)), allow_pickle=False) + return serialize_routed_experts(array, start=start) diff --git a/src/prime_rl/inference/vllm/server.py b/src/prime_rl/inference/vllm/server.py index 984ed47eed..c1dd10274b 100644 --- a/src/prime_rl/inference/vllm/server.py +++ b/src/prime_rl/inference/vllm/server.py @@ -163,28 +163,25 @@ async def custom_init_app_state( args: Namespace, supported_tasks: tuple, ): - """ - Modifies init_app_state: - 1. Call the original init_app_state to set up standard state, including - vLLM 0.20's ``serving_tokens`` for ``/inference/v1/generate``. - 2. Replace ``serving_tokens`` with ``PrimeRlServingTokens`` so DP-rank - routing and ``routed_experts`` export survive the migration off the - legacy ``/v1/generate`` endpoint. - """ + """Initialize vLLM app state and install Prime's token response adapter.""" await init_app_state(engine_client, state, args, supported_tasks) state.liveness_timeout_seconds = args.liveness_timeout_seconds - # Swap in our ServingTokens subclass for /inference/v1/generate so the - # X-data-parallel-rank header and routed_experts response field — both - # used by prime-RL's renderer / router-replay paths — keep working. if "generate" in supported_tasks and state.serving_tokens is not None: from prime_rl.inference.vllm.serving_tokens import PrimeRlServingTokens upstream = state.serving_tokens - prime_serving = object.__new__(PrimeRlServingTokens) - prime_serving.__dict__.update(upstream.__dict__) - state.serving_tokens = prime_serving + state.serving_tokens = PrimeRlServingTokens( + upstream.engine_client, + upstream.models, + upstream.online_renderer, + request_logger=upstream.request_logger, + return_tokens_as_token_ids=upstream.return_tokens_as_token_ids, + force_no_detokenize=upstream.force_no_detokenize, + enable_prompt_tokens_details=True, + enable_log_outputs=upstream.enable_log_outputs, + ) import vllm.entrypoints.openai.api_server diff --git a/src/prime_rl/inference/vllm/serving_tokens.py b/src/prime_rl/inference/vllm/serving_tokens.py index e14a5ac83e..77caea00ce 100644 --- a/src/prime_rl/inference/vllm/serving_tokens.py +++ b/src/prime_rl/inference/vllm/serving_tokens.py @@ -1,53 +1,21 @@ -"""Prime-RL extensions to vLLM's `/inference/v1/generate` handler. - -vLLM ships a generic tokens-in / tokens-out handler at -``vllm.entrypoints.scale_out.token_in_token_out.serving.ServingTokens`` that covers -prefix-cache salting, lora dispatch, multimodal features, prompt logprobs, -priority, ``data_parallel_rank`` header routing and server-side ``max_tokens`` -defaulting. We subclass it for the bits still missing from the upstream handler: - -1. ``data_parallel_rank`` routing — read from the ``X-data-parallel-rank`` - header and forwarded to ``engine_client.generate``. Upstream ``ServingTokens`` - now does this too; we keep the equivalent path for the DP-replicated - inference servers prime-RL runs. - -2. Compact ``routed_experts`` export — when the engine emits routing - decisions, surface them as base64 raw-byte payloads without requiring a vLLM - source fork. - -3. Server-side ``max_tokens`` defaulting — upstream ``ServingTokens`` now applies - this itself (via ``GenerateRequest.is_sampling_param_provided`` + - ``get_max_tokens``); we keep an equivalent guard so callers that omit - ``max_tokens`` don't truncate at vLLM's 16-token ``SamplingParams`` default. - -Everything else (request/response schema, sampling params, error handling) -delegates to upstream so we track future vLLM changes for free. -""" +"""Small Prime extensions to vLLM's canonical token-in/token-out handler.""" from __future__ import annotations -from collections.abc import AsyncGenerator, AsyncIterable -from functools import cached_property +from collections.abc import AsyncGenerator from typing import Any from fastapi import Request -from vllm.entrypoints.openai.engine.protocol import ( - ErrorResponse, - PromptTokenUsageInfo, - RequestResponseMetadata, - UsageInfo, -) +from vllm.entrypoints.openai.engine.protocol import ErrorResponse, RequestResponseMetadata from vllm.entrypoints.scale_out.token_in_token_out.protocol import ( GenerateRequest, GenerateResponse, GenerateResponseChoice, ) from vllm.entrypoints.scale_out.token_in_token_out.serving import ServingTokens -from vllm.entrypoints.serve.utils.api_utils import get_max_tokens from vllm.outputs import RequestOutput -from vllm.sampling_params import RequestOutputKind, SamplingParams -from prime_rl.inference.vllm.routed_experts import RoutedExpertsCapture +from prime_rl.inference.vllm.routed_experts import compact_vllm_routed_experts class PrimeRlGenerateResponseChoice(GenerateResponseChoice): @@ -56,255 +24,26 @@ class PrimeRlGenerateResponseChoice(GenerateResponseChoice): class PrimeRlGenerateResponse(GenerateResponse): choices: list[PrimeRlGenerateResponseChoice] - # Upstream ``GenerateResponse`` doesn't declare a ``usage`` field, so the - # parent ``ServingTokens.serve_tokens_full_generator`` constructs it and - # Pydantic silently drops it on serialization. Declare it here so the - # router can extract per-run token counts (and cached-prefix tokens) for - # platform billing — see https://github.com/PrimeIntellect-ai/router/pull/43. - usage: UsageInfo | None = None - - -class _GenerateRoutedExpertsCapture(RoutedExpertsCapture): - def post_process(self, response: GenerateResponse) -> PrimeRlGenerateResponse: - choices = [ - PrimeRlGenerateResponseChoice( - **choice.model_dump(exclude={"routed_experts"}), - routed_experts=self.routed_experts.get(choice.index), - ) - for choice in response.choices - ] - return PrimeRlGenerateResponse( - request_id=response.request_id, - choices=choices, - prompt_logprobs=response.prompt_logprobs, - kv_transfer_params=response.kv_transfer_params, - ) - - -class _FinalOutputCapture: - """Wraps a ``RequestOutput`` async generator to record the last yielded item. - - Needed so the response builder can construct a ``usage`` block from - ``final_res.prompt_token_ids`` / ``output.token_ids`` / ``num_cached_tokens`` - after delegating iteration to upstream. - """ - - def __init__(self, source: AsyncIterable[RequestOutput]) -> None: - # ``source`` may be any async-iterable — including - # ``_GenerateRoutedExpertsCapture``, which exposes the protocol via - # ``async def __aiter__`` (an async generator function) and has no - # ``__anext__`` method. Drive it through ``async for`` rather than - # poking ``__anext__`` directly so both shapes work. - self._source = source - self.final_res: RequestOutput | None = None - - async def __aiter__(self) -> AsyncGenerator[RequestOutput, None]: - async for item in self._source: - self.final_res = item - yield item - - -def _build_usage(final_res: RequestOutput) -> UsageInfo: - assert final_res.prompt_token_ids is not None - num_prompt_tokens = len(final_res.prompt_token_ids) - if final_res.encoder_prompt_token_ids is not None: - num_prompt_tokens += len(final_res.encoder_prompt_token_ids) - num_generated_tokens = sum(len(output.token_ids) for output in final_res.outputs) - usage = UsageInfo( - prompt_tokens=num_prompt_tokens, - completion_tokens=num_generated_tokens, - total_tokens=num_prompt_tokens + num_generated_tokens, - ) - # Always emit cached tokens when vLLM reports any. Upstream gates this on - # ``enable_prompt_tokens_details`` (default False) for OpenAI-API compat, - # but ``/inference/v1/generate`` is prime-rl internal — the cache-discount - # billing pipeline always wants the cached subset surfaced. - if final_res.num_cached_tokens: - usage.prompt_tokens_details = PromptTokenUsageInfo(cached_tokens=final_res.num_cached_tokens) - return usage - - -async def _client_set_max_tokens(raw_request: Request | None) -> bool: - """Whether the inbound JSON body carried ``sampling_params.max_tokens``. - - ``GenerateRequest.sampling_params`` is parsed into a ``SamplingParams`` - instance, which means an unset ``max_tokens`` is indistinguishable from - an explicit ``max_tokens=16`` once the request reaches the handler — - both surface as ``sampling_params.max_tokens == 16``. We re-read the - cached body to recover that distinction. When we can't (no raw_request, - non-JSON body, or read error), pessimistically assume the client did - set it so we never clobber an explicit value. - """ - if raw_request is None: - return True - try: - body = await raw_request.json() - except Exception: - return True - if not isinstance(body, dict): - return True - sp = body.get("sampling_params") - return isinstance(sp, dict) and "max_tokens" in sp class PrimeRlServingTokens(ServingTokens): - """ServingTokens + DP-rank routing + compact routed experts + max_tokens defaulting.""" - - @cached_property - def _max_tokens_defaults(self) -> tuple[dict, int | None]: - """Server-side ``max_tokens`` defaulting inputs, mirroring upstream ``ServingTokens``. - - Computed lazily because ``custom_init_app_state`` swaps in this - subclass via ``object.__new__`` + ``__dict__.update`` (so our - ``__init__`` never runs). - """ - diff = self.model_config.get_diff_sampling_param() - mc = self.model_config - override = ( - diff.get("max_tokens") - if mc.generation_config not in ("auto", "vllm") - # Upstream uses ``getattr(..., {})`` directly. Defensive ``or {}`` - # in case a downstream caller ever sets the attribute to ``None`` - # (``getattr``'s default only fires when the attribute is missing, - # not when it exists with a ``None`` value). - else (getattr(mc, "override_generation_config", None) or {}).get("max_new_tokens") - ) - return diff, override + """Add KV handoff and Prime's compact routed-expert response encoding.""" async def serve_tokens( self, request: GenerateRequest, raw_request: Request | None = None, - ) -> PrimeRlGenerateResponse | ErrorResponse | AsyncGenerator[str, None]: - # Mirrors upstream ``ServingTokens.serve_tokens``. Diffs: - # (a) inject ``data_parallel_rank`` from the inbound header into - # ``engine_client.generate``; (b) default ``sampling_params.max_tokens`` - # to ``max_model_len - prompt_len`` when the caller didn't set it; and - # (c) dispatch to our overridden response builder so ``routed_experts`` - # makes it into the JSON. - error_check_ret = await self._check_model(request) - if error_check_ret is not None: - return error_check_ret - - if self.engine_client.errored: - raise self.engine_client.dead_error - - lora_request = self._maybe_get_adapters(request, supports_default_mm_loras=True) - model_name = self.models.model_name(lora_request) + ) -> GenerateResponse | ErrorResponse | AsyncGenerator[str, None]: + if request.kv_transfer_params is None: + return await super().serve_tokens(request, raw_request) - request_id = f"generate-tokens-{self._base_request_id(raw_request, request.request_id)}" - request_metadata = RequestResponseMetadata(request_id=request_id) - if raw_request: - raw_request.state.request_metadata = request_metadata + forwarded = request.model_copy(deep=True) + extra_args = dict(forwarded.sampling_params.extra_args or {}) + extra_args["kv_transfer_params"] = forwarded.kv_transfer_params + forwarded.sampling_params.extra_args = extra_args + return await super().serve_tokens(forwarded, raw_request) - # Build the engine input — features-aware (MM) or text-only fallback. - # Identical to upstream so we keep tracking it. - if features := request.features: - from vllm.entrypoints.scale_out.token_in_token_out.mm_serde import decode_mm_kwargs_item - from vllm.inputs import mm_input - from vllm.multimodal.inputs import ( - MultiModalKwargsItem, - MultiModalKwargsItems, - PlaceholderRange, - ) - - mm_placeholders = { - modality: [PlaceholderRange(offset=p.offset, length=p.length) for p in ranges] - for modality, ranges in features.mm_placeholders.items() - } - mm_kwargs: dict[str, list[MultiModalKwargsItem | None]] = {} - if features.kwargs_data is not None: - for modality, items in features.kwargs_data.items(): - mm_kwargs[modality] = [decode_mm_kwargs_item(item) if item is not None else None for item in items] - else: - for modality, hashes in features.mm_hashes.items(): - mm_kwargs[modality] = [None] * len(hashes) - engine_input = mm_input( - prompt_token_ids=request.token_ids, - mm_kwargs=MultiModalKwargsItems(mm_kwargs), - mm_hashes=features.mm_hashes, - mm_placeholders=mm_placeholders, - cache_salt=request.cache_salt, - ) - else: - (engine_input,) = await self.online_renderer.preprocess_completion( - request, - prompt_input=request.token_ids, - prompt_embeds=None, - skip_mm_cache=True, - ) - - sampling_params: SamplingParams = request.sampling_params - - # Upstream ``ServingTokens.serve_tokens`` parses ``request.kv_transfer_params`` - # but never threads it into the engine, so PD disagg never fires on - # ``/inference/v1/generate`` — decode receives an empty NIXL handshake - # and ends up re-prefilling the prompt locally (~100× slower under - # concurrency). Bridge it through ``sampling_params.extra_args`` so the - # engine's KV connector picks the params up. - # - # Upstream fix: https://github.com/vllm-project/vllm/pull/42644 — drop - # this block once we pin a vLLM version that includes it. - if request.kv_transfer_params is not None: - extra = sampling_params.extra_args or {} - extra["kv_transfer_params"] = request.kv_transfer_params - sampling_params.extra_args = extra - - # Server-side ``max_tokens`` defaulting — see module docstring. Upstream - # ``ServingTokens`` now does this too; kept here so callers that omit - # ``max_tokens`` don't get capped at vLLM's 16-token ``SamplingParams`` - # default. - if not await _client_set_max_tokens(raw_request): - diff_sp, override = self._max_tokens_defaults - sampling_params.max_tokens = get_max_tokens( - max_model_len=self.model_config.max_model_len, - max_tokens=None, - input_length=len(request.token_ids), - default_sampling_params=diff_sp, - override_max_tokens=override, - ) - - if self.force_no_detokenize: - sampling_params.detokenize = False - if request.stream: - sampling_params.output_kind = RequestOutputKind.DELTA - - self._log_inputs( - request_id, - engine_input, - params=sampling_params, - lora_request=lora_request, - ) - - trace_headers = None if raw_request is None else await self._get_trace_headers(raw_request.headers) - - result_generator = self.engine_client.generate( - engine_input, - sampling_params, - request_id, - lora_request=lora_request, - trace_headers=trace_headers, - priority=request.priority, - data_parallel_rank=self._get_data_parallel_rank(raw_request), - ) - - if request.stream: - # Streaming path: defer to upstream — prime-RL's renderer client - # only consumes the full response, so adding routed_experts to the - # streaming choice schema is unnecessary churn. - return self.serve_tokens_stream_generator( - request, - result_generator, - request_id, - model_name, - request_metadata, - ) - - return await self.serve_tokens_full_generator( - request, result_generator, request_id, model_name, request_metadata - ) - - async def serve_tokens_full_generator( # type: ignore[override] + async def serve_tokens_full_generator( self, request: GenerateRequest, result_generator: AsyncGenerator[RequestOutput, None], @@ -312,45 +51,25 @@ async def serve_tokens_full_generator( # type: ignore[override] model_name: str, request_metadata: RequestResponseMetadata, ) -> ErrorResponse | GenerateResponse: - # Capture routed_experts as vLLM streams request outputs, then post-process - # the final response into our GenerateResponse subclass so the encoded - # experts surface in the JSON. - capture: _GenerateRoutedExpertsCapture | None = None - if self.model_config.enable_return_routed_experts: - start = request.sampling_params.routed_experts_prompt_start - capture = _GenerateRoutedExpertsCapture( - result_generator, - start=start, - ) - result_generator = capture - - # Always capture the final ``RequestOutput`` so we can attach a - # ``usage`` block to the response. The router parses ``usage`` for - # per-run billing metrics; without it the cache-discount counter - # (``vllm_router_run_cached_prompt_tokens_total``) stays at zero. - final_capture = _FinalOutputCapture(result_generator) - result_generator = final_capture - response = await super().serve_tokens_full_generator( - request, result_generator, request_id, model_name, request_metadata + request, + result_generator, + request_id, + model_name, + request_metadata, ) - - if not isinstance(response, GenerateResponse): + if not isinstance(response, GenerateResponse) or not any( + choice.routed_experts is not None for choice in response.choices + ): return response - - if capture is not None: - response = capture.post_process(response) - elif not isinstance(response, PrimeRlGenerateResponse): - # Upgrade to the prime-rl subclass so the declared ``usage`` field - # actually surfaces in JSON (the parent class would drop it). - response = PrimeRlGenerateResponse( - request_id=response.request_id, - choices=[PrimeRlGenerateResponseChoice(**choice.model_dump()) for choice in response.choices], - prompt_logprobs=response.prompt_logprobs, - kv_transfer_params=response.kv_transfer_params, - ) - - if final_capture.final_res is not None: - response.usage = _build_usage(final_capture.final_res) - - return response + start = request.sampling_params.routed_experts_prompt_start or 0 + return PrimeRlGenerateResponse( + **response.model_dump(exclude={"choices"}), + choices=[ + PrimeRlGenerateResponseChoice( + **choice.model_dump(exclude={"routed_experts"}), + routed_experts=compact_vllm_routed_experts(choice.routed_experts, start=start), + ) + for choice in response.choices + ], + ) diff --git a/tests/unit/inference/test_serving_tokens.py b/tests/unit/inference/test_serving_tokens.py index 951ef5c9d8..71c90e0377 100644 --- a/tests/unit/inference/test_serving_tokens.py +++ b/tests/unit/inference/test_serving_tokens.py @@ -1,66 +1,40 @@ -"""Sanity tests for the prime-RL ``ServingTokens`` subclass. - -The full happy-path is owned upstream by vLLM's -``vllm/entrypoints/serve/disagg`` test suite. We only cover the prime-RL -deltas here: - * ``serialize_routed_experts`` round-trips a compact raw-byte payload. - * The subclass attaches its overrides without monkey-patching the parent. - * ``_client_set_max_tokens`` distinguishes raw-body shapes correctly. -""" - from __future__ import annotations import asyncio +import io import numpy as np import pybase64 -from vllm.entrypoints.openai.engine.protocol import UsageInfo -from vllm.entrypoints.scale_out.token_in_token_out.protocol import GenerateResponse, GenerateResponseChoice +from vllm.entrypoints.openai.engine.protocol import RequestResponseMetadata, UsageInfo +from vllm.entrypoints.scale_out.token_in_token_out.protocol import ( + GenerateRequest, + GenerateResponse, + GenerateResponseChoice, +) +from vllm.entrypoints.scale_out.token_in_token_out.serving import ServingTokens +from vllm.sampling_params import SamplingParams -from prime_rl.inference.vllm.routed_experts import serialize_routed_experts -from prime_rl.inference.vllm.serving_tokens import ( - PrimeRlGenerateResponse, - PrimeRlGenerateResponseChoice, - PrimeRlServingTokens, - _build_usage, - _client_set_max_tokens, - _FinalOutputCapture, - _GenerateRoutedExpertsCapture, +from prime_rl.inference.vllm.routed_experts import ( + compact_vllm_routed_experts, + serialize_routed_experts, ) +from prime_rl.inference.vllm.serving_tokens import PrimeRlServingTokens -def _decode_routed_experts(encoded: dict) -> np.ndarray: +def _decode_compact(encoded: dict) -> np.ndarray: return np.frombuffer( pybase64.b64decode_as_bytearray(encoded["data"]), dtype=np.uint8, ).reshape(encoded["shape"]) -class _FakeRawRequest: - def __init__(self, body): - self._body = body - self._raise = isinstance(body, Exception) - - async def json(self): - if self._raise: - raise self._body - return self._body +def _encode_vllm(array: np.ndarray) -> str: + buffer = io.BytesIO() + np.save(buffer, array) + return pybase64.b64encode(buffer.getvalue()).decode("ascii") -async def _empty_request_outputs(): - if False: - yield - - -def test_subclass_only_overrides_serve_tokens(): - assert PrimeRlServingTokens.serve_tokens is not PrimeRlServingTokens.__mro__[1].serve_tokens - assert ( - PrimeRlServingTokens.serve_tokens_full_generator - is not PrimeRlServingTokens.__mro__[1].serve_tokens_full_generator - ) - - -def test_serialize_routed_experts_uses_compact_raw_payload(): +def test_routed_experts_round_trip_both_wire_formats(): routed_experts = np.array( [ [[1, 2], [3, 4]], @@ -69,201 +43,96 @@ def test_serialize_routed_experts_uses_compact_raw_payload(): dtype=np.int64, ) - encoded = serialize_routed_experts(routed_experts) - assert encoded is not None + compact = serialize_routed_experts(routed_experts, start=2) + converted = compact_vllm_routed_experts(_encode_vllm(routed_experts), start=2) + + assert compact is not None + assert converted is not None + assert converted["start"] == 2 + np.testing.assert_array_equal(_decode_compact(compact), routed_experts) + np.testing.assert_array_equal(_decode_compact(converted), routed_experts) + + +def test_serve_tokens_forwards_kv_transfer_params_without_mutating_request(monkeypatch): + expected = object() + observed_request = None + + async def upstream(_self, request, _raw_request=None): + nonlocal observed_request + observed_request = request + assert request.sampling_params.extra_args == { + "existing": True, + "kv_transfer_params": {"remote": "metadata"}, + } + return expected + + monkeypatch.setattr(ServingTokens, "serve_tokens", upstream) + server = object.__new__(PrimeRlServingTokens) + request = GenerateRequest( + token_ids=[1], + sampling_params=SamplingParams(max_tokens=1, extra_args={"existing": True}), + kv_transfer_params={"remote": "metadata"}, + ) - decoded = _decode_routed_experts(encoded) - assert decoded.dtype == np.uint8 - np.testing.assert_array_equal(decoded, routed_experts) + assert asyncio.run(server.serve_tokens(request)) is expected + assert observed_request is not request + assert request.sampling_params.extra_args == {"existing": True} -def test_generate_response_post_process_replaces_upstream_routed_experts(): - compact_routed_experts = {"data": "AQID", "shape": [1, 1, 3], "start": 0} - capture = _GenerateRoutedExpertsCapture(_empty_request_outputs()) - capture.routed_experts[0] = compact_routed_experts - response = GenerateResponse( - request_id="request-id", +def test_full_generator_preserves_all_upstream_response_fields(monkeypatch): + routed_experts = np.array([[[1, 2, 3]]], dtype=np.uint8) + usage = UsageInfo( + prompt_tokens=3, + completion_tokens=2, + total_tokens=5, + prompt_tokens_details={"cached_tokens": 2}, + ) + upstream_response = GenerateResponse( + request_id="canonical-request-id", + model="served-model", + created=123456789, + usage=usage, choices=[ GenerateResponseChoice( index=0, - token_ids=[1, 2, 3], - routed_experts="upstream-npy-payload", + token_ids=[10, 11], + routed_experts=_encode_vllm(routed_experts), ) ], ) - processed = capture.post_process(response) - - assert processed.choices[0].routed_experts == compact_routed_experts - - -def test_client_set_max_tokens_recognizes_explicit_value(): - body = {"token_ids": [1, 2, 3], "sampling_params": {"max_tokens": 256}} - assert asyncio.run(_client_set_max_tokens(_FakeRawRequest(body))) is True - - -def test_client_set_max_tokens_detects_unset(): - body = {"token_ids": [1, 2, 3], "sampling_params": {}} - assert asyncio.run(_client_set_max_tokens(_FakeRawRequest(body))) is False - - body_without_sp = {"token_ids": [1, 2, 3]} - assert asyncio.run(_client_set_max_tokens(_FakeRawRequest(body_without_sp))) is False - - -class _FakeOutput: - def __init__(self, token_ids): - self.token_ids = token_ids - - -class _FakeRequestOutput: - """Minimal stand-in for ``vllm.outputs.RequestOutput``. - - ``_build_usage`` only touches four attributes; constructing a real - ``RequestOutput`` would require a full ``CompletionOutput`` graph and - isn't worth it for a serialization-shape test. - """ - - def __init__(self, prompt_token_ids, output_token_ids_list, num_cached_tokens=0, encoder_prompt_token_ids=None): - self.prompt_token_ids = prompt_token_ids - self.encoder_prompt_token_ids = encoder_prompt_token_ids - self.outputs = [_FakeOutput(t) for t in output_token_ids_list] - self.num_cached_tokens = num_cached_tokens - - -def test_prime_rl_generate_response_serializes_usage_block(): - # Regression for prime-rl PR #2408: parent ``GenerateResponse`` doesn't - # declare ``usage``, so the field must be declared on the subclass for - # Pydantic to emit it in JSON. Without this the router can't extract - # per-run token / cache counts for billing. - response = PrimeRlGenerateResponse( - request_id="req-1", - choices=[PrimeRlGenerateResponseChoice(index=0, token_ids=[1, 2, 3])], - usage=UsageInfo(prompt_tokens=4, completion_tokens=3, total_tokens=7), - ) - payload = response.model_dump(mode="json") - assert payload["usage"] == { - "prompt_tokens": 4, - "completion_tokens": 3, - "total_tokens": 7, - "prompt_tokens_details": None, - } - - -def test_build_usage_sums_prompt_and_completion_tokens(): - final_res = _FakeRequestOutput( - prompt_token_ids=[1, 2, 3, 4, 5], - output_token_ids_list=[[10, 11], [20, 21, 22]], - ) - usage = _build_usage(final_res) - assert usage.prompt_tokens == 5 - assert usage.completion_tokens == 5 # 2 + 3 - assert usage.total_tokens == 10 - assert usage.prompt_tokens_details is None - - -def test_build_usage_includes_encoder_prompt_tokens(): - final_res = _FakeRequestOutput( - prompt_token_ids=[1, 2, 3], - output_token_ids_list=[[10]], - encoder_prompt_token_ids=[100, 101], - ) - usage = _build_usage(final_res) - assert usage.prompt_tokens == 5 # 3 + 2 - assert usage.total_tokens == 6 - - -def test_build_usage_reports_cached_tokens_unconditionally(): - # Unlike upstream's ``enable_prompt_tokens_details`` gate, prime-rl always - # surfaces cached tokens — the cache-discount billing pipeline needs them. - final_res = _FakeRequestOutput( - prompt_token_ids=[1, 2, 3, 4], - output_token_ids_list=[[10, 11]], - num_cached_tokens=3, + async def upstream(_self, _request, result_generator, *_args): + async for _ in result_generator: + pass + return upstream_response + + monkeypatch.setattr(ServingTokens, "serve_tokens_full_generator", upstream) + server = object.__new__(PrimeRlServingTokens) + server.enable_prompt_tokens_details = True + request = GenerateRequest( + token_ids=[1, 2, 3], + sampling_params=SamplingParams(max_tokens=2, routed_experts_prompt_start=1), ) - usage = _build_usage(final_res) - assert usage.prompt_tokens_details is not None - assert usage.prompt_tokens_details.cached_tokens == 3 - -def test_build_usage_skips_cached_tokens_when_zero(): - # Don't emit a details block with cached=0, which would be misleading - # to the router's billing extractor. - final_res = _FakeRequestOutput( - prompt_token_ids=[1, 2, 3, 4], - output_token_ids_list=[[10, 11]], - num_cached_tokens=0, + async def outputs(): + if False: + yield + + response = asyncio.run( + server.serve_tokens_full_generator( + request, + outputs(), + "input-request-id", + "input-model", + RequestResponseMetadata(request_id="input-request-id"), + ) ) - usage = _build_usage(final_res) - assert usage.prompt_tokens_details is None - - -def test_final_output_capture_records_last_item(): - async def _gen(): - for r in [ - _FakeRequestOutput(prompt_token_ids=[1], output_token_ids_list=[[1]]), - _FakeRequestOutput(prompt_token_ids=[1, 2], output_token_ids_list=[[1, 2]]), - _FakeRequestOutput(prompt_token_ids=[1, 2, 3], output_token_ids_list=[[1, 2, 3]]), - ]: - yield r - - async def _drain(capture): - async for _ in capture: - pass - - capture = _FinalOutputCapture(_gen()) - asyncio.run(_drain(capture)) - assert capture.final_res is not None - assert capture.final_res.prompt_token_ids == [1, 2, 3] - - -def test_final_output_capture_works_over_async_def_aiter_source(): - # ``_GenerateRoutedExpertsCapture`` exposes the async-iterator protocol - # via ``async def __aiter__`` (an async generator function) and has no - # ``__anext__``. The wrapper must drive it through ``async for`` rather - # than poking ``__anext__`` directly, or routed-experts runs raise - # AttributeError before the response is built. - - class _AsyncGenAiterSource: - def __init__(self, items): - self._items = items - - async def __aiter__(self): - for item in self._items: - yield item - - items = [ - _FakeRequestOutput(prompt_token_ids=[1], output_token_ids_list=[[1]]), - _FakeRequestOutput(prompt_token_ids=[1, 2], output_token_ids_list=[[1, 2]]), - ] - capture = _FinalOutputCapture(_AsyncGenAiterSource(items)) - - async def _drain(): - async for _ in capture: - pass - - asyncio.run(_drain()) - assert capture.final_res is not None - assert capture.final_res.prompt_token_ids == [1, 2] - - -def test_final_output_capture_handles_empty_stream(): - capture = _FinalOutputCapture(_empty_request_outputs()) - - async def _drain(): - async for _ in capture: - pass - - asyncio.run(_drain()) - assert capture.final_res is None - - -def test_client_set_max_tokens_assumes_set_when_body_unreadable(): - # No raw_request → can't tell, don't override. - assert asyncio.run(_client_set_max_tokens(None)) is True - - # body read raises → can't tell, don't override. - err = ValueError("bad json") - assert asyncio.run(_client_set_max_tokens(_FakeRawRequest(err))) is True - # non-dict body → can't tell, don't override. - assert asyncio.run(_client_set_max_tokens(_FakeRawRequest([1, 2, 3]))) is True + assert response.request_id == "canonical-request-id" + assert response.model == "served-model" + assert response.created == 123456789 + assert response.usage == usage + encoded = response.choices[0].routed_experts + assert isinstance(encoded, dict) + assert encoded["start"] == 1 + np.testing.assert_array_equal(_decode_compact(encoded), routed_experts) From 582539ffc18d979b27495ec329214cfc008b7d49 Mon Sep 17 00:00:00 2001 From: Biswa Panda Date: Mon, 3 Aug 2026 11:36:53 -0700 Subject: [PATCH 5/9] feat(inference): discover dynamo workers --- .../src/prime_rl/configs/orchestrator.py | 14 +++ .../src/prime_rl/configs/rl.py | 30 +++-- .../src/prime_rl/configs/shared.py | 15 ++- .../src/prime_rl/utils/validation.py | 3 + src/prime_rl/orchestrator/orchestrator.py | 4 + src/prime_rl/orchestrator/utils.py | 3 +- src/prime_rl/utils/client.py | 89 ++++++++++++-- src/prime_rl/utils/dynamo.py | 114 ++++++++++++++++++ .../orchestrator/test_orchestrator_setup.py | 4 + tests/unit/test_configs.py | 110 +++++++++++++++++ tests/unit/utils/test_dynamo.py | 103 ++++++++++++++++ tests/unit/utils/test_dynamo_inmemory.py | 41 +++++++ 12 files changed, 506 insertions(+), 24 deletions(-) create mode 100644 src/prime_rl/utils/dynamo.py create mode 100644 tests/unit/utils/test_dynamo.py create mode 100644 tests/unit/utils/test_dynamo_inmemory.py diff --git a/packages/prime-rl-configs/src/prime_rl/configs/orchestrator.py b/packages/prime-rl-configs/src/prime_rl/configs/orchestrator.py index 74af6c0de8..74270587d2 100644 --- a/packages/prime-rl-configs/src/prime_rl/configs/orchestrator.py +++ b/packages/prime-rl-configs/src/prime_rl/configs/orchestrator.py @@ -349,6 +349,9 @@ class ZeroAdvantageFilterConfig(BaseConfig): class FileSystemWeightBroadcastConfig(BaseConfig): type: Literal["filesystem"] = "filesystem" + inference_world_size: int | None = Field(None, ge=1) + """Expected inference ranks for Dynamo discovery completeness; unused by filesystem transfer itself.""" + class InMemoryWeightBroadcastConfig(BaseConfig): host: str = "localhost" @@ -520,6 +523,17 @@ def auto_setup_session_headers(self): self.model.client.extra_headers_from_state.setdefault("X-Session-ID", "trajectory_id") return self + @model_validator(mode="after") + def validate_dynamo_world_size(self): + if not self.model.client.is_dynamo: + return self + if ( + self.weight_broadcast.inference_world_size is None + or "inference_world_size" not in self.weight_broadcast.model_fields_set + ): + raise ValueError("Dynamo inference requires an explicit weight_broadcast.inference_world_size") + return self + @model_validator(mode="after") def auto_setup_prime_monitor_run_name(self): """Default ``prime_monitor.run_name`` to the W&B run name when monitoring diff --git a/packages/prime-rl-configs/src/prime_rl/configs/rl.py b/packages/prime-rl-configs/src/prime_rl/configs/rl.py index 6d70c3fae1..e8a46137fc 100644 --- a/packages/prime-rl-configs/src/prime_rl/configs/rl.py +++ b/packages/prime-rl-configs/src/prime_rl/configs/rl.py @@ -138,6 +138,9 @@ class SharedNCCLWeightBroadcastConfig(SharedInMemoryWeightBroadcastConfig): quantize_in_weight_transfer: bool = False """Use kernel-format FP8 quantized NCCL transfer for weight updates. When disabled, uses default HF checkpoint-format transfer.""" + inference_world_size: int | None = Field(None, ge=1) + """Expected number of externally managed inference ranks.""" + class SharedNIXLWeightBroadcastConfig(SharedInMemoryWeightBroadcastConfig): type: Literal["nixl"] = "nixl" @@ -148,10 +151,16 @@ class SharedNIXLWeightBroadcastConfig(SharedInMemoryWeightBroadcastConfig): session_id: str = "default" """ModelExpress session ID.""" + inference_world_size: int | None = Field(None, ge=1) + """Expected number of externally managed inference ranks.""" + class SharedFileSystemWeightBroadcastConfig(BaseConfig): type: Literal["filesystem"] = "filesystem" + inference_world_size: int | None = Field(None, ge=1) + """Expected number of externally managed inference ranks.""" + SharedWeightBroadcastConfig: TypeAlias = Annotated[ SharedFileSystemWeightBroadcastConfig | SharedNCCLWeightBroadcastConfig | SharedNIXLWeightBroadcastConfig, @@ -325,12 +334,12 @@ def validate_deployment(self): @model_validator(mode="after") def validate_enough_devices_for_nccl(self): - if self.deployment.type == "single_node": - if self.trainer.weight_broadcast.type == "nccl": - if self.deployment.num_train_gpus + self.deployment.num_infer_gpus < 2: - raise ValueError( - "NCCL weight broadcast requires at least 2 GPUs to build the broadcast process group." - ) + if self.deployment.type != "single_node" or self.trainer.weight_broadcast.type != "nccl": + return self + if self.inference is None and self.weight_broadcast.inference_world_size is not None: + return self + if self.deployment.num_train_gpus + self.deployment.num_infer_gpus < 2: + raise ValueError("NCCL weight broadcast requires at least 2 local GPUs or external inference ranks.") return self @model_validator(mode="after") @@ -396,14 +405,15 @@ def auto_setup_weight_broadcast(self): inference_world_size = ( self.inference.vllm.data_parallel_size * self.inference.vllm.tensor_parallel_size if self.inference - else 1 + else self.weight_broadcast.inference_world_size ) common_config = dict( host=self.weight_broadcast.host, port=self.weight_broadcast.port, timeout=self.weight_broadcast.timeout, - inference_world_size=inference_world_size, ) + if inference_world_size is not None: + common_config["inference_world_size"] = inference_world_size if self.weight_broadcast.type == "nccl": transport_config = dict( quantize_in_weight_transfer=self.weight_broadcast.quantize_in_weight_transfer, @@ -418,7 +428,9 @@ def auto_setup_weight_broadcast(self): self.orchestrator.weight_broadcast = orchestrator_config_type(**common_config, **transport_config) elif self.weight_broadcast.type == "filesystem": self.trainer.weight_broadcast = TrainerFileSystemWeightBroadcastConfig() - self.orchestrator.weight_broadcast = OrchestratorFileSystemWeightBroadcastConfig() + self.orchestrator.weight_broadcast = OrchestratorFileSystemWeightBroadcastConfig( + inference_world_size=self.weight_broadcast.inference_world_size + ) if self.inference is not None: self.inference.weight_broadcast = InferenceWeightBroadcastConfig(type=self.weight_broadcast.type) diff --git a/packages/prime-rl-configs/src/prime_rl/configs/shared.py b/packages/prime-rl-configs/src/prime_rl/configs/shared.py index b7e5d1ab52..3b7d419dfe 100644 --- a/packages/prime-rl-configs/src/prime_rl/configs/shared.py +++ b/packages/prime-rl-configs/src/prime_rl/configs/shared.py @@ -1,6 +1,6 @@ import os from pathlib import Path -from typing import Annotated, Literal, TypeAlias +from typing import Annotated, Literal, Self, TypeAlias from pydantic import AfterValidator, Field, model_validator @@ -133,6 +133,19 @@ class ClientConfig(BaseConfig): admin_base_url: list[str] | None = None """Separate base URLs for admin operations (weight updates, health checks). When set, admin clients bypass routers and hit each server directly — used in multi-replica or disaggregated P/D deployments where the router must not handle admin traffic.""" + dynamo_discovery_url: str | None = None + """Dynamo URL used to discover vLLM admin endpoints and per-engine world sizes.""" + + @model_validator(mode="after") + def validate_dynamo_discovery(self) -> Self: + if self.dynamo_discovery_url is not None and self.admin_base_url is not None: + raise ValueError("dynamo_discovery_url cannot be combined with admin_base_url") + return self + + @property + def is_dynamo(self) -> bool: + return self.dynamo_discovery_url is not None + class LogConfig(BaseConfig): level: str = Field(default_factory=lambda: os.environ.get("PRIME_LOG_LEVEL", "info")) diff --git a/packages/prime-rl-configs/src/prime_rl/utils/validation.py b/packages/prime-rl-configs/src/prime_rl/utils/validation.py index d1afff155d..5fc837c5f9 100644 --- a/packages/prime-rl-configs/src/prime_rl/utils/validation.py +++ b/packages/prime-rl-configs/src/prime_rl/utils/validation.py @@ -128,6 +128,9 @@ def propagate(shared_path: str, *targets: str) -> None: # [rollout_transport] → both sub-configs (host is launcher-injected for zmq multi-node). propagate("rollout_transport", "trainer.rollout_transport", "orchestrator.rollout_transport") + # The orchestrator validates external inference topology during construction. + propagate("weight_broadcast", "orchestrator.weight_broadcast") + # Top-level scalars. propagate("max_steps", "trainer.max_steps", "orchestrator.max_steps") propagate("seq_len", "trainer.model.seq_len", "orchestrator.seq_len") diff --git a/src/prime_rl/orchestrator/orchestrator.py b/src/prime_rl/orchestrator/orchestrator.py index 4a4d99a179..ab24e5504f 100644 --- a/src/prime_rl/orchestrator/orchestrator.py +++ b/src/prime_rl/orchestrator/orchestrator.py @@ -306,6 +306,8 @@ async def setup(self) -> None: config.weight_broadcast.timeout, inference_world_size=config.weight_broadcast.inference_world_size, quantize_in_weight_transfer=config.weight_broadcast.quantize_in_weight_transfer, + engine_world_sizes=self.policy_inference._engine_world_sizes, + use_native_collective_rpc=self.policy_inference._use_native_collective_rpc, ) elif config.weight_broadcast.type == "nixl": await init_nixl_broadcast( @@ -315,6 +317,8 @@ async def setup(self) -> None: config.weight_broadcast.timeout, config.weight_broadcast.inference_world_size, config.weight_broadcast.session_id, + engine_world_sizes=self.policy_inference._engine_world_sizes, + use_native_collective_rpc=self.policy_inference._use_native_collective_rpc, ) self.model_express = ModelExpressSession( client=MxClient(server_url=f"{config.weight_broadcast.host}:{config.weight_broadcast.port}"), diff --git a/src/prime_rl/orchestrator/utils.py b/src/prime_rl/orchestrator/utils.py index 0c32c31023..4e285a3619 100644 --- a/src/prime_rl/orchestrator/utils.py +++ b/src/prime_rl/orchestrator/utils.py @@ -39,12 +39,13 @@ async def setup_policy_inference_pool(*, config: OrchestratorConfig, tokenizer): get_logger().info("Using direct renderer rollout client") else: get_logger().info("No policy-sourced train env — renderer kept for client-side tokenization only") - inference_pool = InferencePool( + inference_pool = await InferencePool.create( client_config, model_name=model_name, train_client_type="renderer", eval_client_type="openai_chat_completions", renderer_config=config.renderer, + expected_inference_world_size=config.weight_broadcast.inference_world_size, ) return renderer, inference_pool diff --git a/src/prime_rl/utils/client.py b/src/prime_rl/utils/client.py index d660b497cb..f18551d6ce 100644 --- a/src/prime_rl/utils/client.py +++ b/src/prime_rl/utils/client.py @@ -50,6 +50,11 @@ def __init__( train_client_type: str = "openai_chat_completions", eval_client_type: str = "openai_chat_completions", renderer_config: RendererConfig | None = None, + expected_inference_world_size: int | None = None, + *, + admin_clients: list[AsyncClient] | None = None, + engine_world_sizes: list[int] | None = None, + use_native_collective_rpc: bool = False, ): renderer_model_name = model_name if train_client_type == "renderer" else None self.train_client = setup_client( @@ -59,7 +64,9 @@ def __init__( renderer_model_name=renderer_model_name, ) self.eval_client = setup_client(client_config, client_type=eval_client_type) - self._admin_clients = setup_admin_clients(client_config) + self._admin_clients = setup_admin_clients(client_config) if admin_clients is None else admin_clients + self._engine_world_sizes = engine_world_sizes + self._use_native_collective_rpc = use_native_collective_rpc # When admin URLs bypass a router, also health-check the client-facing # (router) endpoint - it only starts serving once its workers are healthy. self._router_clients = ( @@ -72,6 +79,32 @@ def __init__( self._scorer = PrefillScorer() self.model_name = model_name + @classmethod + async def create( + cls, + client_config: ClientConfig, + model_name: str, + *, + expected_inference_world_size: int | None = None, + **kwargs, + ) -> "InferencePool": + if not client_config.is_dynamo: + return cls(client_config, model_name, **kwargs) + if expected_inference_world_size is None: + raise ValueError("Dynamo inference requires an explicit inference_world_size") + from prime_rl.utils.dynamo import discover_dynamo_workers, setup_dynamo_admin_clients + + workers = await discover_dynamo_workers(client_config, model_name, expected_inference_world_size) + return cls( + client_config, + model_name, + expected_inference_world_size=expected_inference_world_size, + admin_clients=setup_dynamo_admin_clients(workers), + engine_world_sizes=[worker.world_size for worker in workers], + use_native_collective_rpc=True, + **kwargs, + ) + @property def admin_clients(self) -> list[AsyncClient]: return self._admin_clients @@ -87,7 +120,13 @@ async def wait_for_ready(self, model_name: str, timeout: int | None = None) -> N await maybe_check_has_model(self._admin_clients, model_name, skip_model_check=self._skip_model_check) async def update_weights(self, weight_dir: Path | None, lora_name: str | None = None, step: int = 0) -> None: - await update_weights(self._admin_clients, weight_dir, lora_name=lora_name, step=step) + await update_weights( + self._admin_clients, + weight_dir, + lora_name=lora_name, + step=step, + use_native_collective_rpc=self._use_native_collective_rpc, + ) async def score(self, token_ids: list[int]) -> list[float]: """Prefill-score ``token_ids`` under this pool's model (one logprob per @@ -276,6 +315,8 @@ async def update_weights( weight_dir: Path | None, lora_name: str | None = None, step: int = 0, + *, + use_native_collective_rpc: bool = False, ) -> None: """Update weights on static inference servers. @@ -305,12 +346,18 @@ async def update_weights( nccl_ready_file.touch() logger.debug(f"Created NCCL_READY marker at {nccl_ready_file}") + update_path = "/collective_rpc" if use_native_collective_rpc else "/update_weights" + payload = ( + {"method": "update_weights_from_path", "args": [weight_dir_posix]} + if use_native_collective_rpc + else {"weight_dir": weight_dir_posix} + ) await asyncio.gather( *[ _admin_post( admin_client, - "/update_weights", - json={"weight_dir": weight_dir_posix}, + update_path, + json=payload, timeout_s=UPDATE_WEIGHTS_TIMEOUT_S, ) for admin_client in admin_clients @@ -397,6 +444,7 @@ async def init_nccl_broadcast( quantize_in_weight_transfer: bool = False, *, engine_world_sizes: list[int] | None = None, + use_native_collective_rpc: bool = False, ) -> None: """Initialize NCCL broadcast on all inference servers. @@ -444,13 +492,20 @@ async def _init_nccl_broadcast( } if has_explicit_engine_world_sizes: payload["engine_world_size"] = engine_world_size - try: + if use_native_collective_rpc: + response = await admin_client.post( + "/collective_rpc", + json={"method": "init_broadcaster", "kwargs": payload}, + ) + else: response = await admin_client.post("/init_broadcaster", json=payload) + try: response.raise_for_status() - except httpx.HTTPStatusError as e: - if e.response.status_code == 404: + except httpx.HTTPStatusError as error: + if not use_native_collective_rpc and error.response.status_code == 404: logger.warning("The route /init_broadcaster does not exist. Skipping NCCL broadcast initialization.") return + raise await asyncio.gather( *[ @@ -471,6 +526,7 @@ async def init_nixl_broadcast( session_id: str, *, engine_world_sizes: list[int] | None = None, + use_native_collective_rpc: bool = False, ) -> None: """Configure every vLLM worker for NIXL + ModelExpress pulls.""" has_explicit_engine_world_sizes = engine_world_sizes is not None @@ -494,12 +550,19 @@ async def initialize(admin_client: AsyncClient, rank_offset: int, engine_world_s } if has_explicit_engine_world_sizes: payload["engine_world_size"] = engine_world_size - await _admin_post( - admin_client, - "/init_broadcaster", - timeout_s=max(ADMIN_TIMEOUT_S, timeout), - json=payload, - ) + if use_native_collective_rpc: + response = await admin_client.post( + "/collective_rpc", + json={"method": "init_broadcaster", "kwargs": payload}, + ) + response.raise_for_status() + else: + await _admin_post( + admin_client, + "/init_broadcaster", + timeout_s=max(ADMIN_TIMEOUT_S, timeout), + json=payload, + ) await asyncio.gather( *[ diff --git a/src/prime_rl/utils/dynamo.py b/src/prime_rl/utils/dynamo.py new file mode 100644 index 0000000000..1cd0f05250 --- /dev/null +++ b/src/prime_rl/utils/dynamo.py @@ -0,0 +1,114 @@ +from __future__ import annotations + +import asyncio +from typing import Any, cast + +import httpx +from httpx import AsyncClient +from pydantic import BaseModel, ConfigDict, Field +from tenacity import AsyncRetrying, retry_if_exception, stop_after_delay, wait_exponential + +from prime_rl.configs.shared import ClientConfig + +DYNAMO_RL_DISCOVERY_PROTOCOL_VERSION = 1 +DYNAMO_READINESS_REQUEST_TIMEOUT_S = 30.0 + + +class DiscoveredDynamoWorker(BaseModel): + model_config = ConfigDict(frozen=True, extra="ignore") + + component: str = Field(min_length=1) + instance_id: int = Field(ge=0, strict=True) + model: str + admin_base_url: str = Field(min_length=1) + world_size: int = Field(gt=0, strict=True) + + +class DynamoDiscoverySnapshot(BaseModel): + model_config = ConfigDict(extra="ignore") + + protocol_version: int = Field( + strict=True, + ge=DYNAMO_RL_DISCOVERY_PROTOCOL_VERSION, + le=DYNAMO_RL_DISCOVERY_PROTOCOL_VERSION, + ) + workers: list[dict[str, Any]] + + +class DynamoDiscoveryPending(ValueError): + """A well-formed discovery snapshot that is not ready yet.""" + + +def _is_retryable_dynamo_error(exception: BaseException) -> bool: + if isinstance(exception, httpx.HTTPStatusError): + return exception.response.status_code == 429 or exception.response.status_code >= 500 + return isinstance(exception, (DynamoDiscoveryPending, httpx.TransportError)) + + +def _parse_dynamo_workers(payload: object, model_name: str) -> tuple[DiscoveredDynamoWorker, ...]: + snapshot = DynamoDiscoverySnapshot.model_validate(payload) + workers = [] + for raw_worker in snapshot.workers: + if raw_worker.get("model") not in (None, model_name): + continue + if error := raw_worker.get("error"): + raise DynamoDiscoveryPending(f"Dynamo RL worker probe is not ready: {error}") + workers.append(DiscoveredDynamoWorker.model_validate(raw_worker)) + if not workers: + raise DynamoDiscoveryPending("Dynamo RL discovery returned no workers yet") + + identities = [(worker.component, worker.instance_id) for worker in workers] + admin_urls = [worker.admin_base_url for worker in workers] + if len(set(identities)) != len(identities): + raise ValueError("Dynamo RL discovery returned duplicate worker identities") + if len(set(admin_urls)) != len(admin_urls): + raise ValueError("Dynamo RL discovery returned duplicate admin endpoints") + return tuple(sorted(workers, key=lambda worker: (worker.component, worker.instance_id))) + + +def setup_dynamo_admin_clients(workers: tuple[DiscoveredDynamoWorker, ...]) -> list[AsyncClient]: + return [ + AsyncClient( + base_url=worker.admin_base_url.rstrip("/"), + limits=httpx.Limits(max_connections=4, max_keepalive_connections=1), + timeout=httpx.Timeout(None), + ) + for worker in workers + ] + + +async def discover_dynamo_workers( + client_config: ClientConfig, + model_name: str, + expected_inference_world_size: int, +) -> tuple[DiscoveredDynamoWorker, ...]: + discovery_url = cast(str, client_config.dynamo_discovery_url).rstrip("/").removesuffix("/v1") + timeout = client_config.wait_for_ready_timeout + loop = asyncio.get_running_loop() + deadline = loop.time() + timeout + async with asyncio.timeout(timeout): + async with AsyncClient(timeout=httpx.Timeout(None)) as client: + async for attempt in AsyncRetrying( + stop=stop_after_delay(timeout), + wait=wait_exponential(multiplier=0.1, min=0.1, max=1), + retry=retry_if_exception(_is_retryable_dynamo_error), + reraise=True, + ): + with attempt: + remaining = deadline - loop.time() + if remaining <= 0: + raise TimeoutError + response = await client.get( + f"{discovery_url}/v1/rl/workers", + timeout=httpx.Timeout(min(DYNAMO_READINESS_REQUEST_TIMEOUT_S, remaining)), + ) + response.raise_for_status() + workers = _parse_dynamo_workers(response.json(), model_name) + discovered_world_size = sum(worker.world_size for worker in workers) + if discovered_world_size != expected_inference_world_size: + raise DynamoDiscoveryPending( + f"Dynamo discovery returned inference_world_size={discovered_world_size}; " + f"waiting for expected inference_world_size={expected_inference_world_size}" + ) + return workers + raise TimeoutError("Dynamo worker discovery timed out") diff --git a/tests/unit/orchestrator/test_orchestrator_setup.py b/tests/unit/orchestrator/test_orchestrator_setup.py index d1bcdff0fa..784c6bc705 100644 --- a/tests/unit/orchestrator/test_orchestrator_setup.py +++ b/tests/unit/orchestrator/test_orchestrator_setup.py @@ -18,6 +18,7 @@ async def run() -> None: ), renderer=renderer_settings, any_policy_sourced=True, + weight_broadcast=SimpleNamespace(type="filesystem", inference_world_size=None), ) renderer = object() inference_pool = object() @@ -43,6 +44,7 @@ async def run() -> None: train_client_type="renderer", eval_client_type="openai_chat_completions", renderer_config=renderer_settings, + expected_inference_world_size=None, ) asyncio.run(run()) @@ -64,6 +66,7 @@ async def run() -> None: ), renderer=renderer_settings, any_policy_sourced=False, + weight_broadcast=SimpleNamespace(inference_world_size=8), ) renderer = object() inference_pool = object() @@ -89,6 +92,7 @@ async def run() -> None: train_client_type="renderer", eval_client_type="openai_chat_completions", renderer_config=renderer_settings, + expected_inference_world_size=8, ) asyncio.run(run()) diff --git a/tests/unit/test_configs.py b/tests/unit/test_configs.py index abeac169e8..5c294e66a7 100644 --- a/tests/unit/test_configs.py +++ b/tests/unit/test_configs.py @@ -222,6 +222,116 @@ def test_trainer_enable_token_export_cli_flag(): assert cli(TrainerConfig, args=["--enable-token-export"]).enable_token_export +def test_external_dynamo_world_size_survives_rl_config_resolution(): + config = RLConfig.model_validate( + { + "trainer": {}, + "orchestrator": { + "model": { + "client": { + "base_url": ["http://frontend:8000/v1"], + "dynamo_discovery_url": "http://frontend:8001", + } + } + }, + "inference": None, + "weight_broadcast": { + "type": "nccl", + "host": "trainer.service", + "inference_world_size": 8, + }, + } + ) + + assert config.trainer.weight_broadcast.inference_world_size == 8 + assert config.trainer.weight_broadcast.host == "trainer.service" + assert config.orchestrator.weight_broadcast.inference_world_size == 8 + assert config.orchestrator.weight_broadcast.host == "trainer.service" + + +def test_external_dynamo_nccl_does_not_require_a_local_inference_gpu(): + config = RLConfig.model_validate( + { + "trainer": {}, + "orchestrator": { + "model": { + "client": { + "base_url": ["http://frontend:8000/v1"], + "dynamo_discovery_url": "http://frontend:8001", + } + } + }, + "inference": None, + "deployment": { + "type": "single_node", + "num_train_gpus": 1, + "num_infer_gpus": 0, + }, + "weight_broadcast": { + "type": "nccl", + "host": "trainer.service", + "inference_world_size": 1, + }, + } + ) + + assert config.deployment.num_train_gpus == 1 + assert config.deployment.num_infer_gpus == 0 + assert config.trainer.weight_broadcast.inference_world_size == 1 + + +def test_default_nccl_world_size_does_not_bypass_local_gpu_guard(): + with pytest.raises(ValueError, match="NCCL weight broadcast requires at least 2"): + RLConfig.model_validate( + { + "trainer": {}, + "orchestrator": {}, + "inference": None, + "deployment": { + "type": "single_node", + "num_train_gpus": 1, + "num_infer_gpus": 0, + }, + "weight_broadcast": {"type": "nccl"}, + } + ) + + +def test_external_dynamo_lora_world_size_survives_filesystem_config_resolution(): + config = RLConfig.model_validate( + { + "trainer": {"model": {"lora": {}}}, + "orchestrator": { + "model": { + "client": { + "base_url": ["http://frontend:8000/v1"], + "dynamo_discovery_url": "http://frontend:8001", + } + } + }, + "inference": None, + "weight_broadcast": {"type": "filesystem", "inference_world_size": 8}, + } + ) + + assert config.orchestrator.weight_broadcast.inference_world_size == 8 + + +def test_dynamo_orchestrator_requires_explicit_inference_world_size(): + with pytest.raises(ValueError, match="inference_world_size"): + OrchestratorConfig.model_validate( + { + "model": { + "client": { + "base_url": ["http://frontend:8000/v1"], + "dynamo_discovery_url": "http://frontend:8001", + } + }, + "weight_broadcast": {"type": "filesystem"}, + } + ) + + def test_single_node_auto_inference_ports_follow_server_port(): config = RLConfig.model_validate( { diff --git a/tests/unit/utils/test_dynamo.py b/tests/unit/utils/test_dynamo.py new file mode 100644 index 0000000000..c492c83f63 --- /dev/null +++ b/tests/unit/utils/test_dynamo.py @@ -0,0 +1,103 @@ +import asyncio +from unittest.mock import AsyncMock, MagicMock, patch + +import httpx +import pytest + +from prime_rl.configs.shared import ClientConfig +from prime_rl.utils.dynamo import _parse_dynamo_workers, discover_dynamo_workers + +MODEL = "Qwen/Qwen3-0.6B" + + +def worker(**updates): + value = { + "component": "backend", + "instance_id": 10, + "model": MODEL, + "admin_base_url": "http://decode:8120", + "world_size": 2, + } + return {**value, **updates} + + +def payload(*workers): + return {"protocol_version": 1, "workers": list(workers)} + + +def response(body): + result = MagicMock() + result.raise_for_status = MagicMock() + result.json.return_value = body + return result + + +def test_parse_workers_orders_identity_and_preserves_topology(): + workers = _parse_dynamo_workers( + payload( + worker(component="prefill", instance_id=20, admin_base_url="http://prefill:8121"), + worker(), + ), + MODEL, + ) + assert [(item.component, item.instance_id) for item in workers] == [("backend", 10), ("prefill", 20)] + assert [item.world_size for item in workers] == [2, 2] + + +@pytest.mark.parametrize( + "workers", + [ + [], + [worker(error="probe timed out")], + [worker(admin_base_url=None)], + [worker(world_size=0)], + [worker(model="other/model")], + [worker(), worker(instance_id=11)], + [worker(), worker(component="prefill", instance_id=20, admin_base_url="http://decode:8120")], + ], +) +def test_parse_workers_rejects_incomplete_or_duplicate_snapshots(workers): + with pytest.raises(ValueError): + _parse_dynamo_workers(payload(*workers), MODEL) + + +def test_discovery_config_rejects_static_admin_urls(): + with pytest.raises(ValueError, match="dynamo_discovery_url"): + ClientConfig( + dynamo_discovery_url="http://frontend:8001", + admin_base_url=["http://worker:8120"], + ) + + +def test_discovery_retries_until_expected_world_size_is_complete(): + transient = MagicMock() + transient.raise_for_status.side_effect = httpx.HTTPStatusError( + "Service unavailable", + request=httpx.Request("GET", "http://frontend:8001/v1/rl/workers"), + response=httpx.Response(503), + ) + discovery_client = AsyncMock() + discovery_client.get.side_effect = [ + transient, + response(payload(worker())), + response( + payload( + worker(), + worker(component="prefill", instance_id=20, admin_base_url="http://prefill:8121"), + ) + ), + ] + context = AsyncMock() + context.__aenter__.return_value = discovery_client + + with patch("prime_rl.utils.dynamo.AsyncClient", return_value=context): + workers = asyncio.run( + discover_dynamo_workers( + ClientConfig(dynamo_discovery_url="http://frontend:8001", wait_for_ready_timeout=1), + model_name=MODEL, + expected_inference_world_size=4, + ) + ) + + assert discovery_client.get.await_count == 3 + assert [item.component for item in workers] == ["backend", "prefill"] diff --git a/tests/unit/utils/test_dynamo_inmemory.py b/tests/unit/utils/test_dynamo_inmemory.py new file mode 100644 index 0000000000..8a6f1c08f4 --- /dev/null +++ b/tests/unit/utils/test_dynamo_inmemory.py @@ -0,0 +1,41 @@ +import asyncio +from unittest.mock import AsyncMock, MagicMock + +from prime_rl.utils.client import init_nccl_broadcast, update_weights + + +def test_native_nccl_initialization_uses_collective_rpc(): + clients = [AsyncMock(), AsyncMock()] + for client in clients: + response = MagicMock() + response.raise_for_status = MagicMock() + client.post.return_value = response + + asyncio.run( + init_nccl_broadcast( + clients, + host="127.0.0.1", + port=29519, + timeout=1200, + inference_world_size=4, + engine_world_sizes=[2, 2], + use_native_collective_rpc=True, + ) + ) + assert [client.post.await_args.args[0] for client in clients] == ["/collective_rpc", "/collective_rpc"] + assert [client.post.await_args.kwargs["json"]["kwargs"]["rank_offset"] for client in clients] == [0, 2] + + +def test_native_full_weight_update_uses_positional_path(tmp_path): + client = AsyncMock() + response = MagicMock() + response.raise_for_status = MagicMock() + client.post.return_value = response + + asyncio.run(update_weights([client], tmp_path, step=2, use_native_collective_rpc=True)) + + collective_calls = [call for call in client.post.await_args_list if call.args[0] == "/collective_rpc"] + assert collective_calls[0].kwargs["json"] == { + "method": "update_weights_from_path", + "args": [tmp_path.as_posix()], + } From abac8e606d0c2ee382a6ca11cabf55d621e2649a Mon Sep 17 00:00:00 2001 From: Biswa Panda Date: Mon, 3 Aug 2026 11:44:48 -0700 Subject: [PATCH 6/9] docs(dynamo): add training recipes --- examples/dynamo/README.md | 80 +++++++++++ examples/dynamo/glm52_fp8_r2e/README.md | 128 ++++++++++++++++++ .../dynamo/glm52_fp8_r2e/orchestrator.toml | 60 ++++++++ examples/dynamo/glm52_fp8_r2e/trainer.toml | 47 +++++++ examples/dynamo/qwen3_06b_math/README.md | 45 ++++++ examples/dynamo/qwen3_06b_math/rl.toml | 43 ++++++ examples/dynamo/qwen3_30b_Thinking/README.md | 53 ++++++++ examples/dynamo/qwen3_30b_Thinking/rl.toml | 56 ++++++++ 8 files changed, 512 insertions(+) create mode 100644 examples/dynamo/README.md create mode 100644 examples/dynamo/glm52_fp8_r2e/README.md create mode 100644 examples/dynamo/glm52_fp8_r2e/orchestrator.toml create mode 100644 examples/dynamo/glm52_fp8_r2e/trainer.toml create mode 100644 examples/dynamo/qwen3_06b_math/README.md create mode 100644 examples/dynamo/qwen3_06b_math/rl.toml create mode 100644 examples/dynamo/qwen3_30b_Thinking/README.md create mode 100644 examples/dynamo/qwen3_30b_Thinking/rl.toml diff --git a/examples/dynamo/README.md b/examples/dynamo/README.md new file mode 100644 index 0000000000..067945e157 --- /dev/null +++ b/examples/dynamo/README.md @@ -0,0 +1,80 @@ +# Dynamo native-gRPC deployment requirements + +The recipes in this directory use Dynamo for inference and Prime only for the +trainer/orchestrator. Start from Dynamo's `sidecar_agg.yaml` or +`sidecar_disagg.yaml`, then apply the following RL overlay. The stock manifests +are serving examples and do not enable Prime worker discovery or vLLM's admin +control routes by themselves. + +## Frontend + +Set these variables on the Dynamo frontend and expose both container ports: + +```yaml +env: + - name: DYN_ENABLE_RL + value: "true" + - name: DYN_RL_PORT + value: "8001" +ports: + - name: http + containerPort: 8000 + - name: rl-discovery + containerPort: 8001 +``` + +The frontend Kubernetes Service must also map ports 8000 and 8001. Prime's +`base_url` targets 8000; `dynamo_discovery_url` targets 8001. + +## Every vLLM engine and sidecar pair + +The engine HTTP address published by discovery must be reachable from the +trainer, so bind vLLM to the pod network rather than loopback: + +```text +vllm-rs serve --host 0.0.0.0 --port 8000 --grpc-port 50051 -- \ + --worker-extension-cls prime_rl.inference.vllm.worker.nccl.NCCLWeightUpdateWorker \ + +``` + +Install the matching Prime source in the engine image so Python can import the +worker extension. Set this environment variable on the `vllm-engine` +container to expose `/pause`, `/resume`, and `/collective_rpc`: + +```yaml +env: + - name: VLLM_SERVER_DEV_MODE + value: "1" +``` + +Set the following variables on each `dynamo-vllm-sidecar` container. `POD_IP` +must appear before `VLLM_HTTP_ENDPOINT` so Kubernetes expands it: + +```yaml +env: + - name: POD_IP + valueFrom: + fieldRef: + fieldPath: status.podIP + - name: VLLM_HTTP_ENDPOINT + value: "http://$(POD_IP):8000" + - name: DYN_ENABLE_RL + value: "true" +``` + +Keep the existing `--grpc-endpoint 127.0.0.1:50051`: gRPC stays pod-local, +while discovery publishes the pod-reachable HTTP admin address. No per-worker +Kubernetes Service is required when trainer-to-pod networking is routable. + +After startup, `/v1/rl/workers` must return every expected worker with a +non-null `admin_base_url`, positive `world_size`, and no `error` before Prime is +launched. + +## Recipes + +- [`qwen3_06b_math`](qwen3_06b_math): single-GPU trainer and aggregate Dynamo + inference smoke test. +- [`qwen3_30b_Thinking`](qwen3_30b_Thinking): Qwen3-30B Thinking math with an + external prefill/decode deployment. +- [`glm52_fp8_r2e`](glm52_fp8_r2e): multi-node GLM-5.2 FP8 R2E training with a + separately managed DGD. diff --git a/examples/dynamo/glm52_fp8_r2e/README.md b/examples/dynamo/glm52_fp8_r2e/README.md new file mode 100644 index 0000000000..41932321ef --- /dev/null +++ b/examples/dynamo/glm52_fp8_r2e/README.md @@ -0,0 +1,128 @@ +# GLM-5.2 FP8 R2E with external Dynamo inference + +This three-step smoke recipe runs a distributed Prime trainer and orchestrator +against a separately managed Dynamo deployment serving +`zai-org/GLM-5.2-FP8`. Dynamo owns the frontend, vLLM engines, and one native +gRPC sidecar per engine group; Prime discovers the mutable engine control +surface through `/v1/rl/workers`. + +The recipe is split into trainer and orchestrator files because the external +Dynamo DGD and the multi-node trainer have independent lifecycles. It does not +add a Prime inference configuration or require Prime's launcher to manage the +DGD. + +## Reference topology + +| Component | Shape | GPUs | +|---|---|---:| +| Dynamo prefill | 2 nodes, DP4 x TP2 x PP1 x EP8 | 8 | +| Dynamo decode | 2 nodes, DP4 x TP2 x PP1 x EP8 | 8 | +| Prime trainer | 4 nodes, FSDP16 x CP4 x EP8 | 16 | + +The checked-in configuration targets this topology; it is not a claim that +every cluster can use these parallelism dimensions unchanged. The two +discovery records must report `prefill` and `backend` components with a +combined `world_size` of 16. If the DGD topology changes, update +`weight_broadcast.inference_world_size` in both TOML files to the atomic sum +returned by the same `/v1/rl/workers` response. + +The inference engines must load +`prime_rl.inference.vllm.worker.nccl.NCCLWeightUpdateWorker`, expose vLLM's +admin routes, and run the version-matched `vllm-rs` and +`dynamo-vllm-sidecar` binaries described in [`../README.md`](../README.md). +For mutable GLM weight reloads, launch every vLLM rank with `--enforce-eager`. +When serving a snapshot path, set `--served-model-name +zai-org/GLM-5.2-FP8`. The tested GLM entrypoint also uses the `glm47` tool +parser, `glm45` reasoning parser, the model chat template, and complementary +NIXL `kv_producer`/`kv_consumer` roles for prefill/decode. +The trainer and every inference engine must also use a compatible NCCL +transport. Apply cluster-specific settings such as `NCCL_IB_DISABLE`, +`NCCL_SOCKET_IFNAME`, and the NCCL network plugin consistently on both sides; +do not force Socket on the trainer while allowing inference to select IB. +The filesystem rollout transport requires the orchestrator and all trainer +nodes to mount the same read-write shared output root. + +## Configure + +Initialize the R2E environment submodule and install its workspace package: + +```bash +git submodule update --init -- deps/research-environments +uv sync --package prime-rl --package r2e-gym-v1 +``` + +The taskset intentionally uses the current `r2e-gym-v1` default, +`PrimeIntellect/R2E-Gym-Subset-Verified`, instead of pinning an older dataset +override in this recipe. + +Replace the checked-in service names when the DGD and trainer use different +DNS names: + +- `model.client.base_url`: Dynamo OpenAI frontend on port 8000; +- `model.client.dynamo_discovery_url`: Dynamo RL discovery on port 8001; +- `weight_broadcast.host`: trainer rank zero, reachable from every inference + engine. + +The R2E harness uses Prime sandboxes. Configure the normal Prime credentials, +or replace the runtime with the sandbox backend used by your cluster. Model +and dataset caches should be shared across the trainer and inference nodes. + +Multi-turn affinity is not enabled merely by sending a session header. Start +the Dynamo frontend with `--router-session-affinity-ttl-secs ` (or +`DYN_ROUTER_SESSION_AFFINITY_TTL_SECS`) and choose an idle TTL longer than the +longest expected R2E turn gap. Prime maps each trajectory ID to the canonical +`X-Dynamo-Session-ID` header in `orchestrator.toml`. + +Verify both the model and the complete atomic worker snapshot before starting +Prime. A successful HTTP status alone is insufficient: + +```bash +MODEL=zai-org/GLM-5.2-FP8 +curl -fsS http://dynamo-frontend:8000/v1/models | + jq -e --arg model "$MODEL" '.data | any(.id == $model)' +curl -fsS http://dynamo-frontend:8001/v1/rl/workers | + jq -e --arg model "$MODEL" ' + .protocol_version == 1 and + (.workers | length == 2) and + (all(.workers[]; .model == $model and + ((.error // "") == "") and + (.instance_id != null) and + ((.admin_base_url // "") != ""))) and + ([.workers[].instance_id] | unique | length == 2) and + ([.workers[].admin_base_url] | unique | length == 2) and + ([.workers[] | select(.model == $model) | .component] | sort == ["backend", "prefill"]) and + ([.workers[] | select(.model == $model) | .world_size] | add == 16) + ' +``` + +## Run + +Launch the trainer on four 4-GPU nodes with the cluster's distributed runner. +For example, rank zero's rendezvous address can be passed to `torchrun` while +all ranks consume the same trainer file: + +```bash +uv run torchrun \ + --nnodes=4 --nproc-per-node=4 \ + --rdzv-backend=c10d --rdzv-endpoint="$TRAINER_RANK_ZERO:29501" \ + --node-rank="$NODE_RANK" \ + -m prime_rl.trainer.rl.train \ + @ examples/dynamo/glm52_fp8_r2e/trainer.toml \ + --output-dir /shared/glm52-dynamo-r2e/train +``` + +After trainer rank zero opens port 29500, launch the orchestrator once: + +```bash +uv run orchestrator \ + @ examples/dynamo/glm52_fp8_r2e/orchestrator.toml \ + --output-dir /shared/glm52-dynamo-r2e/train/run_0 +``` + +The gate succeeds when the first optimizer step completes, policy version 1 +settles on all 16 inference ranks through NCCL, and a later multi-turn rollout +completes without changing its Dynamo session assignment. Three steps are +required because finite NCCL runs skip broadcasts once +`step >= max_steps - 1`; this leaves step 1 as the first non-final broadcast +slot. The disabled post-batch zero-advantage filter keeps this small smoke run +from stalling on a homogeneous batch; enable it for a production training run. diff --git a/examples/dynamo/glm52_fp8_r2e/orchestrator.toml b/examples/dynamo/glm52_fp8_r2e/orchestrator.toml new file mode 100644 index 0000000000..74a0c0cf60 --- /dev/null +++ b/examples/dynamo/glm52_fp8_r2e/orchestrator.toml @@ -0,0 +1,60 @@ +max_steps = 3 +batch_size = 2 +group_size = 2 +seq_len = 32768 +max_inflight_episodes = 2 +max_off_policy_steps = 1 +tasks_per_minute = 1 + +[model] +name = "zai-org/GLM-5.2-FP8" + +[model.client] +base_url = ["http://dynamo-frontend:8000/v1"] +dynamo_discovery_url = "http://dynamo-frontend:8001" +wait_for_ready_timeout = 7200 + +[model.client.extra_headers_from_state] +X-Dynamo-Session-ID = "trajectory_id" + +[tokenizer] +name = "zai-org/GLM-5.2-FP8" + +[renderer] +name = "glm-5.1" +clear_thinking = false + +[train.sampling] +temperature = 1.0 +max_completion_tokens = 2048 +extra_body = { chat_template_kwargs = { clear_thinking = false } } + +[[train.source]] +name = "r2e" +group_size = 2 +serve.pool = { type = "static", num_workers = 2 } +env.taskset = { id = "r2e-gym-v1" } +env.agent.max_turns = 64 +env.agent.max_input_tokens = 30720 +env.agent.max_output_tokens = 16384 +env.agent.max_total_tokens = 32768 +env.agent.timeout = { setup = 600, rollout = 1800, finalize = 300, scoring = 600 } +env.agent.harness = { id = "bash", edit = true } +env.agent.runtime = { type = "prime", labels = ["glm52-dynamo"], cpu = 4, creates_per_min = 2 } + +[[post_batch_filters]] +type = "zero_advantage" +enforce = false + +[weight_broadcast] +type = "nccl" +host = "trainer-0.trainer-headless" +port = 29500 +timeout = 12000 +inference_world_size = 16 + +[rollout_transport] +type = "filesystem" + +[log] +level = "debug" diff --git a/examples/dynamo/glm52_fp8_r2e/trainer.toml b/examples/dynamo/glm52_fp8_r2e/trainer.toml new file mode 100644 index 0000000000..0298631e29 --- /dev/null +++ b/examples/dynamo/glm52_fp8_r2e/trainer.toml @@ -0,0 +1,47 @@ +max_steps = 3 +dist_timeout_seconds = 12000 + +[model] +name = "zai-org/GLM-5.2-FP8" +seq_len = 32768 +impl = "custom" +attn = "flash_attention_2" +dp_replicate = 1 +cp = 4 +ep = 8 +optimization_dtype = "bfloat16" +reduce_dtype = "bfloat16" +moe_router_dtype = "float32" +fp8 = false +optim_cpu_offload = true +fused_lm_head_token_chunk_size = 1024 + +[model.ac] +freq = 1 + +[model.ac_offloading] +max_inflight_activations = 1 + +[tokenizer] +name = "zai-org/GLM-5.2-FP8" + +[optim] +type = "sign_sgd" +lr = 1e-6 +weight_decay = 0.0 + +[scheduler] +type = "constant" + +[weight_broadcast] +type = "nccl" +host = "0.0.0.0" +port = 29500 +timeout = 12000 +inference_world_size = 16 + +[rollout_transport] +type = "filesystem" + +[log] +level = "debug" diff --git a/examples/dynamo/qwen3_06b_math/README.md b/examples/dynamo/qwen3_06b_math/README.md new file mode 100644 index 0000000000..dcfae8c576 --- /dev/null +++ b/examples/dynamo/qwen3_06b_math/README.md @@ -0,0 +1,45 @@ +# Qwen3 0.6B math with external Dynamo inference + +This four-step smoke recipe runs the Prime trainer and orchestrator locally while generation is served by an already-running Dynamo frontend, vLLM sidecar, and vLLM engine. Prime does not launch local inference, so the configuration intentionally has no `[inference]` block. + +## Prerequisites + +Use this version-matched native-gRPC source set: + +- Dynamo `feat/dyn-pi-sidecar-v2-review-001` at `836fe81012` +- vLLM `feat/dyn-pi-sidecar-v2-review-001` at `e56ee21b2c` +- Prime `feat/dyn-pi-sidecar-v2-review-001` + +Build `vllm-rs` and `dynamo-vllm-sidecar` from those revisions into the same +runtime image. For Kubernetes, start from Dynamo's +`examples/backends/vllm/deploy/sidecar_agg.yaml`; its adjacent `README.md` +documents the paired-binary image. Then apply the required Prime RL discovery +and admin overlay in [`../README.md`](../README.md). This contract requires both +native gRPC and the Dynamo `/v1/rl/workers` endpoint; a standard Python-only +vLLM worker is not compatible. + +Install the math environment: + +```bash +prime env install primeintellect/math-env +``` + +Start an aggregated DP1 Dynamo deployment for `Qwen/Qwen3-0.6B`. The orchestrator waits for both model publication and worker discovery. These requests are useful diagnostics: + +```bash +curl http://127.0.0.1:8000/v1/models +curl http://127.0.0.1:8001/v1/rl/workers +``` + +The checked-in URLs assume Dynamo is reachable from the trainer through localhost, as in a shared dev pod. For a remote DGD, replace both URLs with its frontend services. Also replace `weight_broadcast.host` with a trainer hostname or IP reachable from every sidecar; localhost is not valid across pods or nodes. + +`inference_world_size` must equal the sum of `world_size` in one `/v1/rl/workers` response. This recipe assumes one aggregated DP1 engine and therefore uses `1`. + +## Run + +```bash +uv run rl @ examples/dynamo/qwen3_06b_math/rl.toml \ + --output-dir outputs/dynamo-qwen3-06b-math +``` + +The run is successful when four optimizer steps complete, the verifier reports math rewards, weight versions advance after each update, and the Dynamo workers remain healthy. diff --git a/examples/dynamo/qwen3_06b_math/rl.toml b/examples/dynamo/qwen3_06b_math/rl.toml new file mode 100644 index 0000000000..d26bf43ba2 --- /dev/null +++ b/examples/dynamo/qwen3_06b_math/rl.toml @@ -0,0 +1,43 @@ +max_steps = 4 +seq_len = 2048 + +[model] +name = "Qwen/Qwen3-0.6B" + +[deployment] +type = "single_node" +num_train_gpus = 1 +num_infer_gpus = 0 + +[weight_broadcast] +type = "nccl" +host = "127.0.0.1" +inference_world_size = 1 + +[trainer] + +[orchestrator] +batch_size = 32 +group_size = 4 + +[orchestrator.model.client] +base_url = ["http://127.0.0.1:8000/v1"] +dynamo_discovery_url = "http://127.0.0.1:8001" +wait_for_ready_timeout = 1800 + +[orchestrator.model.client.extra_headers_from_state] +X-Session-ID = "trajectory_id" +X-Dynamo-Session-ID = "trajectory_id" + +[orchestrator.renderer] +name = "auto" +thinking_retention = "all" + +[orchestrator.train.sampling] +max_completion_tokens = 2048 + +[[orchestrator.train.source]] +name = "math" +env.taskset = { id = "math-env-v1", dataset_name = "openai/gsm8k", dataset_subset = "main" } +env.agent.harness = { id = "null" } +env.agent.runtime = { type = "subprocess" } diff --git a/examples/dynamo/qwen3_30b_Thinking/README.md b/examples/dynamo/qwen3_30b_Thinking/README.md new file mode 100644 index 0000000000..d01b182882 --- /dev/null +++ b/examples/dynamo/qwen3_30b_Thinking/README.md @@ -0,0 +1,53 @@ +# Qwen3 30B Thinking math with external Dynamo inference + +This four-step scale recipe runs a two-GPU Prime trainer against an external 1-prefill/1-decode Dynamo deployment serving `Qwen/Qwen3-30B-A3B-Thinking-2507`. Prime launches no inference process; Dynamo owns the frontend, sidecars, and vLLM engines. + +The two-GPU trainer uses BF16 optimization and reduction plus CPU optimizer +offload. Leaving Prime's FP32 optimization default in place exceeds two GB200s +when Adam allocates its first-step state; larger production recipes use the +existing eight-GPU training shape instead. + +The same model is used by the existing public examples `qwen30b_math`, `qwen30b_swe`, `multinode/rl.toml`, and `multinode/sft.toml`. Those examples remain the source of truth for larger training and non-Dynamo deployment settings; this recipe adds only the external-Dynamo client shape. + +## Prerequisites + +Use this version-matched native-gRPC source set: + +- Dynamo `feat/dyn-pi-sidecar-v2-review-001` at `836fe81012` +- vLLM `feat/dyn-pi-sidecar-v2-review-001` at `e56ee21b2c` +- Prime `feat/dyn-pi-sidecar-v2-review-001` + +Build `vllm-rs` and `dynamo-vllm-sidecar` from those revisions into the same +runtime image. For Kubernetes, start from Dynamo's +`examples/backends/vllm/deploy/sidecar_disagg.yaml` and change the model plus +parallelism/resources for the 30B topology; its adjacent `README.md` documents +the paired-binary image. Then apply the required Prime RL discovery and admin +overlay in [`../README.md`](../README.md). This contract requires both native +gRPC and the Dynamo `/v1/rl/workers` endpoint; a standard Python-only vLLM +worker is not compatible. + +Install the math environment: + +```bash +prime env install primeintellect/math-env +``` + +Start a Dynamo 1P/1D deployment and verify its public endpoints: + +```bash +curl http://127.0.0.1:8000/v1/models +curl http://127.0.0.1:8001/v1/rl/workers +``` + +The checked-in localhost URLs are for a colocated dev-pod run. For a DGD, replace them with the frontend services. Replace `weight_broadcast.host` with a trainer address reachable from every sidecar and allow the configured NCCL port through the network policy. + +`inference_world_size` must equal the sum of `world_size` in one `/v1/rl/workers` response. This recipe assumes 1P/1D with one rank per engine, for a total of `2`; TP, PP, or managed-DP topologies require the corresponding larger sum. + +## Run + +```bash +uv run rl @ examples/dynamo/qwen3_30b_Thinking/rl.toml \ + --output-dir outputs/dynamo-qwen3-30b-thinking-math +``` + +The run is successful when four optimizer steps complete, math rewards are emitted, all worker weight versions advance, and both prefill and decode workers remain healthy. diff --git a/examples/dynamo/qwen3_30b_Thinking/rl.toml b/examples/dynamo/qwen3_30b_Thinking/rl.toml new file mode 100644 index 0000000000..12c6995bf6 --- /dev/null +++ b/examples/dynamo/qwen3_30b_Thinking/rl.toml @@ -0,0 +1,56 @@ +max_steps = 4 +seq_len = 2048 + +[model] +name = "Qwen/Qwen3-30B-A3B-Thinking-2507" + +[deployment] +type = "single_node" +num_train_gpus = 2 +num_infer_gpus = 0 + +[weight_broadcast] +type = "nccl" +host = "127.0.0.1" +timeout = 1800 +inference_world_size = 2 + +[trainer.model] +impl = "custom" +attn = "flash_attention_3" +ep = 2 +optim_cpu_offload = true +optimization_dtype = "bfloat16" +reduce_dtype = "bfloat16" + +[trainer.model.ac] +freq = 1 + +[orchestrator] +batch_size = 2 +group_size = 2 +max_inflight_episodes = 2 +max_off_policy_steps = 0 + +[orchestrator.model.client] +base_url = ["http://127.0.0.1:8000/v1"] +dynamo_discovery_url = "http://127.0.0.1:8001" +wait_for_ready_timeout = 3600 + +[orchestrator.model.client.extra_headers_from_state] +X-Session-ID = "trajectory_id" +X-Dynamo-Session-ID = "trajectory_id" + +[orchestrator.renderer] +name = "qwen3" +enable_thinking = true + +[orchestrator.train.sampling] +temperature = 1.0 +max_completion_tokens = 2048 + +[[orchestrator.train.source]] +name = "math" +env.taskset = { id = "math-env-v1", dataset_name = "PrimeIntellect/Hendrycks-Math", dataset_subset = "default", task = { math_verify_timeout = 60 } } +env.agent.harness = { id = "null" } +env.agent.runtime = { type = "subprocess" } From 0793e2eac73f429a02d6e6110304909eb8ad6458 Mon Sep 17 00:00:00 2001 From: Biswa Panda Date: Mon, 3 Aug 2026 12:10:02 -0700 Subject: [PATCH 7/9] fix(examples): remove obsolete fp8 field --- examples/dynamo/glm52_fp8_r2e/trainer.toml | 1 - 1 file changed, 1 deletion(-) diff --git a/examples/dynamo/glm52_fp8_r2e/trainer.toml b/examples/dynamo/glm52_fp8_r2e/trainer.toml index 0298631e29..d7476a734d 100644 --- a/examples/dynamo/glm52_fp8_r2e/trainer.toml +++ b/examples/dynamo/glm52_fp8_r2e/trainer.toml @@ -12,7 +12,6 @@ ep = 8 optimization_dtype = "bfloat16" reduce_dtype = "bfloat16" moe_router_dtype = "float32" -fp8 = false optim_cpu_offload = true fused_lm_head_token_chunk_size = 1024 From e9d3f8791bbf56065ee905d213feeafa280bd9f6 Mon Sep 17 00:00:00 2001 From: sami jaghouar Date: Thu, 13 Aug 2026 03:42:08 +0000 Subject: [PATCH 8/9] chore(dynamo): add pinned researcher build and smoke tooling --- examples/dynamo/README.md | 27 ++++-- examples/dynamo/glm52_fp8_r2e/README.md | 6 +- .../dynamo/glm52_fp8_r2e/orchestrator.toml | 4 +- examples/dynamo/local/README.md | 95 +++++++++++++++++++ examples/dynamo/local/rl.toml | 45 +++++++++ examples/dynamo/qwen3_06b_math/rl.toml | 4 +- examples/dynamo/qwen3_30b_Thinking/rl.toml | 4 +- .../dynamo/scripts/build_dynamo_artifacts.sh | 27 ++++++ examples/dynamo/scripts/build_vllm_wheel.sh | 27 ++++++ examples/dynamo/scripts/install_artifacts.sh | 19 ++++ skills/install/SKILL.md | 15 +++ src/prime_rl/orchestrator/orchestrator.py | 27 ++---- src/prime_rl/utils/client.py | 54 +++++++++-- src/prime_rl/utils/dynamo.py | 15 ++- .../orchestrator/test_orchestrator_setup.py | 14 +-- tests/unit/test_configs.py | 25 +---- tests/unit/utils/test_dynamo.py | 18 +++- 17 files changed, 350 insertions(+), 76 deletions(-) create mode 100644 examples/dynamo/local/README.md create mode 100644 examples/dynamo/local/rl.toml create mode 100755 examples/dynamo/scripts/build_dynamo_artifacts.sh create mode 100755 examples/dynamo/scripts/build_vllm_wheel.sh create mode 100755 examples/dynamo/scripts/install_artifacts.sh diff --git a/examples/dynamo/README.md b/examples/dynamo/README.md index 067945e157..e6750d634b 100644 --- a/examples/dynamo/README.md +++ b/examples/dynamo/README.md @@ -1,10 +1,23 @@ -# Dynamo native-gRPC deployment requirements +# Dynamo native-gRPC integration -The recipes in this directory use Dynamo for inference and Prime only for the -trainer/orchestrator. Start from Dynamo's `sidecar_agg.yaml` or -`sidecar_disagg.yaml`, then apply the following RL overlay. The stock manifests -are serving examples and do not enable Prime worker discovery or vLLM's admin -control routes by themselves. +Prime-RL uses Dynamo's OpenAI frontend for generation and `/v1/rl/workers` to +discover the direct vLLM admin endpoints used for weight updates. + +## Pinned researcher stack + +The currently validated source set is: + +- vLLM `biswapanda/vllm@e74fc3f` +- Dynamo `ai-dynamo/dynamo@fc556d9` +- Prime-RL changes ported from Biswa's combined integration PR #3181 onto current `main` + +These features are not all present in the public vLLM 0.26 and Dynamo 1.3.0 +wheels. Use the build/install scripts in [`scripts/`](scripts/) or an internally +published image containing those exact revisions. The scripts use seven-character +revision pins as the repository policy requires. + +For a two-GPU researcher smoke test, follow [`local/README.md`](local/README.md). +The other recipes describe larger externally deployed topologies. ## Frontend @@ -72,6 +85,8 @@ launched. ## Recipes +- [`local`](local): single-node 2-GPU smoke test with a real Dynamo stack + (etcd + frontend + sidecar + vLLM engine). - [`qwen3_06b_math`](qwen3_06b_math): single-GPU trainer and aggregate Dynamo inference smoke test. - [`qwen3_30b_Thinking`](qwen3_30b_Thinking): Qwen3-30B Thinking math with an diff --git a/examples/dynamo/glm52_fp8_r2e/README.md b/examples/dynamo/glm52_fp8_r2e/README.md index 41932321ef..b0dcbcf0f6 100644 --- a/examples/dynamo/glm52_fp8_r2e/README.md +++ b/examples/dynamo/glm52_fp8_r2e/README.md @@ -47,11 +47,11 @@ nodes to mount the same read-write shared output root. Initialize the R2E environment submodule and install its workspace package: ```bash -git submodule update --init -- deps/research-environments -uv sync --package prime-rl --package r2e-gym-v1 +git submodule update --init -- deps/prime-envs +uv sync --package prime-rl --package r2e-gym ``` -The taskset intentionally uses the current `r2e-gym-v1` default, +The taskset intentionally uses the current `r2e-gym` default, `PrimeIntellect/R2E-Gym-Subset-Verified`, instead of pinning an older dataset override in this recipe. diff --git a/examples/dynamo/glm52_fp8_r2e/orchestrator.toml b/examples/dynamo/glm52_fp8_r2e/orchestrator.toml index 74a0c0cf60..58e2a319c9 100644 --- a/examples/dynamo/glm52_fp8_r2e/orchestrator.toml +++ b/examples/dynamo/glm52_fp8_r2e/orchestrator.toml @@ -10,7 +10,7 @@ tasks_per_minute = 1 name = "zai-org/GLM-5.2-FP8" [model.client] -base_url = ["http://dynamo-frontend:8000/v1"] +base_url = "http://dynamo-frontend:8000/v1" dynamo_discovery_url = "http://dynamo-frontend:8001" wait_for_ready_timeout = 7200 @@ -33,7 +33,7 @@ extra_body = { chat_template_kwargs = { clear_thinking = false } } name = "r2e" group_size = 2 serve.pool = { type = "static", num_workers = 2 } -env.taskset = { id = "r2e-gym-v1" } +env.taskset = { id = "r2e-gym" } env.agent.max_turns = 64 env.agent.max_input_tokens = 30720 env.agent.max_output_tokens = 16384 diff --git a/examples/dynamo/local/README.md b/examples/dynamo/local/README.md new file mode 100644 index 0000000000..0f945f1b9d --- /dev/null +++ b/examples/dynamo/local/README.md @@ -0,0 +1,95 @@ +# Local Dynamo smoke test + +This recipe runs the Prime-RL trainer on GPU 1 and a Dynamo-managed vLLM engine +on GPU 0. It uses the exact source revisions from the validated integration: + +| Component | Revision | +| --- | --- | +| Prime-RL integration source | Biswa's combined PR #3181, ported onto this branch's `main` base | +| vLLM | `biswapanda/vllm@e74fc3f` | +| Dynamo | `ai-dynamo/dynamo@fc556d9` | + +The public vLLM 0.26 and Dynamo 1.3.0 wheels do not contain this complete +native-gRPC/worker-discovery stack. Build the artifacts below; do not replace +them with `uv sync --extra dynamo`. + +## Build and install the pinned artifacts + +Building vLLM requires a CUDA development environment. Building Dynamo requires +Rust, CMake, Clang, Protobuf, and the system packages listed in Dynamo's source +build guide. + +```bash +examples/dynamo/scripts/build_vllm_wheel.sh +examples/dynamo/scripts/build_dynamo_artifacts.sh +examples/dynamo/scripts/install_artifacts.sh +``` + +The scripts verify the seven-character revisions before building. They place the wheels and `dynamo-vllm-sidecar` executable under +`dist/dynamo/`. The installer creates a separate `.venv-dynamo` because the +custom inference stack's Torch and Pydantic constraints differ from Prime-RL's +trainer environment. + +## Start the stack + +Run each process in its own terminal. Stop old instances before retrying. + +```bash +# 1. etcd +etcd --data-dir /tmp/etcd-data \ + --listen-client-urls http://0.0.0.0:2379 \ + --advertise-client-urls http://127.0.0.1:2379 \ + --listen-peer-urls http://127.0.0.1:2380 +``` + +```bash +# 2. Custom vLLM engine on GPU 0 +CUDA_VISIBLE_DEVICES=0 VLLM_SERVER_DEV_MODE=1 \ + PYTHONPATH="$(pwd)/src" .venv-dynamo/bin/vllm-rs serve Qwen/Qwen3-0.6B \ + --host 0.0.0.0 --port 8002 --grpc-port 50051 \ + --python "$(pwd)/.venv-dynamo/bin/python" -- \ + --worker-extension-cls prime_rl.inference.vllm.worker.nccl.NCCLWeightUpdateWorker +``` + +```bash +# 3. Dynamo frontend: generation on 8000, RL discovery on 8001 +DYN_ENABLE_RL=true DYN_RL_PORT=8001 \ +DYN_VLLM_ENABLE_INFERENCE_V1_GENERATE=true \ + .venv-dynamo/bin/python -m dynamo.frontend \ + --http-host 0.0.0.0 --http-port 8000 \ + --namespace dynamo --discovery-backend etcd \ + --request-plane tcp --event-plane zmq --router-min-initial-workers 1 +``` + +```bash +# 4. Dynamo sidecar: generation uses gRPC 50051; Prime control uses HTTP 8002 +DYN_NAMESPACE=dynamo DYN_DISCOVERY_BACKEND=etcd DYN_ENABLE_RL=true \ +VLLM_HTTP_ENDPOINT=http://127.0.0.1:8002 \ + dist/dynamo/dynamo-vllm-sidecar \ + --vllm-endpoint 127.0.0.1:50051 \ + --admin-endpoint http://127.0.0.1:8002 \ + --model-path Qwen/Qwen3-0.6B \ + --namespace dynamo \ + --rl-discovery-model-name Qwen/Qwen3-0.6B +``` + +## Verify discovery before training + +```bash +curl -fsS http://127.0.0.1:8000/v1/models | jq . +curl -fsS http://127.0.0.1:8001/v1/rl/workers | jq . +``` + +The discovery response must report protocol version 1, model +`Qwen/Qwen3-0.6B`, `admin_base_url=http://127.0.0.1:8002`, and total +`world_size=1`. + +## Run Prime-RL + +```bash +CUDA_VISIBLE_DEVICES=1 uv run rl @ examples/dynamo/local/rl.toml \ + --output-dir outputs/dynamo-local --clean-output-dir +``` + +Success means four optimizer steps complete, policy versions advance, and a +generation request succeeds after each weight update. diff --git a/examples/dynamo/local/rl.toml b/examples/dynamo/local/rl.toml new file mode 100644 index 0000000000..11f9407f70 --- /dev/null +++ b/examples/dynamo/local/rl.toml @@ -0,0 +1,45 @@ +max_steps = 4 +seq_len = 2048 + +[model] +name = "Qwen/Qwen3-0.6B" + +[deployment] +type = "single_node" +num_train_gpus = 1 +num_infer_gpus = 0 + +[weight_broadcast] +type = "nccl" +host = "127.0.0.1" +inference_world_size = 1 + +[trainer] + +[orchestrator] +batch_size = 32 +group_size = 4 + +[orchestrator.model.client] +base_url = "http://127.0.0.1:8000/v1" +dynamo_discovery_url = "http://127.0.0.1:8001" + +[orchestrator.renderer] +name = "auto" + +[orchestrator.train.sampling] +max_completion_tokens = 2048 + +[[orchestrator.train.source]] +name = "math" + +[orchestrator.train.source.env.taskset] +id = "math-env" +dataset_name = "openai/gsm8k" +dataset_subset = "main" + +[orchestrator.train.source.env.agent.harness] +id = "null" + +[orchestrator.train.source.env.agent.runtime] +type = "subprocess" diff --git a/examples/dynamo/qwen3_06b_math/rl.toml b/examples/dynamo/qwen3_06b_math/rl.toml index d26bf43ba2..cc93d06257 100644 --- a/examples/dynamo/qwen3_06b_math/rl.toml +++ b/examples/dynamo/qwen3_06b_math/rl.toml @@ -21,7 +21,7 @@ batch_size = 32 group_size = 4 [orchestrator.model.client] -base_url = ["http://127.0.0.1:8000/v1"] +base_url = "http://127.0.0.1:8000/v1" dynamo_discovery_url = "http://127.0.0.1:8001" wait_for_ready_timeout = 1800 @@ -38,6 +38,6 @@ max_completion_tokens = 2048 [[orchestrator.train.source]] name = "math" -env.taskset = { id = "math-env-v1", dataset_name = "openai/gsm8k", dataset_subset = "main" } +env.taskset = { id = "math-env", dataset_name = "openai/gsm8k", dataset_subset = "main" } env.agent.harness = { id = "null" } env.agent.runtime = { type = "subprocess" } diff --git a/examples/dynamo/qwen3_30b_Thinking/rl.toml b/examples/dynamo/qwen3_30b_Thinking/rl.toml index 12c6995bf6..52f838a81f 100644 --- a/examples/dynamo/qwen3_30b_Thinking/rl.toml +++ b/examples/dynamo/qwen3_30b_Thinking/rl.toml @@ -33,7 +33,7 @@ max_inflight_episodes = 2 max_off_policy_steps = 0 [orchestrator.model.client] -base_url = ["http://127.0.0.1:8000/v1"] +base_url = "http://127.0.0.1:8000/v1" dynamo_discovery_url = "http://127.0.0.1:8001" wait_for_ready_timeout = 3600 @@ -51,6 +51,6 @@ max_completion_tokens = 2048 [[orchestrator.train.source]] name = "math" -env.taskset = { id = "math-env-v1", dataset_name = "PrimeIntellect/Hendrycks-Math", dataset_subset = "default", task = { math_verify_timeout = 60 } } +env.taskset = { id = "math-env", dataset_name = "PrimeIntellect/Hendrycks-Math", dataset_subset = "default", task = { math_verify_timeout = 60 } } env.agent.harness = { id = "null" } env.agent.runtime = { type = "subprocess" } diff --git a/examples/dynamo/scripts/build_dynamo_artifacts.sh b/examples/dynamo/scripts/build_dynamo_artifacts.sh new file mode 100755 index 0000000000..cdedf785cd --- /dev/null +++ b/examples/dynamo/scripts/build_dynamo_artifacts.sh @@ -0,0 +1,27 @@ +#!/usr/bin/env bash +set -euo pipefail + +DYNAMO_REPO=${DYNAMO_REPO:-https://github.com/ai-dynamo/dynamo.git} +DYNAMO_REV=${DYNAMO_REV:-fc556d9} +OUTPUT_DIR=${OUTPUT_DIR:-$PWD/dist/dynamo} +BUILD_DIR=${BUILD_DIR:-$PWD/.build/dynamo-prime-rl} + +mkdir -p "$OUTPUT_DIR" "$(dirname "$BUILD_DIR")" +if [[ ! -d "$BUILD_DIR/.git" ]]; then + git clone "$DYNAMO_REPO" "$BUILD_DIR" +fi +git -C "$BUILD_DIR" fetch origin "$DYNAMO_REV" +git -C "$BUILD_DIR" checkout --detach FETCH_HEAD +actual_rev=$(git -C "$BUILD_DIR" rev-parse --short=7 HEAD) +[[ "$actual_rev" == "$DYNAMO_REV" ]] || { + echo "Expected Dynamo $DYNAMO_REV, got $actual_rev" >&2 + exit 1 +} + +( + cd "$BUILD_DIR" + uv build --wheel --out-dir "$OUTPUT_DIR" + uvx --from 'maturin[patchelf]' maturin build --release --manifest-path lib/bindings/python/Cargo.toml --out "$OUTPUT_DIR" + cargo build --release --locked -p dynamo-vllm-sidecar + install -m 0755 target/release/dynamo-vllm-sidecar "$OUTPUT_DIR/dynamo-vllm-sidecar" +) diff --git a/examples/dynamo/scripts/build_vllm_wheel.sh b/examples/dynamo/scripts/build_vllm_wheel.sh new file mode 100755 index 0000000000..b11182536f --- /dev/null +++ b/examples/dynamo/scripts/build_vllm_wheel.sh @@ -0,0 +1,27 @@ +#!/usr/bin/env bash +set -euo pipefail + +VLLM_REPO=${VLLM_REPO:-https://github.com/biswapanda/vllm.git} +VLLM_REV=${VLLM_REV:-e74fc3f} +OUTPUT_DIR=${OUTPUT_DIR:-$PWD/dist/dynamo} +BUILD_DIR=${BUILD_DIR:-$PWD/.build/vllm-dynamo} +MAX_JOBS=${MAX_JOBS:-$(nproc)} + +mkdir -p "$OUTPUT_DIR" "$(dirname "$BUILD_DIR")" +if [[ ! -d "$BUILD_DIR/.git" ]]; then + git clone "$VLLM_REPO" "$BUILD_DIR" +fi +git -C "$BUILD_DIR" fetch origin "$VLLM_REV" +git -C "$BUILD_DIR" checkout --detach FETCH_HEAD +actual_rev=$(git -C "$BUILD_DIR" rev-parse --short=7 HEAD) +[[ "$actual_rev" == "$VLLM_REV" ]] || { + echo "Expected vLLM $VLLM_REV, got $actual_rev" >&2 + exit 1 +} + +( + cd "$BUILD_DIR" + uv venv --clear --python 3.12 .venv-build + uv pip install --python .venv-build/bin/python -r requirements/build/cuda.txt + MAX_JOBS="$MAX_JOBS" uv build --python .venv-build/bin/python --wheel --no-build-isolation --out-dir "$OUTPUT_DIR" +) diff --git a/examples/dynamo/scripts/install_artifacts.sh b/examples/dynamo/scripts/install_artifacts.sh new file mode 100755 index 0000000000..e2be248e7e --- /dev/null +++ b/examples/dynamo/scripts/install_artifacts.sh @@ -0,0 +1,19 @@ +#!/usr/bin/env bash +set -euo pipefail + +ARTIFACT_DIR=${ARTIFACT_DIR:-$PWD/dist/dynamo} +DYNAMO_ENV=${DYNAMO_ENV:-$PWD/.venv-dynamo} +mapfile -t vllm_wheels < <(find "$ARTIFACT_DIR" -maxdepth 1 -type f -name 'vllm-*.whl' -print) +mapfile -t dynamo_wheels < <(find "$ARTIFACT_DIR" -maxdepth 1 -type f \( -name 'ai_dynamo-*.whl' -o -name 'ai_dynamo_runtime-*.whl' \) -print) + +[[ ${#vllm_wheels[@]} -eq 1 ]] || { echo "Expected one vLLM wheel in $ARTIFACT_DIR" >&2; exit 1; } +[[ ${#dynamo_wheels[@]} -eq 2 ]] || { echo "Expected ai-dynamo and ai-dynamo-runtime wheels in $ARTIFACT_DIR" >&2; exit 1; } +[[ -x "$ARTIFACT_DIR/dynamo-vllm-sidecar" ]] || { echo "Missing dynamo-vllm-sidecar" >&2; exit 1; } + +# Keep the custom inference stack isolated: its Torch and Pydantic constraints +# differ from Prime-RL's trainer environment. +uv venv --clear --python 3.12 "$DYNAMO_ENV" +uv pip install --python "$DYNAMO_ENV/bin/python" "${vllm_wheels[0]}" "${dynamo_wheels[@]}" + +"$DYNAMO_ENV/bin/vllm" --version +"$ARTIFACT_DIR/dynamo-vllm-sidecar" --help >/dev/null diff --git a/skills/install/SKILL.md b/skills/install/SKILL.md index 4fb6694913..779212a22e 100644 --- a/skills/install/SKILL.md +++ b/skills/install/SKILL.md @@ -71,3 +71,18 @@ Binaries land in `third_party/llmd/bin/{epp,envoy,pd-sidecar}` (a shared path, s - `uv.lock` — pinned lockfile (refresh with `uv sync --all-extras`) - `scripts/install.sh` — bootstrap installer - `scripts/install_ep_kernels.sh` — DeepEP build script +## Dynamo research stack + +The Dynamo native-gRPC integration currently requires source-built artifacts; +the public vLLM 0.26 and Dynamo 1.3.0 wheels do not contain the complete tested +stack. Follow [`examples/dynamo/local/README.md`](../../examples/dynamo/local/README.md). +The pinned helper scripts are: + +```bash +examples/dynamo/scripts/build_vllm_wheel.sh +examples/dynamo/scripts/build_dynamo_artifacts.sh +examples/dynamo/scripts/install_artifacts.sh +``` + +They verify the pinned seven-character vLLM and Dynamo revisions before +building the custom artifacts and installing the inference stack in an isolated `.venv-dynamo`. diff --git a/src/prime_rl/orchestrator/orchestrator.py b/src/prime_rl/orchestrator/orchestrator.py index ab24e5504f..20dea6a76b 100644 --- a/src/prime_rl/orchestrator/orchestrator.py +++ b/src/prime_rl/orchestrator/orchestrator.py @@ -76,7 +76,6 @@ from prime_rl.trainer.rl.broadcast.nixl.model_express import ModelExpressSession from prime_rl.transport import setup_micro_batch_sender from prime_rl.utils.async_utils import EventLoopLagMonitor, EventLoopLagStats, safe_cancel -from prime_rl.utils.client import init_nccl_broadcast, init_nixl_broadcast from prime_rl.utils.heartbeat import Heartbeat from prime_rl.utils.logger import format_time, get_logger, setup_logger from prime_rl.utils.monitor import setup_monitor @@ -299,26 +298,20 @@ async def setup(self) -> None: get_logger().info(f"Initializing weight broadcast ({config.weight_broadcast})") if config.weight_broadcast.type == "nccl": - await init_nccl_broadcast( - self.policy_inference.admin_clients, - config.weight_broadcast.host, - config.weight_broadcast.port, - config.weight_broadcast.timeout, + await self.policy_inference.init_nccl_broadcast( + host=config.weight_broadcast.host, + port=config.weight_broadcast.port, + timeout=config.weight_broadcast.timeout, inference_world_size=config.weight_broadcast.inference_world_size, quantize_in_weight_transfer=config.weight_broadcast.quantize_in_weight_transfer, - engine_world_sizes=self.policy_inference._engine_world_sizes, - use_native_collective_rpc=self.policy_inference._use_native_collective_rpc, ) elif config.weight_broadcast.type == "nixl": - await init_nixl_broadcast( - self.policy_inference.admin_clients, - config.weight_broadcast.host, - config.weight_broadcast.port, - config.weight_broadcast.timeout, - config.weight_broadcast.inference_world_size, - config.weight_broadcast.session_id, - engine_world_sizes=self.policy_inference._engine_world_sizes, - use_native_collective_rpc=self.policy_inference._use_native_collective_rpc, + await self.policy_inference.init_nixl_broadcast( + host=config.weight_broadcast.host, + port=config.weight_broadcast.port, + timeout=config.weight_broadcast.timeout, + inference_world_size=config.weight_broadcast.inference_world_size, + session_id=config.weight_broadcast.session_id, ) self.model_express = ModelExpressSession( client=MxClient(server_url=f"{config.weight_broadcast.host}:{config.weight_broadcast.port}"), diff --git a/src/prime_rl/utils/client.py b/src/prime_rl/utils/client.py index f18551d6ce..b8c0c374ba 100644 --- a/src/prime_rl/utils/client.py +++ b/src/prime_rl/utils/client.py @@ -70,10 +70,11 @@ def __init__( # When admin URLs bypass a router, also health-check the client-facing # (router) endpoint - it only starts serving once its workers are healthy. self._router_clients = ( - setup_admin_clients(client_config.model_copy(update={"admin_base_url": None})) - if client_config.admin_base_url + setup_admin_clients(client_config.model_copy(update={"admin_base_url": None, "dynamo_discovery_url": None})) + if client_config.admin_base_url or client_config.is_dynamo else [] ) + self._check_router_model = client_config.is_dynamo self._skip_model_check = client_config.skip_model_check self._wait_for_ready_timeout = client_config.wait_for_ready_timeout self._scorer = PrefillScorer() @@ -99,7 +100,7 @@ async def create( client_config, model_name, expected_inference_world_size=expected_inference_world_size, - admin_clients=setup_dynamo_admin_clients(workers), + admin_clients=setup_dynamo_admin_clients(client_config, workers), engine_world_sizes=[worker.world_size for worker in workers], use_native_collective_rpc=True, **kwargs, @@ -117,7 +118,48 @@ async def wait_for_ready(self, model_name: str, timeout: int | None = None) -> N self._admin_clients + self._router_clients, timeout=timeout if timeout is not None else self._wait_for_ready_timeout, ) - await maybe_check_has_model(self._admin_clients, model_name, skip_model_check=self._skip_model_check) + model_clients = self._admin_clients + (self._router_clients if self._check_router_model else []) + await maybe_check_has_model(model_clients, model_name, skip_model_check=self._skip_model_check) + + async def init_nccl_broadcast( + self, + *, + host: str, + port: int, + timeout: int, + inference_world_size: int | None, + quantize_in_weight_transfer: bool, + ) -> None: + await init_nccl_broadcast( + self._admin_clients, + host, + port, + timeout, + inference_world_size=inference_world_size, + quantize_in_weight_transfer=quantize_in_weight_transfer, + engine_world_sizes=self._engine_world_sizes, + use_native_collective_rpc=self._use_native_collective_rpc, + ) + + async def init_nixl_broadcast( + self, + *, + host: str, + port: int, + timeout: int, + inference_world_size: int, + session_id: str, + ) -> None: + await init_nixl_broadcast( + self._admin_clients, + host, + port, + timeout, + inference_world_size, + session_id, + engine_world_sizes=self._engine_world_sizes, + use_native_collective_rpc=self._use_native_collective_rpc, + ) async def update_weights(self, weight_dir: Path | None, lora_name: str | None = None, step: int = 0) -> None: await update_weights( @@ -164,14 +206,14 @@ def setup_client( ) -def setup_admin_clients(client_config: ClientConfig) -> list[AsyncClient]: +def setup_admin_clients(client_config: ClientConfig, urls: list[str] | None = None) -> list[AsyncClient]: """Create dedicated admin clients for weight update operations. Uses a separate connection pool to avoid queueing behind streaming requests. When admin_base_url is set, uses those URLs instead of base_url, allowing weight updates to bypass routers in disaggregated P/D deployments. """ - urls = client_config.admin_base_url if client_config.admin_base_url else [client_config.base_url] + urls = urls or (client_config.admin_base_url if client_config.admin_base_url else [client_config.base_url]) def _setup_admin_client(base_url: str) -> httpx.AsyncClient: env_headers = { diff --git a/src/prime_rl/utils/dynamo.py b/src/prime_rl/utils/dynamo.py index 1cd0f05250..dc5e402e97 100644 --- a/src/prime_rl/utils/dynamo.py +++ b/src/prime_rl/utils/dynamo.py @@ -9,6 +9,7 @@ from tenacity import AsyncRetrying, retry_if_exception, stop_after_delay, wait_exponential from prime_rl.configs.shared import ClientConfig +from prime_rl.utils.client import setup_admin_clients DYNAMO_RL_DISCOVERY_PROTOCOL_VERSION = 1 DYNAMO_READINESS_REQUEST_TIMEOUT_S = 30.0 @@ -66,15 +67,11 @@ def _parse_dynamo_workers(payload: object, model_name: str) -> tuple[DiscoveredD return tuple(sorted(workers, key=lambda worker: (worker.component, worker.instance_id))) -def setup_dynamo_admin_clients(workers: tuple[DiscoveredDynamoWorker, ...]) -> list[AsyncClient]: - return [ - AsyncClient( - base_url=worker.admin_base_url.rstrip("/"), - limits=httpx.Limits(max_connections=4, max_keepalive_connections=1), - timeout=httpx.Timeout(None), - ) - for worker in workers - ] +def setup_dynamo_admin_clients( + client_config: ClientConfig, + workers: tuple[DiscoveredDynamoWorker, ...], +) -> list[AsyncClient]: + return setup_admin_clients(client_config, [worker.admin_base_url for worker in workers]) async def discover_dynamo_workers( diff --git a/tests/unit/orchestrator/test_orchestrator_setup.py b/tests/unit/orchestrator/test_orchestrator_setup.py index 784c6bc705..2885f78ac6 100644 --- a/tests/unit/orchestrator/test_orchestrator_setup.py +++ b/tests/unit/orchestrator/test_orchestrator_setup.py @@ -1,6 +1,6 @@ import asyncio from types import SimpleNamespace -from unittest.mock import MagicMock, patch +from unittest.mock import AsyncMock, patch from renderers import Qwen3VLRendererConfig @@ -26,8 +26,8 @@ async def run() -> None: with ( patch("renderers.base.create_renderer", return_value=renderer) as create_renderer_mock, patch( - "prime_rl.orchestrator.utils.InferencePool", - new=MagicMock(return_value=inference_pool), + "prime_rl.orchestrator.utils.InferencePool.create", + new=AsyncMock(return_value=inference_pool), ) as setup_pool_mock, ): returned_renderer, returned_pool = await setup_policy_inference_pool( @@ -38,7 +38,7 @@ async def run() -> None: assert returned_renderer is renderer assert returned_pool is inference_pool create_renderer_mock.assert_called_once_with(tokenizer, renderer_settings) - setup_pool_mock.assert_called_once_with( + setup_pool_mock.assert_awaited_once_with( config.model.client, model_name="policy-model", train_client_type="renderer", @@ -74,8 +74,8 @@ async def run() -> None: with ( patch("renderers.base.create_renderer", return_value=renderer) as create_renderer_mock, patch( - "prime_rl.orchestrator.utils.InferencePool", - new=MagicMock(return_value=inference_pool), + "prime_rl.orchestrator.utils.InferencePool.create", + new=AsyncMock(return_value=inference_pool), ) as setup_pool_mock, ): returned_renderer, returned_pool = await setup_policy_inference_pool( @@ -86,7 +86,7 @@ async def run() -> None: assert returned_renderer is renderer assert returned_pool is inference_pool create_renderer_mock.assert_called_once_with(tokenizer, renderer_settings) - setup_pool_mock.assert_called_once_with( + setup_pool_mock.assert_awaited_once_with( config.model.client, model_name="policy-model", train_client_type="renderer", diff --git a/tests/unit/test_configs.py b/tests/unit/test_configs.py index 5c294e66a7..e2561b0c3f 100644 --- a/tests/unit/test_configs.py +++ b/tests/unit/test_configs.py @@ -229,7 +229,7 @@ def test_external_dynamo_world_size_survives_rl_config_resolution(): "orchestrator": { "model": { "client": { - "base_url": ["http://frontend:8000/v1"], + "base_url": "http://frontend:8000/v1", "dynamo_discovery_url": "http://frontend:8001", } } @@ -256,7 +256,7 @@ def test_external_dynamo_nccl_does_not_require_a_local_inference_gpu(): "orchestrator": { "model": { "client": { - "base_url": ["http://frontend:8000/v1"], + "base_url": "http://frontend:8000/v1", "dynamo_discovery_url": "http://frontend:8001", } } @@ -280,23 +280,6 @@ def test_external_dynamo_nccl_does_not_require_a_local_inference_gpu(): assert config.trainer.weight_broadcast.inference_world_size == 1 -def test_default_nccl_world_size_does_not_bypass_local_gpu_guard(): - with pytest.raises(ValueError, match="NCCL weight broadcast requires at least 2"): - RLConfig.model_validate( - { - "trainer": {}, - "orchestrator": {}, - "inference": None, - "deployment": { - "type": "single_node", - "num_train_gpus": 1, - "num_infer_gpus": 0, - }, - "weight_broadcast": {"type": "nccl"}, - } - ) - - def test_external_dynamo_lora_world_size_survives_filesystem_config_resolution(): config = RLConfig.model_validate( { @@ -304,7 +287,7 @@ def test_external_dynamo_lora_world_size_survives_filesystem_config_resolution() "orchestrator": { "model": { "client": { - "base_url": ["http://frontend:8000/v1"], + "base_url": "http://frontend:8000/v1", "dynamo_discovery_url": "http://frontend:8001", } } @@ -323,7 +306,7 @@ def test_dynamo_orchestrator_requires_explicit_inference_world_size(): { "model": { "client": { - "base_url": ["http://frontend:8000/v1"], + "base_url": "http://frontend:8000/v1", "dynamo_discovery_url": "http://frontend:8001", } }, diff --git a/tests/unit/utils/test_dynamo.py b/tests/unit/utils/test_dynamo.py index c492c83f63..c9bc15d146 100644 --- a/tests/unit/utils/test_dynamo.py +++ b/tests/unit/utils/test_dynamo.py @@ -5,7 +5,7 @@ import pytest from prime_rl.configs.shared import ClientConfig -from prime_rl.utils.dynamo import _parse_dynamo_workers, discover_dynamo_workers +from prime_rl.utils.dynamo import _parse_dynamo_workers, discover_dynamo_workers, setup_dynamo_admin_clients MODEL = "Qwen/Qwen3-0.6B" @@ -101,3 +101,19 @@ def test_discovery_retries_until_expected_world_size_is_complete(): assert discovery_client.get.await_count == 3 assert [item.component for item in workers] == ["backend", "prefill"] + + +def test_discovered_admin_clients_preserve_configured_headers(monkeypatch): + monkeypatch.setenv("DYNAMO_TOKEN", "secret") + clients = setup_dynamo_admin_clients( + ClientConfig( + headers={"X-Static": "value"}, + headers_from_env={"X-Token": "DYNAMO_TOKEN"}, + ), + _parse_dynamo_workers(payload(worker()), MODEL), + ) + try: + assert clients[0].headers["X-Static"] == "value" + assert clients[0].headers["X-Token"] == "secret" + finally: + asyncio.run(clients[0].aclose()) From 33d5cdafe8fd118f4d2b80be00be3102971cd6ff Mon Sep 17 00:00:00 2001 From: sami jaghouar Date: Thu, 13 Aug 2026 20:57:39 +0000 Subject: [PATCH 9/9] feat(dynamo): add opt-in Slurm deployment --- examples/dynamo/README.md | 1 + examples/dynamo/slurm-managed/README.md | 99 +++++++++++++++++++ examples/dynamo/slurm-managed/rl.toml | 53 ++++++++++ .../src/prime_rl/configs/inference.py | 35 +++++++ .../src/prime_rl/configs/rl.py | 84 ++++++++++++++++ skills/training/monitor-run/SKILL.md | 14 +++ skills/training/start-run/SKILL.md | 9 ++ src/prime_rl/entrypoints/rl.py | 18 ++++ src/prime_rl/templates/_launch_dynamo.sh.j2 | 64 ++++++++++++ .../templates/multi_node_rl.sbatch.j2 | 49 +++++++++ tests/unit/test_configs.py | 47 +++++++++ 11 files changed, 473 insertions(+) create mode 100644 examples/dynamo/slurm-managed/README.md create mode 100644 examples/dynamo/slurm-managed/rl.toml create mode 100644 src/prime_rl/templates/_launch_dynamo.sh.j2 diff --git a/examples/dynamo/README.md b/examples/dynamo/README.md index e6750d634b..73a248582a 100644 --- a/examples/dynamo/README.md +++ b/examples/dynamo/README.md @@ -87,6 +87,7 @@ launched. - [`local`](local): single-node 2-GPU smoke test with a real Dynamo stack (etcd + frontend + sidecar + vLLM engine). +- [`slurm-managed`](slurm-managed): launcher-managed aggregated Dynamo deployment on dedicated inference nodes. - [`qwen3_06b_math`](qwen3_06b_math): single-GPU trainer and aggregate Dynamo inference smoke test. - [`qwen3_30b_Thinking`](qwen3_30b_Thinking): Qwen3-30B Thinking math with an diff --git a/examples/dynamo/slurm-managed/README.md b/examples/dynamo/slurm-managed/README.md new file mode 100644 index 0000000000..5836a7ab1b --- /dev/null +++ b/examples/dynamo/slurm-managed/README.md @@ -0,0 +1,99 @@ +# Dynamo on SLURM + +This recipe gives Prime-RL one trainer node and one Dynamo inference node. The +SLURM launcher starts, supervises, and cleans up: + +- job-local etcd; +- one Dynamo frontend with the RL discovery listener; +- one native-gRPC vLLM engine per local DP rank; +- one `dynamo-vllm-sidecar` per engine; +- the Prime trainer, orchestrator, and environment server. + +## Prepare artifacts once on the shared checkout + +```bash +git submodule update --init --recursive +uv sync --all-extras --all-packages +examples/dynamo/scripts/build_vllm_wheel.sh +examples/dynamo/scripts/build_dynamo_artifacts.sh +examples/dynamo/scripts/install_artifacts.sh +``` + +The inference environment lives at `.venv-dynamo`; Prime-RL continues to use +`.venv`. Both paths and `dist/dynamo/dynamo-vllm-sidecar` must be visible on all +allocated nodes. + +## Configure and submit + +Edit `partition`, `account`, and any cluster-specific Slurm settings in +`rl.toml`, then run: + +```bash +uv run rl @ examples/dynamo/slurm-managed/rl.toml --output-dir outputs/dynamo-slurm --clean-output-dir +``` + +Preview the rendered job without submitting: + +```bash +uv run rl @ examples/dynamo/slurm-managed/rl.toml --output-dir /tmp/dynamo-slurm --clean-output-dir --dry-run +bash -n /tmp/dynamo-slurm/rl.sbatch +``` + +`[dynamo] enabled = true` replaces the normal global `vllm-router` launch. The +orchestrator uses the Dynamo frontend for generation and receives +`dynamo_discovery_url` from the generated job script. Direct admin endpoints and +world sizes come from `/v1/rl/workers`. + + +## Configuration reference + +The `[dynamo]` block supports: + +```toml +[dynamo] +enabled = true +env_path = ".venv-dynamo" +sidecar_path = "dist/dynamo/dynamo-vllm-sidecar" +etcd_path = "etcd" +discovery_port = 8001 +etcd_port = 2379 +etcd_peer_port = 2380 +engine_port = 8100 +grpc_port = 50051 +system_port = 9000 +namespace = "prime-rl" +``` + +The launcher appends the Slurm job ID to the namespace, starts job-local etcd on +inference node zero, and advertises each engine's node-reachable HTTP admin URL. +The native gRPC and engine HTTP base ports are reused on different nodes; local +DP ranks use consecutive ports. The custom `vllm-rs` splitter keeps frontend +flags on the Rust process and forwards vLLM engine flags to its managed Python +engine. Override these values when your cluster reserves +one of the defaults. + +## Scale inference + +Each inference node starts `gpus_per_node / tensor_parallel_size` independent +engines. Therefore: + +```text +inference_world_size = num_infer_nodes * gpus_per_node +``` + +for the aggregated topology. Keep `weight_broadcast.inference_world_size` equal +to the sum of worker `world_size` values returned by one discovery snapshot. + +This initial Slurm path supports aggregated dense inference. Disaggregated +prefill/decode and cross-engine expert-parallel deployment remain unsupported by +the launcher. + +## Logs + +```text +outputs/dynamo-slurm/logs/inference/dynamo-etcd.log +outputs/dynamo-slurm/logs/inference/dynamo-frontend.log +outputs/dynamo-slurm/logs/inference/node_0.log +outputs/dynamo-slurm/logs/orchestrator.log +outputs/dynamo-slurm/logs/trainer/node_0.log +``` diff --git a/examples/dynamo/slurm-managed/rl.toml b/examples/dynamo/slurm-managed/rl.toml new file mode 100644 index 0000000000..505b4428d9 --- /dev/null +++ b/examples/dynamo/slurm-managed/rl.toml @@ -0,0 +1,53 @@ +max_steps = 4 +seq_len = 2048 + +[model] +name = "Qwen/Qwen3-0.6B" + +[deployment] +type = "multi_node" +num_train_nodes = 1 +num_infer_nodes = 1 +gpus_per_node = 1 + +[slurm] +project_dir = "." +partition = "cluster" +time = "02:00:00" + +[dynamo] +enabled = true + +[weight_broadcast] +type = "nccl" +inference_world_size = 1 + +[trainer] + +[orchestrator] +batch_size = 32 +group_size = 4 + +[orchestrator.renderer] +name = "auto" + +[orchestrator.train.sampling] +max_completion_tokens = 2048 + +[[orchestrator.train.source]] +name = "math" +env.taskset = { id = "math-env", dataset_name = "openai/gsm8k", dataset_subset = "main" } +env.agent.harness = { id = "null" } +env.agent.runtime = { type = "subprocess" } + +[inference] + +[inference.deployment] +type = "multi_node" +num_nodes = 1 + +[inference.vllm] +model = "Qwen/Qwen3-0.6B" +tensor_parallel_size = 1 +max_model_len = 2048 +enforce_eager = true diff --git a/packages/prime-rl-configs/src/prime_rl/configs/inference.py b/packages/prime-rl-configs/src/prime_rl/configs/inference.py index d551ac7243..76db0aca82 100644 --- a/packages/prime-rl-configs/src/prime_rl/configs/inference.py +++ b/packages/prime-rl-configs/src/prime_rl/configs/inference.py @@ -1,4 +1,6 @@ +import base64 import json +import shlex from argparse import Namespace from pathlib import Path from typing import Annotated, Any, Literal, TypeAlias @@ -552,6 +554,39 @@ def auto_setup_slurm_template(self): self.slurm.template_path = templates_dir / "inference.sbatch.j2" return self + @property + def dynamo_vllm_args(self) -> str: + """Shell-quoted managed-engine arguments for the pinned Dynamo vLLM launcher.""" + args: list[str] = [] + values = self.to_namespace() + skip = { + "model", + "host", + "port", + "liveness_timeout_seconds", + "max_model_len", + "tensor_parallel_size", + "data_parallel_size", + "data_parallel_size_local", + "data_parallel_rpc_port", + "api_server_count", + "worker_extension_cls", + } + for key, value in vars(values).items(): + if key in skip or value is None or value is False: + continue + flag = f"--{key.replace('_', '-')}" + if value is True: + args.append(flag) + else: + encoded = json.dumps(value, separators=(",", ":")) if isinstance(value, (dict, list)) else str(value) + args.extend((flag, encoded)) + return " ".join(shlex.quote(arg) for arg in args) + + @property + def dynamo_vllm_args_b64(self) -> str: + return base64.b64encode(self.dynamo_vllm_args.encode()).decode() + def build_kv_transfer_config(self) -> dict[str, Any] | None: """Build the single vLLM ``kv_transfer_config`` from the transfer + offload connectors. diff --git a/packages/prime-rl-configs/src/prime_rl/configs/rl.py b/packages/prime-rl-configs/src/prime_rl/configs/rl.py index e8a46137fc..e09bcb2c37 100644 --- a/packages/prime-rl-configs/src/prime_rl/configs/rl.py +++ b/packages/prime-rl-configs/src/prime_rl/configs/rl.py @@ -225,6 +225,57 @@ def total_infer_nodes(self) -> int: ] +class DynamoSlurmConfig(BaseConfig): + """SLURM-managed Dynamo frontend, sidecar, and native-gRPC vLLM engines.""" + + enabled: bool = False + """Launch Dynamo instead of Prime-RL's vLLM router and inference entrypoint.""" + + env_path: Path = Path(".venv-dynamo") + """Inference environment created by examples/dynamo/scripts/install_artifacts.sh.""" + + sidecar_path: Path = Path("dist/dynamo/dynamo-vllm-sidecar") + """Pinned Dynamo sidecar executable built by examples/dynamo/scripts/build_dynamo_artifacts.sh.""" + + etcd_path: str = "etcd" + """etcd executable available on inference nodes.""" + + discovery_port: int = Field(8001, ge=1, le=65535) + """Dynamo RL worker discovery port.""" + + etcd_port: int = Field(2379, ge=1, le=65535) + """Dynamo etcd client port.""" + + etcd_peer_port: int = Field(2380, ge=1, le=65535) + """Node-local etcd peer port.""" + + engine_port: int = Field(8100, ge=1, le=65535) + """Base vLLM HTTP/admin port; local DP ranks use consecutive ports.""" + + grpc_port: int = Field(50051, ge=1, le=65535) + """Base native vLLM gRPC port; local DP ranks use consecutive ports.""" + + system_port: int = Field(9000, ge=1, le=65535) + """Base Dynamo sidecar system port; local DP ranks use consecutive ports.""" + + namespace: str = "prime-rl" + """Dynamo discovery namespace; the SLURM job ID is appended at launch time.""" + + @model_validator(mode="after") + def validate_ports(self): + ports = [ + self.discovery_port, + self.etcd_port, + self.etcd_peer_port, + self.engine_port, + self.grpc_port, + self.system_port, + ] + if len(set(ports)) != len(ports): + raise ValueError("Dynamo SLURM ports must be distinct") + return self + + class RLConfig(BaseConfig): trainer: TrainerConfig @@ -280,11 +331,40 @@ class RLConfig(BaseConfig): slurm: SlurmConfig | None = None """SLURM configuration. If None, runs locally.""" + dynamo: DynamoSlurmConfig = DynamoSlurmConfig() + """Optional SLURM-managed Dynamo inference stack.""" + dry_run: bool = False """Only validate and dump resolved configs, then exit early.""" ### Validate configs (e.g. raise for unsupported (combinations of) configs) + @model_validator(mode="after") + def validate_dynamo_slurm(self): + if not self.dynamo.enabled: + return self + if self.slurm is None: + raise ValueError("dynamo.enabled requires [slurm]") + if self.deployment.type != "multi_node": + raise ValueError("Dynamo SLURM launch currently requires deployment.type = 'multi_node'") + if not self.slurm.shared_fs: + raise ValueError("Dynamo SLURM launch currently requires slurm.shared_fs = true") + if "template_path" in self.slurm.model_fields_set: + raise ValueError("Dynamo SLURM launch requires the default multi-node template") + if self.inference is None: + raise ValueError("dynamo.enabled requires [inference] for the managed vLLM topology") + if self.inference.deployment.type == "disaggregated": + raise ValueError("Dynamo SLURM launch currently supports aggregated inference only") + if self.deployment.gpus_per_node % self.inference.vllm.tensor_parallel_size: + raise ValueError("deployment.gpus_per_node must be divisible by inference.vllm.tensor_parallel_size") + if self.inference.vllm.enable_expert_parallel: + raise ValueError("Dynamo SLURM launch does not yet support cross-engine expert parallelism") + if self.inference.router is not None and "router" in self.inference.model_fields_set: + raise ValueError("Dynamo replaces inference.router; remove the explicit [inference.router] block") + if self.orchestrator.model.client.admin_base_url is not None: + raise ValueError("Dynamo discovers admin endpoints; remove orchestrator.model.client.admin_base_url") + return self + @model_validator(mode="after") def auto_setup_infer_nodes(self): if self.deployment.type != "multi_node": @@ -768,12 +848,16 @@ def auto_setup_inference_client(self): if self.inference is None: return self client = self.orchestrator.model.client + if self.dynamo.enabled: + client.dynamo_discovery_url = f"http://localhost:{self.dynamo.discovery_port}" + client.admin_base_url = None if not self.orchestrator.any_policy_sourced and "base_url" not in client.model_fields_set: host = self.inference.server.host or "localhost" port = self.inference.server.port client.base_url = f"http://{host}:{port}/v1" if ( self.deployment.type == "single_node" + and not self.dynamo.enabled and self.inference.router is not None and "admin_base_url" not in client.model_fields_set ): diff --git a/skills/training/monitor-run/SKILL.md b/skills/training/monitor-run/SKILL.md index 084d1610dd..8f2069e1c4 100644 --- a/skills/training/monitor-run/SKILL.md +++ b/skills/training/monitor-run/SKILL.md @@ -189,3 +189,17 @@ PRIME-RL::Launcher ``` For multi-node runs, trainer and inference processes are on separate nodes — use `srun` or `ssh` to inspect them. + +### Dynamo SLURM inference logs + +For `[dynamo] enabled = true`, the global router log is replaced by: + +```text +logs/inference/dynamo-etcd.log +logs/inference/dynamo-frontend.log +logs/inference/node_.log +``` + +The node log contains both the native-gRPC vLLM engine and its sidecar. Check +`/v1/models` on the frontend port and `/v1/rl/workers` on the configured +Dynamo discovery port before diagnosing Prime weight-transfer startup. diff --git a/skills/training/start-run/SKILL.md b/skills/training/start-run/SKILL.md index 8356bc20f3..f3e5642a2d 100644 --- a/skills/training/start-run/SKILL.md +++ b/skills/training/start-run/SKILL.md @@ -92,3 +92,12 @@ curl http://localhost:8000/v1/chat/completions \ - `packages/prime-rl-configs/src/prime_rl/configs/` — all config classes - `configs/debug/` — minimal debug configs - `examples/` — full example configs (e.g. `reverse-text/`) + +## Dynamo on SLURM + +Set `[dynamo] enabled = true` in a multi-node RL config to replace the normal +SLURM `vllm-router` deployment with a launcher-managed aggregated Dynamo stack. +Build the pinned inference artifacts first, then use +`examples/dynamo/slurm-managed/rl.toml`; the full workflow is documented in +`examples/dynamo/slurm-managed/README.md`. Dynamo currently requires dedicated inference +nodes and does not support the launcher's disaggregated prefill/decode mode. diff --git a/src/prime_rl/entrypoints/rl.py b/src/prime_rl/entrypoints/rl.py index 4dab3d9911..dec80ff8fe 100644 --- a/src/prime_rl/entrypoints/rl.py +++ b/src/prime_rl/entrypoints/rl.py @@ -472,6 +472,9 @@ def write_slurm_script(config: RLConfig, config_dir: Path, script_path: Path) -> config_path=config_dir / RL_TOML, output_dir=config.output_dir, gpus_per_node=config.deployment.gpus_per_node, + dynamo=config.dynamo, + inference=config.inference, + inference_env_vars=inference_env_vars, ) elif config.inference is not None and config.inference.deployment.type == "disaggregated": infer_deploy = config.inference.deployment @@ -514,6 +517,7 @@ def write_slurm_script(config: RLConfig, config_dir: Path, script_path: Path) -> orchestrator_on_inference=config.deployment.orchestrator_on_inference, train_env_names=train_env_names, eval_env_names=eval_env_names, + dynamo=config.dynamo, ) else: script = template.render( @@ -550,6 +554,8 @@ def write_slurm_script(config: RLConfig, config_dir: Path, script_path: Path) -> inference_env_vars=inference_env_vars, train_env_names=train_env_names, eval_env_names=eval_env_names, + dynamo=config.dynamo, + inference=config.inference, ) script_path.parent.mkdir(parents=True, exist_ok=True) @@ -566,6 +572,18 @@ def rl_slurm(config: RLConfig): config_dir = config.output_dir / "configs" log_dir = get_log_dir(config.output_dir) + if config.dynamo.enabled and not config.dry_run: + dynamo_python = config.slurm.project_dir / config.dynamo.env_path / "bin" / "python" + vllm_rs = config.slurm.project_dir / config.dynamo.env_path / "bin" / "vllm-rs" + sidecar = config.slurm.project_dir / config.dynamo.sidecar_path + missing = [path for path in (dynamo_python, vllm_rs, sidecar) if not path.exists()] + if missing: + paths = ", ".join(map(str, missing)) + raise FileNotFoundError( + f"Missing Dynamo SLURM artifacts: {paths}. Run examples/dynamo/scripts/build_vllm_wheel.sh, " + "build_dynamo_artifacts.sh, and install_artifacts.sh first." + ) + if config.deployment.type == "single_node": write_config(config, config_dir, exclude={"slurm", "dry_run", "clean_output_dir"}) logger.info(f"Wrote config to {config_dir / RL_TOML}") diff --git a/src/prime_rl/templates/_launch_dynamo.sh.j2 b/src/prime_rl/templates/_launch_dynamo.sh.j2 new file mode 100644 index 0000000000..cb55036799 --- /dev/null +++ b/src/prime_rl/templates/_launch_dynamo.sh.j2 @@ -0,0 +1,64 @@ +{# Launch helpers for an aggregated Dynamo topology. One frontend + etcd run on + inference node 0. Every local DP rank runs one native-gRPC vLLM engine and + one Dynamo sidecar. #} +launch_dynamo_control_plane() { + local log="$OUTPUT_DIR/logs/inference/dynamo-frontend.log" + local etcd_data="/tmp/prime-rl-dynamo-etcd-$SLURM_JOB_ID" + echo "Starting Dynamo control plane on $LOCAL_IP:$DYNAMO_FRONTEND_PORT (discovery $DYNAMO_DISCOVERY_PORT)" | tee "$log" + rm -rf "$etcd_data" + "$DYNAMO_ETCD_BIN" --name "dynamo-$SLURM_JOB_ID" --data-dir "$etcd_data" \ + --listen-client-urls "http://0.0.0.0:$DYNAMO_ETCD_PORT" \ + --advertise-client-urls "http://$LOCAL_IP:$DYNAMO_ETCD_PORT" \ + --listen-peer-urls "http://127.0.0.1:$DYNAMO_ETCD_PEER_PORT" \ + --initial-advertise-peer-urls "http://127.0.0.1:$DYNAMO_ETCD_PEER_PORT" \ + --initial-cluster "dynamo-$SLURM_JOB_ID=http://127.0.0.1:$DYNAMO_ETCD_PEER_PORT" \ + --initial-cluster-state new >> "$OUTPUT_DIR/logs/inference/dynamo-etcd.log" 2>&1 & + local deadline=$((SECONDS + 60)) + until curl -fsS "$DYNAMO_ETCD_ENDPOINT/health" >/dev/null; do + [ "$SECONDS" -lt "$deadline" ] || { echo "Timed out waiting for Dynamo etcd" >&2; return 1; } + sleep 1 + done + DYN_ENABLE_RL=true DYN_RL_PORT="$DYNAMO_DISCOVERY_PORT" \ + DYN_NAMESPACE="$DYNAMO_NAMESPACE" DYN_DISCOVERY_BACKEND=etcd \ + ETCD_ENDPOINTS="$DYNAMO_ETCD_ENDPOINT" \ + DYN_VLLM_ENABLE_INFERENCE_V1_GENERATE=true \ + "$DYNAMO_PYTHON" -m dynamo.frontend \ + --http-host 0.0.0.0 --http-port "$DYNAMO_FRONTEND_PORT" \ + --namespace "$DYNAMO_NAMESPACE" --discovery-backend etcd \ + --request-plane tcp --event-plane zmq \ + --router-min-initial-workers "$DYNAMO_TOTAL_ENGINES" >> "$log" 2>&1 & +} + +# launch_dynamo_rank +launch_dynamo_rank() { + local local_rank="$1" global_rank="$2" http_port="$3" grpc_port="$4" gpus="$5" log="$6" + local system_port=$((DYNAMO_SYSTEM_PORT + global_rank)) + local rpcbase="/tmp/vllm-rpc-$USER-$SLURM_JOB_ID-$http_port" + mkdir -p "$rpcbase" + echo "Starting Dynamo engine global_rank=$global_rank http=$http_port grpc=$grpc_port GPUs=$gpus" | tee -a "$log" + local -a vllm_args=() + if [ -n "$DYNAMO_VLLM_ARGS" ]; then + eval "vllm_args=($DYNAMO_VLLM_ARGS)" + fi + local -a managed_args=() + [ -n "$DYNAMO_MAX_MODEL_LEN" ] && managed_args+=(--max-model-len "$DYNAMO_MAX_MODEL_LEN") + + CUDA_VISIBLE_DEVICES="$gpus" VLLM_SERVER_DEV_MODE=1 \ + PYTHONPATH="$PROJECT_DIR/src${PYTHONPATH:+:$PYTHONPATH}" \ + VLLM_RPC_BASE_PATH="$rpcbase" \ + "$DYNAMO_VLLM_RS" serve "$DYNAMO_MODEL" \ + --host 0.0.0.0 --port "$http_port" --grpc-port "$grpc_port" \ + --python "$DYNAMO_PYTHON" --data-parallel-size 1 "${managed_args[@]}" \ + --tensor-parallel-size "$INFERENCE_TP" \ + --worker-extension-cls prime_rl.inference.vllm.worker.nccl.NCCLWeightUpdateWorker \ + "${vllm_args[@]}" >> "$log" 2>&1 & + + DYN_NAMESPACE="$DYNAMO_NAMESPACE" DYN_DISCOVERY_BACKEND=etcd \ + ETCD_ENDPOINTS="$DYNAMO_ETCD_ENDPOINT" DYN_ENABLE_RL=true \ + DYN_SYSTEM_PORT="$system_port" VLLM_HTTP_ENDPOINT="http://$LOCAL_IP:$http_port" \ + "$DYNAMO_SIDECAR" --vllm-endpoint "$LOCAL_IP:$grpc_port" \ + --admin-endpoint "http://$LOCAL_IP:$http_port" \ + --model-path "$DYNAMO_MODEL" --namespace "$DYNAMO_NAMESPACE" \ + --component "backend-$global_rank" \ + --rl-discovery-model-name "$DYNAMO_MODEL" >> "$log" 2>&1 & +} diff --git a/src/prime_rl/templates/multi_node_rl.sbatch.j2 b/src/prime_rl/templates/multi_node_rl.sbatch.j2 index 33da3e7a00..9ff04efe3c 100755 --- a/src/prime_rl/templates/multi_node_rl.sbatch.j2 +++ b/src/prime_rl/templates/multi_node_rl.sbatch.j2 @@ -58,6 +58,26 @@ export INFERENCE_ENABLE_EXPERT_PARALLEL={{ "1" if inference_enable_expert_parall export PROJECT_DIR={{ project_dir }} export CONFIG_DIR={{ config_dir }} export OUTPUT_DIR={{ output_dir }} +{% if dynamo.enabled -%} +export DYNAMO_HOME="$PROJECT_DIR/{{ dynamo.env_path }}" +export DYNAMO_PYTHON="$DYNAMO_HOME/bin/python" +export DYNAMO_VLLM_RS="$DYNAMO_HOME/bin/vllm-rs" +export DYNAMO_SIDECAR="$PROJECT_DIR/{{ dynamo.sidecar_path }}" +export DYNAMO_ETCD_BIN="{{ dynamo.etcd_path }}" +export DYNAMO_FRONTEND_PORT={{ router_port }} +export DYNAMO_DISCOVERY_PORT={{ dynamo.discovery_port }} +export DYNAMO_ETCD_PORT={{ dynamo.etcd_port }} +export DYNAMO_ETCD_PEER_PORT={{ dynamo.etcd_peer_port }} +export DYNAMO_ENGINE_PORT={{ dynamo.engine_port }} +export DYNAMO_GRPC_PORT={{ dynamo.grpc_port }} +export DYNAMO_SYSTEM_PORT={{ dynamo.system_port }} +export DYNAMO_NAMESPACE="{{ dynamo.namespace }}-$SLURM_JOB_ID" +export DYNAMO_MODEL="{{ inference.vllm.model }}" +export DYNAMO_MAX_MODEL_LEN="{{ inference.vllm.max_model_len or '' }}" +export DYNAMO_TOTAL_ENGINES=$((NUM_INFER_NODES * INFERENCE_DP_LOCAL)) +export DYNAMO_VLLM_ARGS_B64="{{ inference.dynamo_vllm_args_b64 }}" +export DYNAMO_VLLM_ARGS="$(printf %s "$DYNAMO_VLLM_ARGS_B64" | base64 -d)" +{%- endif %} mkdir -p $OUTPUT_DIR/logs/trainer $OUTPUT_DIR/logs/inference rm -f $OUTPUT_DIR/logs/inference/*.log ln -sfn trainer/node_0.log $OUTPUT_DIR/logs/trainer.log @@ -77,6 +97,11 @@ INFER_HOSTS=( ${HOSTNAMES[@]:0:$NUM_INFER_NODES} ) # Single global router on inference node 0 fronts every per-rank endpoint; # admin URLs bypass it and hit each rank directly. INFER_URLS="http://${INFER_HOSTS[0]}:${ROUTER_PORT}/v1" +{% if dynamo.enabled -%} +DYNAMO_DISCOVERY_URL="http://${INFER_HOSTS[0]}:${DYNAMO_DISCOVERY_PORT}" +DYNAMO_ETCD_ENDPOINT="http://${INFER_HOSTS[0]}:${DYNAMO_ETCD_PORT}" +export DYNAMO_DISCOVERY_URL DYNAMO_ETCD_ENDPOINT +{%- endif %} ADMIN_URLS="" ROUTER_ARGS="" {% if is_disaggregated -%} @@ -169,6 +194,11 @@ srun bash -s <<'CLEANUP_SH' pkill -9 -f "[p]ython.*prime_rl" 2>/dev/null pkill -9 -f "[t]orchrun" 2>/dev/null pkill -9 -f "[v]llm-router" 2>/dev/null +{% if dynamo.enabled -%} + pkill -9 -f "[d]ynamo.frontend" 2>/dev/null + pkill -9 -f "[d]ynamo-vllm-sidecar" 2>/dev/null + pkill -9 -f "[e]tcd.*prime-rl-dynamo" 2>/dev/null +{%- endif %} # llm-d router procs (EPP + Envoy + decode pd-sidecar) left over from a prior run. pkill -9 -f "[e]pp.*pool-name" 2>/dev/null pkill -9 -f "[e]nvoy.*envoy.yaml" 2>/dev/null @@ -245,6 +275,9 @@ if [ "$SLURM_PROCID" -lt "$NUM_INFER_NODES" ]; then CONFIG_PATH="$CONFIG_DIR/inference.toml" {% include '_launch_rank.sh.j2' %} {% include '_launch_router.sh.j2' %} +{% if dynamo.enabled -%} +{% include '_launch_dynamo.sh.j2' %} +{%- endif %} {% if is_disaggregated -%} # PD disaggregated: NIXL KV transfer + role-based inference @@ -349,6 +382,17 @@ if [ "$SLURM_PROCID" -lt "$NUM_INFER_NODES" ]; then launch_inference_rank "$((PORT + k))" "$RANK_GPUS" "$DP" "$INFERENCE_DATA_PARALLEL_RPC_PORT" "$RANK_EXTRA" "$((5600 + k))" "$(rank_log "$INFER_NODE_RANK" "$k")" done {%- else -%} +{% if dynamo.enabled -%} + if [ "$INFER_NODE_RANK" -eq 0 ]; then + launch_dynamo_control_plane + fi + for ((d=0; d