support non disturbing remote instance weight loader v2 (#14997)
Signed-off-by: Anqi Shen <amy.saq@antgroup.com>
This commit is contained in:
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user