Revert several PRs (#14958)
Co-authored-by: fzyzcjy <ch271828n@outlook.com>
This commit is contained in:
@@ -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
|
||||
```
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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."""
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user