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.graph_shared_output import GraphSharedOutput
|
||||||
from sglang.srt.model_executor.hook_manager import register_forward_hooks
|
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 (
|
from sglang.srt.model_executor.model_runner_components.weight_exporter import (
|
||||||
WeightExporter,
|
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.loader import get_model_loader
|
||||||
from sglang.srt.model_loader.remote_instance_weight_loader_utils import (
|
from sglang.srt.model_loader.remote_instance_weight_loader_utils import (
|
||||||
RemoteInstanceWeightLoaderBackend,
|
RemoteInstanceWeightLoaderBackend,
|
||||||
register_memory_region,
|
|
||||||
trigger_init_weights_send_group_for_remote_instance_request,
|
trigger_init_weights_send_group_for_remote_instance_request,
|
||||||
)
|
)
|
||||||
from sglang.srt.model_loader.utils import resolve_language_model
|
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.draft_model_idx = draft_model_idx
|
||||||
self.enable_hisparse = server_args.enable_hisparse
|
self.enable_hisparse = server_args.enable_hisparse
|
||||||
|
|
||||||
self.remote_instance_transfer_engine = None
|
self.init_remote_instance_weight_transporter()
|
||||||
self.remote_instance_transfer_engine_session_id = ""
|
|
||||||
self.remote_instance_transfer_engine_weight_info = None
|
|
||||||
|
|
||||||
self.msprobe_debugger = None
|
self.msprobe_debugger = None
|
||||||
if server_args.msprobe_dump_config is not None:
|
if server_args.msprobe_dump_config is not None:
|
||||||
@@ -550,6 +550,14 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
get_model=lambda: self.model,
|
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):
|
def init_msprobe(self):
|
||||||
# Init the msprobe
|
# Init the msprobe
|
||||||
try:
|
try:
|
||||||
@@ -587,7 +595,7 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
)
|
)
|
||||||
|
|
||||||
if self.server_args.remote_instance_weight_loader_use_transfer_engine():
|
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:
|
if not self.is_draft_worker:
|
||||||
set_global_expert_location_metadata(
|
set_global_expert_location_metadata(
|
||||||
@@ -650,21 +658,7 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
else None
|
else None
|
||||||
)
|
)
|
||||||
|
|
||||||
if (
|
self.remote_instance_weight_transporter.maybe_register_and_publish_weight_info()
|
||||||
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()
|
|
||||||
|
|
||||||
# For MTP models like DeepSeek-V3 or GLM-4.5, the MTP layer(s) are used separately as draft
|
# 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
|
# 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."
|
"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):
|
def check_quantized_moe_compatibility(self):
|
||||||
if (
|
if (
|
||||||
quantization_config := getattr(
|
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_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_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_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=self.remote_instance_weight_transporter.engine,
|
||||||
remote_instance_weight_loader_transfer_engine_session_id=self.remote_instance_transfer_engine_session_id,
|
remote_instance_weight_loader_transfer_engine_session_id=self.remote_instance_weight_transporter.session_id,
|
||||||
modelexpress_url=self.server_args.modelexpress_url,
|
modelexpress_url=self.server_args.modelexpress_url,
|
||||||
modelexpress_transport=self.server_args.modelexpress_transport,
|
modelexpress_transport=self.server_args.modelexpress_transport,
|
||||||
modelopt_config=modelopt_config,
|
modelopt_config=modelopt_config,
|
||||||
@@ -1265,7 +1193,7 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
device_config=DeviceConfig(self.device, self.gpu_id),
|
device_config=DeviceConfig(self.device, self.gpu_id),
|
||||||
)
|
)
|
||||||
if hasattr(self.loader, "remote_instance_transfer_engine_weight_info"):
|
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
|
self.loader.remote_instance_transfer_engine_weight_info
|
||||||
)
|
)
|
||||||
# Cache needs to be cleared after loading model weights (in the self.loader.load_model function).
|
# 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