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
+49
View File
@@ -0,0 +1,49 @@
# 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. |
| remote-instance-weight-loader-start-seed-via-transfer-engine | set to start seed service that supports TransferEngine as backend. It is needed for seed instances when using `transfer_engine` as backend. |
### NCCL as backend
seed instance:
```shell
python -m sglang.launch_server [args]
```
client instance:
```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
seed instance:
```shell
python -m sglang.launch_server [args] \
--remote-instance-weight-loader-start-seed-via-transfer-engine
```
```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
```
+3 -1
View File
@@ -2,7 +2,7 @@
import enum import enum
import logging import logging
from dataclasses import dataclass, field from dataclasses import dataclass, field
from typing import List, Optional, Union from typing import Any, List, Optional, Union
import orjson import orjson
@@ -73,6 +73,8 @@ class LoadConfig:
remote_instance_weight_loader_seed_instance_ip: Optional[str] = None 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_seed_instance_service_port: Optional[int] = None
remote_instance_weight_loader_send_weights_group_ports: Optional[List[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-specific loading options
modelopt_checkpoint_restore_path: Optional[str] = None 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.scheduler import run_scheduler_process
from sglang.srt.managers.template_manager import TemplateManager from sglang.srt.managers.template_manager import TemplateManager
from sglang.srt.managers.tokenizer_manager import TokenizerManager 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.server_args import PortArgs, ServerArgs
from sglang.srt.tracing.trace import process_tracing_init, trace_set_thread_info from sglang.srt.tracing.trace import process_tracing_init, trace_set_thread_info
from sglang.srt.utils import ( from sglang.srt.utils import (
@@ -173,10 +176,9 @@ def _launch_subprocesses(
scheduler_infos.append(data) scheduler_infos.append(data)
# Get back some info from scheduler to tokenizer_manager # Get back some info from scheduler to tokenizer_manager
scheduler_info = scheduler_infos[0] tokenizer_manager.max_req_input_len = scheduler_infos[0]["max_req_input_len"]
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_infos, port_args
class Engine(EngineBase): class Engine(EngineBase):
@@ -221,13 +223,21 @@ class Engine(EngineBase):
atexit.register(self.shutdown) atexit.register(self.shutdown)
# Launch subprocesses # 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.launch_subprocesses_func(server_args=server_args)
) )
self.tokenizer_manager = tokenizer_manager self.tokenizer_manager = tokenizer_manager
self.template_manager = template_manager self.template_manager = template_manager
scheduler_info = scheduler_infos[0]
self.scheduler_info = scheduler_info self.scheduler_info = scheduler_info
self.port_args = port_args 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 # Initialize ZMQ sockets
context = zmq.Context(2) 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.template_manager import TemplateManager
from sglang.srt.managers.tokenizer_manager import ServerStatus, TokenizerManager from sglang.srt.managers.tokenizer_manager import ServerStatus, TokenizerManager
from sglang.srt.metrics.func_timer import enable_func_timer 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.parser.reasoning_parser import ReasoningParser
from sglang.srt.server_args import PortArgs, ServerArgs from sglang.srt.server_args import PortArgs, ServerArgs
from sglang.srt.tracing.trace import process_tracing_init, trace_set_thread_info 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] tokenizer_manager: Union[TokenizerManager, MultiTokenizerRouter, TokenizerWorker]
template_manager: TemplateManager template_manager: TemplateManager
scheduler_info: Dict 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 _global_state: Optional[_GlobalState] = None
@@ -825,6 +837,30 @@ async def send_weights_to_remote_instance(
return ORJSONResponse(content, status_code=HTTPStatus.BAD_REQUEST) 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") @app.post("/init_weights_update_group")
async def init_weights_update_group( async def init_weights_update_group(
obj: InitWeightsUpdateGroupReqInput, request: Request obj: InitWeightsUpdateGroupReqInput, request: Request
@@ -1615,15 +1651,24 @@ def launch_server(
1. The HTTP server, Engine, and TokenizerManager all run in the main process. 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. 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) 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( set_global_state(
_GlobalState( _GlobalState(
tokenizer_manager=tokenizer_manager, tokenizer_manager=tokenizer_manager,
template_manager=template_manager, template_manager=template_manager,
scheduler_info=scheduler_info, 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) self.send_to_detokenizer.send_output(recv_req, recv_req)
return None return None
def get_remote_instance_transfer_engine_info(self):
return self.tp_worker.get_remote_instance_transfer_engine_info()
class IdleSleeper: class IdleSleeper:
""" """
@@ -2769,14 +2772,25 @@ def run_scheduler_process(
pp_rank, pp_rank,
dp_rank, dp_rank,
) )
pipe_writer.send( result_dict = {
{ "status": "ready",
"status": "ready", "max_total_num_tokens": scheduler.max_total_num_tokens,
"max_total_num_tokens": scheduler.max_total_num_tokens, "max_req_input_len": scheduler.max_req_input_len,
"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 disaggregation_mode: DisaggregationMode = scheduler.disaggregation_mode
if disaggregation_mode == DisaggregationMode.NULL: if disaggregation_mode == DisaggregationMode.NULL:
if scheduler.enable_pdmux: if scheduler.enable_pdmux:
+6
View File
@@ -366,6 +366,12 @@ class TpModelWorker(BaseTpWorker):
can_run_cuda_graph=can_run_cuda_graph, 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( def forward_batch_generation(
self, self,
model_worker_batch: ModelWorkerBatch, 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 ( from sglang.srt.model_executor.piecewise_cuda_graph_runner import (
PiecewiseCudaGraphRunner, 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.loader import DefaultModelLoader, get_model_loader
from sglang.srt.model_loader.remote_instance_weight_loader_utils import ( from sglang.srt.model_loader.remote_instance_weight_loader_utils import (
RemoteInstanceWeightLoaderBackend,
register_memory_region,
trigger_init_weights_send_group_for_remote_instance_request, trigger_init_weights_send_group_for_remote_instance_request,
) )
from sglang.srt.model_loader.utils import set_default_torch_dtype 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_available_gpu_memory,
get_bool_env_var, get_bool_env_var,
get_cpu_ids_by_node, get_cpu_ids_by_node,
get_local_ip_auto,
init_custom_process_group, init_custom_process_group,
is_cuda, is_cuda,
is_float4_e2m1fn_x2, is_float4_e2m1fn_x2,
@@ -319,6 +321,10 @@ class ModelRunner:
self.forward_pass_id = 0 self.forward_pass_id = 0
self.init_new_workspace = False 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 # Apply the rank zero filter to logger
if server_args.show_time_cost: if server_args.show_time_cost:
enable_show_time_cost() enable_show_time_cost()
@@ -393,6 +399,9 @@ class ModelRunner:
enable=self.server_args.enable_memory_saver 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: if not self.is_draft_worker:
set_global_expert_location_metadata( set_global_expert_location_metadata(
compute_initial_expert_location_metadata( compute_initial_expert_location_metadata(
@@ -433,6 +442,15 @@ class ModelRunner:
self.sampler = Sampler() self.sampler = Sampler()
self.load_model() 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 # Check if the model is using hybrid SWA
if ( if (
not self.server_args.disable_hybrid_swa_memory not self.server_args.disable_hybrid_swa_memory
@@ -547,6 +565,23 @@ class ModelRunner:
# Initialize piecewise CUDA graph # Initialize piecewise CUDA graph
self.init_piecewise_cuda_graphs() 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): def model_specific_adjustment(self):
server_args = self.server_args 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_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_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_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, modelopt_config=modelopt_config,
rl_quant_profile=self.server_args.rl_quant_profile, rl_quant_profile=self.server_args.rl_quant_profile,
) )
@@ -772,7 +809,11 @@ class ModelRunner:
self.model_config, self.load_config, self.tp_size 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: if self.tp_rank == 0:
instance_ip = socket.gethostbyname(socket.gethostname()) instance_ip = socket.gethostbyname(socket.gethostname())
t = threading.Thread( t = threading.Thread(
@@ -797,11 +838,18 @@ class ModelRunner:
GPU_MEMORY_TYPE_WEIGHTS, GPU_MEMORY_TYPE_WEIGHTS,
enable_cpu_backup=enable_cpu_backup, enable_cpu_backup=enable_cpu_backup,
): ):
self.model = get_model( self.loader = get_model_loader(
model_config=self.model_config,
load_config=self.load_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), 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) monkey_patch_vllm_parallel_state(reverse=True)
get_offloader().post_init() get_offloader().post_init()
+104 -4
View File
@@ -34,6 +34,11 @@ import huggingface_hub
import numpy as np import numpy as np
import torch 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 from sglang.srt.server_args import get_global_server_args
# Try to import accelerate (optional dependency) # Try to import accelerate (optional dependency)
@@ -1987,6 +1992,7 @@ class RemoteInstanceModelLoader(BaseModelLoader):
f"Model loader extra config is not supported for " f"Model loader extra config is not supported for "
f"load format {load_config.load_format}" f"load format {load_config.load_format}"
) )
self.remote_instance_transfer_engine_weight_info = None
def download_model(self, model_config: ModelConfig) -> None: def download_model(self, model_config: ModelConfig) -> None:
raise NotImplementedError raise NotImplementedError
@@ -2005,16 +2011,19 @@ class RemoteInstanceModelLoader(BaseModelLoader):
f"load format {load_config.load_format}" 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 set_default_torch_dtype(model_config.dtype):
with torch.device(device_config.device): with torch.device(device_config.device):
model = _initialize_model(model_config, self.load_config) 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: with create_remote_connector(model_weights, device_config.device) as client:
connector_type = get_connector_type(client) connector_type = get_connector_type(client)
if connector_type == ConnectorType.INSTANCE: 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 model, client, model_config, device_config
) )
else: else:
@@ -2022,9 +2031,43 @@ class RemoteInstanceModelLoader(BaseModelLoader):
f"Unsupported connector type {connector_type} for " f"Unsupported connector type {connector_type} for "
f"remote tensor model loading." 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() 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 self, model, client, model_config: ModelConfig, device_config: DeviceConfig
) -> nn.Module: ) -> nn.Module:
load_config = self.load_config load_config = self.load_config
@@ -2075,6 +2118,63 @@ class RemoteInstanceModelLoader(BaseModelLoader):
) )
torch.cuda.empty_cache() 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): class RemoteModelLoader(BaseModelLoader):
"""Model loader that can load Tensors from remote database.""" """Model loader that can load Tensors from remote database."""
@@ -1,6 +1,10 @@
# SPDX-License-Identifier: Apache-2.0 # SPDX-License-Identifier: Apache-2.0
import enum
import importlib
import importlib.util
import logging import logging
import time
from typing import List from typing import List
import requests import requests
@@ -8,6 +12,11 @@ import requests
logger = logging.getLogger(__name__) 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( def trigger_init_weights_send_group_for_remote_instance_request(
remote_instance_weight_loader_seed_instance_ip: str, remote_instance_weight_loader_seed_instance_ip: str,
remote_instance_weight_loader_seed_instance_service_port: int, remote_instance_weight_loader_seed_instance_service_port: int,
@@ -67,3 +76,133 @@ def trigger_transferring_weights_request(
except Exception as e: except Exception as e:
logger.error(f"Failed to trigger send weights to remote instance request: {e}") logger.error(f"Failed to trigger send weights to remote instance request: {e}")
raise 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 argparse
import dataclasses import dataclasses
import importlib import importlib
import importlib.util
import json import json
import logging import logging
import os import os
@@ -614,6 +615,8 @@ class ServerArgs:
remote_instance_weight_loader_seed_instance_ip: Optional[str] = None 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_seed_instance_service_port: Optional[int] = None
remote_instance_weight_loader_send_weights_group_ports: Optional[List[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 # For PD-Multiplexing
enable_pdmux: bool = False enable_pdmux: bool = False
@@ -690,6 +693,9 @@ class ServerArgs:
# Handle speculative decoding logic. # Handle speculative decoding logic.
self._handle_speculative_decoding() self._handle_speculative_decoding()
# Handle remote instance weight loader.
self._handle_remote_instance_weight_loader_start_seed_via_transfer_engine()
# Handle model loading format. # Handle model loading format.
self._handle_load_format() self._handle_load_format()
@@ -2107,8 +2113,26 @@ class ServerArgs:
if ( if (
self.remote_instance_weight_loader_seed_instance_ip is None 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_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" self.load_format = "auto"
def _handle_encoder_disaggregation(self): def _handle_encoder_disaggregation(self):
@@ -2366,19 +2390,12 @@ class ServerArgs:
self.disable_cuda_graph = True self.disable_cuda_graph = True
self.skip_server_warmup = True self.skip_server_warmup = True
def _handle_remote_instance_weight_loader_support_transfer_engine(self): def _handle_remote_instance_weight_loader_start_seed_via_transfer_engine(self):
if importlib.util.find_spec("mooncake.engine") is None: # Check whether TransferEngine can be used when users want to start seed service that supports TransferEngine backend.
logger.warning( if self.remote_instance_weight_loader_start_seed_via_transfer_engine:
f"Failed to import mooncake.engine. Does not support using TransferEngine as remote instance weight loader backend." 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 @staticmethod
def add_cli_args(parser: argparse.ArgumentParser): def add_cli_args(parser: argparse.ArgumentParser):
@@ -4277,6 +4294,18 @@ class ServerArgs:
default=ServerArgs.remote_instance_weight_loader_send_weights_group_ports, default=ServerArgs.remote_instance_weight_loader_send_weights_group_ports,
help="The communication group ports for loading weights from remote instance.", 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 # For PD-Multiplexing
parser.add_argument( parser.add_argument(
@@ -4782,6 +4811,33 @@ class ServerArgs:
original_server_arg_mem_fraction * final_overall_factor 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. # NOTE: This is a global variable to hold the server args for scheduler.
_global_server_args: Optional[ServerArgs] = None _global_server_args: Optional[ServerArgs] = None
@@ -62,6 +62,7 @@ def init_process(
seed_instance_group_base_port, seed_instance_group_base_port,
event_seed_ready, event_seed_ready,
event_dst_ready_list, event_dst_ready_list,
remote_instance_loader_backend,
): ):
torch.cuda.set_device(rank) torch.cuda.set_device(rank)
@@ -90,6 +91,7 @@ def init_process(
tp_size, tp_size,
event_seed_ready, event_seed_ready,
event_dst_ready_list, event_dst_ready_list,
remote_instance_loader_backend,
) )
@@ -122,6 +124,7 @@ def init_process_seed(
str(rank), str(rank),
"--tp-size", "--tp-size",
str(tp_size), str(tp_size),
"--remote-instance-weight-loader-start-seed-via-transfer-engine",
), ),
) )
torch.cuda.synchronize() torch.cuda.synchronize()
@@ -159,6 +162,7 @@ def init_process_dst(
tp_size, tp_size,
event_seed_ready, event_seed_ready,
event_dst_ready_list, event_dst_ready_list,
remote_instance_loader_backend,
): ):
torch.cuda.set_device(rank * tp_size) torch.cuda.set_device(rank * tp_size)
torch.cuda.synchronize() torch.cuda.synchronize()
@@ -186,6 +190,10 @@ def init_process_dst(
remote_instance_weight_loader_seed_instance_service_port=seed_instance_service_port, remote_instance_weight_loader_seed_instance_service_port=seed_instance_service_port,
remote_instance_weight_loader_send_weights_group_ports=ports, remote_instance_weight_loader_send_weights_group_ports=ports,
load_format="remote_instance", load_format="remote_instance",
remote_instance_weight_loader_backend=remote_instance_loader_backend,
remote_instance_weight_loader_start_seed_via_transfer_engine=(
remote_instance_loader_backend == "transfer_engine"
),
) )
else: else:
host, _, port = DEFAULT_URL_FOR_TEST.rpartition(":") host, _, port = DEFAULT_URL_FOR_TEST.rpartition(":")
@@ -213,6 +221,9 @@ def init_process_dst(
f"[{','.join(str(port) for port in ports)}]", f"[{','.join(str(port) for port in ports)}]",
"--load-format", "--load-format",
"remote_instance", "remote_instance",
"--remote-instance-weight-loader-backend",
remote_instance_loader_backend,
"--remote-instance-weight-loader-start-seed-via-transfer-engine",
), ),
) )
torch.cuda.synchronize() torch.cuda.synchronize()
@@ -250,9 +261,10 @@ def test_load_weights_from_remote_instance(
seed_instance_ip, seed_instance_ip,
seed_instance_service_port, seed_instance_service_port,
seed_instance_group_base_port, seed_instance_group_base_port,
remote_instance_loader_backend,
): ):
print( 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() param_queue = mp.Queue()
results = {} results = {}
@@ -276,6 +288,7 @@ def test_load_weights_from_remote_instance(
seed_instance_group_base_port, seed_instance_group_base_port,
event_seed_ready, event_seed_ready,
event_dst_ready_list, event_dst_ready_list,
remote_instance_loader_backend,
), ),
nprocs=1 + dp_size, nprocs=1 + dp_size,
join=False, join=False,
@@ -340,14 +353,42 @@ class TestLoadWeightsFromRemoteInstance(CustomTestCase):
# test_suits : tp, dp, model_name, backend, dst_instance_id # test_suits : tp, dp, model_name, backend, dst_instance_id
if is_in_ci(): if is_in_ci():
mode = random.choice(["Engine", "Server"]) mode = random.choice(["Engine", "Server"])
remote_instance_loader_backend = random.choice(["nccl", "transfer_engine"])
test_suits = [ 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: else:
test_suits = [ test_suits = [
(1, 1, DEFAULT_SMALL_MODEL_NAME_FOR_TEST, ["Engine"]), (1, 1, DEFAULT_SMALL_MODEL_NAME_FOR_TEST, ["Engine"], "nccl"),
(1, 1, DEFAULT_SMALL_MODEL_NAME_FOR_TEST, ["Sever"]), (1, 1, DEFAULT_SMALL_MODEL_NAME_FOR_TEST, ["Server"], "nccl"),
(2, 2, DEFAULT_SMALL_MODEL_NAME_FOR_TEST, ["Engine", "Server"]), (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,
["Server"],
"transfer_engine",
),
(
2,
2,
DEFAULT_SMALL_MODEL_NAME_FOR_TEST,
["Engine", "Server"],
"transfer_engine",
),
] ]
truncate_size = 10 truncate_size = 10
@@ -365,7 +406,13 @@ class TestLoadWeightsFromRemoteInstance(CustomTestCase):
"model.norm.weight", "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( test_load_weights_from_remote_instance(
tp_size, tp_size,
dp_size, dp_size,
@@ -376,6 +423,7 @@ class TestLoadWeightsFromRemoteInstance(CustomTestCase):
"127.0.0.1", "127.0.0.1",
DEFAULT_PORT_FOR_SRT_TEST_RUNNER + 1000, DEFAULT_PORT_FOR_SRT_TEST_RUNNER + 1000,
60000, 60000,
remote_instance_loader_backend,
) )