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
+3 -1
View File
@@ -2,7 +2,7 @@
import enum
import logging
from dataclasses import dataclass, field
from typing import List, Optional, Union
from typing import Any, List, Optional, Union
import orjson
@@ -73,6 +73,8 @@ class LoadConfig:
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: Optional[str] = None
remote_instance_weight_loader_transfer_engine: Optional[Any] = None
# ModelOpt-specific loading options
modelopt_checkpoint_restore_path: Optional[str] = None
+14 -4
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)
+46 -1
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,
)
)
+21 -7
View File
@@ -2656,6 +2656,9 @@ class Scheduler(
self.send_to_detokenizer.send_output(recv_req, recv_req)
return None
def get_remote_instance_transfer_engine_info(self):
return self.tp_worker.get_remote_instance_transfer_engine_info()
class IdleSleeper:
"""
@@ -2769,14 +2772,25 @@ def run_scheduler_process(
pp_rank,
dp_rank,
)
pipe_writer.send(
{
"status": "ready",
"max_total_num_tokens": scheduler.max_total_num_tokens,
"max_req_input_len": scheduler.max_req_input_len,
}
)
result_dict = {
"status": "ready",
"max_total_num_tokens": scheduler.max_total_num_tokens,
"max_req_input_len": scheduler.max_req_input_len,
}
if server_args.remote_instance_weight_loader_use_transfer_engine():
(
remote_instance_transfer_engine_session_id,
remote_instance_transfer_engine_weights_info_dict,
) = scheduler.get_remote_instance_transfer_engine_info()
result_dict.update(
{
"tp_rank": tp_rank,
"remote_instance_transfer_engine_session_id": remote_instance_transfer_engine_session_id,
"remote_instance_transfer_engine_weights_info_dict": remote_instance_transfer_engine_weights_info_dict,
}
)
pipe_writer.send(result_dict)
disaggregation_mode: DisaggregationMode = scheduler.disaggregation_mode
if disaggregation_mode == DisaggregationMode.NULL:
if scheduler.enable_pdmux:
+6
View File
@@ -366,6 +366,12 @@ class TpModelWorker(BaseTpWorker):
can_run_cuda_graph=can_run_cuda_graph,
)
def get_remote_instance_transfer_engine_info(self):
return (
self.model_runner.remote_instance_transfer_engine_session_id,
self.model_runner.remote_instance_transfer_engine_weight_info,
)
def forward_batch_generation(
self,
model_worker_batch: ModelWorkerBatch,
@@ -136,9 +136,10 @@ from sglang.srt.model_executor.input_buffers import GraphInputBuffers
from sglang.srt.model_executor.piecewise_cuda_graph_runner import (
PiecewiseCudaGraphRunner,
)
from sglang.srt.model_loader import get_model
from sglang.srt.model_loader.loader import DefaultModelLoader, get_model_loader
from sglang.srt.model_loader.remote_instance_weight_loader_utils import (
RemoteInstanceWeightLoaderBackend,
register_memory_region,
trigger_init_weights_send_group_for_remote_instance_request,
)
from sglang.srt.model_loader.utils import set_default_torch_dtype
@@ -158,6 +159,7 @@ from sglang.srt.utils import (
get_available_gpu_memory,
get_bool_env_var,
get_cpu_ids_by_node,
get_local_ip_auto,
init_custom_process_group,
is_cuda,
is_float4_e2m1fn_x2,
@@ -319,6 +321,10 @@ class ModelRunner:
self.forward_pass_id = 0
self.init_new_workspace = False
self.remote_instance_transfer_engine = None
self.remote_instance_transfer_engine_session_id = ""
self.remote_instance_transfer_engine_weight_info = None
# Apply the rank zero filter to logger
if server_args.show_time_cost:
enable_show_time_cost()
@@ -393,6 +399,9 @@ class ModelRunner:
enable=self.server_args.enable_memory_saver
)
if self.server_args.remote_instance_weight_loader_use_transfer_engine():
self.remote_instance_init_transfer_engine()
if not self.is_draft_worker:
set_global_expert_location_metadata(
compute_initial_expert_location_metadata(
@@ -433,6 +442,15 @@ class ModelRunner:
self.sampler = Sampler()
self.load_model()
if (
self.server_args.remote_instance_weight_loader_use_transfer_engine()
and self.remote_instance_transfer_engine is not None
and self.remote_instance_transfer_engine_weight_info is None
):
self.remote_instance_transfer_engine_weight_info = register_memory_region(
self.model, self.remote_instance_transfer_engine
)
# Check if the model is using hybrid SWA
if (
not self.server_args.disable_hybrid_swa_memory
@@ -547,6 +565,23 @@ class ModelRunner:
# Initialize piecewise CUDA graph
self.init_piecewise_cuda_graphs()
def remote_instance_init_transfer_engine(self):
try:
from mooncake.engine import TransferEngine
except ImportError as e:
logger.warning(
"Please install mooncake for using remote instance transfer engine: pip install mooncake"
)
return
self.remote_instance_transfer_engine = TransferEngine()
local_ip = get_local_ip_auto()
self.remote_instance_transfer_engine.initialize(
local_ip, "P2PHANDSHAKE", "rdma", envs.MOONCAKE_DEVICE.value
)
self.remote_instance_transfer_engine_session_id = (
f"{local_ip}:{self.remote_instance_transfer_engine.get_rpc_port()}"
)
def model_specific_adjustment(self):
server_args = self.server_args
@@ -764,6 +799,8 @@ class ModelRunner:
remote_instance_weight_loader_seed_instance_ip=self.server_args.remote_instance_weight_loader_seed_instance_ip,
remote_instance_weight_loader_seed_instance_service_port=self.server_args.remote_instance_weight_loader_seed_instance_service_port,
remote_instance_weight_loader_send_weights_group_ports=self.server_args.remote_instance_weight_loader_send_weights_group_ports,
remote_instance_weight_loader_backend=self.server_args.remote_instance_weight_loader_backend,
remote_instance_weight_loader_transfer_engine=self.remote_instance_transfer_engine,
modelopt_config=modelopt_config,
rl_quant_profile=self.server_args.rl_quant_profile,
)
@@ -772,7 +809,11 @@ class ModelRunner:
self.model_config, self.load_config, self.tp_size
)
if self.server_args.load_format == LoadFormat.REMOTE_INSTANCE:
if (
self.server_args.load_format == LoadFormat.REMOTE_INSTANCE
and self.server_args.remote_instance_weight_loader_backend
== RemoteInstanceWeightLoaderBackend.NCCL
):
if self.tp_rank == 0:
instance_ip = socket.gethostbyname(socket.gethostname())
t = threading.Thread(
@@ -797,11 +838,18 @@ class ModelRunner:
GPU_MEMORY_TYPE_WEIGHTS,
enable_cpu_backup=enable_cpu_backup,
):
self.model = get_model(
model_config=self.model_config,
self.loader = get_model_loader(
load_config=self.load_config,
model_config=self.model_config,
)
self.model = self.loader.load_model(
model_config=self.model_config,
device_config=DeviceConfig(self.device, self.gpu_id),
)
if hasattr(self.loader, "remote_instance_transfer_engine_weight_info"):
self.remote_instance_transfer_engine_weight_info = (
self.loader.remote_instance_transfer_engine_weight_info
)
monkey_patch_vllm_parallel_state(reverse=True)
get_offloader().post_init()
+104 -4
View File
@@ -34,6 +34,11 @@ import huggingface_hub
import numpy as np
import torch
from sglang.srt.model_loader.remote_instance_weight_loader_utils import (
RemoteInstanceWeightLoaderBackend,
get_remote_instance_transfer_engine_info_per_rank,
register_memory_region,
)
from sglang.srt.server_args import get_global_server_args
# Try to import accelerate (optional dependency)
@@ -1987,6 +1992,7 @@ class RemoteInstanceModelLoader(BaseModelLoader):
f"Model loader extra config is not supported for "
f"load format {load_config.load_format}"
)
self.remote_instance_transfer_engine_weight_info = None
def download_model(self, model_config: ModelConfig) -> None:
raise NotImplementedError
@@ -2005,16 +2011,19 @@ class RemoteInstanceModelLoader(BaseModelLoader):
f"load format {load_config.load_format}"
)
model_weights = f"instance://{load_config.remote_instance_weight_loader_seed_instance_ip}:{load_config.remote_instance_weight_loader_send_weights_group_ports[load_config.tp_rank]}"
with set_default_torch_dtype(model_config.dtype):
with torch.device(device_config.device):
model = _initialize_model(model_config, self.load_config)
if (
load_config.remote_instance_weight_loader_backend
== RemoteInstanceWeightLoaderBackend.NCCL
):
model_weights = f"instance://{load_config.remote_instance_weight_loader_seed_instance_ip}:{load_config.remote_instance_weight_loader_send_weights_group_ports[load_config.tp_rank]}"
with create_remote_connector(model_weights, device_config.device) as client:
connector_type = get_connector_type(client)
if connector_type == ConnectorType.INSTANCE:
self.load_model_from_remote_instance(
self.load_model_from_remote_instance_by_nccl(
model, client, model_config, device_config
)
else:
@@ -2022,9 +2031,43 @@ class RemoteInstanceModelLoader(BaseModelLoader):
f"Unsupported connector type {connector_type} for "
f"remote tensor model loading."
)
elif (
load_config.remote_instance_weight_loader_backend
== RemoteInstanceWeightLoaderBackend.TRANSFER_ENGINE
):
if load_config.remote_instance_weight_loader_transfer_engine is None:
raise RuntimeError(
"Transfer engine is not initialized for remote instance "
"model loader with `transfer_engine` backend. "
)
logger.info(
"TransferEngine registering memory regions (this may take a few seconds)..."
)
# register memory region
self.remote_instance_transfer_engine_weight_info = register_memory_region(
model, load_config.remote_instance_weight_loader_transfer_engine
)
logger.info(
"TransferEngine memory regions have been successfully registered."
)
# transfer weights
success = self.load_model_from_remote_instance_by_transfer_engine(
model,
load_config.remote_instance_weight_loader_transfer_engine,
f"http://{load_config.remote_instance_weight_loader_seed_instance_ip}:{load_config.remote_instance_weight_loader_seed_instance_service_port}",
load_config.tp_rank,
)
if not success:
raise RuntimeError(
"Failed to load weights from remote instance via transfer engine."
)
else:
raise ValueError("Invalid remote instance weight loader backend.")
return model.eval()
def load_model_from_remote_instance(
def load_model_from_remote_instance_by_nccl(
self, model, client, model_config: ModelConfig, device_config: DeviceConfig
) -> nn.Module:
load_config = self.load_config
@@ -2075,6 +2118,63 @@ class RemoteInstanceModelLoader(BaseModelLoader):
)
torch.cuda.empty_cache()
def load_model_from_remote_instance_by_transfer_engine(
self, model, transfer_engine, seed_url, tp_rank
) -> bool:
# get remote weights metadata from source instance
seed_transfer_engine_session_id, seed_transfer_engine_weight_info = (
get_remote_instance_transfer_engine_info_per_rank(seed_url, tp_rank)
)
if (
seed_transfer_engine_session_id is None
or seed_transfer_engine_weight_info is None
):
logger.error("Cannot get transfer engine session or weight info.")
return False
# prepare local/remote RDMA keys
seed_ptr_list = []
client_ptr_list = []
client_len_list = []
for name, tensor in model.named_parameters():
weight_info = seed_transfer_engine_weight_info.get(name, None)
if weight_info is None:
logger.error(f"Cannot find weight info for {name}.")
return False
seed_ptr, seed_numel, seed_element_size = weight_info
if (
seed_numel != tensor.numel()
or seed_element_size != tensor.element_size()
):
logger.error(
f"Weight info does not match for {name}, "
f"expected ({seed_numel}, {seed_element_size}), "
f"got ({tensor.numel()}, {tensor.element_size()})"
)
return False
client_ptr = tensor.data_ptr()
client_len = tensor.numel() * tensor.element_size()
seed_ptr_list.append(seed_ptr)
client_ptr_list.append(client_ptr)
client_len_list.append(client_len)
# load weights from source instance through TransferEngine
ret = transfer_engine.batch_transfer_sync_read(
seed_transfer_engine_session_id,
client_ptr_list,
seed_ptr_list,
client_len_list,
)
if ret < 0:
logger.error(f"batch transfer failed, error: {ret}")
return False
if hasattr(model, "post_load_weights"):
model.post_load_weights()
return True
class RemoteModelLoader(BaseModelLoader):
"""Model loader that can load Tensors from remote database."""
@@ -1,6 +1,10 @@
# SPDX-License-Identifier: Apache-2.0
import enum
import importlib
import importlib.util
import logging
import time
from typing import List
import requests
@@ -8,6 +12,11 @@ import requests
logger = logging.getLogger(__name__)
class RemoteInstanceWeightLoaderBackend(str, enum.Enum):
NCCL = "nccl"
TRANSFER_ENGINE = "transfer_engine"
def trigger_init_weights_send_group_for_remote_instance_request(
remote_instance_weight_loader_seed_instance_ip: str,
remote_instance_weight_loader_seed_instance_service_port: int,
@@ -67,3 +76,133 @@ def trigger_transferring_weights_request(
except Exception as e:
logger.error(f"Failed to trigger send weights to remote instance request: {e}")
raise
def get_remote_instance_transfer_engine_info_per_rank(seed_url: str, rank: int):
try:
response = requests.get(
f"{seed_url}/get_remote_instance_transfer_engine_info",
params={
"rank": rank,
},
)
if response.status_code == 200:
data = response.json()
if "remote_instance_transfer_engine_info" in data:
return data["remote_instance_transfer_engine_info"]
else:
logger.error(
"Failed to get `remote_instance_transfer_engine_info` in response."
)
return None, None
else:
logger.error(f"request.get failed: {response.status_code}")
return None, None
except Exception as e:
logger.error(f"Exception: {e}")
return None, None
def parse_remote_instance_transfer_engine_info_from_scheduler_infos(scheduler_infos):
remote_instance_transfer_engine_info = {}
for data in scheduler_infos:
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"],
)
return remote_instance_transfer_engine_info
def register_memory_region(model, transfer_engine):
if importlib.util.find_spec("torch") is None:
return register_memory_region_v1(model, transfer_engine)
else:
return register_memory_region_v2(model, transfer_engine)
def register_memory_region_v1(model, transfer_engine):
start_tic = time.time()
weight_mr_dict = {}
for name, weight in model.named_parameters():
ret = transfer_engine.register_memory(
weight.data_ptr(), weight.numel() * weight.element_size()
)
if ret != 0:
raise RuntimeError(
f"register memory failed for weight {name}, error: {ret}"
)
weight_mr_dict[name] = (
weight.data_ptr(),
weight.numel(),
weight.element_size(),
)
end_tic = time.time()
logger.debug(f"Register memory region time: {(end_tic - start_tic):.4f}s")
return weight_mr_dict
def register_memory_region_v2(model, transfer_engine):
start_tic = time.time()
weight_mr_dict = {}
weight_addr_set = set()
for name, weight in model.named_parameters():
weight_mr_dict[name] = (
weight.data_ptr(),
weight.numel(),
weight.element_size(),
)
weight_addr_set.add(weight.data_ptr())
import torch
memory_snapshot = torch.cuda.memory.memory_snapshot()
weight_blocks_for_reg_mr = []
# Blocks in each segment have continuous physical addresses,
# so they can be merged for memory registration.
for segment in memory_snapshot:
current_weight_block = None
blocks = segment.get("blocks", [])
for block in blocks:
address = block.get("address", -1)
size = block.get("size", -1)
state = block.get("state", "")
if address < 0 or size < 0 or state == "":
continue
# Only register active allocated memory blocks that hold weights.
if state == "active_allocated":
if address in weight_addr_set:
if current_weight_block is None:
current_weight_block = (address, size)
elif current_weight_block[0] + current_weight_block[1] == address:
current_weight_block = (
current_weight_block[0],
current_weight_block[1] + size,
)
else:
weight_blocks_for_reg_mr.append(current_weight_block)
current_weight_block = (address, size)
if current_weight_block is not None:
weight_blocks_for_reg_mr.append(current_weight_block)
# Register merged memory blocks that hold weights.
for weight_block in weight_blocks_for_reg_mr:
address, size = weight_block
ret = transfer_engine.register_memory(address, size)
if ret != 0:
raise RuntimeError(
f"register memory failed for weight block at address {address} with size {size}, error: {ret}"
)
end_tic = time.time()
logger.debug(f"Register memory region v2 time: {(end_tic - start_tic):.4f}s")
return weight_mr_dict
+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