diff --git a/python/sglang/srt/configs/load_config.py b/python/sglang/srt/configs/load_config.py index e642bdc27..6072b2ad5 100644 --- a/python/sglang/srt/configs/load_config.py +++ b/python/sglang/srt/configs/load_config.py @@ -85,6 +85,7 @@ class LoadConfig: modelexpress_ep_size: Optional[int] = None modelexpress_dtype: Optional[str] = None modelexpress_quantization: Optional[str] = None + modelexpress_transport: str = "transfer_engine" # ModelOpt-specific loading options modelopt_checkpoint_restore_path: Optional[str] = None diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index f69cab8d2..2b261fa83 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -884,7 +884,12 @@ class ModelRunner(ModelRunnerKVCacheMixin): ) 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: from modelexpress import p2p_pb2 from modelexpress.client import MxClient @@ -898,15 +903,7 @@ class ModelRunner(ModelRunnerKVCacheMixin): self.server_args.modelexpress_model_name or self.server_args.model_path ) mx_url = self.server_args.modelexpress_url - 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 + transport = self.server_args.modelexpress_transport # Build SourceIdentity for this instance identity = p2p_pb2.SourceIdentity( @@ -919,23 +916,12 @@ class ModelRunner(ModelRunnerKVCacheMixin): quantization=self.server_args.quantization or "", ) - # Build tensor descriptors from weight_info dict - 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, - ) + if transport == "nixl": + worker, tensor_count = self._build_nixl_worker_metadata(p2p_pb2) + else: + worker, tensor_count = self._build_transfer_engine_worker_metadata(p2p_pb2) + if worker is None: + return # Generate a unique worker_id for this running instance worker_id = str(uuid.uuid4()) @@ -943,12 +929,12 @@ class ModelRunner(ModelRunnerKVCacheMixin): mx_client = MxClient(server_url=mx_url) try: logger.info( - "ModelExpress source: publishing metadata for model=%s, " - "tp_rank=%d, session=%s, %d tensors, worker_id=%s", + "ModelExpress source [%s]: publishing metadata for model=%s, " + "tp_rank=%d, %d tensors, worker_id=%s", + transport, model_name, self.tp_rank, - session_id, - len(tensors), + tensor_count, worker_id, ) mx_source_id = mx_client.publish_metadata(identity, worker, worker_id) @@ -968,6 +954,86 @@ class ModelRunner(ModelRunnerKVCacheMixin): finally: 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): server_args = self.server_args @@ -1261,6 +1327,7 @@ class ModelRunner(ModelRunnerKVCacheMixin): modelexpress_ep_size=self.server_args.ep_size, modelexpress_dtype=self.server_args.dtype, modelexpress_quantization=self.server_args.quantization or "", + modelexpress_transport=self.server_args.modelexpress_transport, modelopt_config=modelopt_config, rl_quant_profile=self.server_args.rl_quant_profile, draft_model_idx=self.draft_model_idx, diff --git a/python/sglang/srt/model_loader/loader.py b/python/sglang/srt/model_loader/loader.py index 87dee392d..6b4450899 100644 --- a/python/sglang/srt/model_loader/loader.py +++ b/python/sglang/srt/model_loader/loader.py @@ -2303,7 +2303,12 @@ class RemoteInstanceModelLoader(BaseModelLoader): load_config: LoadConfig, 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: import grpc from modelexpress import p2p_pb2 @@ -2314,14 +2319,11 @@ class RemoteInstanceModelLoader(BaseModelLoader): "Install it with: pip install modelexpress" ) 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 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) for _, module in model.named_modules(): quant_method = getattr(module, "quant_method", None) @@ -2329,14 +2331,23 @@ class RemoteInstanceModelLoader(BaseModelLoader): with device_loading_context(module, target_device): quant_method.process_weights_after_loading(module) - 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 - ) + # Register local memory for the chosen transport + if transport == "nixl": + nixl_mgr = self._init_nixl_for_target(model, load_config, device_config) + else: + 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( model_name=model_name, backend_framework=p2p_pb2.BACKEND_FRAMEWORK_SGLANG, @@ -2347,11 +2358,11 @@ class RemoteInstanceModelLoader(BaseModelLoader): 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) try: logger.info( - "ModelExpress: looking for seed (model=%s, rank=%d)...", + "ModelExpress [%s]: looking for seed (model=%s, rank=%d)...", + transport, model_name, tp_rank, ) @@ -2380,7 +2391,6 @@ class RemoteInstanceModelLoader(BaseModelLoader): f"Ensure the seed instance is running and has published metadata." ) - # Fetch full metadata for the discovered worker response = mx_client.get_metadata( mx_source_id=source_ref.mx_source_id, worker_id=source_ref.worker_id, @@ -2393,31 +2403,46 @@ class RemoteInstanceModelLoader(BaseModelLoader): ) 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: 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 = [] client_ptr_list = [] client_len_list = [] @@ -2440,7 +2465,7 @@ class RemoteInstanceModelLoader(BaseModelLoader): client_len_list.append(local_size) logger.info( - "ModelExpress: starting RDMA transfer of %d tensors...", + "ModelExpress: starting TransferEngine RDMA of %d tensors...", len(seed_ptr_list), ) ret = transfer_engine.batch_transfer_sync_read( @@ -2454,10 +2479,88 @@ class RemoteInstanceModelLoader(BaseModelLoader): f"ModelExpress: batch_transfer_sync_read failed, error={ret}" ) - if hasattr(model, "post_load_weights"): - model.post_load_weights() + def _init_nixl_for_target(self, model, load_config, device_config): + """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): diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index 168d5f316..6997298e3 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -7106,18 +7106,24 @@ class ServerArgs: def modelexpress_source(self) -> bool: 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): # Use TransferEngine as seed backend. if self.remote_instance_weight_loader_start_seed_via_transfer_engine: return True - # ModelExpress source mode also needs TransferEngine init. - if self.modelexpress_source: + # ModelExpress source mode needs TransferEngine init only if transport is transfer_engine. + if self.modelexpress_source and self.modelexpress_transport == "transfer_engine": return True # Use TransferEngine as client backend. elif ( self.load_format == "remote_instance" and self.remote_instance_weight_loader_backend in ("transfer_engine", "modelexpress") + and self.modelexpress_transport == "transfer_engine" ): return True else: