diff --git a/docs/advanced_features/rfork.md b/docs/advanced_features/rfork.md deleted file mode 100644 index 67da3a843..000000000 --- a/docs/advanced_features/rfork.md +++ /dev/null @@ -1,36 +0,0 @@ -# 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 77b96461a..4619dc286 100644 --- a/python/sglang/srt/configs/load_config.py +++ b/python/sglang/srt/configs/load_config.py @@ -73,8 +73,6 @@ 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 377495d12..912d1b9ec 100644 --- a/python/sglang/srt/entrypoints/engine.py +++ b/python/sglang/srt/entrypoints/engine.py @@ -127,18 +127,13 @@ class Engine(EngineBase): atexit.register(self.shutdown) # Launch subprocesses - ( - tokenizer_manager, - template_manager, - scheduler_info, - port_args, - remote_instance_transfer_engine_info, - ) = _launch_subprocesses(server_args=server_args) + tokenizer_manager, template_manager, scheduler_info, port_args = ( + _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) @@ -915,7 +910,6 @@ 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() @@ -932,24 +926,9 @@ 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, - remote_instance_transfer_engine_info, - ) + return tokenizer_manager, template_manager, scheduler_info, port_args diff --git a/python/sglang/srt/entrypoints/http_server.py b/python/sglang/srt/entrypoints/http_server.py index 671de3aee..5eb908d62 100644 --- a/python/sglang/srt/entrypoints/http_server.py +++ b/python/sglang/srt/entrypoints/http_server.py @@ -144,15 +144,6 @@ 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 @@ -822,24 +813,6 @@ 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 @@ -1413,20 +1386,15 @@ 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, - remote_instance_transfer_engine_info, - ) = _launch_subprocesses(server_args=server_args) + tokenizer_manager, template_manager, scheduler_info, port_args = ( + _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 4a3885b30..094e1aca3 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -2609,9 +2609,6 @@ 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: """ @@ -2725,29 +2722,13 @@ def run_scheduler_process( pp_rank, dp_rank, ) - 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, - } - ) + 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 b3d933df4..8853f5ba1 100644 --- a/python/sglang/srt/managers/tp_worker.py +++ b/python/sglang/srt/managers/tp_worker.py @@ -366,12 +366,6 @@ 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 a3fc18e57..4e82119e9 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -136,10 +136,9 @@ 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 @@ -159,7 +158,6 @@ 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,10 +317,6 @@ 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() @@ -397,9 +391,6 @@ 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( @@ -440,16 +431,6 @@ 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 @@ -564,23 +545,6 @@ 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 @@ -798,8 +762,6 @@ 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, ) @@ -808,11 +770,7 @@ class ModelRunner: self.model_config, self.load_config, self.tp_size ) - if ( - self.server_args.load_format == LoadFormat.REMOTE_INSTANCE - and self.server_args.remote_instance_weight_loader_backend - == RemoteInstanceWeightLoaderBackend.NCCL - ): + if self.server_args.load_format == LoadFormat.REMOTE_INSTANCE: if self.tp_rank == 0: instance_ip = socket.gethostbyname(socket.gethostname()) t = threading.Thread( @@ -837,18 +795,11 @@ class ModelRunner: GPU_MEMORY_TYPE_WEIGHTS, enable_cpu_backup=enable_cpu_backup, ): - self.loader = get_model_loader( + self.model = get_model( + model_config=self.model_config, 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 a981bcf1e..cd9cdcde2 100644 --- a/python/sglang/srt/model_loader/loader.py +++ b/python/sglang/srt/model_loader/loader.py @@ -34,11 +34,6 @@ 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) @@ -1992,7 +1987,6 @@ 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 @@ -2011,19 +2005,16 @@ 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_by_nccl( + self.load_model_from_remote_instance( model, client, model_config, device_config ) else: @@ -2031,45 +2022,9 @@ 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_by_nccl( + def load_model_from_remote_instance( self, model, client, model_config: ModelConfig, device_config: DeviceConfig ) -> nn.Module: load_config = self.load_config @@ -2120,63 +2075,6 @@ 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 91d318a21..5974bba20 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,21 +1,13 @@ # 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, @@ -75,110 +67,3 @@ 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 bc17db6f8..f3942976a 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -18,7 +18,6 @@ from __future__ import annotations import argparse import dataclasses import importlib -import importlib.util import json import logging import os @@ -596,8 +595,6 @@ 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 @@ -710,9 +707,6 @@ 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"} @@ -1971,26 +1965,8 @@ 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): @@ -2226,25 +2202,18 @@ class ServerArgs: self.skip_server_warmup = True def _handle_remote_instance_weight_loader_support_transfer_engine(self): - try: - importlib.import_module("mooncake") - 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 - except ImportError: + if importlib.util.find_spec("mooncake.engine") is None: logger.warning( - f"Failed to import mooncake. Does not support using TransferEngine as remote instance weight loader backend." + 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): @@ -4087,18 +4056,6 @@ 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 e2a6f8516..bd722a920 100644 --- a/test/srt/test_load_weights_from_remote_instance.py +++ b/test/srt/test_load_weights_from_remote_instance.py @@ -62,7 +62,6 @@ def init_process( seed_instance_group_base_port, event_seed_ready, event_dst_ready_list, - remote_instance_loader_backend, ): torch.cuda.set_device(rank) @@ -91,7 +90,6 @@ def init_process( tp_size, event_seed_ready, event_dst_ready_list, - remote_instance_loader_backend, ) @@ -161,7 +159,6 @@ 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() @@ -189,7 +186,6 @@ 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(":") @@ -217,8 +213,6 @@ 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() @@ -256,10 +250,9 @@ 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} remote_instance_loader_backend: {remote_instance_loader_backend}" + f"Testing model: {model_name} tp_size: {tp_size}, dp_size: {dp_size} backend: {backends}" ) param_queue = mp.Queue() results = {} @@ -283,7 +276,6 @@ 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, @@ -348,36 +340,14 @@ 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], - remote_instance_loader_backend, - ), + (1, 1, DEFAULT_SMALL_MODEL_NAME_FOR_TEST, [mode]), ] else: test_suits = [ - (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", - ), + (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"]), ] truncate_size = 10 @@ -395,13 +365,7 @@ class TestLoadWeightsFromRemoteInstance(CustomTestCase): "model.norm.weight", ] - for ( - tp_size, - dp_size, - model_name, - backends, - remote_instance_loader_backend, - ) in test_suits: + for tp_size, dp_size, model_name, backends in test_suits: test_load_weights_from_remote_instance( tp_size, dp_size, @@ -412,7 +376,6 @@ class TestLoadWeightsFromRemoteInstance(CustomTestCase): "127.0.0.1", DEFAULT_PORT_FOR_SRT_TEST_RUNNER + 1000, 60000, - remote_instance_loader_backend, )