Introduce RemoteInstanceWeightTransporter component (#31153)
This commit is contained in:
@@ -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).
|
||||
|
||||
+114
@@ -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}"
|
||||
)
|
||||
Reference in New Issue
Block a user