Revert several PRs (#14958)

Co-authored-by: fzyzcjy <ch271828n@outlook.com>
This commit is contained in:
Yineng Zhang
2025-12-12 11:25:12 -08:00
committed by GitHub
parent ec242f516e
commit 4b7b5af36a
11 changed files with 38 additions and 500 deletions

View File

@@ -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 **<a href=https://lmsys.org/blog/2025-12-10-rfork/> R-Fork blog </a>**
## 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
```

View File

@@ -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

View File

@@ -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

View File

@@ -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,
)
)

View File

@@ -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:

View File

@@ -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,

View File

@@ -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()

View File

@@ -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."""

View File

@@ -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

View File

@@ -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(

View File

@@ -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,
)