Introduce RemoteInstanceWeightTransporter component (#31153)

This commit is contained in:
fzyzcjy
2026-07-14 15:57:43 +08:00
committed by GitHub
parent caa85ea022
commit 0f20f52e5e
2 changed files with 131 additions and 89 deletions
@@ -129,6 +129,9 @@ from sglang.srt.model_executor.forward_context import (
)
from sglang.srt.model_executor.graph_shared_output import GraphSharedOutput
from sglang.srt.model_executor.hook_manager import register_forward_hooks
from sglang.srt.model_executor.model_runner_components.remote_instance_weight_transporter import (
RemoteInstanceWeightTransporter,
)
from sglang.srt.model_executor.model_runner_components.weight_exporter import (
WeightExporter,
)
@@ -150,7 +153,6 @@ from sglang.srt.model_executor.runner import (
from sglang.srt.model_loader.loader import get_model_loader
from sglang.srt.model_loader.remote_instance_weight_loader_utils import (
RemoteInstanceWeightLoaderBackend,
register_memory_region,
trigger_init_weights_send_group_for_remote_instance_request,
)
from sglang.srt.model_loader.utils import resolve_language_model
@@ -311,9 +313,7 @@ class ModelRunner(ModelRunnerKVCacheMixin):
self.draft_model_idx = draft_model_idx
self.enable_hisparse = server_args.enable_hisparse
self.remote_instance_transfer_engine = None
self.remote_instance_transfer_engine_session_id = ""
self.remote_instance_transfer_engine_weight_info = None
self.init_remote_instance_weight_transporter()
self.msprobe_debugger = None
if server_args.msprobe_dump_config is not None:
@@ -550,6 +550,14 @@ class ModelRunner(ModelRunnerKVCacheMixin):
get_model=lambda: self.model,
)
def init_remote_instance_weight_transporter(self):
self.remote_instance_weight_transporter = RemoteInstanceWeightTransporter(
server_args=self.server_args,
get_model=lambda: self.model,
tp_rank=self.tp_rank,
gpu_id=self.gpu_id,
)
def init_msprobe(self):
# Init the msprobe
try:
@@ -587,7 +595,7 @@ class ModelRunner(ModelRunnerKVCacheMixin):
)
if self.server_args.remote_instance_weight_loader_use_transfer_engine():
self.remote_instance_init_transfer_engine()
self.remote_instance_weight_transporter.init_engine()
if not self.is_draft_worker:
set_global_expert_location_metadata(
@@ -650,21 +658,7 @@ class ModelRunner(ModelRunnerKVCacheMixin):
else None
)
if (
self.server_args.remote_instance_weight_loader_use_transfer_engine()
# ModelExpress owns TransferEngine memory registration and metadata
# publishing for backend=modelexpress. Re-registering here would
# overlap the same weight buffers.
and self.server_args.remote_instance_weight_loader_backend
!= RemoteInstanceWeightLoaderBackend.MODELEXPRESS
and self.remote_instance_transfer_engine is not None
and self.remote_instance_transfer_engine_weight_info is None
):
# Register memory and upstream the transfer engine info to the bootstrap server
self.remote_instance_transfer_engine_weight_info = register_memory_region(
self.model, self.remote_instance_transfer_engine
)
self._register_to_engine_info_bootstrap()
self.remote_instance_weight_transporter.maybe_register_and_publish_weight_info()
# For MTP models like DeepSeek-V3 or GLM-4.5, the MTP layer(s) are used separately as draft
# models for speculative decoding. In those cases, `num_nextn_predict_layers` is used to
@@ -993,72 +987,6 @@ class ModelRunner(ModelRunnerKVCacheMixin):
"one of which is required for DFLASH/DSPARK."
)
def remote_instance_init_transfer_engine(self):
try:
from mooncake.engine import TransferEngine
except ImportError:
logger.warning(
"Please install mooncake for using remote instance transfer engine: pip install mooncake-transfer-engine"
)
return
self.remote_instance_transfer_engine = TransferEngine()
local_ip = get_local_ip_auto()
self.remote_instance_transfer_engine.initialize(
local_ip,
"P2PHANDSHAKE",
envs.MOONCAKE_PROTOCOL.get(),
envs.MOONCAKE_DEVICE.get(),
)
self.remote_instance_transfer_engine_session_id = NetworkAddress(
local_ip, self.remote_instance_transfer_engine.get_rpc_port()
).to_host_port_str()
def _register_to_engine_info_bootstrap(self):
"""Register transfer engine info with the EngineInfoBootstrapServer via HTTP PUT.
The bootstrap server runs on node_rank==0. For multi-node setups, the
host is derived from dist_init_addr. For single-node, use 127.0.0.1.
"""
import requests as http_requests
if self.server_args.dist_init_addr:
# Multi-node: bootstrap server is on the head node (node_rank==0).
# Derive host from dist_init_addr (shared across all nodes).
bootstrap_host = (
NetworkAddress.parse(self.server_args.dist_init_addr).resolved().host
)
else:
bootstrap_host = "127.0.0.1"
bootstrap_port = self.server_args.engine_info_bootstrap_port
bootstrap_na = NetworkAddress(bootstrap_host, bootstrap_port)
url = f"{bootstrap_na.to_url()}/register_transfer_engine_info"
payload = {
"tp_rank": self.tp_rank,
"transfer_engine_info": {
"session_id": self.remote_instance_transfer_engine_session_id,
"weights_info_dict": self.remote_instance_transfer_engine_weight_info,
},
}
try:
resp = http_requests.put(url, json=payload, timeout=5)
if resp.status_code == 200:
logger.info(
f"Registered transfer engine info for tp_rank={self.tp_rank} "
f"with bootstrap server at {bootstrap_na}"
)
else:
logger.error(
f"Failed to register transfer engine info for tp_rank={self.tp_rank}: "
f"{resp.status_code}, {resp.text}"
)
except Exception as e:
logger.error(
f"Failed to register transfer engine info for tp_rank={self.tp_rank}: {e}"
)
def check_quantized_moe_compatibility(self):
if (
quantization_config := getattr(
@@ -1214,8 +1142,8 @@ class ModelRunner(ModelRunnerKVCacheMixin):
remote_instance_weight_loader_seed_instance_service_port=self.server_args.remote_instance_weight_loader_seed_instance_service_port,
remote_instance_weight_loader_send_weights_group_ports=self.server_args.remote_instance_weight_loader_send_weights_group_ports,
remote_instance_weight_loader_backend=self.server_args.remote_instance_weight_loader_backend,
remote_instance_weight_loader_transfer_engine=self.remote_instance_transfer_engine,
remote_instance_weight_loader_transfer_engine_session_id=self.remote_instance_transfer_engine_session_id,
remote_instance_weight_loader_transfer_engine=self.remote_instance_weight_transporter.engine,
remote_instance_weight_loader_transfer_engine_session_id=self.remote_instance_weight_transporter.session_id,
modelexpress_url=self.server_args.modelexpress_url,
modelexpress_transport=self.server_args.modelexpress_transport,
modelopt_config=modelopt_config,
@@ -1265,7 +1193,7 @@ class ModelRunner(ModelRunnerKVCacheMixin):
device_config=DeviceConfig(self.device, self.gpu_id),
)
if hasattr(self.loader, "remote_instance_transfer_engine_weight_info"):
self.remote_instance_transfer_engine_weight_info = (
self.remote_instance_weight_transporter.weight_info = (
self.loader.remote_instance_transfer_engine_weight_info
)
# Cache needs to be cleared after loading model weights (in the self.loader.load_model function).
@@ -0,0 +1,114 @@
from __future__ import annotations
import logging
from dataclasses import dataclass
from typing import Any, Callable, Optional
import torch
from sglang.srt.environ import envs
from sglang.srt.model_loader.remote_instance_weight_loader_utils import (
RemoteInstanceWeightLoaderBackend,
register_memory_region,
)
from sglang.srt.server_args import ServerArgs
from sglang.srt.utils.network import NetworkAddress, get_local_ip_auto
logger = logging.getLogger(__name__)
@dataclass(slots=True, kw_only=True)
class RemoteInstanceWeightTransporter:
server_args: ServerArgs
get_model: Callable[[], torch.nn.Module]
tp_rank: int
gpu_id: int
engine: Optional[Any] = None
session_id: str = ""
weight_info: Optional[dict[str, tuple[int, int, int]]] = None
_nixl_manager: Optional[Any] = None
@property
def model(self) -> torch.nn.Module:
return self.get_model()
def init_engine(self):
try:
from mooncake.engine import TransferEngine
except ImportError:
logger.warning(
"Please install mooncake for using remote instance transfer engine: pip install mooncake-transfer-engine"
)
return
self.engine = TransferEngine()
local_ip = get_local_ip_auto()
self.engine.initialize(
local_ip,
"P2PHANDSHAKE",
envs.MOONCAKE_PROTOCOL.get(),
envs.MOONCAKE_DEVICE.get(),
)
self.session_id = NetworkAddress(
local_ip, self.engine.get_rpc_port()
).to_host_port_str()
def maybe_register_and_publish_weight_info(self) -> None:
if (
self.server_args.remote_instance_weight_loader_use_transfer_engine()
# ModelExpress owns TransferEngine memory registration and metadata
# publishing for backend=modelexpress. Re-registering here would
# overlap the same weight buffers.
and self.server_args.remote_instance_weight_loader_backend
!= RemoteInstanceWeightLoaderBackend.MODELEXPRESS
and self.engine is not None
and self.weight_info is None
):
# Register memory and upstream the transfer engine info to the bootstrap server
self.weight_info = register_memory_region(self.model, self.engine)
self._register_to_engine_info_bootstrap()
def _register_to_engine_info_bootstrap(self: RemoteInstanceWeightTransporter):
"""Register transfer engine info with the EngineInfoBootstrapServer via HTTP PUT.
The bootstrap server runs on node_rank==0. For multi-node setups, the
host is derived from dist_init_addr. For single-node, use 127.0.0.1.
"""
import requests as http_requests
if self.server_args.dist_init_addr:
# Multi-node: bootstrap server is on the head node (node_rank==0).
# Derive host from dist_init_addr (shared across all nodes).
bootstrap_host = (
NetworkAddress.parse(self.server_args.dist_init_addr).resolved().host
)
else:
bootstrap_host = "127.0.0.1"
bootstrap_port = self.server_args.engine_info_bootstrap_port
bootstrap_na = NetworkAddress(bootstrap_host, bootstrap_port)
url = f"{bootstrap_na.to_url()}/register_transfer_engine_info"
payload = {
"tp_rank": self.tp_rank,
"transfer_engine_info": {
"session_id": self.session_id,
"weights_info_dict": self.weight_info,
},
}
try:
resp = http_requests.put(url, json=payload, timeout=5)
if resp.status_code == 200:
logger.info(
f"Registered transfer engine info for tp_rank={self.tp_rank} "
f"with bootstrap server at {bootstrap_na}"
)
else:
logger.error(
f"Failed to register transfer engine info for tp_rank={self.tp_rank}: "
f"{resp.status_code}, {resp.text}"
)
except Exception as e:
logger.error(
f"Failed to register transfer engine info for tp_rank={self.tp_rank}: {e}"
)