support non disturbing remote instance weight loader v2 (#14997)

Signed-off-by: Anqi Shen <amy.saq@antgroup.com>
This commit is contained in:
amysaq2023
2025-12-17 06:39:56 +08:00
committed by GitHub
parent a4c762811a
commit ccc8f3b266
11 changed files with 557 additions and 40 deletions

View File

@@ -63,6 +63,9 @@ from sglang.srt.managers.multi_tokenizer_mixin import MultiTokenizerRouter
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.model_loader.remote_instance_weight_loader_utils import (
parse_remote_instance_transfer_engine_info_from_scheduler_infos,
)
from sglang.srt.server_args import PortArgs, ServerArgs
from sglang.srt.tracing.trace import process_tracing_init, trace_set_thread_info
from sglang.srt.utils import (
@@ -173,10 +176,9 @@ def _launch_subprocesses(
scheduler_infos.append(data)
# Get back some info from scheduler to tokenizer_manager
scheduler_info = scheduler_infos[0]
tokenizer_manager.max_req_input_len = scheduler_info["max_req_input_len"]
tokenizer_manager.max_req_input_len = scheduler_infos[0]["max_req_input_len"]
return tokenizer_manager, template_manager, scheduler_info, port_args
return tokenizer_manager, template_manager, scheduler_infos, port_args
class Engine(EngineBase):
@@ -221,13 +223,21 @@ class Engine(EngineBase):
atexit.register(self.shutdown)
# Launch subprocesses
tokenizer_manager, template_manager, scheduler_info, port_args = (
tokenizer_manager, template_manager, scheduler_infos, port_args = (
self.launch_subprocesses_func(server_args=server_args)
)
self.tokenizer_manager = tokenizer_manager
self.template_manager = template_manager
scheduler_info = scheduler_infos[0]
self.scheduler_info = scheduler_info
self.port_args = port_args
self.remote_instance_transfer_engine_info = (
parse_remote_instance_transfer_engine_info_from_scheduler_infos(
scheduler_infos
)
)
# Initialize ZMQ sockets
context = zmq.Context(2)

View File

@@ -123,6 +123,9 @@ 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.metrics.func_timer import enable_func_timer
from sglang.srt.model_loader.remote_instance_weight_loader_utils import (
parse_remote_instance_transfer_engine_info_from_scheduler_infos,
)
from sglang.srt.parser.reasoning_parser import ReasoningParser
from sglang.srt.server_args import PortArgs, ServerArgs
from sglang.srt.tracing.trace import process_tracing_init, trace_set_thread_info
@@ -152,6 +155,15 @@ 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
@@ -825,6 +837,30 @@ async def send_weights_to_remote_instance(
return ORJSONResponse(content, status_code=HTTPStatus.BAD_REQUEST)
@app.get("/get_remote_instance_transfer_engine_info")
async def get_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)
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)
@app.post("/init_weights_update_group")
async def init_weights_update_group(
obj: InitWeightsUpdateGroupReqInput, request: Request
@@ -1615,15 +1651,24 @@ def launch_server(
1. The HTTP server, Engine, and TokenizerManager all run in the main process.
2. Inter-process communication is done through IPC (each process uses a different port) via the ZMQ library.
"""
tokenizer_manager, template_manager, scheduler_info, port_args = (
tokenizer_manager, template_manager, scheduler_infos, port_args = (
launch_subprocesses_func(server_args=server_args)
)
scheduler_info = scheduler_infos[0]
remote_instance_transfer_engine_info = None
if server_args.remote_instance_weight_loader_use_transfer_engine():
remote_instance_transfer_engine_info = (
parse_remote_instance_transfer_engine_info_from_scheduler_infos(
scheduler_infos
)
)
set_global_state(
_GlobalState(
tokenizer_manager=tokenizer_manager,
template_manager=template_manager,
scheduler_info=scheduler_info,
remote_instance_transfer_engine_info=remote_instance_transfer_engine_info,
)
)