diff --git a/docs/advanced_features/rfork.md b/docs/advanced_features/rfork.md new file mode 100644 index 000000000..67da3a843 --- /dev/null +++ b/docs/advanced_features/rfork.md @@ -0,0 +1,36 @@ +# R-Fork + +R-Fork (Tensor Remote Fork) is a novel weight loading methodology that leverages efficient inter-node GPU-to-GPU data transfer path to load tensors from a running SGLang instance to a new instance with zero-copy. It can significantly optimize the SGLang instance boot-up time by reducing model weights loading from several minutes to mere seconds. + +To learn more details about R-Fork, please check ** R-Fork blog ** + +## Usage + +| Argument | Usage | +|--------------|--------------------------------------------| +| load-format | set to `remote_instance` to enable R-Fork. | +| remote-instance-weight-loader-backend | `nccl` or `transfer_engine`, default value is `nccl` | +| remote-instance-weight-loader-seed-instance-ip | IP address of the seed instance who will provide the model weight | +| remote-instance-weight-loader-seed-instance-service-port | the port that the seed instance's HTTP server is listening on | +| remote-instance-weight-loader-send-weights-group-ports | the list of available ports on the seed instance that will be used to build NCCL communication groups between seed and client instance. This argument is only needed by `nccl` backend. | + +### NCCL as backend + +```shell +python -m sglang.launch_server [args] \ + --load-format remote_instance \ + --remote-instance-weight-loader-seed-instance-ip [seed_instance_ip] \ + --remote-instance-weight-loader-seed-instance-service-port [seed_instance_service_port] \ + --remote-instance-weight-loader-send-weights-group-ports [send_weights_nccl_group_ports_list] \ + --remote-instance-weight-loader-backend nccl +``` + +### TransferEngine as backend + +```shell +python -m sglang.launch_server [args] \ + --load-format remote_instance \ + --remote-instance-weight-loader-seed-instance-ip [seed_instance_ip] \ + --remote-instance-weight-loader-seed-instance-service-port [seed_instance_service_port] \ + --remote-instance-weight-loader-backend transfer_engine +``` diff --git a/python/sglang/srt/configs/load_config.py b/python/sglang/srt/configs/load_config.py index 4619dc286..77b96461a 100644 --- a/python/sglang/srt/configs/load_config.py +++ b/python/sglang/srt/configs/load_config.py @@ -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 diff --git a/python/sglang/srt/entrypoints/engine.py b/python/sglang/srt/entrypoints/engine.py index 912d1b9ec..377495d12 100644 --- a/python/sglang/srt/entrypoints/engine.py +++ b/python/sglang/srt/entrypoints/engine.py @@ -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, + ) diff --git a/python/sglang/srt/entrypoints/http_server.py b/python/sglang/srt/entrypoints/http_server.py index 5eb908d62..671de3aee 100644 --- a/python/sglang/srt/entrypoints/http_server.py +++ b/python/sglang/srt/entrypoints/http_server.py @@ -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, ) ) diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index 152d9717b..9696facef 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -2573,6 +2573,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: """ @@ -2686,13 +2689,29 @@ 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, - } - ) + if server_args.remote_instance_weight_loader_support_transfer_engine: + ( + remote_instance_transfer_engine_session_id, + remote_instance_transfer_engine_weights_info_dict, + ) = scheduler.get_remote_instance_transfer_engine_info() + pipe_writer.send( + { + "status": "ready", + "max_total_num_tokens": scheduler.max_total_num_tokens, + "max_req_input_len": scheduler.max_req_input_len, + "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, + } + ) + else: + pipe_writer.send( + { + "status": "ready", + "max_total_num_tokens": scheduler.max_total_num_tokens, + "max_req_input_len": scheduler.max_req_input_len, + } + ) disaggregation_mode: DisaggregationMode = scheduler.disaggregation_mode if disaggregation_mode == DisaggregationMode.NULL: diff --git a/python/sglang/srt/managers/tp_worker.py b/python/sglang/srt/managers/tp_worker.py index 758f0ffc9..c1e73f85b 100644 --- a/python/sglang/srt/managers/tp_worker.py +++ b/python/sglang/srt/managers/tp_worker.py @@ -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, diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index 26569c5fd..d65e5ef42 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -65,6 +65,7 @@ from sglang.srt.distributed import ( ) from sglang.srt.distributed.parallel_state import monkey_patch_vllm_parallel_state from sglang.srt.elastic_ep.elastic_ep import ElasticEPStateManager +from sglang.srt.environ import envs from sglang.srt.eplb.eplb_manager import EPLBManager from sglang.srt.eplb.expert_distribution import ( ExpertDistributionRecorder, @@ -135,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_v2, trigger_init_weights_send_group_for_remote_instance_request, ) from sglang.srt.model_loader.utils import set_default_torch_dtype @@ -157,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 +322,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 +400,9 @@ class ModelRunner: enable=self.server_args.enable_memory_saver ) + if self.server_args.remote_instance_weight_loader_support_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 +443,16 @@ class ModelRunner: self.sampler = Sampler() self.load_model() + if ( + self.server_args.remote_instance_weight_loader_support_transfer_engine + and self.remote_instance_transfer_engine_weight_info is None + ): + self.remote_instance_transfer_engine_weight_info = ( + register_memory_region_v2( + 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 +567,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 +801,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 +811,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 +840,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() diff --git a/python/sglang/srt/model_loader/loader.py b/python/sglang/srt/model_loader/loader.py index 4b8e5b084..5e41aea5f 100644 --- a/python/sglang/srt/model_loader/loader.py +++ b/python/sglang/srt/model_loader/loader.py @@ -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_v2, +) from sglang.srt.server_args import get_global_server_args # Try to import accelerate (optional dependency) @@ -1974,6 +1979,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 @@ -1992,16 +1998,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: @@ -2009,9 +2018,45 @@ 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_v2( + 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 @@ -2062,6 +2107,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.""" diff --git a/python/sglang/srt/model_loader/remote_instance_weight_loader_utils.py b/python/sglang/srt/model_loader/remote_instance_weight_loader_utils.py index 5974bba20..91d318a21 100644 --- a/python/sglang/srt/model_loader/remote_instance_weight_loader_utils.py +++ b/python/sglang/srt/model_loader/remote_instance_weight_loader_utils.py @@ -1,13 +1,21 @@ # SPDX-License-Identifier: Apache-2.0 +import enum import logging +import time from typing import List import requests +import torch 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 +75,110 @@ 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 + + +# DEPRECATED. Use register_memory_region_v2 instead. +def register_memory_region(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()) + + 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 diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index a0a943f6d..ea1470300 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -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( diff --git a/test/srt/test_load_weights_from_remote_instance.py b/test/srt/test_load_weights_from_remote_instance.py index bd722a920..e2a6f8516 100644 --- a/test/srt/test_load_weights_from_remote_instance.py +++ b/test/srt/test_load_weights_from_remote_instance.py @@ -62,6 +62,7 @@ def init_process( seed_instance_group_base_port, event_seed_ready, event_dst_ready_list, + remote_instance_loader_backend, ): torch.cuda.set_device(rank) @@ -90,6 +91,7 @@ def init_process( tp_size, event_seed_ready, event_dst_ready_list, + remote_instance_loader_backend, ) @@ -159,6 +161,7 @@ def init_process_dst( tp_size, event_seed_ready, event_dst_ready_list, + remote_instance_loader_backend, ): torch.cuda.set_device(rank * tp_size) torch.cuda.synchronize() @@ -186,6 +189,7 @@ def init_process_dst( remote_instance_weight_loader_seed_instance_service_port=seed_instance_service_port, remote_instance_weight_loader_send_weights_group_ports=ports, load_format="remote_instance", + remote_instance_weight_loader_backend=remote_instance_loader_backend, ) else: host, _, port = DEFAULT_URL_FOR_TEST.rpartition(":") @@ -213,6 +217,8 @@ def init_process_dst( f"[{','.join(str(port) for port in ports)}]", "--load-format", "remote_instance", + "--remote-instance-weight-loader-backend", + remote_instance_loader_backend, ), ) torch.cuda.synchronize() @@ -250,9 +256,10 @@ def test_load_weights_from_remote_instance( seed_instance_ip, seed_instance_service_port, seed_instance_group_base_port, + remote_instance_loader_backend, ): print( - f"Testing model: {model_name} tp_size: {tp_size}, dp_size: {dp_size} backend: {backends}" + f"Testing model: {model_name} tp_size: {tp_size}, dp_size: {dp_size} backend: {backends} remote_instance_loader_backend: {remote_instance_loader_backend}" ) param_queue = mp.Queue() results = {} @@ -276,6 +283,7 @@ def test_load_weights_from_remote_instance( seed_instance_group_base_port, event_seed_ready, event_dst_ready_list, + remote_instance_loader_backend, ), nprocs=1 + dp_size, join=False, @@ -340,14 +348,36 @@ class TestLoadWeightsFromRemoteInstance(CustomTestCase): # test_suits : tp, dp, model_name, backend, dst_instance_id if is_in_ci(): mode = random.choice(["Engine", "Server"]) + remote_instance_loader_backend = random.choice(["nccl", "transfer_engine"]) test_suits = [ - (1, 1, DEFAULT_SMALL_MODEL_NAME_FOR_TEST, [mode]), + ( + 1, + 1, + DEFAULT_SMALL_MODEL_NAME_FOR_TEST, + [mode], + remote_instance_loader_backend, + ), ] else: test_suits = [ - (1, 1, DEFAULT_SMALL_MODEL_NAME_FOR_TEST, ["Engine"]), - (1, 1, DEFAULT_SMALL_MODEL_NAME_FOR_TEST, ["Sever"]), - (2, 2, DEFAULT_SMALL_MODEL_NAME_FOR_TEST, ["Engine", "Server"]), + (1, 1, DEFAULT_SMALL_MODEL_NAME_FOR_TEST, ["Engine"], "nccl"), + (1, 1, DEFAULT_SMALL_MODEL_NAME_FOR_TEST, ["Sever"], "nccl"), + (2, 2, DEFAULT_SMALL_MODEL_NAME_FOR_TEST, ["Engine", "Server"], "nccl"), + ( + 1, + 1, + DEFAULT_SMALL_MODEL_NAME_FOR_TEST, + ["Engine"], + "transfer_engine", + ), + (1, 1, DEFAULT_SMALL_MODEL_NAME_FOR_TEST, ["Sever"], "transfer_engine"), + ( + 2, + 2, + DEFAULT_SMALL_MODEL_NAME_FOR_TEST, + ["Engine", "Server"], + "transfer_engine", + ), ] truncate_size = 10 @@ -365,7 +395,13 @@ class TestLoadWeightsFromRemoteInstance(CustomTestCase): "model.norm.weight", ] - for tp_size, dp_size, model_name, backends in test_suits: + for ( + tp_size, + dp_size, + model_name, + backends, + remote_instance_loader_backend, + ) in test_suits: test_load_weights_from_remote_instance( tp_size, dp_size, @@ -376,6 +412,7 @@ class TestLoadWeightsFromRemoteInstance(CustomTestCase): "127.0.0.1", DEFAULT_PORT_FOR_SRT_TEST_RUNNER + 1000, 60000, + remote_instance_loader_backend, )