Fix remote weight info nnode>1 and dp>1 (#17389)

This commit is contained in:
JD
2026-03-31 21:17:18 +08:00
committed by GitHub
parent ca2b2130ba
commit 20d07c4384
10 changed files with 434 additions and 82 deletions
+29 -8
View File
@@ -49,6 +49,9 @@ import uvloop
import zmq
from sglang.srt.elastic_ep.expert_backup_manager import run_expert_backup_manager
from sglang.srt.entrypoints.engine_info_bootstrap_server import (
EngineInfoBootstrapServer,
)
from sglang.srt.entrypoints.EngineBase import EngineBase
from sglang.srt.managers.data_parallel_controller import (
run_data_parallel_controller_process,
@@ -80,9 +83,6 @@ from sglang.srt.managers.scheduler import run_scheduler_process
from sglang.srt.managers.template_manager import TemplateManager
from sglang.srt.managers.tokenizer_manager import TokenizerManager
from sglang.srt.managers.tokenizer_manager_multiitem_mixin import ScoreResult
from sglang.srt.model_loader.remote_instance_weight_loader_utils import (
parse_remote_instance_transfer_engine_info_from_scheduler_infos,
)
from sglang.srt.observability.trace import process_tracing_init, trace_set_thread_info
from sglang.srt.server_args import PortArgs, ServerArgs
from sglang.srt.utils import (
@@ -98,7 +98,7 @@ from sglang.srt.utils import (
set_prometheus_multiproc_dir,
set_ulimit,
)
from sglang.srt.utils.network import get_zmq_socket
from sglang.srt.utils.network import get_zmq_socket, is_port_available
from sglang.srt.utils.torch_memory_saver_adapter import TorchMemorySaverAdapter
from sglang.srt.utils.watchdog import SubprocessWatchdog
from sglang.version import __version__
@@ -116,6 +116,7 @@ class SchedulerInitResult:
scheduler_infos: List[Dict[str, Any]]
wait_for_ready: Callable[[], None] = lambda: None
wait_for_completion: Callable[[], None] = lambda: None
engine_info_bootstrap_server: Optional[Any] = None
def init_tokenizer_manager(
@@ -201,11 +202,11 @@ class Engine(EngineBase):
if tokenizer_manager is not None:
tokenizer_manager._subprocess_watchdog = subprocess_watchdog
self.port_args = port_args
self.remote_instance_transfer_engine_info = (
parse_remote_instance_transfer_engine_info_from_scheduler_infos(
scheduler_init_result.scheduler_infos
# Access transfer engine info if bootstrap server is started.
if scheduler_init_result.engine_info_bootstrap_server is not None:
self.remote_instance_transfer_engine_info = (
scheduler_init_result.engine_info_bootstrap_server.transfer_engine_info
)
)
# Initialize ZMQ sockets
context = zmq.Context(2)
@@ -642,10 +643,30 @@ class Engine(EngineBase):
port_args = PortArgs.init_new(server_args)
logger.info(f"{server_args=}")
# Start the engine info bootstrap server if per-rank info is needed.
engine_info_bootstrap_server = None
if (
server_args.remote_instance_weight_loader_start_seed_via_transfer_engine
and server_args.node_rank == 0
):
bootstrap_port = server_args.engine_info_bootstrap_port
if not is_port_available(bootstrap_port):
raise RuntimeError(
f"engine_info_bootstrap_port {bootstrap_port} is already in use. "
f"When running multiple instances on the same node, each instance must use a "
f"different --engine-info-bootstrap-port."
)
engine_info_bootstrap_server = EngineInfoBootstrapServer(
host=server_args.host, port=bootstrap_port
)
# Launch scheduler processes
scheduler_init_result, scheduler_procs = cls._launch_scheduler_processes(
server_args, port_args, run_scheduler_process_func
)
scheduler_init_result.engine_info_bootstrap_server = (
engine_info_bootstrap_server
)
if (
server_args.enable_elastic_expert_backup
@@ -0,0 +1,105 @@
# Copyright 2023-2024 SGLang Team
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
# ==============================================================================
import logging
import threading
from typing import Dict, Optional, Tuple
import uvicorn
from fastapi import FastAPI, HTTPException
from fastapi.responses import PlainTextResponse
logger = logging.getLogger(__name__)
class EngineInfoBootstrapServer:
"""Lightweight HTTP server for per-rank model info registration.
Runs in a daemon thread on node_rank==0. Each ModelRunner registers its
info via HTTP PUT after model initialization. The Engine
accesses the collected info directly in-process; external consumers can
query via HTTP GET.
Currently supports transfer engine memory registration info.
"""
def __init__(self, host: str, port: int):
self.host = host
self.port = port
# Storage: {tp_rank: (session_id, weights_info_dict)}
self.transfer_engine_info: Dict[int, Tuple] = {}
self.lock = threading.Lock()
app = FastAPI()
@app.get("/health")
def health():
return PlainTextResponse("OK")
@app.put("/register_transfer_engine_info")
def register_transfer_engine_info(data: dict):
try:
tp_rank = data["tp_rank"]
info = data["transfer_engine_info"]
session_id = info["session_id"]
weights_info_dict = info["weights_info_dict"]
with self.lock:
self.transfer_engine_info[tp_rank] = (
session_id,
weights_info_dict,
)
logger.info(
f"Registered transfer engine info for tp_rank={tp_rank}, "
f"session_id={session_id}"
)
return PlainTextResponse("OK")
except Exception as e:
logger.error(f"Failed to register engine info: {e}")
raise HTTPException(status_code=400, detail=str(e))
@app.get("/get_transfer_engine_info")
def get_transfer_engine_info(rank: int):
if rank < 0:
raise HTTPException(status_code=400, detail="Invalid rank parameter")
with self.lock:
info = self.transfer_engine_info.get(rank)
if info is None:
raise HTTPException(
status_code=404,
detail=f"No transfer engine info for rank {rank}",
)
return {"rank": rank, "remote_instance_transfer_engine_info": list(info)}
config = uvicorn.Config(app, host=host, port=port, log_level="warning")
self._server = uvicorn.Server(config)
self._thread = threading.Thread(
target=self._server.run,
daemon=True,
)
self._thread.start()
logger.info(f"EngineInfoBootstrapServer started on {host}:{port}")
def close(self):
self._server.should_exit = True
self._thread.join(timeout=5)
def get_transfer_engine_info(self, rank: int) -> Optional[Tuple]:
"""Direct in-process access for co-located HTTP server (no HTTP round-trip)."""
return self.transfer_engine_info.get(rank)
+30 -35
View File
@@ -153,9 +153,6 @@ from sglang.srt.managers.multi_tokenizer_mixin import (
)
from sglang.srt.managers.template_manager import TemplateManager
from sglang.srt.managers.tokenizer_manager import ServerStatus, TokenizerManager
from sglang.srt.model_loader.remote_instance_weight_loader_utils import (
parse_remote_instance_transfer_engine_info_from_scheduler_infos,
)
from sglang.srt.observability.func_timer import enable_func_timer
from sglang.srt.observability.trace import (
process_tracing_init,
@@ -196,15 +193,6 @@ class _GlobalState:
tokenizer_manager: Union[TokenizerManager, MultiTokenizerRouter, TokenizerWorker]
template_manager: TemplateManager
scheduler_info: Dict
# Dict{
# rank: Tuple(
# session_id,
# Dict{
# name: Tuple (d_ptr, numel, element_size)
# }
# )
# }
remote_instance_transfer_engine_info: Optional[Dict] = None
_global_state: Optional[_GlobalState] = None
@@ -1030,26 +1018,39 @@ async def send_weights_to_remote_instance(
@app.get("/get_remote_instance_transfer_engine_info")
@auth_level(AuthLevel.ADMIN_OPTIONAL)
async def get_remote_instance_transfer_engine_info(rank: int = None):
"""Get the server information (deprecated - use /remote_instance_transfer_engine_info instead)."""
logger.warning(
"Endpoint '/get_remote_instance_transfer_engine_info' is deprecated and will be removed in a future version. "
"Please use '/remote_instance_transfer_engine_info' instead."
)
return await remote_instance_transfer_engine_info(rank=rank)
@app.get("/remote_instance_transfer_engine_info")
@auth_level(AuthLevel.ADMIN_OPTIONAL)
async def remote_instance_transfer_engine_info(rank: int = None):
if rank is None or rank < 0:
return Response(status_code=HTTPStatus.BAD_REQUEST)
if (
_global_state.remote_instance_transfer_engine_info is None
or len(_global_state.remote_instance_transfer_engine_info) == 0
):
return Response(status_code=HTTPStatus.BAD_REQUEST)
return ORJSONResponse(
{"error": {"message": "Missing or invalid rank parameter"}},
status_code=HTTPStatus.BAD_REQUEST,
)
server_args = _global_state.tokenizer_manager.server_args
try:
result = {
"rank": rank,
"remote_instance_transfer_engine_info": _global_state.remote_instance_transfer_engine_info[
rank
],
}
return result
except Exception as e:
logger.error(f"Exception: {e}")
return Response(status_code=HTTPStatus.BAD_REQUEST)
resp = requests.get(
f"{server_args.engine_info_bootstrap_url}/get_transfer_engine_info",
params={"rank": rank},
timeout=5,
)
if resp.status_code == 200:
return resp.json()
except (requests.exceptions.RequestException, ValueError) as e:
logger.warning(f"Failed to get transfer engine info for rank {rank}: {e}")
return ORJSONResponse(
{"error": {"message": f"Failed to get transfer engine info for rank {rank}"}},
status_code=HTTPStatus.BAD_REQUEST,
)
@app.post("/init_weights_update_group")
@@ -1993,18 +1994,12 @@ def _setup_and_run_http_server(
Called by launch_server after subprocesses have been launched.
"""
# Parse info got from the schedulers
remote_instance_transfer_engine_info = (
parse_remote_instance_transfer_engine_info_from_scheduler_infos(scheduler_infos)
)
# Set global states
set_global_state(
_GlobalState(
tokenizer_manager=tokenizer_manager,
template_manager=template_manager,
scheduler_info=scheduler_infos[0],
remote_instance_transfer_engine_info=remote_instance_transfer_engine_info,
)
)
-16
View File
@@ -1255,19 +1255,6 @@ class Scheduler(
"max_req_input_len": self.max_req_input_len,
}
if self.server_args.remote_instance_weight_loader_use_transfer_engine():
(
remote_instance_transfer_engine_session_id,
remote_instance_transfer_engine_weights_info_dict,
) = self.get_remote_instance_transfer_engine_info()
result_dict.update(
{
"tp_rank": self.tp_rank,
"remote_instance_transfer_engine_session_id": remote_instance_transfer_engine_session_id,
"remote_instance_transfer_engine_weights_info_dict": remote_instance_transfer_engine_weights_info_dict,
}
)
return result_dict
def run_event_loop(self) -> None:
@@ -3384,9 +3371,6 @@ class Scheduler(
):
pass
def get_remote_instance_transfer_engine_info(self):
return self.tp_worker.get_remote_instance_transfer_engine_info()
class IdleSleeper:
"""
-6
View File
@@ -441,12 +441,6 @@ class TpModelWorker(BaseTpWorker):
can_run_cuda_graph=can_run_cuda_graph,
)
def get_remote_instance_transfer_engine_info(self):
return (
self.model_runner.remote_instance_transfer_engine_session_id,
self.model_runner.remote_instance_transfer_engine_weight_info,
)
def forward_batch_generation(
self,
model_worker_batch: ModelWorkerBatch,
@@ -520,9 +520,11 @@ class ModelRunner(ModelRunnerKVCacheMixin):
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
# models for speculative decoding. In those cases, `num_nextn_predict_layers` is used to
@@ -700,6 +702,52 @@ class ModelRunner(ModelRunnerKVCacheMixin):
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 _publish_modelexpress_metadata(self):
"""Publish TransferEngine metadata to ModelExpress server (seed mode)."""
try:
@@ -106,21 +106,6 @@ def get_remote_instance_transfer_engine_info_per_rank(seed_url: str, rank: int):
return None, None
def parse_remote_instance_transfer_engine_info_from_scheduler_infos(scheduler_infos):
remote_instance_transfer_engine_info = {}
for data in scheduler_infos:
if (
"tp_rank" in data
and "remote_instance_transfer_engine_session_id" in data
and "remote_instance_transfer_engine_weights_info_dict" in data
):
remote_instance_transfer_engine_info[data["tp_rank"]] = (
data["remote_instance_transfer_engine_session_id"],
data["remote_instance_transfer_engine_weights_info_dict"],
)
return remote_instance_transfer_engine_info
def register_memory_region(model, transfer_engine):
if importlib.util.find_spec("torch") is None:
return register_memory_region_v1(model, transfer_engine)
+16 -2
View File
@@ -714,6 +714,7 @@ class ServerArgs:
"transfer_engine", "nccl", "modelexpress"
] = "nccl"
remote_instance_weight_loader_start_seed_via_transfer_engine: bool = False
engine_info_bootstrap_port: int = 6789
modelexpress_config: Optional[str] = None
# For PD-Multiplexing
@@ -5812,6 +5813,13 @@ class ServerArgs:
action="store_true",
help="Start seed server via transfer engine backend for remote instance weight loader.",
)
parser.add_argument(
"--engine-info-bootstrap-port",
type=int,
default=ServerArgs.engine_info_bootstrap_port,
help="Port for the engine info bootstrap server. Default is 6789. "
"Must be set explicitly when running multiple instances on the same node.",
)
parser.add_argument(
"--modelexpress-config",
type=str,
@@ -5931,7 +5939,7 @@ class ServerArgs:
attrs = [attr.name for attr in dataclasses.fields(cls)]
return cls(**{attr: getattr(args, attr) for attr in attrs})
def url(self):
def url(self, port: Optional[int] = None):
scheme = "https" if self.ssl_certfile else "http"
# When binding to all interfaces, use loopback for internal requests.
host = self.host
@@ -5939,7 +5947,13 @@ class ServerArgs:
host = "127.0.0.1"
elif host == "::":
host = "::1"
return NetworkAddress(host, self.port).to_url(scheme)
return NetworkAddress(host, port if port is not None else self.port).to_url(
scheme
)
@property
def engine_info_bootstrap_url(self):
return self.url(port=self.engine_info_bootstrap_port)
def ssl_verify(self):
"""Return the value for the requests library's ``verify=`` parameter.