Clean logging under --weight-loader-prefetch-checkpoints (#33930)
Co-authored-by: Brayden Zhong <brayden@radixark.ai> Co-authored-by: Mohammad Miadh Angkad <176301910+mmangkad@users.noreply.github.com>
This commit is contained in:
co-authored by
Brayden Zhong
Mohammad Miadh Angkad
parent
0645398a32
commit
92a4d8b5ee
@@ -95,10 +95,6 @@ import { GLM5Deployment } from '/src/snippets/autoregressive/glm-5-deployment.js
|
||||
- For other configuration tips (MTP, DSA kernel, Context Parallel, HiSparse, NVFP4, Index Cache), see the [DeepSeek-V3.2 cookbook page](../DeepSeek/DeepSeek-V3_2). GLM-5 and DeepSeek-V3.2 share the same model structure, so the optimization techniques are common.
|
||||
- Use `--json-model-override-args '{"index_topk_pattern": "FFSFSSSFSSFFFSSSFFFSFSSSSSSFFSFFSFFSSFFFFFFSFFFFFSFFSSSSSSFSFFFSFSSSFSFFSFFSSS"}'` for GLM-5-FP8 if you want to enable the [IndexCache](https://github.com/THUDM/IndexCache) method. This feature is supported through [this PR](https://github.com/sgl-project/sglang/pull/21405) and introduces only a small accuracy loss. However, if you are running rigorous accuracy evaluations, it is not recommended to enable this feature.
|
||||
|
||||
<Warning>
|
||||
**FP8 KV Cache**: `--kv-cache-dtype fp8_e4m3` quantizes the KV cache to FP8 at runtime. Since these FP8 model checkpoints do not include pre-calibrated KV cache scaling factors, SGLang defaults to a scale of 1.0, which may cause noticeable accuracy degradation on reasoning-heavy tasks. It is not included in the generated commands above; add it manually only if memory constraints require the trade-off.
|
||||
</Warning>
|
||||
|
||||
## 4. Model Invocation
|
||||
|
||||
Deploy GLM-5 with the following command (FP8 on H200, all features enabled):
|
||||
|
||||
@@ -322,10 +322,7 @@ class ModelOptFp4Config(ModelOptQuantConfig):
|
||||
super().__init__(exclude_modules, packed_modules_mapping)
|
||||
self.is_checkpoint_nvfp4_serialized = is_checkpoint_nvfp4_serialized
|
||||
if is_checkpoint_nvfp4_serialized:
|
||||
logger.warning(
|
||||
"Detected nvfp4 checkpoint. Please note that the "
|
||||
"format is experimental and subject to change."
|
||||
)
|
||||
logger.info("Detected nvfp4 checkpoint.")
|
||||
self.group_size = group_size
|
||||
self.checkpoint_uses_packed_qkv = checkpoint_uses_packed_qkv
|
||||
self.swap_weight_nibbles = swap_weight_nibbles
|
||||
|
||||
@@ -60,7 +60,6 @@ class CpDecodeAttnTpContext:
|
||||
else:
|
||||
self.decode_tp_rank = None
|
||||
self.decode_tp_size = None
|
||||
logger.info("Disable CP decode attention TP")
|
||||
|
||||
self.use_decode_attn_tp = False
|
||||
self._slice_cache: Dict = {}
|
||||
|
||||
@@ -83,7 +83,6 @@ from sglang.srt.utils import (
|
||||
is_cpu,
|
||||
is_hip,
|
||||
is_npu,
|
||||
print_info_once,
|
||||
round_up,
|
||||
)
|
||||
from sglang.srt.utils.custom_op import register_custom_op
|
||||
@@ -474,7 +473,7 @@ class FusedMoE(torch.nn.Module):
|
||||
global _deferred_finalize_info_logged
|
||||
if not _deferred_finalize_info_logged:
|
||||
_deferred_finalize_info_logged = True
|
||||
logging.getLogger(__name__).info(
|
||||
logging.getLogger(__name__).debug(
|
||||
"FlashInfer TRTLLM MoE deferred finalize is "
|
||||
f"{'enabled' if self.supports_deferred_finalize else 'disabled'} "
|
||||
f"(moe_runner_backend={get_exec().moe.moe_runner_backend}, "
|
||||
@@ -516,10 +515,6 @@ class FusedMoE(torch.nn.Module):
|
||||
get_moe_runner_backend().is_flashinfer_trtllm_routed()
|
||||
or get_moe_runner_backend().is_flashinfer_trtllm()
|
||||
):
|
||||
if self.moe_runner_config.inplace:
|
||||
print_info_once(
|
||||
"Setting inplace to False for FlashInfer TRTLLM MoE backend."
|
||||
)
|
||||
self.moe_runner_config.inplace = False
|
||||
|
||||
self.should_fuse_routed_scaling_factor_in_topk = (
|
||||
|
||||
@@ -1429,10 +1429,7 @@ class ModelOptFp4Config(ModelOptQuantConfig):
|
||||
super().__init__(kv_cache_quant_algo, exclude_modules, packed_modules_mapping)
|
||||
self.is_checkpoint_nvfp4_serialized = is_checkpoint_nvfp4_serialized
|
||||
if is_checkpoint_nvfp4_serialized:
|
||||
logger.warning(
|
||||
"Detected nvfp4 checkpoint. Please note that the "
|
||||
"format is experimental and subject to change."
|
||||
)
|
||||
logger.info("Detected nvfp4 checkpoint.")
|
||||
self.is_awq = is_awq
|
||||
self.is_w4a16 = False
|
||||
self.group_size = group_size
|
||||
|
||||
@@ -45,10 +45,7 @@ class PetitNvFp4Config(QuantizationConfig):
|
||||
) -> None:
|
||||
self.is_checkpoint_nvfp4_serialized = is_checkpoint_nvfp4_serialized
|
||||
if is_checkpoint_nvfp4_serialized:
|
||||
logger.warning(
|
||||
"Detected nvfp4 checkpoint. Please note that the "
|
||||
"format is experimental and subject to change."
|
||||
)
|
||||
logger.info("Detected nvfp4 checkpoint.")
|
||||
self.group_size = group_size
|
||||
self.kv_cache_quant_algo = kv_cache_quant_algo
|
||||
self.exclude_modules = exclude_modules
|
||||
|
||||
@@ -287,6 +287,7 @@ def capture_prefill_graph(
|
||||
"""Initialize a prefill graph and return its startup resource usage."""
|
||||
|
||||
memory_phase = "draft_prefill" if model_runner.is_draft_worker else "prefill"
|
||||
role = "draft" if model_runner.is_draft_worker else "target"
|
||||
|
||||
def result(
|
||||
runner: Optional[BaseRunner],
|
||||
@@ -430,8 +431,10 @@ def capture_prefill_graph(
|
||||
layer_model = layer_model.model
|
||||
|
||||
if not hasattr(layer_model, "layers"):
|
||||
logger.warning(
|
||||
"Disable prefill CUDA graph because the model does not have a 'layers' attribute"
|
||||
log_info_on_rank0(
|
||||
logger,
|
||||
f"Disable {role} prefill CUDA graph because the {role} model does "
|
||||
"not have a 'layers' attribute",
|
||||
)
|
||||
return result(None)
|
||||
|
||||
@@ -467,7 +470,6 @@ def capture_prefill_graph(
|
||||
|
||||
tic = time.perf_counter()
|
||||
before_mem = get_available_gpu_memory(model_runner.device, model_runner.gpu_id)
|
||||
role = "draft" if model_runner.is_draft_worker else "target"
|
||||
capture_name = f"{role} prefill"
|
||||
logger.info(
|
||||
f"Capture {capture_name} CUDA graph begin. "
|
||||
|
||||
@@ -133,7 +133,6 @@ def load_kv_cache_scales(*, model, kv_cache_dtype: str) -> None:
|
||||
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!"
|
||||
)
|
||||
|
||||
|
||||
|
||||
+6
@@ -23,6 +23,7 @@ buffers to keep break-point tensors at stable addresses.
|
||||
"""
|
||||
|
||||
import threading
|
||||
import warnings
|
||||
from contextvars import ContextVar
|
||||
from typing import Any, Callable, Optional
|
||||
|
||||
@@ -394,6 +395,11 @@ class BreakableCUDAGraphCapture:
|
||||
forked.clear()
|
||||
graph = self._current_graph
|
||||
assert graph is not None
|
||||
# A segment that enqueued no kernels (back-to-back breaks, or a segment
|
||||
# whose ops all ran eagerly) captures an empty graph, which replays as a
|
||||
# no-op. Torch warns about it on every such capture_end; expected here.
|
||||
with warnings.catch_warnings():
|
||||
warnings.filterwarnings("ignore", message="The CUDA Graph is empty")
|
||||
graph.capture_end()
|
||||
self.cuda_graph._append_segment(graph, self._current_graph_needs_instantiate)
|
||||
self._current_graph = None
|
||||
|
||||
@@ -640,7 +640,7 @@ class DefaultModelLoader(BaseModelLoader):
|
||||
{"enable_multithread_load", "num_threads"} & extra_config.keys()
|
||||
)
|
||||
):
|
||||
logger.warning(
|
||||
logger.debug(
|
||||
"Checkpoint prefetching is active; falling "
|
||||
"back to single-threaded weight loading to avoid I/O "
|
||||
"oversubscription with the prefetch threads. Set "
|
||||
|
||||
@@ -1021,7 +1021,7 @@ def _prefetch_all_checkpoints(
|
||||
succeeded_event = threading.Event()
|
||||
errors: List[Tuple[str, Exception]] = []
|
||||
|
||||
logger.info(
|
||||
logger.debug(
|
||||
"Rank %d: prefetching %d/%d checkpoint shards into page cache "
|
||||
"(background, %d local ranks sharing the work, %d threads per rank)...",
|
||||
local_rank,
|
||||
@@ -1042,7 +1042,7 @@ def _prefetch_all_checkpoints(
|
||||
if total_for_rank > 0 and next_log_pct <= 100:
|
||||
pct = 100 * completed / total_for_rank
|
||||
while pct >= next_log_pct and next_log_pct <= 100:
|
||||
logger.info(
|
||||
logger.debug(
|
||||
"Rank %d: prefetching checkpoint files: %d%% (%d/%d)",
|
||||
local_rank,
|
||||
next_log_pct,
|
||||
@@ -1093,7 +1093,7 @@ def _prefetch_all_checkpoints(
|
||||
start = time.perf_counter()
|
||||
_prefetch_all()
|
||||
succeeded_event.set()
|
||||
logger.info(
|
||||
logger.debug(
|
||||
"Rank %d: prefetching checkpoint files into page cache finished in %.2fs",
|
||||
local_rank,
|
||||
time.perf_counter() - start,
|
||||
|
||||
@@ -80,7 +80,7 @@ class BailingMoEModelNextN(nn.Module):
|
||||
config.for_nextn_model = True
|
||||
|
||||
if quant_config is not None and quant_config.get_name() == "modelopt_fp4":
|
||||
logger.warning(
|
||||
logger.debug(
|
||||
"Overriding DeepseekV3ForCausalLMNextN quant config for modelopt_fp4 Deepseek model."
|
||||
)
|
||||
quant_config = None
|
||||
|
||||
@@ -116,7 +116,7 @@ class DeepseekModelNextN(nn.Module):
|
||||
moe_quant_config_override = None
|
||||
|
||||
if quant_config is not None and quant_config.get_name() == "modelopt_fp4":
|
||||
logger.warning(
|
||||
logger.debug(
|
||||
"Overriding DeepseekV3ForCausalLMNextN quant config for modelopt_fp4 Deepseek model."
|
||||
)
|
||||
quant_config = None
|
||||
|
||||
@@ -50,7 +50,7 @@ class Glm4MoeLiteModelNextN(nn.Module):
|
||||
) -> None:
|
||||
super().__init__()
|
||||
if quant_config is not None and quant_config.get_name() == "modelopt_fp4":
|
||||
logger.warning(
|
||||
logger.debug(
|
||||
"Overriding Glm4MoeLiteForCausalLMNextN quant config for modelopt_fp4 "
|
||||
"GLM-4.7-Flash model."
|
||||
)
|
||||
|
||||
@@ -47,7 +47,7 @@ class Glm4MoeModelNextN(nn.Module):
|
||||
) -> None:
|
||||
super().__init__()
|
||||
if quant_config is not None and quant_config.get_name() == "modelopt_fp4":
|
||||
logger.warning(
|
||||
logger.debug(
|
||||
"Overriding Glm4MoeForCausalLMNextN quant config for modelopt_fp4 GLM-4.5 / GLM-4.6 / GLM-4.7 model."
|
||||
)
|
||||
quant_config = None
|
||||
|
||||
@@ -49,7 +49,7 @@ class GlmOcrModelNextN(nn.Module):
|
||||
) -> None:
|
||||
super().__init__()
|
||||
if quant_config is not None and quant_config.get_name() == "modelopt_fp4":
|
||||
logger.warning(
|
||||
logger.debug(
|
||||
"Overriding GlmOcrModelNextN quant config for modelopt_fp4 GLM-OCR model."
|
||||
)
|
||||
quant_config = None
|
||||
|
||||
@@ -58,6 +58,7 @@ from importlib.metadata import PackageNotFoundError, version
|
||||
from importlib.util import find_spec
|
||||
from io import BytesIO
|
||||
from json import JSONDecodeError
|
||||
from multiprocessing import parent_process
|
||||
from multiprocessing.reduction import ForkingPickler
|
||||
from pathlib import Path
|
||||
from typing import (
|
||||
@@ -2409,6 +2410,14 @@ def configure_logger(server_args, prefix: str = ""):
|
||||
for name in ("httpx", "httpcore"):
|
||||
logging.getLogger(name).setLevel(logging.WARNING)
|
||||
|
||||
# Server-sent hub warnings (e.g. the unauthenticated-request / HF_TOKEN
|
||||
# hint) are deduplicated per process, so a TP-N launch repeats each one N
|
||||
# times. Keep them only in the launching process -- every worker (scheduler,
|
||||
# detokenizer, DP controller, ...) is spawned via multiprocessing, whether
|
||||
# or not it passes a log prefix -- where they are printed exactly once.
|
||||
if parent_process() is not None:
|
||||
logging.getLogger("huggingface_hub.utils._http").setLevel(logging.ERROR)
|
||||
|
||||
if is_flashinfer_available():
|
||||
from flashinfer.jit.core import logger as flashinfer_logger
|
||||
|
||||
@@ -4106,7 +4115,7 @@ def freeze_gc(context: str):
|
||||
g0_before, g1_before, g2_before = gc_object_counts()
|
||||
gc.freeze()
|
||||
g0_after, g1_after, g2_after = gc_object_counts()
|
||||
logger.info(
|
||||
logger.debug(
|
||||
f"Freezing GC in {context} process. "
|
||||
f"gen0: {g0_before}->{g0_after}, "
|
||||
f"gen1: {g1_before}->{g1_after}, "
|
||||
|
||||
@@ -53,6 +53,8 @@ def apply_all():
|
||||
return
|
||||
_applied = True
|
||||
|
||||
_mute_diffusers_torchao_probe()
|
||||
|
||||
# v5.4 patches
|
||||
_patch_flash_attn_availability()
|
||||
_patch_rope_parameters_validation()
|
||||
@@ -71,6 +73,23 @@ def apply_all():
|
||||
logger.debug("transformers compatibility patches applied")
|
||||
|
||||
|
||||
def _mute_diffusers_torchao_probe():
|
||||
"""Silence diffusers' torchao-Tensor-subclass probe warning.
|
||||
|
||||
diffusers lazily imports its torchao quantizer and warns when the installed
|
||||
torchao has moved the optional Tensor subclasses it probes for. It only
|
||||
affects loading torchao-serialized diffusers checkpoints, which no sglang
|
||||
path does. Set here rather than in ``configure_logger`` because the import
|
||||
can land before logging is configured, and the level sticks whenever the
|
||||
lazy import happens.
|
||||
"""
|
||||
import logging
|
||||
|
||||
logging.getLogger("diffusers.quantizers.torchao.torchao_quantizer").setLevel(
|
||||
logging.ERROR
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Public API: on-demand helpers (called explicitly by other modules)
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -203,15 +222,22 @@ def _patch_removed_symbols():
|
||||
|
||||
# Importing modeling_llama triggers a deep import chain:
|
||||
# modeling_llama -> modeling_utils -> quantizers -> torchao
|
||||
# torchao emits a noisy warning about incompatible torch versions
|
||||
# that is irrelevant here — suppress it during this import.
|
||||
_torchao_logger = logging.getLogger("torchao")
|
||||
_prev_level = _torchao_logger.level
|
||||
_torchao_logger.setLevel(logging.ERROR)
|
||||
# torchao emits a noisy warning about incompatible torch versions, and
|
||||
# its register_as_pytree_constant() calls on Enum types make
|
||||
# torch.utils._pytree log a deprecation warning once per Enum and per
|
||||
# rank. Neither is actionable here — suppress both during this import.
|
||||
_muted = [
|
||||
logging.getLogger("torchao"),
|
||||
logging.getLogger("torch.utils._pytree"),
|
||||
]
|
||||
_prev_levels = [lg.level for lg in _muted]
|
||||
for lg in _muted:
|
||||
lg.setLevel(logging.ERROR)
|
||||
try:
|
||||
from transformers.models.llama import modeling_llama
|
||||
finally:
|
||||
_torchao_logger.setLevel(_prev_level)
|
||||
for lg, level in zip(_muted, _prev_levels):
|
||||
lg.setLevel(level)
|
||||
|
||||
if not hasattr(modeling_llama, "LlamaFlashAttention2"):
|
||||
if hasattr(modeling_llama, "LlamaAttention"):
|
||||
|
||||
@@ -211,13 +211,13 @@ class TestPrefetchCheckpoints(CustomTestCase):
|
||||
patch("concurrent.futures.ThreadPoolExecutor", _InlineExecutor),
|
||||
patch("concurrent.futures.wait", side_effect=_wait_all),
|
||||
patch("sglang.srt.model_loader.weight_utils._prefetch_checkpoint_file"),
|
||||
patch("sglang.srt.model_loader.weight_utils.logger.info") as log_info,
|
||||
patch("sglang.srt.model_loader.weight_utils.logger.debug") as log_debug,
|
||||
):
|
||||
_prefetch_all_checkpoints(paths, num_threads=1)
|
||||
|
||||
progress_pcts = [
|
||||
call.args[2]
|
||||
for call in log_info.call_args_list
|
||||
for call in log_debug.call_args_list
|
||||
if call.args
|
||||
and call.args[0] == "Rank %d: prefetching checkpoint files: %d%% (%d/%d)"
|
||||
]
|
||||
@@ -426,14 +426,24 @@ class TestPrefetchDispatch(CustomTestCase):
|
||||
"sglang.srt.model_loader.loader.safetensors_weights_iterator",
|
||||
return_value=iter([]),
|
||||
),
|
||||
patch("sglang.srt.model_loader.loader.logger.warning"),
|
||||
patch("sglang.srt.model_loader.loader.logger.debug"),
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _override_notices(mock_log):
|
||||
"""The single-thread override notice among the captured log calls."""
|
||||
return [
|
||||
call
|
||||
for call in mock_log.call_args_list
|
||||
if call.args
|
||||
and "falling back to single-threaded weight loading" in call.args[0]
|
||||
]
|
||||
|
||||
def test_prefetch_uses_single_thread_for_default_config(self):
|
||||
"""Prefetch on + no explicit multithread config -> single-threaded,
|
||||
and the opt-out warning fires once."""
|
||||
and the opt-out notice fires once."""
|
||||
loader = self._make_loader({})
|
||||
p_prep, p_model, p_buffered, p_single, p_warn = self._patch_dispatch(
|
||||
p_prep, p_model, p_buffered, p_single, p_log = self._patch_dispatch(
|
||||
prefetch=True
|
||||
)
|
||||
with (
|
||||
@@ -441,18 +451,18 @@ class TestPrefetchDispatch(CustomTestCase):
|
||||
p_model,
|
||||
p_buffered as mock_buffered,
|
||||
p_single as mock_single,
|
||||
p_warn as mock_warning,
|
||||
p_log as mock_log,
|
||||
):
|
||||
self._run(loader)
|
||||
mock_single.assert_called_once()
|
||||
mock_buffered.assert_not_called()
|
||||
mock_warning.assert_called_once()
|
||||
self.assertEqual(len(self._override_notices(mock_log)), 1)
|
||||
|
||||
def test_explicit_enable_multithread_keeps_buffered_with_prefetch(self):
|
||||
"""Explicit enable_multithread_load=true is the escape hatch; the
|
||||
override and its warning must not fire."""
|
||||
loader = self._make_loader({"enable_multithread_load": True})
|
||||
p_prep, p_model, p_buffered, p_single, p_warn = self._patch_dispatch(
|
||||
p_prep, p_model, p_buffered, p_single, p_log = self._patch_dispatch(
|
||||
prefetch=True
|
||||
)
|
||||
with (
|
||||
@@ -460,19 +470,19 @@ class TestPrefetchDispatch(CustomTestCase):
|
||||
p_model,
|
||||
p_buffered as mock_buffered,
|
||||
p_single as mock_single,
|
||||
p_warn as mock_warning,
|
||||
p_log as mock_log,
|
||||
):
|
||||
self._run(loader)
|
||||
mock_buffered.assert_called_once()
|
||||
mock_single.assert_not_called()
|
||||
mock_warning.assert_not_called()
|
||||
self.assertEqual(self._override_notices(mock_log), [])
|
||||
|
||||
def test_num_threads_only_keeps_buffered_with_prefetch(self):
|
||||
"""num_threads alone (relying on the enable_multithread_load=True
|
||||
default) also signals multi-thread intent, so the override must not
|
||||
fire and num_threads stays live."""
|
||||
loader = self._make_loader({"num_threads": 64})
|
||||
p_prep, p_model, p_buffered, p_single, p_warn = self._patch_dispatch(
|
||||
p_prep, p_model, p_buffered, p_single, p_log = self._patch_dispatch(
|
||||
prefetch=True
|
||||
)
|
||||
with (
|
||||
@@ -480,20 +490,20 @@ class TestPrefetchDispatch(CustomTestCase):
|
||||
p_model,
|
||||
p_buffered as mock_buffered,
|
||||
p_single as mock_single,
|
||||
p_warn as mock_warning,
|
||||
p_log as mock_log,
|
||||
):
|
||||
self._run(loader)
|
||||
mock_buffered.assert_called_once()
|
||||
# num_threads is forwarded as max_workers to the buffered iterator.
|
||||
self.assertEqual(mock_buffered.call_args.kwargs["max_workers"], 64)
|
||||
mock_single.assert_not_called()
|
||||
mock_warning.assert_not_called()
|
||||
self.assertEqual(self._override_notices(mock_log), [])
|
||||
|
||||
def test_no_prefetch_uses_multithread(self):
|
||||
"""Prefetch off -> multi-threaded iterator is used (default), no
|
||||
override warning."""
|
||||
loader = self._make_loader({})
|
||||
p_prep, p_model, p_buffered, p_single, p_warn = self._patch_dispatch(
|
||||
p_prep, p_model, p_buffered, p_single, p_log = self._patch_dispatch(
|
||||
prefetch=False
|
||||
)
|
||||
with (
|
||||
@@ -501,12 +511,12 @@ class TestPrefetchDispatch(CustomTestCase):
|
||||
p_model,
|
||||
p_buffered as mock_buffered,
|
||||
p_single as mock_single,
|
||||
p_warn as mock_warning,
|
||||
p_log as mock_log,
|
||||
):
|
||||
self._run(loader)
|
||||
mock_buffered.assert_called_once()
|
||||
mock_single.assert_not_called()
|
||||
mock_warning.assert_not_called()
|
||||
self.assertEqual(self._override_notices(mock_log), [])
|
||||
|
||||
def test_startup_prefetch_reuses_existing_background_handle(self):
|
||||
"""Startup commit reuses resolved shards and the active prefetch handle."""
|
||||
@@ -518,7 +528,7 @@ class TestPrefetchDispatch(CustomTestCase):
|
||||
weight_files=("f.safetensors",),
|
||||
use_safetensors=True,
|
||||
)
|
||||
p_prep, p_model, p_buffered, p_single, p_warn = self._patch_dispatch(
|
||||
p_prep, p_model, p_buffered, p_single, p_log = self._patch_dispatch(
|
||||
prefetch=False
|
||||
)
|
||||
with (
|
||||
@@ -526,7 +536,7 @@ class TestPrefetchDispatch(CustomTestCase):
|
||||
p_model,
|
||||
p_buffered as mock_buffered,
|
||||
p_single as mock_single,
|
||||
p_warn as mock_warning,
|
||||
p_log as mock_log,
|
||||
):
|
||||
list(
|
||||
loader._get_weights_iterator(
|
||||
@@ -541,7 +551,7 @@ class TestPrefetchDispatch(CustomTestCase):
|
||||
mock_single.assert_called_once()
|
||||
self.assertFalse(mock_single.call_args.kwargs["prefetch"])
|
||||
mock_buffered.assert_not_called()
|
||||
mock_warning.assert_called_once()
|
||||
self.assertEqual(len(self._override_notices(mock_log)), 1)
|
||||
|
||||
def test_completed_startup_prefetch_restores_multithread_loader(self):
|
||||
loader = self._make_loader({})
|
||||
@@ -552,7 +562,7 @@ class TestPrefetchDispatch(CustomTestCase):
|
||||
weight_files=("f.safetensors",),
|
||||
use_safetensors=True,
|
||||
)
|
||||
p_prep, p_model, p_buffered, p_single, p_warn = self._patch_dispatch(
|
||||
p_prep, p_model, p_buffered, p_single, p_log = self._patch_dispatch(
|
||||
prefetch=False
|
||||
)
|
||||
with (
|
||||
@@ -560,7 +570,7 @@ class TestPrefetchDispatch(CustomTestCase):
|
||||
p_model,
|
||||
p_buffered as mock_buffered,
|
||||
p_single as mock_single,
|
||||
p_warn as mock_warning,
|
||||
p_log as mock_log,
|
||||
):
|
||||
list(
|
||||
loader._get_weights_iterator(
|
||||
@@ -575,7 +585,7 @@ class TestPrefetchDispatch(CustomTestCase):
|
||||
mock_buffered.assert_called_once()
|
||||
self.assertFalse(mock_buffered.call_args.kwargs["prefetch"])
|
||||
mock_single.assert_not_called()
|
||||
mock_warning.assert_not_called()
|
||||
self.assertEqual(self._override_notices(mock_log), [])
|
||||
|
||||
def test_completed_startup_prefetch_is_not_started_twice(self):
|
||||
loader = self._make_loader({})
|
||||
@@ -586,7 +596,7 @@ class TestPrefetchDispatch(CustomTestCase):
|
||||
weight_files=("f.safetensors",),
|
||||
use_safetensors=True,
|
||||
)
|
||||
p_prep, p_model, p_buffered, p_single, p_warn = self._patch_dispatch(
|
||||
p_prep, p_model, p_buffered, p_single, p_log = self._patch_dispatch(
|
||||
prefetch=True
|
||||
)
|
||||
with (
|
||||
@@ -594,7 +604,7 @@ class TestPrefetchDispatch(CustomTestCase):
|
||||
p_model,
|
||||
p_buffered as mock_buffered,
|
||||
p_single as mock_single,
|
||||
p_warn as mock_warning,
|
||||
p_log as mock_log,
|
||||
):
|
||||
list(
|
||||
loader._get_weights_iterator(
|
||||
@@ -608,13 +618,13 @@ class TestPrefetchDispatch(CustomTestCase):
|
||||
mock_buffered.assert_called_once()
|
||||
self.assertFalse(mock_buffered.call_args.kwargs["prefetch"])
|
||||
mock_single.assert_not_called()
|
||||
mock_warning.assert_not_called()
|
||||
self.assertEqual(self._override_notices(mock_log), [])
|
||||
|
||||
def test_prefetch_does_not_override_when_mmap_disabled(self):
|
||||
"""Prefetch is a no-op without mmap, so the override and its warning
|
||||
must not fire."""
|
||||
loader = self._make_loader({})
|
||||
p_prep, p_model, p_buffered, p_single, p_warn = self._patch_dispatch(
|
||||
p_prep, p_model, p_buffered, p_single, p_log = self._patch_dispatch(
|
||||
prefetch=True, disable_mmap=True
|
||||
)
|
||||
with (
|
||||
@@ -622,18 +632,18 @@ class TestPrefetchDispatch(CustomTestCase):
|
||||
p_model,
|
||||
p_buffered as mock_buffered,
|
||||
p_single as mock_single,
|
||||
p_warn as mock_warning,
|
||||
p_log as mock_log,
|
||||
):
|
||||
self._run(loader)
|
||||
mock_buffered.assert_called_once()
|
||||
mock_single.assert_not_called()
|
||||
mock_warning.assert_not_called()
|
||||
self.assertEqual(self._override_notices(mock_log), [])
|
||||
|
||||
def test_prefetch_does_not_override_for_fastsafetensors(self):
|
||||
"""FASTSAFETENSORS ignores both flags; override + warning must not
|
||||
fire."""
|
||||
loader = self._make_loader({}, load_format=LoadFormat.FASTSAFETENSORS)
|
||||
p_prep, p_model, p_buffered, p_single, p_warn = self._patch_dispatch(
|
||||
p_prep, p_model, p_buffered, p_single, p_log = self._patch_dispatch(
|
||||
prefetch=True
|
||||
)
|
||||
with (
|
||||
@@ -645,7 +655,7 @@ class TestPrefetchDispatch(CustomTestCase):
|
||||
p_model,
|
||||
p_buffered as mock_buffered,
|
||||
p_single as mock_single,
|
||||
p_warn as mock_warning,
|
||||
p_log as mock_log,
|
||||
):
|
||||
self._run(loader)
|
||||
mock_fast.assert_called_once_with(
|
||||
@@ -655,13 +665,13 @@ class TestPrefetchDispatch(CustomTestCase):
|
||||
)
|
||||
mock_buffered.assert_not_called()
|
||||
mock_single.assert_not_called()
|
||||
mock_warning.assert_not_called()
|
||||
self.assertEqual(self._override_notices(mock_log), [])
|
||||
|
||||
def test_fastsafetensors_gds_can_be_disabled(self):
|
||||
loader = self._make_loader(
|
||||
{"enable_gds": False}, load_format=LoadFormat.FASTSAFETENSORS
|
||||
)
|
||||
p_prep, p_model, p_buffered, p_single, p_warn = self._patch_dispatch(
|
||||
p_prep, p_model, p_buffered, p_single, p_log = self._patch_dispatch(
|
||||
prefetch=False,
|
||||
drop_cache=True,
|
||||
)
|
||||
@@ -674,7 +684,7 @@ class TestPrefetchDispatch(CustomTestCase):
|
||||
p_model,
|
||||
p_buffered,
|
||||
p_single,
|
||||
p_warn,
|
||||
p_log,
|
||||
):
|
||||
self._run(loader)
|
||||
mock_fast.assert_called_once_with(
|
||||
|
||||
Reference in New Issue
Block a user