Reland serving-time Triton load diagnostics (#33908)

Co-authored-by: kangwangamd <100359556+kangwangamd@users.noreply.github.com>
This commit is contained in:
cctry
2026-08-07 14:51:43 -07:00
committed by GitHub
co-authored by kangwangamd
parent 8fc66d1a62
commit df3aa20d89
3 changed files with 133 additions and 0 deletions
+6
View File
@@ -721,6 +721,12 @@ class Envs:
# (parity with flash-attn's ragged-aware launch). The feature checks _is_hip # (parity with flash-attn's ragged-aware launch). The feature checks _is_hip
# explicitly in code; this env var allows override (0=force off, 1=force on). # explicitly in code; this env var allows override (0=force off, 1=force on).
SGLANG_TRITON_COMPACT_EXTEND_ATTENTION = EnvBool(True) SGLANG_TRITON_COMPACT_EXTEND_ATTENTION = EnvBool(True)
# Raise if Triton loads a kernel after the engine starts serving. This
# verifies that startup warmup covers every kernel specialization used at
# serving time.
SGLANG_CRASH_ON_TRITON_LOAD_AFTER_READY = EnvBool(False)
SGLANG_TRITON_SLOW_COMPILE_THRESHOLD_SECS = EnvFloat(1.0)
SGLANG_TRITON_LOAD_WARNING_THRESHOLD_GB = EnvFloat(1.0)
# Torch Compile # Torch Compile
SGLANG_ENABLE_TORCH_COMPILE = EnvBool(False) SGLANG_ENABLE_TORCH_COMPILE = EnvBool(False)
+6
View File
@@ -301,6 +301,7 @@ from sglang.srt.utils import (
set_gpu_proc_affinity, set_gpu_proc_affinity,
set_random_seed, set_random_seed,
suppress_other_loggers, suppress_other_loggers,
triton_load_watch,
) )
from sglang.srt.utils.common import is_npu from sglang.srt.utils.common import is_npu
from sglang.srt.utils.hf_transformers_utils import ( from sglang.srt.utils.hf_transformers_utils import (
@@ -1629,6 +1630,11 @@ class Scheduler(
Sets up the schedule stream and dispatches to the appropriate event loop. Sets up the schedule stream and dispatches to the appropriate event loop.
The event loop blocks until shutdown. The event loop blocks until shutdown.
""" """
# Engine init (graph capture, warmups) is done; from here on any
# Triton kernel device-load is a lazy first-use at serving time.
triton_load_watch.install()
triton_load_watch.mark_serving_started()
if use_mlx(): if use_mlx():
# MLX overlap uses mx.async_eval for CPU/GPU overlap, # MLX overlap uses mx.async_eval for CPU/GPU overlap,
# not PyTorch MPS streams. # not PyTorch MPS streams.
@@ -0,0 +1,121 @@
"""Detect Triton kernel device-loads after the engine starts serving.
Triton loads each kernel specialization's cubin onto the GPU at its first
launch (``CompiledKernel._init_handles`` -> ``cuModuleLoadData``). That load
needs free device memory *outside* the torch caching allocator. Engines size
their pools to leave little post-init headroom, and the allocator's high-water
mark consumes the rest during early serving — so a specialization first used
mid-serving (e.g. a new adaptive speculative draft length, or a rare batch-size
bucket) can die in ``cuModuleLoadData`` with CUDA OOM, minutes or hours in.
Once ``mark_serving_started()`` has been called, this module warns when an
uncached Triton compilation takes at least one second or a device-load starts
with less than 1 GiB of free device memory. Set
``SGLANG_CRASH_ON_TRITON_LOAD_AFTER_READY=1`` to raise on every late load
instead — for CI recipes that assert full startup warmup coverage. The hooks
only run for compilation and first-use loads, so steady-state cost is zero.
Note: request-driven warmup (``--warmups``, the server warmup request) runs
*after* ``mark_serving_started()`` and is subject to the same diagnostics;
crash mode is only meant for deployments whose kernels are fully pre-loaded at
engine init.
"""
from __future__ import annotations
import logging
import torch
from sglang.srt.environ import envs
from sglang.srt.utils.common import get_available_gpu_memory
logger = logging.getLogger(__name__)
_serving_started = False
_prev_compile_listener = None
_installed = False
def install() -> None:
"""Install the diagnostics (idempotent; chains pre-existing hooks)."""
global _installed, _prev_compile_listener
if _installed:
return
try:
import triton.knobs as knobs
except ImportError:
return
# Triton 3.4 builds used by ROCm expose triton.knobs without these hooks.
runtime = getattr(knobs, "runtime", None)
compilation = getattr(knobs, "compilation", None)
hook = getattr(runtime, "kernel_load_start_hook", None)
if hook is None or compilation is None or not hasattr(compilation, "listener"):
return
_prev_compile_listener = compilation.listener
hook.add(_on_kernel_load)
compilation.listener = _on_compilation
_installed = True
def mark_serving_started() -> None:
"""Arm diagnostics for subsequent Triton compilations and device-loads."""
global _serving_started
_serving_started = True
def _on_compilation(*, src, metadata, metadata_group, times, cache_hit) -> None:
if _prev_compile_listener is not None:
_prev_compile_listener(
src=src,
metadata=metadata,
metadata_group=metadata_group,
times=times,
cache_hit=cache_hit,
)
if not _serving_started or cache_hit:
return
compile_time_secs = times.total / 1e6
if compile_time_secs < envs.SGLANG_TRITON_SLOW_COMPILE_THRESHOLD_SECS.get():
return
logger.warning(
"Triton kernel '%s' took %.2f s to compile after serving started. "
"Serving-time compilation can stall the engine; pre-compile it during "
"engine init.",
src.name,
compile_time_secs,
)
def _on_kernel_load(module, function, name, metadata_group, hash) -> None:
if not _serving_started:
return
free_gb = None
if torch.cuda.is_available():
try:
free_gb = get_available_gpu_memory(
"cuda", torch.cuda.current_device(), empty_cache=False
)
except RuntimeError:
logger.debug("Unable to query free device memory", exc_info=True)
should_crash = envs.SGLANG_CRASH_ON_TRITON_LOAD_AFTER_READY.get()
if not should_crash and (
free_gb is None or free_gb >= envs.SGLANG_TRITON_LOAD_WARNING_THRESHOLD_GB.get()
):
return
free_memory = f"{free_gb:.2f} GiB" if free_gb is not None else "unknown"
msg = (
f"Triton kernel '{name}' device-loaded after serving started "
f"(free device mem: {free_memory}). Pre-load it during engine init "
f"to avoid CUDA OOM."
)
if should_crash:
raise RuntimeError(msg)
logger.warning(msg)