Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
14 changes: 14 additions & 0 deletions packages/prime-rl-configs/src/prime_rl/configs/orchestrator.py
Original file line number Diff line number Diff line change
Expand Up @@ -403,6 +403,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"
Expand Down Expand Up @@ -572,6 +575,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
Expand Down
45 changes: 32 additions & 13 deletions packages/prime-rl-configs/src/prime_rl/configs/rl.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 inference ranks when inference is managed externally."""


class SharedNIXLWeightBroadcastConfig(SharedInMemoryWeightBroadcastConfig):
type: Literal["nixl"] = "nixl"
Expand All @@ -148,10 +151,16 @@ class SharedNIXLWeightBroadcastConfig(SharedInMemoryWeightBroadcastConfig):
session_id: str = "default"
"""ModelExpress session ID."""

inference_world_size: int | None = Field(None, ge=1)
"""Expected inference ranks when inference is managed externally."""


class SharedFileSystemWeightBroadcastConfig(BaseConfig):
type: Literal["filesystem"] = "filesystem"

inference_world_size: int | None = Field(None, ge=1)
"""Expected inference ranks when inference is managed externally (e.g. Dynamo LoRA over filesystem)."""


SharedWeightBroadcastConfig: TypeAlias = Annotated[
SharedFileSystemWeightBroadcastConfig | SharedNCCLWeightBroadcastConfig | SharedNIXLWeightBroadcastConfig,
Expand Down Expand Up @@ -323,16 +332,6 @@ def validate_deployment(self):
)
return 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."
)
return self

@model_validator(mode="after")
def validate_quantize_in_weight_transfer(self):
if not isinstance(self.weight_broadcast, SharedNCCLWeightBroadcastConfig):
Expand Down Expand Up @@ -393,13 +392,18 @@ def auto_setup_weight_broadcast(self):
"Set weight_broadcast.type = 'filesystem'."
)
if self.weight_broadcast.type in ("nccl", "nixl"):
inference_world_size = self.inference.parallel.dp * self.inference.parallel.tp if self.inference else 1
inference_world_size = (
self.inference.parallel.dp * self.inference.parallel.tp
if self.inference
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,
Expand All @@ -414,7 +418,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)

Expand Down Expand Up @@ -444,6 +450,19 @@ def auto_setup_rollout_transport(self):
self.rollout_transport = self.trainer.rollout_transport
return self

@model_validator(mode="after")
def validate_enough_devices_for_nccl(self):
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
local_inference_gpus = self.deployment.num_infer_gpus if self.inference is not None else 0
if self.deployment.num_train_gpus + local_inference_gpus < 2:
raise ValueError(
"NCCL weight broadcast requires at least 2 local GPUs or an explicit external inference_world_size."
)
return self

@model_validator(mode="after")
def validate_eplb_requires_quantized_weight_transfer(self):
if self.inference is None or not self.inference.enable_eplb:
Expand Down
18 changes: 17 additions & 1 deletion packages/prime-rl-configs/src/prime_rl/configs/shared.py
Original file line number Diff line number Diff line change
@@ -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

Expand Down Expand Up @@ -144,17 +144,33 @@ 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 disaggregated P/D deployments where the router must not handle admin traffic."""

dynamo_discovery_url: str | None = None
"""Dynamo discovery URL. When set, Prime discovers vLLM admin endpoints and per-engine world sizes from ``/v1/rl/workers`` instead of requiring ``admin_base_url`` entries."""

elastic: ElasticConfig | None = None
"""Elastic inference pool config for DNS-based service discovery. When set, ``base_url`` is ignored and inference servers are discovered dynamically via DNS."""

router_url: str | None = None
"""vllm-router URL for load-aware inference routing. With elastic mode, inference requests go through the router while admin ops still hit discovered pods directly."""

@model_validator(mode="after")
def validate_pool_mode(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")
if self.dynamo_discovery_url is not None and self.elastic is not None:
raise ValueError("dynamo_discovery_url cannot be combined with elastic discovery")
return self

@property
def is_elastic(self) -> bool:
"""Check if elastic mode is enabled."""
return self.elastic is not None

@property
def is_dynamo(self) -> bool:
"""Check if Dynamo worker discovery is enabled."""
return self.dynamo_discovery_url is not None


class LogConfig(BaseConfig):
level: str = Field(default_factory=lambda: os.environ.get("PRIME_LOG_LEVEL", "info"))
Expand Down
3 changes: 3 additions & 0 deletions packages/prime-rl-configs/src/prime_rl/utils/validation.py
Original file line number Diff line number Diff line change
Expand Up @@ -129,6 +129,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")
Expand Down
39 changes: 39 additions & 0 deletions src/prime_rl/inference/vllm/ranks.py
Original file line number Diff line number Diff line change
@@ -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
12 changes: 11 additions & 1 deletion src/prime_rl/inference/vllm/server.py
Original file line number Diff line number Diff line change
Expand Up @@ -143,11 +143,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"}

Expand Down
61 changes: 39 additions & 22 deletions src/prime_rl/inference/vllm/worker/nccl.py
Original file line number Diff line number Diff line change
Expand Up @@ -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_weights_checkpoint_layerwise,
load_weights_kernel,
Expand Down Expand Up @@ -78,15 +79,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):
Expand All @@ -101,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,
Expand All @@ -144,15 +162,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,
)
20 changes: 18 additions & 2 deletions src/prime_rl/inference/vllm/worker/nixl.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 (
Expand Down Expand Up @@ -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))
Expand Down
Loading