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-17 06:39:56 +08:00
committed by GitHub
parent a4c762811a
commit ccc8f3b266
11 changed files with 557 additions and 40 deletions

View File

@@ -34,6 +34,11 @@ import huggingface_hub
import numpy as np
import torch
from sglang.srt.model_loader.remote_instance_weight_loader_utils import (
RemoteInstanceWeightLoaderBackend,
get_remote_instance_transfer_engine_info_per_rank,
register_memory_region,
)
from sglang.srt.server_args import get_global_server_args
# Try to import accelerate (optional dependency)
@@ -1987,6 +1992,7 @@ class RemoteInstanceModelLoader(BaseModelLoader):
f"Model loader extra config is not supported for "
f"load format {load_config.load_format}"
)
self.remote_instance_transfer_engine_weight_info = None
def download_model(self, model_config: ModelConfig) -> None:
raise NotImplementedError
@@ -2005,16 +2011,19 @@ class RemoteInstanceModelLoader(BaseModelLoader):
f"load format {load_config.load_format}"
)
model_weights = f"instance://{load_config.remote_instance_weight_loader_seed_instance_ip}:{load_config.remote_instance_weight_loader_send_weights_group_ports[load_config.tp_rank]}"
with set_default_torch_dtype(model_config.dtype):
with torch.device(device_config.device):
model = _initialize_model(model_config, self.load_config)
if (
load_config.remote_instance_weight_loader_backend
== RemoteInstanceWeightLoaderBackend.NCCL
):
model_weights = f"instance://{load_config.remote_instance_weight_loader_seed_instance_ip}:{load_config.remote_instance_weight_loader_send_weights_group_ports[load_config.tp_rank]}"
with create_remote_connector(model_weights, device_config.device) as client:
connector_type = get_connector_type(client)
if connector_type == ConnectorType.INSTANCE:
self.load_model_from_remote_instance(
self.load_model_from_remote_instance_by_nccl(
model, client, model_config, device_config
)
else:
@@ -2022,9 +2031,43 @@ 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(
model, load_config.remote_instance_weight_loader_transfer_engine
)
logger.info(
"TransferEngine memory regions have been successfully registered."
)
# transfer weights
success = self.load_model_from_remote_instance_by_transfer_engine(
model,
load_config.remote_instance_weight_loader_transfer_engine,
f"http://{load_config.remote_instance_weight_loader_seed_instance_ip}:{load_config.remote_instance_weight_loader_seed_instance_service_port}",
load_config.tp_rank,
)
if not success:
raise RuntimeError(
"Failed to load weights from remote instance via transfer engine."
)
else:
raise ValueError("Invalid remote instance weight loader backend.")
return model.eval()
def load_model_from_remote_instance(
def load_model_from_remote_instance_by_nccl(
self, model, client, model_config: ModelConfig, device_config: DeviceConfig
) -> nn.Module:
load_config = self.load_config
@@ -2075,6 +2118,63 @@ class RemoteInstanceModelLoader(BaseModelLoader):
)
torch.cuda.empty_cache()
def load_model_from_remote_instance_by_transfer_engine(
self, model, transfer_engine, seed_url, tp_rank
) -> bool:
# get remote weights metadata from source instance
seed_transfer_engine_session_id, seed_transfer_engine_weight_info = (
get_remote_instance_transfer_engine_info_per_rank(seed_url, tp_rank)
)
if (
seed_transfer_engine_session_id is None
or seed_transfer_engine_weight_info is None
):
logger.error("Cannot get transfer engine session or weight info.")
return False
# prepare local/remote RDMA keys
seed_ptr_list = []
client_ptr_list = []
client_len_list = []
for name, tensor in model.named_parameters():
weight_info = seed_transfer_engine_weight_info.get(name, None)
if weight_info is None:
logger.error(f"Cannot find weight info for {name}.")
return False
seed_ptr, seed_numel, seed_element_size = weight_info
if (
seed_numel != tensor.numel()
or seed_element_size != tensor.element_size()
):
logger.error(
f"Weight info does not match for {name}, "
f"expected ({seed_numel}, {seed_element_size}), "
f"got ({tensor.numel()}, {tensor.element_size()})"
)
return False
client_ptr = tensor.data_ptr()
client_len = tensor.numel() * tensor.element_size()
seed_ptr_list.append(seed_ptr)
client_ptr_list.append(client_ptr)
client_len_list.append(client_len)
# load weights from source instance through TransferEngine
ret = transfer_engine.batch_transfer_sync_read(
seed_transfer_engine_session_id,
client_ptr_list,
seed_ptr_list,
client_len_list,
)
if ret < 0:
logger.error(f"batch transfer failed, error: {ret}")
return False
if hasattr(model, "post_load_weights"):
model.post_load_weights()
return True
class RemoteModelLoader(BaseModelLoader):
"""Model loader that can load Tensors from remote database."""

View File

@@ -1,6 +1,10 @@
# SPDX-License-Identifier: Apache-2.0
import enum
import importlib
import importlib.util
import logging
import time
from typing import List
import requests
@@ -8,6 +12,11 @@ import requests
logger = logging.getLogger(__name__)
class RemoteInstanceWeightLoaderBackend(str, enum.Enum):
NCCL = "nccl"
TRANSFER_ENGINE = "transfer_engine"
def trigger_init_weights_send_group_for_remote_instance_request(
remote_instance_weight_loader_seed_instance_ip: str,
remote_instance_weight_loader_seed_instance_service_port: int,
@@ -67,3 +76,133 @@ 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
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