Skip to content
Merged
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
4 changes: 4 additions & 0 deletions docs/scaling.md
Original file line number Diff line number Diff line change
Expand Up @@ -245,6 +245,10 @@ The generated job starts a job-scoped ModelExpress server and Redis backend on t

The launcher requires the CUDA and InfiniBand transports from `third_party/ucx`. Each NIXL process selects the active InfiniBand port nearest its GPU; an explicitly configured `UCX_NET_DEVICES` takes precedence. Inference ranks start their pulls at different trainer ranks so concurrent workers distribute traffic across all available source rails.

ModelExpress exchanges peer metadata during startup. Weight updates reuse prepared NIXL requests, post every trainer-rank read in a transfer group concurrently, and use versioned NIXL notifications for source readiness and buffer credits.

By default, the trainer and inference worker each allocate one transfer arena. Set `weight_broadcast.overlap_transfer_and_replay = true` to allocate two arenas on both sides and replay one weight group while receiving the next. The additional arena is the size of the largest transfer group per GPU; allocation errors are reported instead of silently disabling overlap.

### Custom Templates

For unusual partitions, module loads, or environment setup, supply your own Jinja2 template:
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -408,6 +408,9 @@ class NIXLWeightBroadcastConfig(InMemoryWeightBroadcastConfig):
session_id: str = "default"
"""ModelExpress session ID."""

overlap_transfer_and_replay: bool = False
"""Allocate two transfer arenas so inference can replay one weight group while receiving the next."""


WeightBroadcastConfig: TypeAlias = Annotated[
FileSystemWeightBroadcastConfig | NCCLWeightBroadcastConfig | NIXLWeightBroadcastConfig,
Expand Down
8 changes: 7 additions & 1 deletion packages/prime-rl-configs/src/prime_rl/configs/rl.py
Original file line number Diff line number Diff line change
Expand Up @@ -158,6 +158,9 @@ class SharedNIXLWeightBroadcastConfig(SharedInMemoryWeightBroadcastConfig):
session_id: str = "default"
"""ModelExpress session ID."""

overlap_transfer_and_replay: bool = False
"""Allocate two transfer arenas so inference can replay one weight group while receiving the next."""


class SharedFileSystemWeightBroadcastConfig(BaseConfig):
type: Literal["filesystem"] = "filesystem"
Expand Down Expand Up @@ -490,7 +493,10 @@ def auto_setup_weight_broadcast(self):
trainer_config_type = TrainerNCCLWeightBroadcastConfig
orchestrator_config_type = OrchestratorNCCLWeightBroadcastConfig
else:
transport_config = dict(session_id=self.weight_broadcast.session_id)
transport_config = dict(
session_id=self.weight_broadcast.session_id,
overlap_transfer_and_replay=self.weight_broadcast.overlap_transfer_and_replay,
)
trainer_config_type = TrainerNIXLWeightBroadcastConfig
orchestrator_config_type = OrchestratorNIXLWeightBroadcastConfig
self.trainer.weight_broadcast = trainer_config_type(**common_config, **transport_config)
Expand Down
3 changes: 3 additions & 0 deletions packages/prime-rl-configs/src/prime_rl/configs/trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -565,6 +565,9 @@ class NIXLWeightBroadcastConfig(InMemoryWeightBroadcastConfig):
session_id: str = "default"
"""ModelExpress session ID."""

overlap_transfer_and_replay: bool = False
"""Allocate two staging arenas so inference can replay one weight group while receiving the next."""


WeightBroadcastConfig: TypeAlias = Annotated[
FileSystemWeightBroadcastConfig | NCCLWeightBroadcastConfig | NIXLWeightBroadcastConfig,
Expand Down
150 changes: 46 additions & 104 deletions src/prime_rl/inference/vllm/worker/nixl.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,17 +13,21 @@

import torch
import torch.nn as nn
from modelexpress import p2p_pb2
from modelexpress.client import MxClient
from vllm.config import set_current_vllm_config
from vllm.logger import init_logger

from prime_rl.inference.vllm.worker.weight_transfer import update_mla_absorbed_weights
from prime_rl.transports.weights.nixl.agent import MemDesc, NixlAgent, make_agent_name, set_ucx_env_defaults
from prime_rl.transports.weights.nixl.cuda_malloc_memory import (
size_cuda_buffers,
use_cuda_malloc_pool,
from prime_rl.transports.weights.nixl.agent import (
MemDesc,
NixlAgent,
NixlPeer,
PreparedRead,
group_notification,
make_agent_name,
set_ucx_env_defaults,
)
from prime_rl.transports.weights.nixl.cuda_malloc_memory import use_cuda_malloc_pool
from prime_rl.transports.weights.nixl.graph import (
Destination,
OperationChain,
Expand All @@ -44,7 +48,6 @@
Worker = object

logger = init_logger("vllm.inference.vllm.worker_nixl")
_BUFFER_POLL_INTERVAL = 0.01


@dataclass
Expand All @@ -69,13 +72,14 @@ def destination_names(self) -> set[str]:
class WeightTransferGroup:
name: str
layers: list[LayerWeightTransferPlan]
pulls: list[tuple[Any, Any, list[int]]]
pulls: list[PreparedRead]


@dataclass
class WeightTransferPlan:
receive_arenas: dict[torch.dtype, torch.Tensor]
receive_buffer_count: int
trainer_peers: list[NixlPeer]
groups: list[WeightTransferGroup]


Expand Down Expand Up @@ -109,6 +113,7 @@ def init_broadcaster(
session_id=session_id,
worker_id=f"inference-{global_rank}",
)
self.model_express.publish(nixl_metadata=self.nixl_agent.get_metadata())
self.weight_transfer_timeout = timeout
self.weight_transfer_plan: WeightTransferPlan | None = None
logger.info(
Expand All @@ -126,30 +131,13 @@ def initialize_transfer(self) -> WeightTransferPlan:
trainer_ref = self.model_express.wait_for(
"trainer",
count=1,
status=None,
timeout=self.weight_transfer_timeout,
)[0]
table = TrainerTensorTable.decode(self.model_express.fetch(trainer_ref).nixl_metadata)
copies = self.trace_weight_loads(table)
plan = self.build_transfer_plan(table, copies)
self.buffer_sessions = []
for buffer_index in range(table.staging_buffer_count):
session = ModelExpressSession(
client=self.model_express.client,
role="inference",
rank=self.model_express.rank,
session_id=f"{self.model_express.session_id}:layers:{buffer_index}",
worker_id=f"inference-buffer-{self.model_express.rank}-{buffer_index}",
)
session.publish()
session.set_status(p2p_pb2.SOURCE_STATUS_INITIALIZING)
self.buffer_sessions.append(session)
# Join the current generation directly. Publishing a transient READY
# before the first pull would let the trainer mistake initialization
# for a completed acknowledgement.
self.model_express.publish(nixl_metadata=self.nixl_agent.get_metadata())
self.model_express.set_status(p2p_pb2.SOURCE_STATUS_INITIALIZING)
self.weight_transfer_plan = plan
self.group_generations = [0] * len(plan.groups)
logger.info(
"Initialized NIXL transfer plan on rank %d with %d groups",
self.model_express.rank,
Expand Down Expand Up @@ -255,25 +243,29 @@ def build_transfer_plan(
copies,
replay_plans,
)
receive_buffer_count = self.choose_receive_buffer_count(
receive_buffer_elements,
table.staging_buffer_count,
)
receive_buffer_count = table.staging_buffer_count
receive_arenas = self.allocate_receive_arenas(
receive_buffer_elements,
receive_buffer_count,
)
peers: list[NixlPeer] = []
for agent in table.agents:
peer = self.nixl_agent.add_remote_agent(agent.metadata)
self.nixl_agent.make_connection(peer)
peers.append(peer)
groups = self.build_transfer_groups(
table,
copies,
replay_plans,
receive_buffer_elements,
receive_arenas,
receive_buffer_count,
peers,
)
return WeightTransferPlan(
receive_arenas=receive_arenas,
receive_buffer_count=receive_buffer_count,
trainer_peers=peers,
groups=groups,
)

Expand Down Expand Up @@ -310,31 +302,6 @@ def calculate_receive_buffer_elements(
group_elements[source_dtype][tensor_groups[source.name]] += prod(replay_plans[id(copy)].source_shape)
return {dtype: max(elements, default=0) for dtype, elements in group_elements.items()}

def choose_receive_buffer_count(
self,
receive_buffer_elements: dict[torch.dtype, int],
staging_buffer_count: int,
) -> int:
receive_buffer_bytes = max(
1,
sum(elements * dtype.itemsize for dtype, elements in receive_buffer_elements.items()),
)
allocated_bytes = torch.cuda.memory_allocated(self.device)
peak_growth_bytes = max(
0,
torch.cuda.max_memory_allocated(self.device) - allocated_bytes,
)
free_bytes, _ = torch.cuda.mem_get_info(self.device)
max_receive_buffers = min(2, staging_buffer_count) if peak_growth_bytes else 1
if peak_growth_bytes or free_bytes < receive_buffer_bytes:
torch.cuda.empty_cache()
return size_cuda_buffers(
receive_buffer_bytes,
max_receive_buffers,
self.device,
extra_headroom_bytes=receive_buffer_bytes + peak_growth_bytes,
)

def allocate_receive_arenas(
self,
receive_buffer_elements: dict[torch.dtype, int],
Expand Down Expand Up @@ -362,6 +329,7 @@ def build_transfer_groups(
receive_buffer_elements: dict[torch.dtype, int],
receive_arenas: dict[torch.dtype, torch.Tensor],
receive_buffer_count: int,
peers: list[NixlPeer],
) -> list[WeightTransferGroup]:
tensors = {tensor.name: tensor for group in table.groups for tensor in group.tensors}
tensor_groups = {
Expand Down Expand Up @@ -390,7 +358,6 @@ def build_transfer_groups(
)

agent_devices = {agent_index: agent.device_id for agent_index, agent in enumerate(table.agents)}
peer_names: dict[int, str] = {}
transfer_groups: list[WeightTransferGroup] = []

for group_index, group in enumerate(table.groups):
Expand Down Expand Up @@ -432,10 +399,9 @@ def build_transfer_groups(
persistent_plans_by_layer,
),
pulls=self.prepare_group_pulls(
table,
local_descs,
remote_descs,
peer_names,
peers,
),
)
)
Expand Down Expand Up @@ -486,44 +452,37 @@ def build_layer_transfer_plans(

def prepare_group_pulls(
self,
table: TrainerTensorTable,
local_descs: dict[int, list[MemDesc]],
remote_descs: dict[int, list[MemDesc]],
peer_names: dict[int, str],
) -> list[tuple[Any, Any, list[int]]]:
pulls: list[tuple[Any, Any, list[int]]] = []
peers: list[NixlPeer],
) -> list[PreparedRead]:
pulls: list[PreparedRead] = []
agent_indices = sorted(remote_descs)
rotation = self.model_express.rank % len(agent_indices) if agent_indices else 0
agent_indices = agent_indices[rotation:] + agent_indices[:rotation]
for agent_index in agent_indices:
remote = remote_descs[agent_index]
peer_name = peer_names.get(agent_index)
if peer_name is None:
peer_name = self.nixl_agent.add_remote_agent(table.agents[agent_index].metadata)
self.nixl_agent.make_connection(peer_name)
peer_names[agent_index] = peer_name
peer = peers[agent_index]
local_prepared = self.nixl_agent.prepare_xfer_dlist(local_descs[agent_index])
remote_prepared = self.nixl_agent.prepare_xfer_dlist(remote, agent_name=peer_name)
pulls.append((local_prepared, remote_prepared, list(range(len(remote)))))
remote_prepared = self.nixl_agent.prepare_xfer_dlist(remote, peer=peer)
pulls.append(
self.nixl_agent.prepare_read(
local_prepared,
list(range(len(remote))),
remote_prepared,
)
)
return pulls

@torch.no_grad()
def update_weights_from_path(self, weight_dir: str | None = None) -> None:
del weight_dir
plan = self.initialize_transfer()
self.model_express.set_status(p2p_pb2.SOURCE_STATUS_INITIALIZING)
self.model_express.wait_for(
"trainer",
count=1,
status=p2p_pb2.SOURCE_STATUS_READY,
timeout=self.weight_transfer_timeout,
)

started = time.perf_counter()
self.apply_transfer_plan(plan)
update_mla_absorbed_weights(self.raw_model)
torch.cuda.synchronize(self.device)
self.model_express.set_status(p2p_pb2.SOURCE_STATUS_READY)
logger.info(
"Applied NIXL policy update on rank %d in %.2fs",
self.model_express.rank,
Expand All @@ -546,44 +505,30 @@ def apply_transfer_plan(self, plan: WeightTransferPlan) -> None:

def pull_group(group_index: int) -> WeightTransferGroup:
transfer_group = plan.groups[group_index]
session = self.buffer_sessions[group_index % len(self.buffer_sessions)]
session.wait_for(
"trainer",
count=1,
status=p2p_pb2.SOURCE_STATUS_READY,
notification = group_notification(group_index, self.group_generations[group_index])
self.nixl_agent.wait_for_notification(
plan.trainer_peers,
notification,
timeout=self.weight_transfer_timeout,
poll_interval=_BUFFER_POLL_INTERVAL,
cancelled=cancelled.is_set,
)

for local, remote, indices in transfer_group.pulls:
handle = self.nixl_agent.post_read(local, indices, remote)
for read in transfer_group.pulls:
self.nixl_agent.post_read(read, notification)
for read in transfer_group.pulls:
self.nixl_agent.wait(
handle,
read,
context=f"weight pull for {transfer_group.name}",
timeout=self.weight_transfer_timeout,
cancelled=cancelled.is_set,
)
return transfer_group

def acknowledge_group(group_index: int) -> None:
session = self.buffer_sessions[group_index % len(self.buffer_sessions)]
session.set_status(p2p_pb2.SOURCE_STATUS_READY)
session.wait_for(
"trainer",
count=1,
status=p2p_pb2.SOURCE_STATUS_INITIALIZING,
timeout=self.weight_transfer_timeout,
poll_interval=_BUFFER_POLL_INTERVAL,
cancelled=cancelled.is_set,
)
session.set_status(p2p_pb2.SOURCE_STATUS_INITIALIZING)
self.group_generations[group_index] += 1
return transfer_group
Comment thread
cursor[bot] marked this conversation as resolved.

def prefetch_group(group_index: int) -> WeightTransferGroup:
torch.cuda.set_device(self.device)
transfer_group = pull_group(group_index)
acknowledge_group(group_index)
return transfer_group
return pull_group(group_index)

def replay_group(transfer_group: WeightTransferGroup) -> None:
for layer_plan in transfer_group.layers:
Expand Down Expand Up @@ -629,9 +574,6 @@ def replay_group(transfer_group: WeightTransferGroup) -> None:

replay_group(transfer_group)
torch.cuda.synchronize(self.device)

if not pipelined:
acknowledge_group(group_index)
finally:
cancelled.set()
if executor is not None:
Expand Down
Loading