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
Original file line number Diff line number Diff line change
@@ -1,4 +1,4 @@
"""Shared NIXL pixel-payload puller for TokenSpeed multimodal inputs."""
"""Engine-neutral NIXL pixel-payload puller for gateway-exported multimodal inputs."""

from __future__ import annotations

Expand All @@ -15,7 +15,8 @@
_DESCRIPTOR_MAGIC = b"SMGRDMA1"
_GEN_BYTES = 8
_FRAME_BYTES = 2 * _GEN_BYTES
_GATEWAY_AGENT_NAME = "smg-gateway-encode"
# The gateway's fixed NIXL agent name (see RDMA_GATEWAY_AGENT_NAME gateway-side).
DEFAULT_GATEWAY_AGENT_NAME = "smg-gateway-encode"
_LOCAL_IP_CACHE: dict[str, bool] = {}


Expand Down Expand Up @@ -102,8 +103,15 @@ def _parse_descriptor(td, explicit_room: int | None) -> _RemotePixelDescriptor:
class RdmaPixelPuller:
"""Persistent NIXL READ agent for gateway-exported multimodal pixel buffers."""

def __init__(self, *, agent_name: str, log_prefix: str):
def __init__(
self,
*,
agent_name: str,
log_prefix: str,
gateway_agent_name: str = DEFAULT_GATEWAY_AGENT_NAME,
):
self._log_prefix = log_prefix
self._gateway_agent_name = gateway_agent_name
self._nixl_agent = None
self._rdma_md_ready = set()
self._rdma_md_lock = threading.Lock()
Expand Down Expand Up @@ -171,7 +179,7 @@ def _ensure_remote_ready(self, ip: str, port: int, remote, room: int) -> None:
return

agent = self._nixl_agent
agent.fetch_remote_metadata(_GATEWAY_AGENT_NAME, ip, port)
agent.fetch_remote_metadata(self._gateway_agent_name, ip, port)

send_md_env = os.environ.get("SMG_RDMA_SEND_MD")
if send_md_env in ("1", "true"):
Expand All @@ -194,23 +202,16 @@ def _ensure_remote_ready(self, ip: str, port: int, remote, room: int) -> None:

ready = False
for _ in range(5000):
if agent.check_remote_metadata(_GATEWAY_AGENT_NAME, remote):
if agent.check_remote_metadata(self._gateway_agent_name, remote):
ready = True
break
time.sleep(0.001)
if not ready:
raise RuntimeError(f"NIXL remote metadata not ready room={room}")
self._rdma_md_ready.add(key)

def feature_from_remote(
self,
td,
*,
explicit_room: int | None,
cast_to,
publish_shm: bool = False,
):
"""Read a remote TensorData payload and return ``(feature, content_hash)``."""
def feature_from_remote(self, td, *, explicit_room: int | None, cast_to):
"""Read a remote TensorData payload and return the feature tensor."""
import numpy as np
import torch

Expand Down Expand Up @@ -271,7 +272,7 @@ def feature_from_remote(
"READ",
local,
remote,
_GATEWAY_AGENT_NAME,
self._gateway_agent_name,
str(descriptor.room).encode(),
)
read_deadline = time.monotonic() + float(os.environ.get("SMG_RDMA_READ_TIMEOUT_S", 60))
Expand Down Expand Up @@ -327,13 +328,7 @@ def feature_from_remote(
tensor = tensor.to(cast_to)
copied = True

if publish_shm:
from tokenspeed.runtime.multimodal.hash import hash_feature
from tokenspeed.runtime.multimodal.shm_transport import ShmTensorHandle

feat_hash = hash_feature(tensor)
return ShmTensorHandle.publish(tensor), feat_hash

return (tensor if copied else tensor.clone()), None
# Copy out of the landing slot before it returns to the free ring.
return tensor if copied else tensor.clone()

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P2 Badge Avoid cloning RDMA tensors before SHM publishing

When EPD_PIXEL_SHM is left at its default true and the remote tensor already has the model dtype, this return clones the full RDMA landing buffer before TokenSpeedEncoderServicer._items_from_proto immediately hashes it and calls ShmTensorHandle.publish(feature). For the large encoder-input tensors that use RDMA, that adds an extra full-size copy and temporary allocation on every EPD encode item; keep the publish/hash inside the landing-slot lifetime (for example via a caller-supplied materializer) so the SHM path can copy directly from the landing view.

Useful? React with 👍 / 👎.

finally:
self._landing_free.put(slot)
29 changes: 13 additions & 16 deletions grpc_servicer/smg_grpc_servicer/tokenspeed/encoder_servicer.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,7 +22,7 @@
tokenspeed_encoder_pb2_grpc,
)

from smg_grpc_servicer.tokenspeed.rdma_pixel import RdmaPixelPuller
from smg_grpc_servicer.mm_rdma import RdmaPixelPuller
from smg_grpc_servicer.tokenspeed.servicer import TokenSpeedSchedulerServicer

if TYPE_CHECKING:
Expand Down Expand Up @@ -125,27 +125,24 @@ def _items_from_proto(self, mm_inputs, bootstrap_room: int = 0):
# before the swap and pre-set on the item.
td = item_proto.encoder_input
if td.WhichOneof("payload") == "remote":
# EPD RDMA: pull the encoder input from exported NIXL memory.
# With EPD_PIXEL_SHM, the received slot is published directly to
# scheduler SHM so the scheduler ingest path still avoids pickle
# copies of the full encoder tensor.
feature, feat_hash = self._rdma_pixel_puller.feature_from_remote(
feature = self._rdma_pixel_puller.feature_from_remote(
td,
explicit_room=bootstrap_room,
cast_to=model_dtype,
publish_shm=self._pixel_shm,
)
else:
feature = TokenSpeedSchedulerServicer._tensor_from_proto(td, cast_to=model_dtype)
feat_hash = None
if self._pixel_shm:
from tokenspeed.runtime.multimodal.hash import hash_feature
from tokenspeed.runtime.multimodal.shm_transport import (
ShmTensorHandle,
)

feat_hash = hash_feature(feature)
feature = ShmTensorHandle.publish(feature)

# EPD_PIXEL_SHM publishes the feature to scheduler SHM so the ZMQ hop
# pickles ~KB instead of the full encoder tensor; the content hash is
# taken on the real bytes first.
feat_hash = None
if self._pixel_shm:
from tokenspeed.runtime.multimodal.hash import hash_feature
from tokenspeed.runtime.multimodal.shm_transport import ShmTensorHandle

feat_hash = hash_feature(feature)
feature = ShmTensorHandle.publish(feature)

item = MultimodalDataItem(
modality=item_modality,
Expand Down
5 changes: 2 additions & 3 deletions grpc_servicer/smg_grpc_servicer/tokenspeed/servicer.py
Original file line number Diff line number Diff line change
Expand Up @@ -39,9 +39,9 @@
from tokenspeed.runtime.pd.kv_events import KVEventBatch

from smg_grpc_servicer.kv_events import endpoint_for_rank, stream_kv_events
from smg_grpc_servicer.mm_rdma import RdmaPixelPuller
from smg_grpc_servicer.tokenspeed.health_servicer import TokenSpeedHealthServicer
from smg_grpc_servicer.tokenspeed.kv_events import resolve_kv_events_config
from smg_grpc_servicer.tokenspeed.rdma_pixel import RdmaPixelPuller

if TYPE_CHECKING:
# Type-only — keeps these out of the cold-path graph when the servicer is
Expand Down Expand Up @@ -1072,11 +1072,10 @@ def _mm_inputs_from_itemized_proto(
feature_started = time.perf_counter() if LOG_MM_TIMING else None
encoder_input = item_proto.encoder_input
if encoder_input.WhichOneof("payload") == "remote":
feature, _ = self._rdma_pixel_puller.feature_from_remote(
feature = self._rdma_pixel_puller.feature_from_remote(
encoder_input,
explicit_room=None,
cast_to=model_dtype,
publish_shm=False,
)
else:
feature = self._feature_from_proto(encoder_input, cast_to=model_dtype)
Expand Down
8 changes: 2 additions & 6 deletions model_gateway/src/routers/grpc/client.rs
Original file line number Diff line number Diff line change
Expand Up @@ -534,9 +534,7 @@ impl GrpcClient {
}
Self::TokenSpeed(client) => {
let tokenspeed_mm = options.multimodal_inputs.map(|mm| match mm {
MultimodalData::TokenSpeed(data) => {
data.try_export_encoder_inputs_nixl_remote().into_proto()
}
MultimodalData::TokenSpeed(data) => data.into_proto(true),
_ => unreachable!("caller guarantees matching variant"),
});
finish_tokenspeed_request(tokenspeed_mm, |mm| {
Expand Down Expand Up @@ -628,9 +626,7 @@ impl GrpcClient {
}
Self::TokenSpeed(client) => {
let tokenspeed_mm = options.multimodal_inputs.map(|mm| match mm {
MultimodalData::TokenSpeed(data) => {
data.try_export_encoder_inputs_nixl_remote().into_proto()
}
MultimodalData::TokenSpeed(data) => data.into_proto(true),
_ => unreachable!("caller guarantees matching variant"),
});
finish_tokenspeed_request(tokenspeed_mm, |mm| {
Expand Down
15 changes: 10 additions & 5 deletions model_gateway/src/routers/grpc/epd_encode.rs
Original file line number Diff line number Diff line change
Expand Up @@ -15,11 +15,11 @@ use uuid::Uuid;
use super::{
client::GrpcClient,
context::{ClientSelection, WorkerSelection},
multimodal::{assemble_tokenspeed_for_encode, MultimodalIntermediate},
multimodal::{assemble_tokenspeed_for_encode, mm_rdma_exporter, MultimodalIntermediate},
proto_wrapper::{
cleanup_mm_shm_handles, cleanup_tokenspeed_items_encoder_shm,
collect_tokenspeed_multimodal_inputs_shm_handles, EncodeItemBootstrapInfo,
TokenSpeedMultimodalData, TokenSpeedMultimodalItem,
collect_tokenspeed_multimodal_inputs_shm_handles, stage_tokenspeed_tensor_rdma,
EncodeItemBootstrapInfo, TokenSpeedMultimodalData, TokenSpeedMultimodalItem,
},
};
use crate::worker::DEFAULT_BOOTSTRAP_PORT;
Expand Down Expand Up @@ -108,7 +108,12 @@ impl PreparedEncodeItem {
.take()
.ok_or_else(|| "encode item was already dispatched".to_string())?;
*cleanup_on_drop = false;
item.encoder_input = item.encoder_input.try_export_nixl_remote(bootstrap_room);
// Stage with bootstrap_room as the slot key (load-bearing), so
// into_proto(false) below does not re-stage with a random key.
if let Some(exporter) = mm_rdma_exporter() {
item.encoder_input =
stage_tokenspeed_tensor_rdma(exporter, bootstrap_room, item.encoder_input);
}
let request = tokenspeed_encoder::EncodeRequest {
request_id: format!("encode-{}", Uuid::now_v7()),
mm_inputs: Some(
Expand All @@ -117,7 +122,7 @@ impl PreparedEncodeItem {
shm_enabled: *shm_enabled,
shm_min_bytes: *shm_min_bytes,
}
.into_proto(),
.into_proto(false),
),
items: vec![tokenspeed_encoder::EncodeItemAssignment { bootstrap_room }],
};
Expand Down
2 changes: 2 additions & 0 deletions model_gateway/src/routers/grpc/multimodal/assemble.rs
Original file line number Diff line number Diff line change
Expand Up @@ -230,6 +230,8 @@ fn assemble_vllm(
modality,
shm_enabled: resolve_mm_shm_enabled(workers, false),
shm_min_bytes: resolve_mm_shm_min_bytes(workers),
// vLLM workers cannot pull RDMA payloads yet.
rdma_enabled: false,
})
}

Expand Down
Loading
Loading