feat: Support modelexpress p2p RDMA transfer (#23105)

This commit is contained in:
Zhongdongming Dai
2026-04-29 12:57:40 -07:00
committed by GitHub
parent db84a8ebbb
commit 7389743d85
4 changed files with 252 additions and 75 deletions
+1
View File
@@ -85,6 +85,7 @@ class LoadConfig:
modelexpress_ep_size: Optional[int] = None modelexpress_ep_size: Optional[int] = None
modelexpress_dtype: Optional[str] = None modelexpress_dtype: Optional[str] = None
modelexpress_quantization: Optional[str] = None modelexpress_quantization: Optional[str] = None
modelexpress_transport: str = "transfer_engine"
# ModelOpt-specific loading options # ModelOpt-specific loading options
modelopt_checkpoint_restore_path: Optional[str] = None modelopt_checkpoint_restore_path: Optional[str] = None
@@ -884,7 +884,12 @@ class ModelRunner(ModelRunnerKVCacheMixin):
) )
def _publish_modelexpress_metadata(self): def _publish_modelexpress_metadata(self):
"""Publish TransferEngine metadata to ModelExpress server (seed mode).""" """Publish metadata to ModelExpress server (seed mode).
Supports two transport backends:
- transfer_engine: publishes TransferEngine session_id (Mooncake)
- nixl: creates NIXL agent, registers tensors, publishes nixl_metadata
"""
try: try:
from modelexpress import p2p_pb2 from modelexpress import p2p_pb2
from modelexpress.client import MxClient from modelexpress.client import MxClient
@@ -898,15 +903,7 @@ class ModelRunner(ModelRunnerKVCacheMixin):
self.server_args.modelexpress_model_name or self.server_args.model_path self.server_args.modelexpress_model_name or self.server_args.model_path
) )
mx_url = self.server_args.modelexpress_url mx_url = self.server_args.modelexpress_url
session_id = self.remote_instance_transfer_engine_session_id transport = self.server_args.modelexpress_transport
weight_info = self.remote_instance_transfer_engine_weight_info
if not session_id or weight_info is None:
logger.warning(
"ModelExpress source: skipping publish -- "
"TransferEngine not initialized or no weight info"
)
return
# Build SourceIdentity for this instance # Build SourceIdentity for this instance
identity = p2p_pb2.SourceIdentity( identity = p2p_pb2.SourceIdentity(
@@ -919,23 +916,12 @@ class ModelRunner(ModelRunnerKVCacheMixin):
quantization=self.server_args.quantization or "", quantization=self.server_args.quantization or "",
) )
# Build tensor descriptors from weight_info dict if transport == "nixl":
tensors = [] worker, tensor_count = self._build_nixl_worker_metadata(p2p_pb2)
for name, (addr, numel, element_size) in weight_info.items(): else:
tensors.append( worker, tensor_count = self._build_transfer_engine_worker_metadata(p2p_pb2)
p2p_pb2.TensorDescriptor( if worker is None:
name=name, return
addr=addr,
size=numel * element_size,
device_id=self.gpu_id,
)
)
worker = p2p_pb2.WorkerMetadata(
worker_rank=self.tp_rank,
transfer_engine_session_id=session_id,
tensors=tensors,
)
# Generate a unique worker_id for this running instance # Generate a unique worker_id for this running instance
worker_id = str(uuid.uuid4()) worker_id = str(uuid.uuid4())
@@ -943,12 +929,12 @@ class ModelRunner(ModelRunnerKVCacheMixin):
mx_client = MxClient(server_url=mx_url) mx_client = MxClient(server_url=mx_url)
try: try:
logger.info( logger.info(
"ModelExpress source: publishing metadata for model=%s, " "ModelExpress source [%s]: publishing metadata for model=%s, "
"tp_rank=%d, session=%s, %d tensors, worker_id=%s", "tp_rank=%d, %d tensors, worker_id=%s",
transport,
model_name, model_name,
self.tp_rank, self.tp_rank,
session_id, tensor_count,
len(tensors),
worker_id, worker_id,
) )
mx_source_id = mx_client.publish_metadata(identity, worker, worker_id) mx_source_id = mx_client.publish_metadata(identity, worker, worker_id)
@@ -968,6 +954,86 @@ class ModelRunner(ModelRunnerKVCacheMixin):
finally: finally:
mx_client.close() mx_client.close()
def _build_transfer_engine_worker_metadata(self, p2p_pb2):
"""Build WorkerMetadata using TransferEngine session_id."""
session_id = self.remote_instance_transfer_engine_session_id
weight_info = self.remote_instance_transfer_engine_weight_info
if not session_id or weight_info is None:
logger.warning(
"ModelExpress source: skipping publish -- "
"TransferEngine not initialized or no weight info"
)
return None, 0
tensors = []
for name, (addr, numel, element_size) in weight_info.items():
tensors.append(
p2p_pb2.TensorDescriptor(
name=name,
addr=addr,
size=numel * element_size,
device_id=self.gpu_id,
)
)
worker = p2p_pb2.WorkerMetadata(
worker_rank=self.tp_rank,
transfer_engine_session_id=session_id,
tensors=tensors,
)
return worker, len(tensors)
def _build_nixl_worker_metadata(self, p2p_pb2):
"""Build WorkerMetadata using NIXL agent for RDMA transfers."""
from modelexpress.nixl_transfer import NixlTransferManager
agent_name = f"sglang-seed-rank{self.tp_rank}-{uuid.uuid4().hex[:8]}"
nixl_mgr = NixlTransferManager(agent_name, self.gpu_id)
nixl_mgr.initialize()
# Collect model tensors for NIXL registration
model_tensors = {}
for name, param in self.model.named_parameters():
t = param.data
if t.is_contiguous():
model_tensors[name] = t
else:
# Non-contiguous tensors: register underlying storage as byte view
sv = torch.empty(0, dtype=torch.uint8, device=t.device).set_(
t.untyped_storage()
)
if sv.data_ptr() not in {
v.data_ptr() for v in model_tensors.values()
}:
model_tensors[f"{name}.__storage"] = sv
nixl_metadata = nixl_mgr.register_tensors(model_tensors)
# Build tensor descriptors from registered tensors
tensors = []
for td in nixl_mgr.tensor_descriptors:
tensors.append(
p2p_pb2.TensorDescriptor(
name=td.name,
addr=td.addr,
size=td.size,
device_id=td.device_id,
dtype=td.dtype,
)
)
worker = p2p_pb2.WorkerMetadata(
worker_rank=self.tp_rank,
nixl_metadata=nixl_metadata,
tensors=tensors,
)
# Keep reference alive so NIXL agent isn't garbage collected
self._nixl_manager = nixl_mgr
return worker, len(tensors)
def model_specific_adjustment(self): def model_specific_adjustment(self):
server_args = self.server_args server_args = self.server_args
@@ -1261,6 +1327,7 @@ class ModelRunner(ModelRunnerKVCacheMixin):
modelexpress_ep_size=self.server_args.ep_size, modelexpress_ep_size=self.server_args.ep_size,
modelexpress_dtype=self.server_args.dtype, modelexpress_dtype=self.server_args.dtype,
modelexpress_quantization=self.server_args.quantization or "", modelexpress_quantization=self.server_args.quantization or "",
modelexpress_transport=self.server_args.modelexpress_transport,
modelopt_config=modelopt_config, modelopt_config=modelopt_config,
rl_quant_profile=self.server_args.rl_quant_profile, rl_quant_profile=self.server_args.rl_quant_profile,
draft_model_idx=self.draft_model_idx, draft_model_idx=self.draft_model_idx,
+145 -42
View File
@@ -2303,7 +2303,12 @@ class RemoteInstanceModelLoader(BaseModelLoader):
load_config: LoadConfig, load_config: LoadConfig,
device_config: DeviceConfig, device_config: DeviceConfig,
): ):
"""Load weights via ModelExpress coordination + TransferEngine RDMA.""" """Load weights via ModelExpress coordination + RDMA transfer.
Supports two transport backends:
- transfer_engine: Mooncake TransferEngine (default)
- nixl: NIXL UCX-based RDMA
"""
try: try:
import grpc import grpc
from modelexpress import p2p_pb2 from modelexpress import p2p_pb2
@@ -2314,14 +2319,11 @@ class RemoteInstanceModelLoader(BaseModelLoader):
"Install it with: pip install modelexpress" "Install it with: pip install modelexpress"
) from exc ) from exc
transfer_engine = load_config.remote_instance_weight_loader_transfer_engine
if transfer_engine is None:
raise RuntimeError(
"TransferEngine is not initialized for modelexpress backend."
)
tp_rank = load_config.tp_rank tp_rank = load_config.tp_rank
model_name = load_config.modelexpress_model_name model_name = load_config.modelexpress_model_name
transport = load_config.modelexpress_transport
# Process quantized weights to establish final tensor layout
target_device = torch.device(device_config.device) target_device = torch.device(device_config.device)
for _, module in model.named_modules(): for _, module in model.named_modules():
quant_method = getattr(module, "quant_method", None) quant_method = getattr(module, "quant_method", None)
@@ -2329,14 +2331,23 @@ class RemoteInstanceModelLoader(BaseModelLoader):
with device_loading_context(module, target_device): with device_loading_context(module, target_device):
quant_method.process_weights_after_loading(module) quant_method.process_weights_after_loading(module)
logger.info( # Register local memory for the chosen transport
"ModelExpress: registering memory regions for tp_rank=%d...", tp_rank if transport == "nixl":
) nixl_mgr = self._init_nixl_for_target(model, load_config, device_config)
self.remote_instance_transfer_engine_weight_info = register_memory_region( else:
model, transfer_engine transfer_engine = load_config.remote_instance_weight_loader_transfer_engine
) if transfer_engine is None:
raise RuntimeError(
"TransferEngine is not initialized for modelexpress backend."
)
logger.info(
"ModelExpress: registering memory regions for tp_rank=%d...", tp_rank
)
self.remote_instance_transfer_engine_weight_info = register_memory_region(
model, transfer_engine
)
# Build SourceIdentity matching the seed's identity # --- Shared MX discovery logic ---
identity = p2p_pb2.SourceIdentity( identity = p2p_pb2.SourceIdentity(
model_name=model_name, model_name=model_name,
backend_framework=p2p_pb2.BACKEND_FRAMEWORK_SGLANG, backend_framework=p2p_pb2.BACKEND_FRAMEWORK_SGLANG,
@@ -2347,11 +2358,11 @@ class RemoteInstanceModelLoader(BaseModelLoader):
quantization=load_config.modelexpress_quantization or "", quantization=load_config.modelexpress_quantization or "",
) )
# Query MX server for a READY source matching our identity and rank
mx_client = MxClient(server_url=load_config.modelexpress_url) mx_client = MxClient(server_url=load_config.modelexpress_url)
try: try:
logger.info( logger.info(
"ModelExpress: looking for seed (model=%s, rank=%d)...", "ModelExpress [%s]: looking for seed (model=%s, rank=%d)...",
transport,
model_name, model_name,
tp_rank, tp_rank,
) )
@@ -2380,7 +2391,6 @@ class RemoteInstanceModelLoader(BaseModelLoader):
f"Ensure the seed instance is running and has published metadata." f"Ensure the seed instance is running and has published metadata."
) )
# Fetch full metadata for the discovered worker
response = mx_client.get_metadata( response = mx_client.get_metadata(
mx_source_id=source_ref.mx_source_id, mx_source_id=source_ref.mx_source_id,
worker_id=source_ref.worker_id, worker_id=source_ref.worker_id,
@@ -2393,31 +2403,46 @@ class RemoteInstanceModelLoader(BaseModelLoader):
) )
source_worker = response.worker source_worker = response.worker
# Extract session_id from oneof backend_metadata
backend_field = source_worker.WhichOneof("backend_metadata")
if backend_field == "transfer_engine_session_id":
seed_session_id = source_worker.transfer_engine_session_id
else:
raise RuntimeError(
f"ModelExpress: expected transfer_engine_session_id, "
f"got backend_metadata={backend_field}"
)
# Build {name: (addr, size_bytes)} from seed tensor descriptors
seed_weight_info = {}
for td in source_worker.tensors:
seed_weight_info[td.name] = (td.addr, td.size)
logger.info(
"ModelExpress: got %d tensor descriptors from seed (session=%s)",
len(seed_weight_info),
seed_session_id,
)
finally: finally:
mx_client.close() mx_client.close()
# Transfer weights via TransferEngine RDMA # --- Transport-specific transfer ---
if transport == "nixl":
self._transfer_via_nixl(
model, nixl_mgr, source_worker, tp_rank
)
else:
self._transfer_via_transfer_engine(
model, transfer_engine, source_worker, tp_rank
)
if hasattr(model, "post_load_weights"):
model.post_load_weights()
logger.info("ModelExpress: weight transfer complete for tp_rank=%d", tp_rank)
def _transfer_via_transfer_engine(
self, model, transfer_engine, source_worker, tp_rank
):
"""Execute weight transfer using Mooncake TransferEngine."""
backend_field = source_worker.WhichOneof("backend_metadata")
if backend_field != "transfer_engine_session_id":
raise RuntimeError(
f"ModelExpress: expected transfer_engine_session_id, "
f"got backend_metadata={backend_field}"
)
seed_session_id = source_worker.transfer_engine_session_id
seed_weight_info = {}
for td in source_worker.tensors:
seed_weight_info[td.name] = (td.addr, td.size)
logger.info(
"ModelExpress: got %d tensor descriptors from seed (session=%s)",
len(seed_weight_info),
seed_session_id,
)
seed_ptr_list = [] seed_ptr_list = []
client_ptr_list = [] client_ptr_list = []
client_len_list = [] client_len_list = []
@@ -2440,7 +2465,7 @@ class RemoteInstanceModelLoader(BaseModelLoader):
client_len_list.append(local_size) client_len_list.append(local_size)
logger.info( logger.info(
"ModelExpress: starting RDMA transfer of %d tensors...", "ModelExpress: starting TransferEngine RDMA of %d tensors...",
len(seed_ptr_list), len(seed_ptr_list),
) )
ret = transfer_engine.batch_transfer_sync_read( ret = transfer_engine.batch_transfer_sync_read(
@@ -2454,10 +2479,88 @@ class RemoteInstanceModelLoader(BaseModelLoader):
f"ModelExpress: batch_transfer_sync_read failed, error={ret}" f"ModelExpress: batch_transfer_sync_read failed, error={ret}"
) )
if hasattr(model, "post_load_weights"): def _init_nixl_for_target(self, model, load_config, device_config):
model.post_load_weights() """Initialize NIXL agent and register local tensors for the target."""
import uuid
logger.info("ModelExpress: weight transfer complete for tp_rank=%d", tp_rank) from modelexpress.nixl_transfer import NixlTransferManager
tp_rank = load_config.tp_rank
device_id = device_config.gpu_id
agent_name = f"sglang-target-rank{tp_rank}-{uuid.uuid4().hex[:8]}"
nixl_mgr = NixlTransferManager(agent_name, device_id)
nixl_mgr.initialize()
# Collect local tensors, handling non-contiguous via storage views
local_tensors = {}
seen_ptrs = set()
for name, param in model.named_parameters():
t = param.data
if t.is_contiguous():
ptr = t.data_ptr()
if ptr in seen_ptrs:
continue
seen_ptrs.add(ptr)
local_tensors[name] = t
else:
sv = torch.empty(0, dtype=torch.uint8, device=t.device).set_(
t.untyped_storage()
)
ptr = sv.data_ptr()
if ptr in seen_ptrs:
continue
seen_ptrs.add(ptr)
local_tensors[f"{name}.__storage"] = sv
nixl_mgr.register_tensors(local_tensors)
logger.info(
"ModelExpress [nixl]: registered %d tensors for tp_rank=%d",
len(local_tensors),
tp_rank,
)
return nixl_mgr
def _transfer_via_nixl(self, model, nixl_mgr, source_worker, tp_rank):
"""Execute weight transfer using NIXL RDMA."""
from modelexpress.types import TensorDescriptor
backend_field = source_worker.WhichOneof("backend_metadata")
if backend_field != "nixl_metadata":
raise RuntimeError(
f"ModelExpress: expected nixl_metadata, "
f"got backend_metadata={backend_field}"
)
source_tensors = [
TensorDescriptor(
name=td.name,
addr=td.addr,
size=td.size,
device_id=td.device_id,
dtype=td.dtype,
)
for td in source_worker.tensors
]
logger.info(
"ModelExpress [nixl]: starting RDMA transfer of %d tensors...",
len(source_tensors),
)
total_bytes, matched, duration = nixl_mgr.receive_from_source(
source_metadata=source_worker.nixl_metadata,
source_tensors=source_tensors,
coalesce_transfers=False,
)
logger.info(
"ModelExpress [nixl]: transferred %d tensors, "
"%.2f GB in %.2fs",
matched,
total_bytes / 1e9,
duration,
)
class RemoteModelLoader(BaseModelLoader): class RemoteModelLoader(BaseModelLoader):
+8 -2
View File
@@ -7106,18 +7106,24 @@ class ServerArgs:
def modelexpress_source(self) -> bool: def modelexpress_source(self) -> bool:
return self._parsed_modelexpress_config.get("source", False) return self._parsed_modelexpress_config.get("source", False)
@property
def modelexpress_transport(self) -> str:
"""Transport backend for modelexpress: 'transfer_engine' (default) or 'nixl'."""
return self._parsed_modelexpress_config.get("transport", "transfer_engine")
def remote_instance_weight_loader_use_transfer_engine(self): def remote_instance_weight_loader_use_transfer_engine(self):
# Use TransferEngine as seed backend. # Use TransferEngine as seed backend.
if self.remote_instance_weight_loader_start_seed_via_transfer_engine: if self.remote_instance_weight_loader_start_seed_via_transfer_engine:
return True return True
# ModelExpress source mode also needs TransferEngine init. # ModelExpress source mode needs TransferEngine init only if transport is transfer_engine.
if self.modelexpress_source: if self.modelexpress_source and self.modelexpress_transport == "transfer_engine":
return True return True
# Use TransferEngine as client backend. # Use TransferEngine as client backend.
elif ( elif (
self.load_format == "remote_instance" self.load_format == "remote_instance"
and self.remote_instance_weight_loader_backend and self.remote_instance_weight_loader_backend
in ("transfer_engine", "modelexpress") in ("transfer_engine", "modelexpress")
and self.modelexpress_transport == "transfer_engine"
): ):
return True return True
else: else: