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-16 14:39:56 -08:00
committed by GitHub
parent a4c762811a
commit ccc8f3b266
11 changed files with 557 additions and 40 deletions
+69 -13
View File
@@ -18,6 +18,7 @@ from __future__ import annotations
import argparse
import dataclasses
import importlib
import importlib.util
import json
import logging
import os
@@ -614,6 +615,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_start_seed_via_transfer_engine: bool = False
# For PD-Multiplexing
enable_pdmux: bool = False
@@ -690,6 +693,9 @@ class ServerArgs:
# Handle speculative decoding logic.
self._handle_speculative_decoding()
# Handle remote instance weight loader.
self._handle_remote_instance_weight_loader_start_seed_via_transfer_engine()
# Handle model loading format.
self._handle_load_format()
@@ -2107,8 +2113,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 (
not self.validate_transfer_engine()
and self.remote_instance_weight_loader_backend == "transfer_engine"
):
logger.warning(
"Fallback load_format to 'auto' due to 'transfer_engine' backend is not supported."
)
self.load_format = "auto"
def _handle_encoder_disaggregation(self):
@@ -2366,19 +2390,12 @@ 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."
def _handle_remote_instance_weight_loader_start_seed_via_transfer_engine(self):
# Check whether TransferEngine can be used when users want to start seed service that supports TransferEngine backend.
if self.remote_instance_weight_loader_start_seed_via_transfer_engine:
self.remote_instance_weight_loader_start_seed_via_transfer_engine = (
self.validate_transfer_engine()
)
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):
@@ -4277,6 +4294,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-start-seed-via-transfer-engine",
action="store_true",
help="Start seed server via transfer engine backend for remote instance weight loader.",
)
# For PD-Multiplexing
parser.add_argument(
@@ -4782,6 +4811,33 @@ class ServerArgs:
original_server_arg_mem_fraction * final_overall_factor
)
def validate_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."
)
return 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."
)
return False
else:
return True
def remote_instance_weight_loader_use_transfer_engine(self):
# Use TransferEngine as seed backend.
if self.remote_instance_weight_loader_start_seed_via_transfer_engine:
return True
# Use TransferEngine as client backend.
elif (
self.load_format == "remote_instance"
and self.remote_instance_weight_loader_backend == "transfer_engine"
):
return True
else:
return False
# NOTE: This is a global variable to hold the server args for scheduler.
_global_server_args: Optional[ServerArgs] = None