support non disturbing remote instance weight loader v2 (#14997)
Signed-off-by: Anqi Shen <amy.saq@antgroup.com>
This commit is contained in:
@@ -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."""
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user