Extract load_model helpers into a load_model_utils module (#31155)
This commit is contained in:
@@ -16,22 +16,16 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import contextlib
|
import contextlib
|
||||||
import datetime
|
|
||||||
import inspect
|
import inspect
|
||||||
import logging
|
import logging
|
||||||
import os
|
|
||||||
import socket
|
|
||||||
import threading
|
|
||||||
import time
|
import time
|
||||||
from collections import defaultdict
|
from collections import defaultdict
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from typing import Optional, Union
|
from typing import Optional, Union
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
import torch.distributed as dist
|
|
||||||
|
|
||||||
from sglang.srt.configs.device_config import DeviceConfig
|
from sglang.srt.configs.load_config import LoadConfig
|
||||||
from sglang.srt.configs.load_config import LoadConfig, LoadFormat
|
|
||||||
from sglang.srt.configs.model_config import (
|
from sglang.srt.configs.model_config import (
|
||||||
AttentionArch,
|
AttentionArch,
|
||||||
ModelConfig,
|
ModelConfig,
|
||||||
@@ -41,20 +35,14 @@ from sglang.srt.configs.model_config import (
|
|||||||
is_deepseek_dsa,
|
is_deepseek_dsa,
|
||||||
)
|
)
|
||||||
from sglang.srt.configs.update_config import adjust_config_with_unaligned_cpu_tp
|
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.dumper import dumper
|
||||||
from sglang.srt.debug_utils.tensor_dump_forward_hook import (
|
|
||||||
register_forward_hook_for_model,
|
|
||||||
)
|
|
||||||
from sglang.srt.distributed import (
|
from sglang.srt.distributed import (
|
||||||
bootstrap,
|
bootstrap,
|
||||||
get_tp_group,
|
|
||||||
get_world_group,
|
get_world_group,
|
||||||
)
|
)
|
||||||
from sglang.srt.distributed.device_communicators.pynccl_allocator import (
|
from sglang.srt.distributed.device_communicators.pynccl_allocator import (
|
||||||
prealloc_symmetric_memory_pool,
|
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.dllm.config import DllmConfig
|
||||||
from sglang.srt.elastic_ep.elastic_ep import (
|
from sglang.srt.elastic_ep.elastic_ep import (
|
||||||
ElasticEPStateManager,
|
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.graph_shared_output import GraphSharedOutput
|
||||||
from sglang.srt.model_executor.hook_manager import register_forward_hooks
|
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 (
|
from sglang.srt.model_executor.model_runner_components.ngram_embedding_manager import (
|
||||||
NgramEmbeddingManager,
|
NgramEmbeddingManager,
|
||||||
)
|
)
|
||||||
@@ -150,11 +149,6 @@ from sglang.srt.model_executor.runner import (
|
|||||||
PrefillCudaGraphRunner,
|
PrefillCudaGraphRunner,
|
||||||
get_batch_sizes_to_capture,
|
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.model_loader.utils import resolve_language_model
|
||||||
from sglang.srt.platforms import current_platform
|
from sglang.srt.platforms import current_platform
|
||||||
from sglang.srt.runtime_context import get_flags, get_server_args
|
from sglang.srt.runtime_context import get_flags, get_server_args
|
||||||
@@ -196,7 +190,7 @@ from sglang.srt.utils import (
|
|||||||
set_cuda_arch,
|
set_cuda_arch,
|
||||||
slow_rank_detector,
|
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_pytorch_hooks import PytHooks
|
||||||
from sglang.srt.utils.nvtx_utils import profile_range
|
from sglang.srt.utils.nvtx_utils import profile_range
|
||||||
from sglang.srt.utils.offloader import (
|
from sglang.srt.utils.offloader import (
|
||||||
@@ -222,7 +216,6 @@ elif current_platform.is_out_of_tree():
|
|||||||
current_platform.init_backend()
|
current_platform.init_backend()
|
||||||
|
|
||||||
# Detect stragger ranks in model loading
|
# Detect stragger ranks in model loading
|
||||||
UNBALANCED_MODEL_LOADING_TIMEOUT_S = 480 # leave more time for post data processing
|
|
||||||
|
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
@@ -1115,49 +1108,17 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
if self.device != "cpu":
|
if self.device != "cpu":
|
||||||
torch.set_num_threads(1)
|
torch.set_num_threads(1)
|
||||||
if self.device == "cuda":
|
if self.device == "cuda":
|
||||||
if torch.cuda.get_device_capability()[0] < 8:
|
maybe_downgrade_dtype_for_legacy_gpu(
|
||||||
logger.info(
|
server_args=self.server_args, model_config=self.model_config
|
||||||
"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.")
|
|
||||||
|
|
||||||
set_cuda_arch()
|
set_cuda_arch()
|
||||||
|
|
||||||
# Prepare the model config
|
self.load_config = build_load_config(
|
||||||
from sglang.srt.configs.modelopt_config import ModelOptConfig
|
server_args=self.server_args,
|
||||||
|
|
||||||
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,
|
|
||||||
tp_rank=self.tp_rank,
|
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_transporter_engine=self.remote_instance_weight_transporter.engine,
|
||||||
remote_instance_weight_loader_seed_instance_service_port=self.server_args.remote_instance_weight_loader_seed_instance_service_port,
|
remote_instance_weight_transporter_session_id=self.remote_instance_weight_transporter.session_id,
|
||||||
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,
|
|
||||||
draft_model_idx=self.draft_model_idx,
|
draft_model_idx=self.draft_model_idx,
|
||||||
)
|
)
|
||||||
if self.device == "cpu":
|
if self.device == "cpu":
|
||||||
@@ -1165,52 +1126,25 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
self.model_config, self.load_config, self.tp_size
|
self.model_config, self.load_config, self.tp_size
|
||||||
)
|
)
|
||||||
|
|
||||||
if (
|
maybe_trigger_remote_instance_nccl_send_group(
|
||||||
self.server_args.load_format == LoadFormat.REMOTE_INSTANCE
|
server_args=self.server_args, tp_rank=self.tp_rank
|
||||||
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
|
loaded = load_model_with_memory_saver(
|
||||||
# Remove monkey_patch when linear.py quant remove dependencies with vllm
|
server_args=self.server_args,
|
||||||
monkey_patch_vllm_parallel_state()
|
model_config=self.model_config,
|
||||||
|
|
||||||
enable_cpu_backup = self.server_args.enable_weights_cpu_backup or (
|
|
||||||
self.is_draft_worker and self.server_args.enable_draft_weights_cpu_backup
|
|
||||||
)
|
|
||||||
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,
|
load_config=self.load_config,
|
||||||
model_config=self.model_config,
|
device=self.device,
|
||||||
|
gpu_id=self.gpu_id,
|
||||||
|
memory_saver_adapter=self.memory_saver_adapter,
|
||||||
|
is_draft_worker=self.is_draft_worker,
|
||||||
)
|
)
|
||||||
self.model = self.loader.load_model(
|
self.loader = loaded.loader
|
||||||
model_config=self.model_config,
|
self.model = loaded.model
|
||||||
device_config=DeviceConfig(self.device, self.gpu_id),
|
if loaded.remote_instance_weight_info is not None:
|
||||||
)
|
|
||||||
if hasattr(self.loader, "remote_instance_transfer_engine_weight_info"):
|
|
||||||
self.remote_instance_weight_transporter.weight_info = (
|
self.remote_instance_weight_transporter.weight_info = (
|
||||||
self.loader.remote_instance_transfer_engine_weight_info
|
loaded.remote_instance_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:
|
if not self.is_draft_worker:
|
||||||
get_offloader().post_init()
|
get_offloader().post_init()
|
||||||
@@ -1220,43 +1154,10 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
pyt_hooks = PytHooks()
|
pyt_hooks = PytHooks()
|
||||||
pyt_hooks.register_hooks(self.model, module_prefix="model")
|
pyt_hooks.register_hooks(self.model, module_prefix="model")
|
||||||
|
|
||||||
if self.server_args.kv_cache_dtype == "fp8_e4m3":
|
load_kv_cache_scales(model=self.model, server_args=self.server_args)
|
||||||
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!"
|
|
||||||
)
|
|
||||||
|
|
||||||
# Parse other args
|
self.sliding_window_size = resolve_sliding_window_size(
|
||||||
self.sliding_window_size = None
|
self.model, self.model_config
|
||||||
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.prefill_aware_swa = (
|
self.prefill_aware_swa = (
|
||||||
@@ -1281,33 +1182,16 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
f"mem usage={self.weight_load_mem_usage:.2f} GB."
|
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.
|
report_online_quantization(model=self.model, server_args=self.server_args)
|
||||||
# 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}"
|
|
||||||
)
|
|
||||||
|
|
||||||
if self.server_args.debug_tensor_dump_output_folder is not None:
|
maybe_register_debug_tensor_dump_hook(
|
||||||
dump_folder = self.server_args.debug_tensor_dump_output_folder
|
model=self.model,
|
||||||
if self.spec_algorithm.is_eagle():
|
server_args=self.server_args,
|
||||||
role = "draft" if self.is_draft_worker else "target"
|
spec_algorithm=self.spec_algorithm,
|
||||||
dump_folder = os.path.join(dump_folder, role)
|
is_draft_worker=self.is_draft_worker,
|
||||||
register_forward_hook_for_model(
|
tp_size=self.tp_size,
|
||||||
self.model,
|
tp_rank=self.tp_rank,
|
||||||
dump_folder,
|
pp_rank=self.pp_rank,
|
||||||
self.server_args.debug_tensor_dump_layers,
|
|
||||||
self.tp_size,
|
|
||||||
self.tp_rank,
|
|
||||||
self.pp_rank,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
if dumper.may_enable:
|
if dumper.may_enable:
|
||||||
@@ -1322,23 +1206,10 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
logger,
|
logger,
|
||||||
)
|
)
|
||||||
|
|
||||||
if self.server_args.elastic_ep_backend == "mooncake":
|
dist_barrier_after_load(
|
||||||
# Mooncake does not support `monitored_barrier`
|
elastic_ep_backend=self.server_args.elastic_ep_backend,
|
||||||
dist.barrier(group=get_tp_group().cpu_group)
|
tp_rank=self.tp_rank,
|
||||||
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
|
|
||||||
|
|
||||||
def _prepare_moe_topk(self):
|
def _prepare_moe_topk(self):
|
||||||
balancer_cls = None
|
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