Extract load_model helpers into a load_model_utils module (#31155)
This commit is contained in:
@@ -16,22 +16,16 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import contextlib
|
||||
import datetime
|
||||
import inspect
|
||||
import logging
|
||||
import os
|
||||
import socket
|
||||
import threading
|
||||
import time
|
||||
from collections import defaultdict
|
||||
from dataclasses import dataclass
|
||||
from typing import Optional, Union
|
||||
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
|
||||
from sglang.srt.configs.device_config import DeviceConfig
|
||||
from sglang.srt.configs.load_config import LoadConfig, LoadFormat
|
||||
from sglang.srt.configs.load_config import LoadConfig
|
||||
from sglang.srt.configs.model_config import (
|
||||
AttentionArch,
|
||||
ModelConfig,
|
||||
@@ -41,20 +35,14 @@ from sglang.srt.configs.model_config import (
|
||||
is_deepseek_dsa,
|
||||
)
|
||||
from sglang.srt.configs.update_config import adjust_config_with_unaligned_cpu_tp
|
||||
from sglang.srt.constants import GPU_MEMORY_TYPE_WEIGHTS
|
||||
from sglang.srt.debug_utils.dumper import dumper
|
||||
from sglang.srt.debug_utils.tensor_dump_forward_hook import (
|
||||
register_forward_hook_for_model,
|
||||
)
|
||||
from sglang.srt.distributed import (
|
||||
bootstrap,
|
||||
get_tp_group,
|
||||
get_world_group,
|
||||
)
|
||||
from sglang.srt.distributed.device_communicators.pynccl_allocator import (
|
||||
prealloc_symmetric_memory_pool,
|
||||
)
|
||||
from sglang.srt.distributed.parallel_state import monkey_patch_vllm_parallel_state
|
||||
from sglang.srt.dllm.config import DllmConfig
|
||||
from sglang.srt.elastic_ep.elastic_ep import (
|
||||
ElasticEPStateManager,
|
||||
@@ -129,6 +117,17 @@ from sglang.srt.model_executor.forward_context import (
|
||||
)
|
||||
from sglang.srt.model_executor.graph_shared_output import GraphSharedOutput
|
||||
from sglang.srt.model_executor.hook_manager import register_forward_hooks
|
||||
from sglang.srt.model_executor.model_runner_components.load_model_utils import (
|
||||
build_load_config,
|
||||
dist_barrier_after_load,
|
||||
load_kv_cache_scales,
|
||||
load_model_with_memory_saver,
|
||||
maybe_downgrade_dtype_for_legacy_gpu,
|
||||
maybe_register_debug_tensor_dump_hook,
|
||||
maybe_trigger_remote_instance_nccl_send_group,
|
||||
report_online_quantization,
|
||||
resolve_sliding_window_size,
|
||||
)
|
||||
from sglang.srt.model_executor.model_runner_components.ngram_embedding_manager import (
|
||||
NgramEmbeddingManager,
|
||||
)
|
||||
@@ -150,11 +149,6 @@ from sglang.srt.model_executor.runner import (
|
||||
PrefillCudaGraphRunner,
|
||||
get_batch_sizes_to_capture,
|
||||
)
|
||||
from sglang.srt.model_loader.loader import get_model_loader
|
||||
from sglang.srt.model_loader.remote_instance_weight_loader_utils import (
|
||||
RemoteInstanceWeightLoaderBackend,
|
||||
trigger_init_weights_send_group_for_remote_instance_request,
|
||||
)
|
||||
from sglang.srt.model_loader.utils import resolve_language_model
|
||||
from sglang.srt.platforms import current_platform
|
||||
from sglang.srt.runtime_context import get_flags, get_server_args
|
||||
@@ -196,7 +190,7 @@ from sglang.srt.utils import (
|
||||
set_cuda_arch,
|
||||
slow_rank_detector,
|
||||
)
|
||||
from sglang.srt.utils.network import NetworkAddress, get_local_ip_auto
|
||||
from sglang.srt.utils.network import get_local_ip_auto
|
||||
from sglang.srt.utils.nvtx_pytorch_hooks import PytHooks
|
||||
from sglang.srt.utils.nvtx_utils import profile_range
|
||||
from sglang.srt.utils.offloader import (
|
||||
@@ -222,7 +216,6 @@ elif current_platform.is_out_of_tree():
|
||||
current_platform.init_backend()
|
||||
|
||||
# Detect stragger ranks in model loading
|
||||
UNBALANCED_MODEL_LOADING_TIMEOUT_S = 480 # leave more time for post data processing
|
||||
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -1115,49 +1108,17 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
||||
if self.device != "cpu":
|
||||
torch.set_num_threads(1)
|
||||
if self.device == "cuda":
|
||||
if torch.cuda.get_device_capability()[0] < 8:
|
||||
logger.info(
|
||||
"Compute capability below sm80. Use float16 due to lack of bfloat16 support."
|
||||
)
|
||||
from sglang.srt.arg_groups.overrides import (
|
||||
declare_load_time_override,
|
||||
)
|
||||
|
||||
declare_load_time_override(
|
||||
"ModelRunner._sm80_dtype_fallback", {"dtype": "float16"}
|
||||
)
|
||||
self.model_config.dtype = torch.float16
|
||||
if torch.cuda.get_device_capability()[1] < 5:
|
||||
raise RuntimeError("SGLang only supports sm75 and above.")
|
||||
maybe_downgrade_dtype_for_legacy_gpu(
|
||||
server_args=self.server_args, model_config=self.model_config
|
||||
)
|
||||
|
||||
set_cuda_arch()
|
||||
|
||||
# Prepare the model config
|
||||
from sglang.srt.configs.modelopt_config import ModelOptConfig
|
||||
|
||||
modelopt_config = ModelOptConfig(
|
||||
quant=self.server_args.modelopt_quant,
|
||||
checkpoint_restore_path=self.server_args.modelopt_checkpoint_restore_path,
|
||||
checkpoint_save_path=self.server_args.modelopt_checkpoint_save_path,
|
||||
export_path=self.server_args.modelopt_export_path,
|
||||
quantize_and_serve=self.server_args.quantize_and_serve,
|
||||
)
|
||||
|
||||
self.load_config = LoadConfig(
|
||||
load_format=self.server_args.load_format,
|
||||
download_dir=self.server_args.download_dir,
|
||||
model_loader_extra_config=self.server_args.model_loader_extra_config,
|
||||
self.load_config = build_load_config(
|
||||
server_args=self.server_args,
|
||||
tp_rank=self.tp_rank,
|
||||
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_weight_transporter.engine,
|
||||
remote_instance_weight_loader_transfer_engine_session_id=self.remote_instance_weight_transporter.session_id,
|
||||
modelexpress_url=self.server_args.modelexpress_url,
|
||||
modelexpress_transport=self.server_args.modelexpress_transport,
|
||||
modelopt_config=modelopt_config,
|
||||
rl_quant_profile=self.server_args.rl_quant_profile,
|
||||
remote_instance_weight_transporter_engine=self.remote_instance_weight_transporter.engine,
|
||||
remote_instance_weight_transporter_session_id=self.remote_instance_weight_transporter.session_id,
|
||||
draft_model_idx=self.draft_model_idx,
|
||||
)
|
||||
if self.device == "cpu":
|
||||
@@ -1165,52 +1126,25 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
||||
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.tp_rank == 0:
|
||||
instance_ip = NetworkAddress.resolve_host(socket.gethostname())
|
||||
t = threading.Thread(
|
||||
target=trigger_init_weights_send_group_for_remote_instance_request,
|
||||
args=(
|
||||
self.server_args.remote_instance_weight_loader_seed_instance_ip,
|
||||
self.server_args.remote_instance_weight_loader_seed_instance_service_port,
|
||||
self.server_args.remote_instance_weight_loader_send_weights_group_ports,
|
||||
instance_ip,
|
||||
),
|
||||
)
|
||||
t.start()
|
||||
|
||||
# Load the model
|
||||
# Remove monkey_patch when linear.py quant remove dependencies with vllm
|
||||
monkey_patch_vllm_parallel_state()
|
||||
|
||||
enable_cpu_backup = self.server_args.enable_weights_cpu_backup or (
|
||||
self.is_draft_worker and self.server_args.enable_draft_weights_cpu_backup
|
||||
maybe_trigger_remote_instance_nccl_send_group(
|
||||
server_args=self.server_args, tp_rank=self.tp_rank
|
||||
)
|
||||
with self.memory_saver_adapter.region(
|
||||
GPU_MEMORY_TYPE_WEIGHTS,
|
||||
enable_cpu_backup=enable_cpu_backup,
|
||||
):
|
||||
self.loader = get_model_loader(
|
||||
load_config=self.load_config,
|
||||
model_config=self.model_config,
|
||||
|
||||
loaded = load_model_with_memory_saver(
|
||||
server_args=self.server_args,
|
||||
model_config=self.model_config,
|
||||
load_config=self.load_config,
|
||||
device=self.device,
|
||||
gpu_id=self.gpu_id,
|
||||
memory_saver_adapter=self.memory_saver_adapter,
|
||||
is_draft_worker=self.is_draft_worker,
|
||||
)
|
||||
self.loader = loaded.loader
|
||||
self.model = loaded.model
|
||||
if loaded.remote_instance_weight_info is not None:
|
||||
self.remote_instance_weight_transporter.weight_info = (
|
||||
loaded.remote_instance_weight_info
|
||||
)
|
||||
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_weight_transporter.weight_info = (
|
||||
self.loader.remote_instance_transfer_engine_weight_info
|
||||
)
|
||||
# Cache needs to be cleared after loading model weights (in the self.loader.load_model function).
|
||||
# To avoid conflict with memory_saver_adapter.region, empty_cache operation is now moved here.
|
||||
if _is_npu:
|
||||
torch.npu.empty_cache()
|
||||
monkey_patch_vllm_parallel_state(reverse=True)
|
||||
|
||||
if not self.is_draft_worker:
|
||||
get_offloader().post_init()
|
||||
@@ -1220,44 +1154,11 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
||||
pyt_hooks = PytHooks()
|
||||
pyt_hooks.register_hooks(self.model, module_prefix="model")
|
||||
|
||||
if self.server_args.kv_cache_dtype == "fp8_e4m3":
|
||||
if self.server_args.quantization_param_path is not None:
|
||||
if callable(getattr(self.model, "load_kv_cache_scales", None)):
|
||||
self.model.load_kv_cache_scales(
|
||||
self.server_args.quantization_param_path
|
||||
)
|
||||
logger.info(
|
||||
"Loaded KV cache scaling factors from %s",
|
||||
self.server_args.quantization_param_path,
|
||||
)
|
||||
else:
|
||||
raise RuntimeError(
|
||||
"Using FP8 KV cache and scaling factors provided but "
|
||||
"model %s does not support loading scaling factors.",
|
||||
self.model.__class__,
|
||||
)
|
||||
else:
|
||||
logger.warning(
|
||||
"Using FP8 KV cache but no scaling factors "
|
||||
"provided. Defaulting to scaling factors of 1.0. "
|
||||
"This may lead to less accurate results!"
|
||||
)
|
||||
load_kv_cache_scales(model=self.model, server_args=self.server_args)
|
||||
|
||||
# Parse other args
|
||||
self.sliding_window_size = None
|
||||
if hasattr(self.model, "get_attention_sliding_window_size"):
|
||||
self.sliding_window_size = self.model.get_attention_sliding_window_size()
|
||||
elif (
|
||||
self.model_config.is_hybrid_swa
|
||||
and self.model_config.sliding_window_size is not None
|
||||
):
|
||||
# sliding window field in model config may have different meaning for different kinds of models (e.g., dllm), here we only consider the sliding window in SWA model
|
||||
self.sliding_window_size = self.model_config.sliding_window_size
|
||||
elif self.model_config.attention_chunk_size is not None:
|
||||
self.sliding_window_size = self.model_config.attention_chunk_size
|
||||
logger.info(
|
||||
f"Setting sliding_window_size to be attention_chunk_size: {self.sliding_window_size}"
|
||||
)
|
||||
self.sliding_window_size = resolve_sliding_window_size(
|
||||
self.model, self.model_config
|
||||
)
|
||||
|
||||
self.prefill_aware_swa = (
|
||||
hasattr(self.model, "is_prefill_aware_swa")
|
||||
@@ -1281,34 +1182,17 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
||||
f"mem usage={self.weight_load_mem_usage:.2f} GB."
|
||||
)
|
||||
|
||||
# TODO: Make sure all models have `quant_config` attribute, and all online quantization methods register which layers they actually quantize.
|
||||
# TODO: Move this online-quantization reporting out of ModelRunner.
|
||||
quantized_layers = getattr(
|
||||
getattr(self.model, "quant_config", None), "quantized_layers", None
|
||||
)
|
||||
if (
|
||||
self.server_args.quantization is not None
|
||||
and isinstance(quantized_layers, tuple)
|
||||
and len(quantized_layers) == 2
|
||||
):
|
||||
layer_types, quantized_layers_count = quantized_layers
|
||||
logger.info(
|
||||
f"Online {self.server_args.quantization} quantization: quantized {quantized_layers_count} layers of types: {layer_types}"
|
||||
)
|
||||
report_online_quantization(model=self.model, server_args=self.server_args)
|
||||
|
||||
if self.server_args.debug_tensor_dump_output_folder is not None:
|
||||
dump_folder = self.server_args.debug_tensor_dump_output_folder
|
||||
if self.spec_algorithm.is_eagle():
|
||||
role = "draft" if self.is_draft_worker else "target"
|
||||
dump_folder = os.path.join(dump_folder, role)
|
||||
register_forward_hook_for_model(
|
||||
self.model,
|
||||
dump_folder,
|
||||
self.server_args.debug_tensor_dump_layers,
|
||||
self.tp_size,
|
||||
self.tp_rank,
|
||||
self.pp_rank,
|
||||
)
|
||||
maybe_register_debug_tensor_dump_hook(
|
||||
model=self.model,
|
||||
server_args=self.server_args,
|
||||
spec_algorithm=self.spec_algorithm,
|
||||
is_draft_worker=self.is_draft_worker,
|
||||
tp_size=self.tp_size,
|
||||
tp_rank=self.tp_rank,
|
||||
pp_rank=self.pp_rank,
|
||||
)
|
||||
|
||||
if dumper.may_enable:
|
||||
dumper.apply_source_patches()
|
||||
@@ -1322,23 +1206,10 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
||||
logger,
|
||||
)
|
||||
|
||||
if self.server_args.elastic_ep_backend == "mooncake":
|
||||
# Mooncake does not support `monitored_barrier`
|
||||
dist.barrier(group=get_tp_group().cpu_group)
|
||||
else:
|
||||
# Handle the case where some ranks do not finish loading.
|
||||
try:
|
||||
dist.monitored_barrier(
|
||||
group=get_tp_group().cpu_group,
|
||||
timeout=datetime.timedelta(
|
||||
seconds=UNBALANCED_MODEL_LOADING_TIMEOUT_S
|
||||
),
|
||||
wait_all_ranks=True,
|
||||
)
|
||||
except RuntimeError:
|
||||
raise ValueError(
|
||||
f"TP rank {self.tp_rank} could finish the model loading, but there are other ranks that didn't finish loading. It is likely due to unexpected failures (e.g., OOM) or a slow node."
|
||||
) from None
|
||||
dist_barrier_after_load(
|
||||
elastic_ep_backend=self.server_args.elastic_ep_backend,
|
||||
tp_rank=self.tp_rank,
|
||||
)
|
||||
|
||||
def _prepare_moe_topk(self):
|
||||
balancer_cls = None
|
||||
|
||||
@@ -0,0 +1,266 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import datetime
|
||||
import logging
|
||||
import os
|
||||
import socket
|
||||
import threading
|
||||
from typing import TYPE_CHECKING, Any, Optional
|
||||
|
||||
import msgspec
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
|
||||
from sglang.srt.configs.device_config import DeviceConfig
|
||||
from sglang.srt.configs.load_config import LoadConfig, LoadFormat
|
||||
from sglang.srt.constants import GPU_MEMORY_TYPE_WEIGHTS
|
||||
from sglang.srt.debug_utils.tensor_dump_forward_hook import (
|
||||
register_forward_hook_for_model,
|
||||
)
|
||||
from sglang.srt.distributed import get_tp_group
|
||||
from sglang.srt.distributed.parallel_state import monkey_patch_vllm_parallel_state
|
||||
from sglang.srt.model_loader.loader import get_model_loader
|
||||
from sglang.srt.model_loader.remote_instance_weight_loader_utils import (
|
||||
RemoteInstanceWeightLoaderBackend,
|
||||
trigger_init_weights_send_group_for_remote_instance_request,
|
||||
)
|
||||
from sglang.srt.utils.common import is_npu
|
||||
from sglang.srt.utils.network import NetworkAddress
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.configs.model_config import ModelConfig
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_is_npu = is_npu()
|
||||
|
||||
|
||||
UNBALANCED_MODEL_LOADING_TIMEOUT_S = 480 # leave more time for post data processing
|
||||
|
||||
|
||||
class LoadedModel(msgspec.Struct, frozen=True, kw_only=True):
|
||||
loader: Any
|
||||
model: Any
|
||||
remote_instance_weight_info: Optional[Any]
|
||||
|
||||
|
||||
def maybe_downgrade_dtype_for_legacy_gpu(
|
||||
*, server_args: ServerArgs, model_config: ModelConfig
|
||||
) -> None:
|
||||
if torch.cuda.get_device_capability()[0] < 8:
|
||||
logger.info(
|
||||
"Compute capability below sm80. Use float16 due to lack of bfloat16 support."
|
||||
)
|
||||
from sglang.srt.arg_groups.overrides import declare_load_time_override
|
||||
|
||||
declare_load_time_override(
|
||||
"ModelRunner._sm80_dtype_fallback", {"dtype": "float16"}
|
||||
)
|
||||
model_config.dtype = torch.float16
|
||||
if torch.cuda.get_device_capability()[1] < 5:
|
||||
raise RuntimeError("SGLang only supports sm75 and above.")
|
||||
|
||||
|
||||
def maybe_trigger_remote_instance_nccl_send_group(
|
||||
*, server_args: ServerArgs, tp_rank: int
|
||||
) -> None:
|
||||
if (
|
||||
server_args.load_format == LoadFormat.REMOTE_INSTANCE
|
||||
and server_args.remote_instance_weight_loader_backend
|
||||
== RemoteInstanceWeightLoaderBackend.NCCL
|
||||
):
|
||||
if tp_rank == 0:
|
||||
instance_ip = NetworkAddress.resolve_host(socket.gethostname())
|
||||
t = threading.Thread(
|
||||
target=trigger_init_weights_send_group_for_remote_instance_request,
|
||||
args=(
|
||||
server_args.remote_instance_weight_loader_seed_instance_ip,
|
||||
server_args.remote_instance_weight_loader_seed_instance_service_port,
|
||||
server_args.remote_instance_weight_loader_send_weights_group_ports,
|
||||
instance_ip,
|
||||
),
|
||||
)
|
||||
t.start()
|
||||
|
||||
|
||||
def load_kv_cache_scales(*, model, server_args: ServerArgs) -> None:
|
||||
if server_args.kv_cache_dtype == "fp8_e4m3":
|
||||
if server_args.quantization_param_path is not None:
|
||||
if callable(getattr(model, "load_kv_cache_scales", None)):
|
||||
model.load_kv_cache_scales(server_args.quantization_param_path)
|
||||
logger.info(
|
||||
"Loaded KV cache scaling factors from %s",
|
||||
server_args.quantization_param_path,
|
||||
)
|
||||
else:
|
||||
raise RuntimeError(
|
||||
"Using FP8 KV cache and scaling factors provided but "
|
||||
"model %s does not support loading scaling factors.",
|
||||
model.__class__,
|
||||
)
|
||||
else:
|
||||
logger.warning(
|
||||
"Using FP8 KV cache but no scaling factors "
|
||||
"provided. Defaulting to scaling factors of 1.0. "
|
||||
"This may lead to less accurate results!"
|
||||
)
|
||||
|
||||
|
||||
def resolve_sliding_window_size(model, model_config: ModelConfig) -> Optional[int]:
|
||||
# Parse other args
|
||||
sliding_window_size = None
|
||||
if hasattr(model, "get_attention_sliding_window_size"):
|
||||
sliding_window_size = model.get_attention_sliding_window_size()
|
||||
elif model_config.is_hybrid_swa and model_config.sliding_window_size is not None:
|
||||
# sliding window field in model config may have different meaning for different kinds of models (e.g., dllm), here we only consider the sliding window in SWA model
|
||||
sliding_window_size = model_config.sliding_window_size
|
||||
elif model_config.attention_chunk_size is not None:
|
||||
sliding_window_size = model_config.attention_chunk_size
|
||||
logger.info(
|
||||
f"Setting sliding_window_size to be attention_chunk_size: {sliding_window_size}"
|
||||
)
|
||||
return sliding_window_size
|
||||
|
||||
|
||||
def report_online_quantization(*, model, server_args: ServerArgs) -> None:
|
||||
# TODO: Make sure all models have `quant_config` attribute, and all online quantization methods register which layers they actually quantize.
|
||||
quantized_layers = getattr(
|
||||
getattr(model, "quant_config", None), "quantized_layers", None
|
||||
)
|
||||
if (
|
||||
server_args.quantization is not None
|
||||
and isinstance(quantized_layers, tuple)
|
||||
and len(quantized_layers) == 2
|
||||
):
|
||||
layer_types, quantized_layers_count = quantized_layers
|
||||
logger.info(
|
||||
f"Online {server_args.quantization} quantization: quantized {quantized_layers_count} layers of types: {layer_types}"
|
||||
)
|
||||
|
||||
|
||||
def maybe_register_debug_tensor_dump_hook(
|
||||
*,
|
||||
model,
|
||||
server_args: ServerArgs,
|
||||
spec_algorithm: SpeculativeAlgorithm,
|
||||
is_draft_worker: bool,
|
||||
tp_size: int,
|
||||
tp_rank: int,
|
||||
pp_rank: int,
|
||||
) -> None:
|
||||
if server_args.debug_tensor_dump_output_folder is not None:
|
||||
dump_folder = server_args.debug_tensor_dump_output_folder
|
||||
if spec_algorithm.is_eagle():
|
||||
role = "draft" if is_draft_worker else "target"
|
||||
dump_folder = os.path.join(dump_folder, role)
|
||||
register_forward_hook_for_model(
|
||||
model,
|
||||
dump_folder,
|
||||
server_args.debug_tensor_dump_layers,
|
||||
tp_size,
|
||||
tp_rank,
|
||||
pp_rank,
|
||||
)
|
||||
|
||||
|
||||
def build_load_config(
|
||||
*,
|
||||
server_args: ServerArgs,
|
||||
tp_rank: int,
|
||||
remote_instance_weight_transporter_engine: Any,
|
||||
remote_instance_weight_transporter_session_id: str,
|
||||
draft_model_idx: Optional[int],
|
||||
) -> LoadConfig:
|
||||
from sglang.srt.configs.modelopt_config import ModelOptConfig
|
||||
|
||||
modelopt_config = ModelOptConfig(
|
||||
quant=server_args.modelopt_quant,
|
||||
checkpoint_restore_path=server_args.modelopt_checkpoint_restore_path,
|
||||
checkpoint_save_path=server_args.modelopt_checkpoint_save_path,
|
||||
export_path=server_args.modelopt_export_path,
|
||||
quantize_and_serve=server_args.quantize_and_serve,
|
||||
)
|
||||
|
||||
return LoadConfig(
|
||||
load_format=server_args.load_format,
|
||||
download_dir=server_args.download_dir,
|
||||
model_loader_extra_config=server_args.model_loader_extra_config,
|
||||
tp_rank=tp_rank,
|
||||
remote_instance_weight_loader_seed_instance_ip=server_args.remote_instance_weight_loader_seed_instance_ip,
|
||||
remote_instance_weight_loader_seed_instance_service_port=server_args.remote_instance_weight_loader_seed_instance_service_port,
|
||||
remote_instance_weight_loader_send_weights_group_ports=server_args.remote_instance_weight_loader_send_weights_group_ports,
|
||||
remote_instance_weight_loader_backend=server_args.remote_instance_weight_loader_backend,
|
||||
remote_instance_weight_loader_transfer_engine=remote_instance_weight_transporter_engine,
|
||||
remote_instance_weight_loader_transfer_engine_session_id=remote_instance_weight_transporter_session_id,
|
||||
modelexpress_url=server_args.modelexpress_url,
|
||||
modelexpress_transport=server_args.modelexpress_transport,
|
||||
modelopt_config=modelopt_config,
|
||||
rl_quant_profile=server_args.rl_quant_profile,
|
||||
draft_model_idx=draft_model_idx,
|
||||
)
|
||||
|
||||
|
||||
def load_model_with_memory_saver(
|
||||
*,
|
||||
server_args: ServerArgs,
|
||||
model_config: ModelConfig,
|
||||
load_config: LoadConfig,
|
||||
device: str,
|
||||
gpu_id: int,
|
||||
memory_saver_adapter: Any,
|
||||
is_draft_worker: bool,
|
||||
) -> LoadedModel:
|
||||
# Remove monkey_patch when linear.py quant remove dependencies with vllm
|
||||
monkey_patch_vllm_parallel_state()
|
||||
|
||||
enable_cpu_backup = server_args.enable_weights_cpu_backup or (
|
||||
is_draft_worker and server_args.enable_draft_weights_cpu_backup
|
||||
)
|
||||
remote_instance_weight_info = None
|
||||
with memory_saver_adapter.region(
|
||||
GPU_MEMORY_TYPE_WEIGHTS,
|
||||
enable_cpu_backup=enable_cpu_backup,
|
||||
):
|
||||
loader = get_model_loader(
|
||||
load_config=load_config,
|
||||
model_config=model_config,
|
||||
)
|
||||
model = loader.load_model(
|
||||
model_config=model_config,
|
||||
device_config=DeviceConfig(device, gpu_id),
|
||||
)
|
||||
if hasattr(loader, "remote_instance_transfer_engine_weight_info"):
|
||||
remote_instance_weight_info = (
|
||||
loader.remote_instance_transfer_engine_weight_info
|
||||
)
|
||||
# Cache needs to be cleared after loading model weights (in the loader.load_model function).
|
||||
# To avoid conflict with memory_saver_adapter.region, empty_cache operation is now moved here.
|
||||
if _is_npu:
|
||||
torch.npu.empty_cache()
|
||||
monkey_patch_vllm_parallel_state(reverse=True)
|
||||
|
||||
return LoadedModel(
|
||||
loader=loader,
|
||||
model=model,
|
||||
remote_instance_weight_info=remote_instance_weight_info,
|
||||
)
|
||||
|
||||
|
||||
def dist_barrier_after_load(*, elastic_ep_backend: Optional[str], tp_rank: int) -> None:
|
||||
if elastic_ep_backend == "mooncake":
|
||||
# Mooncake does not support `monitored_barrier`
|
||||
dist.barrier(group=get_tp_group().cpu_group)
|
||||
else:
|
||||
# Handle the case where some ranks do not finish loading.
|
||||
try:
|
||||
dist.monitored_barrier(
|
||||
group=get_tp_group().cpu_group,
|
||||
timeout=datetime.timedelta(seconds=UNBALANCED_MODEL_LOADING_TIMEOUT_S),
|
||||
wait_all_ranks=True,
|
||||
)
|
||||
except RuntimeError:
|
||||
raise ValueError(
|
||||
f"TP rank {tp_rank} could finish the model loading, but there are other ranks that didn't finish loading. It is likely due to unexpected failures (e.g., OOM) or a slow node."
|
||||
) from None
|
||||
Reference in New Issue
Block a user