Support loading weights from remote instance (#8215)
Signed-off-by: Anqi Shen <amy.saq@antgroup.com> Co-authored-by: Chayenne <74843776+zhaochenyang20@users.noreply.github.com>
This commit is contained in:
@@ -12,6 +12,9 @@ import json
|
||||
import logging
|
||||
import math
|
||||
import os
|
||||
import re
|
||||
import socket
|
||||
import threading
|
||||
import time
|
||||
from abc import ABC, abstractmethod
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
@@ -27,9 +30,11 @@ from typing import (
|
||||
Tuple,
|
||||
cast,
|
||||
)
|
||||
from urllib.parse import urlparse
|
||||
|
||||
import huggingface_hub
|
||||
import numpy as np
|
||||
import requests
|
||||
import safetensors.torch
|
||||
import torch
|
||||
from huggingface_hub import HfApi, hf_hub_download
|
||||
@@ -56,6 +61,7 @@ from sglang.srt.model_loader.utils import (
|
||||
)
|
||||
from sglang.srt.model_loader.weight_utils import (
|
||||
_BAR_FORMAT,
|
||||
default_weight_loader,
|
||||
download_safetensors_index_file_from_hf,
|
||||
download_weights_from_hf,
|
||||
filter_duplicate_safetensors_files,
|
||||
@@ -71,6 +77,9 @@ from sglang.srt.model_loader.weight_utils import (
|
||||
safetensors_weights_iterator,
|
||||
set_runai_streamer_env,
|
||||
)
|
||||
from sglang.srt.remote_instance_weight_loader_utils import (
|
||||
trigger_transferring_weights_request,
|
||||
)
|
||||
from sglang.srt.utils import (
|
||||
get_bool_env_var,
|
||||
get_device_capability,
|
||||
@@ -1380,6 +1389,104 @@ class GGUFModelLoader(BaseModelLoader):
|
||||
return model
|
||||
|
||||
|
||||
class RemoteInstanceModelLoader(BaseModelLoader):
|
||||
"""Model loader that can load Tensors from remote sglang instance."""
|
||||
|
||||
def __init__(self, load_config: LoadConfig):
|
||||
super().__init__(load_config)
|
||||
if load_config.model_loader_extra_config:
|
||||
raise ValueError(
|
||||
f"Model loader extra config is not supported for "
|
||||
f"load format {load_config.load_format}"
|
||||
)
|
||||
|
||||
def download_model(self, model_config: ModelConfig) -> None:
|
||||
raise NotImplementedError
|
||||
|
||||
def load_model(
|
||||
self,
|
||||
*,
|
||||
model_config: ModelConfig,
|
||||
device_config: DeviceConfig,
|
||||
) -> nn.Module:
|
||||
logger.info("Loading weights from remote instance ...")
|
||||
load_config = self.load_config
|
||||
|
||||
assert load_config.load_format == LoadFormat.REMOTE_INSTANCE, (
|
||||
f"Model loader {self.load_config.load_format} is not supported for "
|
||||
f"load format {load_config.load_format}"
|
||||
)
|
||||
|
||||
model_weights = f"instance://{model_config.remote_instance_weight_loader_seed_instance_ip}:{model_config.remote_instance_weight_loader_send_weights_group_ports[model_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)
|
||||
|
||||
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(
|
||||
model, client, model_config, device_config
|
||||
)
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Unsupported connector type {connector_type} for "
|
||||
f"remote tensor model loading."
|
||||
)
|
||||
return model.eval()
|
||||
|
||||
def load_model_from_remote_instance(
|
||||
self, model, client, model_config: ModelConfig, device_config: DeviceConfig
|
||||
) -> nn.Module:
|
||||
instance_ip = socket.gethostbyname(socket.gethostname())
|
||||
start_build_group_tic = time.time()
|
||||
client.build_group(
|
||||
gpu_id=device_config.gpu_id,
|
||||
tp_rank=model_config.tp_rank,
|
||||
instance_ip=instance_ip,
|
||||
)
|
||||
torch.cuda.synchronize()
|
||||
end_build_group_tic = time.time()
|
||||
logger.debug(
|
||||
f"finish building group for remote instance, time used: {(end_build_group_tic - start_build_group_tic):.4f}s"
|
||||
)
|
||||
|
||||
if model_config.tp_rank == 0:
|
||||
t = threading.Thread(
|
||||
target=trigger_transferring_weights_request,
|
||||
args=(
|
||||
model_config.remote_instance_weight_loader_seed_instance_ip,
|
||||
model_config.remote_instance_weight_loader_seed_instance_service_port,
|
||||
model_config.remote_instance_weight_loader_send_weights_group_ports,
|
||||
instance_ip,
|
||||
),
|
||||
)
|
||||
t.start()
|
||||
|
||||
start_get_weights_tic = time.time()
|
||||
with set_default_torch_dtype(model_config.dtype):
|
||||
for _, tensor in model.named_parameters():
|
||||
torch.distributed.broadcast(
|
||||
tensor.data,
|
||||
src=0,
|
||||
group=client._model_update_group,
|
||||
)
|
||||
torch.cuda.synchronize()
|
||||
|
||||
if hasattr(model, "post_load_weights"):
|
||||
model.post_load_weights()
|
||||
end_get_weights_tic = time.time()
|
||||
logger.debug(
|
||||
f"finish getting all weights from remote instance, time used: {(end_get_weights_tic - start_get_weights_tic):.4f}s"
|
||||
)
|
||||
# destroy the process group after loading weights
|
||||
torch.distributed.distributed_c10d.destroy_process_group(
|
||||
client._model_update_group
|
||||
)
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
|
||||
class RemoteModelLoader(BaseModelLoader):
|
||||
"""Model loader that can load Tensors from remote database."""
|
||||
|
||||
@@ -1581,4 +1688,7 @@ def get_model_loader(load_config: LoadConfig) -> BaseModelLoader:
|
||||
if load_config.load_format == LoadFormat.REMOTE:
|
||||
return RemoteModelLoader(load_config)
|
||||
|
||||
if load_config.load_format == LoadFormat.REMOTE_INSTANCE:
|
||||
return RemoteInstanceModelLoader(load_config)
|
||||
|
||||
return DefaultModelLoader(load_config)
|
||||
|
||||
Reference in New Issue
Block a user