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:
Brayden Zhong
2026-09-04 20:05:53 -07:00
committed by GitHub
co-authored by Brayden Zhong Mohammad Miadh Angkad
parent 0645398a32
commit 92a4d8b5ee
19 changed files with 111 additions and 78 deletions
@@ -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!"
)
@@ -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
+1 -1
View File
@@ -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
+1 -1
View File
@@ -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."
)
+1 -1
View File
@@ -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
+1 -1
View File
@@ -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
+10 -1
View File
@@ -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(