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-12 08:45:32 +08:00
committed by GitHub
parent fd1ebbb0d6
commit 70758d457e
11 changed files with 500 additions and 29 deletions

View File

@@ -17,6 +17,8 @@ from __future__ import annotations
import argparse
import dataclasses
import importlib
import importlib.util
import json
import logging
import os
@@ -593,6 +595,8 @@ class ServerArgs:
remote_instance_weight_loader_seed_instance_ip: Optional[str] = None
remote_instance_weight_loader_seed_instance_service_port: Optional[int] = None
remote_instance_weight_loader_send_weights_group_ports: Optional[List[int]] = None
remote_instance_weight_loader_backend: Literal["transfer_engine", "nccl"] = "nccl"
remote_instance_weight_loader_support_transfer_engine: bool = False
# For PD-Multiplexing
enable_pdmux: bool = False
@@ -704,6 +708,9 @@ class ServerArgs:
# Handle elastic expert parallelism.
self._handle_elastic_ep()
# Handle remote instance weight loader.
self._handle_remote_instance_weight_loader_support_transfer_engine()
def _handle_deprecated_args(self):
# handle deprecated tool call parsers
deprecated_tool_call_parsers = {"qwen25": "qwen", "glm45": "glm"}
@@ -1882,8 +1889,26 @@ class ServerArgs:
if (
self.remote_instance_weight_loader_seed_instance_ip is None
or self.remote_instance_weight_loader_seed_instance_service_port is None
or self.remote_instance_weight_loader_send_weights_group_ports is None
):
logger.warning(
"Fallback load_format to 'auto' due to incomplete remote instance weight loader settings."
)
self.load_format = "auto"
elif (
self.remote_instance_weight_loader_send_weights_group_ports is None
and self.remote_instance_weight_loader_backend == "nccl"
):
logger.warning(
"Fallback load_format to 'auto' due to incomplete remote instance weight loader NCCL group ports settings."
)
self.load_format = "auto"
elif (
self.enable_memory_saver
and self.remote_instance_weight_loader_backend == "transfer_engine"
):
logger.warning(
"Fallback load_format to 'auto' due to incompatible remote instance weight loader transfer engine backend with memory saver."
)
self.load_format = "auto"
def _handle_disaggregation(self):
@@ -2118,6 +2143,20 @@ class ServerArgs:
self.disable_cuda_graph = True
self.skip_server_warmup = True
def _handle_remote_instance_weight_loader_support_transfer_engine(self):
if importlib.util.find_spec("mooncake.engine") is None:
logger.warning(
f"Failed to import mooncake.engine. Does not support using TransferEngine as remote instance weight loader backend."
)
self.remote_instance_weight_loader_support_transfer_engine = False
elif self.enable_memory_saver:
logger.warning(
"Memory saver is enabled, which is not compatible with TransferEngine. Does not support using TransferEngine as remote instance weight loader backend."
)
self.remote_instance_weight_loader_support_transfer_engine = False
else:
self.remote_instance_weight_loader_support_transfer_engine = True
@staticmethod
def add_cli_args(parser: argparse.ArgumentParser):
@@ -3947,6 +3986,18 @@ class ServerArgs:
default=ServerArgs.remote_instance_weight_loader_send_weights_group_ports,
help="The communication group ports for loading weights from remote instance.",
)
parser.add_argument(
"--remote-instance-weight-loader-backend",
type=str,
choices=["transfer_engine", "nccl"],
default=ServerArgs.remote_instance_weight_loader_backend,
help="The backend for loading weights from remote instance. Can be 'transfer_engine' or 'nccl'. Default is 'nccl'.",
)
parser.add_argument(
"--remote-instance-weight-loader-support-transfer-engine",
action="store_true",
help="Enable transfer engine support for remote instance weight loader.",
)
# For PD-Multiplexing
parser.add_argument(