support non-disturbing remote-instance-weight-loader (#13125)

Signed-off-by: Anqi Shen <amy.saq@antgroup.com>
This commit is contained in:
amysaq2023
2025-12-11 16:45:32 -08:00
committed by GitHub
parent fd1ebbb0d6
commit 70758d457e
11 changed files with 500 additions and 29 deletions
+25 -4
View File
@@ -127,13 +127,18 @@ class Engine(EngineBase):
atexit.register(self.shutdown)
# Launch subprocesses
tokenizer_manager, template_manager, scheduler_info, port_args = (
_launch_subprocesses(server_args=server_args)
)
(
tokenizer_manager,
template_manager,
scheduler_info,
port_args,
remote_instance_transfer_engine_info,
) = _launch_subprocesses(server_args=server_args)
self.tokenizer_manager = tokenizer_manager
self.template_manager = template_manager
self.scheduler_info = scheduler_info
self.port_args = port_args
self.remote_instance_transfer_engine_info = remote_instance_transfer_engine_info
# Initialize ZMQ sockets
context = zmq.Context(2)
@@ -910,6 +915,7 @@ def _launch_subprocesses(
# Wait for the model to finish loading
scheduler_infos = []
remote_instance_transfer_engine_info = {}
for i in range(len(scheduler_pipe_readers)):
try:
data = scheduler_pipe_readers[i].recv()
@@ -926,9 +932,24 @@ def _launch_subprocesses(
"Initialization failed. Please see the error messages above."
)
scheduler_infos.append(data)
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"],
)
# Assume all schedulers have the same scheduler_info
scheduler_info = scheduler_infos[0]
tokenizer_manager.max_req_input_len = scheduler_info["max_req_input_len"]
return tokenizer_manager, template_manager, scheduler_info, port_args
return (
tokenizer_manager,
template_manager,
scheduler_info,
port_args,
remote_instance_transfer_engine_info,
)
+35 -3
View File
@@ -144,6 +144,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
@@ -813,6 +822,24 @@ 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)
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
@@ -1386,15 +1413,20 @@ 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 = (
_launch_subprocesses(server_args=server_args)
)
(
tokenizer_manager,
template_manager,
scheduler_info,
port_args,
remote_instance_transfer_engine_info,
) = _launch_subprocesses(server_args=server_args)
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,
)
)