diff --git a/python/sglang/jit_kernel/all_reduce.py b/python/sglang/jit_kernel/all_reduce.py index 76eca0581..f913aa2b3 100644 --- a/python/sglang/jit_kernel/all_reduce.py +++ b/python/sglang/jit_kernel/all_reduce.py @@ -12,6 +12,7 @@ from sglang.jit_kernel.utils import ( load_jit, make_cpp_args, ) +from sglang.kernel_api_logging import debug_kernel_api class ConfigResult(NamedTuple): @@ -158,6 +159,7 @@ def get_custom_all_reduce_cls() -> type[CustomAllReduceObj]: def world_size(self) -> int: return self._world_size + @debug_kernel_api def all_reduce( self, input: torch.Tensor, diff --git a/python/sglang/jit_kernel/awq_marlin_repack.py b/python/sglang/jit_kernel/awq_marlin_repack.py index 5c0614629..d51c1fd51 100644 --- a/python/sglang/jit_kernel/awq_marlin_repack.py +++ b/python/sglang/jit_kernel/awq_marlin_repack.py @@ -4,8 +4,8 @@ from typing import TYPE_CHECKING import torch -from sglang.jit_kernel.debug_utils import maybe_wrap_jit_kernel_debug from sglang.jit_kernel.utils import cache_once, load_jit +from sglang.kernel_api_logging import debug_kernel_api if TYPE_CHECKING: from tvm_ffi.module import Module @@ -20,7 +20,7 @@ def _jit_awq_marlin_repack_module() -> Module: ) -@maybe_wrap_jit_kernel_debug +@debug_kernel_api def awq_marlin_repack( b_q_weight: torch.Tensor, size_k: int, @@ -39,7 +39,7 @@ def awq_marlin_repack( return out -@maybe_wrap_jit_kernel_debug +@debug_kernel_api def awq_marlin_moe_repack( b_q_weight: torch.Tensor, perm: torch.Tensor, diff --git a/python/sglang/jit_kernel/debug_utils.py b/python/sglang/jit_kernel/debug_utils.py deleted file mode 100644 index d65ef2feb..000000000 --- a/python/sglang/jit_kernel/debug_utils.py +++ /dev/null @@ -1,45 +0,0 @@ -import os -from typing import Any, Callable, TypeVar, cast, overload - -F = TypeVar("F", bound=Callable[..., Any]) - - -def _wrap_jit_kernel_debug(func: F, op_name: str | None = None) -> F: - try: - if int(os.environ.get("SGLANG_KERNEL_API_LOGLEVEL", "0")) == 0: - return func - except Exception: - return func - - try: - from sglang.kernel_api_logging import debug_kernel_api - except Exception: - return func - - if getattr(func, "_debug_kernel_wrapped", False): - return func - - wrapped = debug_kernel_api(func, op_name=op_name) - setattr(wrapped, "_debug_kernel_wrapped", True) - return cast(F, wrapped) - - -@overload -def maybe_wrap_jit_kernel_debug(func: F) -> F: ... - - -@overload -def maybe_wrap_jit_kernel_debug(func: F, op_name: str) -> F: ... - - -@overload -def maybe_wrap_jit_kernel_debug(*, op_name: str | None = None) -> Callable[[F], F]: ... - - -def maybe_wrap_jit_kernel_debug( - func: F | None = None, op_name: str | None = None -) -> F | Callable[[F], F]: - if func is None: - return lambda wrapped_func: _wrap_jit_kernel_debug(wrapped_func, op_name) - - return _wrap_jit_kernel_debug(func, op_name) diff --git a/python/sglang/jit_kernel/diffusion/triton/rmsnorm_onepass.py b/python/sglang/jit_kernel/diffusion/triton/rmsnorm_onepass.py index 8c776d7a9..801027a11 100644 --- a/python/sglang/jit_kernel/diffusion/triton/rmsnorm_onepass.py +++ b/python/sglang/jit_kernel/diffusion/triton/rmsnorm_onepass.py @@ -2,7 +2,7 @@ import torch import triton # type: ignore import triton.language as tl # type: ignore -from sglang.jit_kernel.debug_utils import maybe_wrap_jit_kernel_debug +from sglang.kernel_api_logging import debug_kernel_api from sglang.multimodal_gen.runtime.platforms import current_platform from sglang.srt.utils.custom_op import register_custom_op @@ -37,7 +37,6 @@ def _rms_norm_tiled_onepass( tl.store(y_blk, x * rstd * w, mask=mask) -@maybe_wrap_jit_kernel_debug @register_custom_op(op_name="triton_one_pass_rms_norm_cuda", out_shape="x") def _triton_one_pass_rms_norm_cuda( x: torch.Tensor, w: torch.Tensor, eps: float = 1e-6 @@ -73,6 +72,6 @@ def triton_one_pass_rms_norm(x: torch.Tensor, w: torch.Tensor, eps: float = 1e-6 if current_platform.is_mps(): from .mps_fallback import triton_one_pass_rms_norm_native - @maybe_wrap_jit_kernel_debug + @debug_kernel_api def triton_one_pass_rms_norm(x: torch.Tensor, w: torch.Tensor, eps: float = 1e-6): return triton_one_pass_rms_norm_native(x, w, eps) diff --git a/python/sglang/jit_kernel/flash_attention_v4.py b/python/sglang/jit_kernel/flash_attention_v4.py index e889cdda3..dcd5f2334 100644 --- a/python/sglang/jit_kernel/flash_attention_v4.py +++ b/python/sglang/jit_kernel/flash_attention_v4.py @@ -4,7 +4,7 @@ from typing import Callable, Optional, Tuple, Union import torch -from sglang.jit_kernel.debug_utils import maybe_wrap_jit_kernel_debug +from sglang.kernel_api_logging import debug_kernel_api try: from flash_attn.cute import flash_attn_varlen_func as _flash_attn_varlen_func @@ -19,7 +19,7 @@ def _maybe_contiguous(x: Optional[torch.Tensor]) -> Optional[torch.Tensor]: return x.contiguous() if x is not None and x.stride(-1) != 1 else x -@maybe_wrap_jit_kernel_debug +@debug_kernel_api def flash_attn_varlen_func( q: torch.Tensor, k: torch.Tensor, @@ -92,7 +92,7 @@ def flash_attn_varlen_func( return result -@maybe_wrap_jit_kernel_debug +@debug_kernel_api def flash_attn_with_kvcache( q: torch.Tensor, k_cache: torch.Tensor, diff --git a/python/sglang/jit_kernel/fused_store_index_cache.py b/python/sglang/jit_kernel/fused_store_index_cache.py index b1d9af897..dc50e21b5 100644 --- a/python/sglang/jit_kernel/fused_store_index_cache.py +++ b/python/sglang/jit_kernel/fused_store_index_cache.py @@ -13,13 +13,13 @@ from typing import TYPE_CHECKING import torch -from sglang.jit_kernel.debug_utils import maybe_wrap_jit_kernel_debug from sglang.jit_kernel.utils import ( cache_once, is_arch_support_pdl, load_jit, make_cpp_args, ) +from sglang.kernel_api_logging import debug_kernel_api if TYPE_CHECKING: from tvm_ffi.module import Module @@ -65,7 +65,7 @@ def can_use_nsa_fused_store( return False -@maybe_wrap_jit_kernel_debug +@debug_kernel_api def fused_store_index_k_cache( key: torch.Tensor, index_k_with_scale: torch.Tensor, diff --git a/python/sglang/jit_kernel/gptq_marlin.py b/python/sglang/jit_kernel/gptq_marlin.py index 8add63f5f..d3bde5336 100644 --- a/python/sglang/jit_kernel/gptq_marlin.py +++ b/python/sglang/jit_kernel/gptq_marlin.py @@ -4,8 +4,8 @@ from typing import TYPE_CHECKING, Optional import torch -from sglang.jit_kernel.debug_utils import maybe_wrap_jit_kernel_debug from sglang.jit_kernel.utils import cache_once, load_jit, make_cpp_args +from sglang.kernel_api_logging import debug_kernel_api if TYPE_CHECKING: from sgl_kernel.scalar_type import ScalarType @@ -32,7 +32,7 @@ def _or_empty( return t if t is not None else torch.empty(0, device=device, dtype=dtype) -@maybe_wrap_jit_kernel_debug +@debug_kernel_api def gptq_marlin_gemm( a: torch.Tensor, c: Optional[torch.Tensor], diff --git a/python/sglang/jit_kernel/gptq_marlin_repack.py b/python/sglang/jit_kernel/gptq_marlin_repack.py index 8f7cafc63..ea7fe9908 100644 --- a/python/sglang/jit_kernel/gptq_marlin_repack.py +++ b/python/sglang/jit_kernel/gptq_marlin_repack.py @@ -4,8 +4,8 @@ from typing import TYPE_CHECKING import torch -from sglang.jit_kernel.debug_utils import maybe_wrap_jit_kernel_debug from sglang.jit_kernel.utils import cache_once, load_jit +from sglang.kernel_api_logging import debug_kernel_api if TYPE_CHECKING: from tvm_ffi.module import Module @@ -23,7 +23,7 @@ def _jit_gptq_marlin_repack_module() -> Module: ) -@maybe_wrap_jit_kernel_debug +@debug_kernel_api def gptq_marlin_repack( b_q_weight: torch.Tensor, perm: torch.Tensor, diff --git a/python/sglang/jit_kernel/hicache.py b/python/sglang/jit_kernel/hicache.py index 41861c377..7a3577901 100644 --- a/python/sglang/jit_kernel/hicache.py +++ b/python/sglang/jit_kernel/hicache.py @@ -3,8 +3,8 @@ from __future__ import annotations import logging from typing import TYPE_CHECKING -from sglang.jit_kernel.debug_utils import maybe_wrap_jit_kernel_debug from sglang.jit_kernel.utils import cache_once, load_jit, make_cpp_args +from sglang.kernel_api_logging import debug_kernel_api if TYPE_CHECKING: import torch @@ -67,7 +67,7 @@ def _default_unroll(element_size: int) -> int: return 1 -@maybe_wrap_jit_kernel_debug +@debug_kernel_api def transfer_hicache_one_layer( k_cache_dst: torch.Tensor, v_cache_dst: torch.Tensor, @@ -103,7 +103,7 @@ def transfer_hicache_one_layer( ) -@maybe_wrap_jit_kernel_debug +@debug_kernel_api def transfer_hicache_all_layer( k_ptr_dst: torch.Tensor, v_ptr_dst: torch.Tensor, diff --git a/python/sglang/jit_kernel/kvcache.py b/python/sglang/jit_kernel/kvcache.py index 46a14612b..542d1866e 100644 --- a/python/sglang/jit_kernel/kvcache.py +++ b/python/sglang/jit_kernel/kvcache.py @@ -11,6 +11,7 @@ from sglang.jit_kernel.utils import ( load_jit, make_cpp_args, ) +from sglang.srt.utils.custom_op import register_custom_op if TYPE_CHECKING: from tvm_ffi.module import Module @@ -46,6 +47,7 @@ def can_use_store_cache(size: int) -> bool: return False +@register_custom_op(mutates_args=["k_cache", "v_cache"]) def store_cache( k: torch.Tensor, v: torch.Tensor, diff --git a/python/sglang/jit_kernel/moe_wna16_marlin.py b/python/sglang/jit_kernel/moe_wna16_marlin.py index 76b2c90dc..e9a8cd253 100644 --- a/python/sglang/jit_kernel/moe_wna16_marlin.py +++ b/python/sglang/jit_kernel/moe_wna16_marlin.py @@ -4,8 +4,8 @@ from typing import TYPE_CHECKING, Optional import torch -from sglang.jit_kernel.debug_utils import maybe_wrap_jit_kernel_debug from sglang.jit_kernel.utils import cache_once, load_jit, make_cpp_args +from sglang.kernel_api_logging import debug_kernel_api if TYPE_CHECKING: from sgl_kernel.scalar_type import ScalarType @@ -37,7 +37,7 @@ def _or_empty( return t if t is not None else torch.empty(0, device=device, dtype=dtype) -@maybe_wrap_jit_kernel_debug +@debug_kernel_api def moe_wna16_marlin_gemm( a: torch.Tensor, c_or_none: Optional[torch.Tensor], diff --git a/python/sglang/jit_kernel/ngram_embedding.py b/python/sglang/jit_kernel/ngram_embedding.py index 600880f12..f07937787 100644 --- a/python/sglang/jit_kernel/ngram_embedding.py +++ b/python/sglang/jit_kernel/ngram_embedding.py @@ -2,8 +2,8 @@ from __future__ import annotations from typing import TYPE_CHECKING -from sglang.jit_kernel.debug_utils import maybe_wrap_jit_kernel_debug from sglang.jit_kernel.utils import cache_once, load_jit +from sglang.kernel_api_logging import debug_kernel_api if TYPE_CHECKING: import torch @@ -22,7 +22,7 @@ def _jit_ngram_embedding_module() -> Module: ) -@maybe_wrap_jit_kernel_debug +@debug_kernel_api def compute_n_gram_ids( ne_n: int, ne_k: int, @@ -68,7 +68,7 @@ def compute_n_gram_ids( ) -@maybe_wrap_jit_kernel_debug +@debug_kernel_api def update_token_table( tokens: torch.Tensor, ne_token_table: torch.Tensor, diff --git a/python/sglang/jit_kernel/norm.py b/python/sglang/jit_kernel/norm.py index 4aef33c20..606358dd1 100644 --- a/python/sglang/jit_kernel/norm.py +++ b/python/sglang/jit_kernel/norm.py @@ -5,13 +5,13 @@ from typing import TYPE_CHECKING, Optional import torch -from sglang.jit_kernel.debug_utils import maybe_wrap_jit_kernel_debug from sglang.jit_kernel.utils import ( cache_once, is_arch_support_pdl, load_jit, make_cpp_args, ) +from sglang.kernel_api_logging import debug_kernel_api if TYPE_CHECKING: from tvm_ffi.module import Module @@ -99,7 +99,7 @@ def can_use_fused_inplace_qknorm(head_dim: int, dtype: torch.dtype) -> bool: return False -@maybe_wrap_jit_kernel_debug +@debug_kernel_api def fused_inplace_qknorm( q: torch.Tensor, k: torch.Tensor, @@ -114,7 +114,7 @@ def fused_inplace_qknorm( module.qknorm(q, k, q_weight, k_weight, eps) -@maybe_wrap_jit_kernel_debug +@debug_kernel_api def rmsnorm( input: torch.Tensor, weight: torch.Tensor, @@ -133,7 +133,7 @@ def rmsnorm( module.rmsnorm(input, weight, output, eps) -@maybe_wrap_jit_kernel_debug +@debug_kernel_api def fused_add_rmsnorm( input: torch.Tensor, residual: torch.Tensor, @@ -144,7 +144,7 @@ def fused_add_rmsnorm( module.fused_add_rmsnorm(input, residual, weight, eps) -@maybe_wrap_jit_kernel_debug +@debug_kernel_api def fused_inplace_qknorm_across_heads( q: torch.Tensor, k: torch.Tensor, diff --git a/python/sglang/jit_kernel/nvfp4.py b/python/sglang/jit_kernel/nvfp4.py index 7a061a9cc..6ca86c47e 100644 --- a/python/sglang/jit_kernel/nvfp4.py +++ b/python/sglang/jit_kernel/nvfp4.py @@ -8,8 +8,8 @@ from typing import TYPE_CHECKING, Optional, Tuple import torch -from sglang.jit_kernel.debug_utils import maybe_wrap_jit_kernel_debug from sglang.jit_kernel.utils import cache_once, load_jit +from sglang.kernel_api_logging import debug_kernel_api from sglang.srt.utils.custom_op import register_custom_op if TYPE_CHECKING: @@ -196,7 +196,7 @@ def _jit_nvfp4_blockwise_moe_module() -> Module: ) -@maybe_wrap_jit_kernel_debug +@debug_kernel_api def cutlass_scaled_fp4_mm( a: torch.Tensor, b: torch.Tensor, @@ -213,7 +213,7 @@ def cutlass_scaled_fp4_mm( return out -@maybe_wrap_jit_kernel_debug +@debug_kernel_api def cutlass_fp4_group_mm( a_fp4: torch.Tensor, b_fp4: torch.Tensor, @@ -293,7 +293,7 @@ def _scaled_fp4_quant_custom_op( module.scaled_fp4_quant(output, input, output_scale, input_global_scale) -@maybe_wrap_jit_kernel_debug +@debug_kernel_api def scaled_fp4_quant( input: torch.Tensor, input_global_scale: torch.Tensor ) -> Tuple[torch.Tensor, torch.Tensor]: @@ -363,7 +363,7 @@ def _scaled_fp4_experts_quant_custom_op( ) -@maybe_wrap_jit_kernel_debug +@debug_kernel_api def scaled_fp4_experts_quant( input_tensor: torch.Tensor, input_global_scale: torch.Tensor, @@ -448,7 +448,7 @@ def _scaled_fp4_grouped_quant_custom_op( ) -@maybe_wrap_jit_kernel_debug +@debug_kernel_api def scaled_fp4_grouped_quant( input_tensor: torch.Tensor, input_global_scale: torch.Tensor, @@ -509,7 +509,7 @@ def _silu_and_mul_scaled_fp4_grouped_quant_custom_op( ) -@maybe_wrap_jit_kernel_debug +@debug_kernel_api def silu_and_mul_scaled_fp4_grouped_quant( input_tensor: torch.Tensor, input_global_scale: torch.Tensor, diff --git a/python/sglang/jit_kernel/per_tensor_quant_fp8.py b/python/sglang/jit_kernel/per_tensor_quant_fp8.py index 89cafb7ce..9225aa45d 100644 --- a/python/sglang/jit_kernel/per_tensor_quant_fp8.py +++ b/python/sglang/jit_kernel/per_tensor_quant_fp8.py @@ -4,7 +4,6 @@ from typing import TYPE_CHECKING import torch -from sglang.jit_kernel.debug_utils import maybe_wrap_jit_kernel_debug from sglang.jit_kernel.utils import cache_once, load_jit, make_cpp_args from sglang.srt.utils.custom_op import register_custom_op @@ -23,7 +22,6 @@ def _jit_per_tensor_quant_fp8_module(is_static: bool, dtype: torch.dtype) -> Mod ) -@maybe_wrap_jit_kernel_debug @register_custom_op( op_name="per_tensor_quant_fp8", mutates_args=["output_q", "output_s"], diff --git a/python/sglang/jit_kernel/per_token_group_quant_8bit.py b/python/sglang/jit_kernel/per_token_group_quant_8bit.py index 529bb6a99..6df31c520 100644 --- a/python/sglang/jit_kernel/per_token_group_quant_8bit.py +++ b/python/sglang/jit_kernel/per_token_group_quant_8bit.py @@ -4,8 +4,8 @@ from typing import TYPE_CHECKING import torch -from sglang.jit_kernel.debug_utils import maybe_wrap_jit_kernel_debug from sglang.jit_kernel.utils import cache_once, load_jit, make_cpp_args +from sglang.kernel_api_logging import debug_kernel_api from sglang.srt.utils.custom_op import register_custom_op if TYPE_CHECKING: @@ -73,7 +73,7 @@ def _per_token_group_quant_8bit_custom_op( return None -@maybe_wrap_jit_kernel_debug +@debug_kernel_api def per_token_group_quant_8bit( input: torch.Tensor, output_q: torch.Tensor, diff --git a/python/sglang/jit_kernel/rope.py b/python/sglang/jit_kernel/rope.py index ff9470027..d9cbe0a8b 100644 --- a/python/sglang/jit_kernel/rope.py +++ b/python/sglang/jit_kernel/rope.py @@ -5,7 +5,6 @@ from typing import TYPE_CHECKING, Optional import torch -from sglang.jit_kernel.debug_utils import maybe_wrap_jit_kernel_debug from sglang.jit_kernel.utils import ( cache_once, is_arch_support_pdl, @@ -177,7 +176,6 @@ def apply_rope_inplace_with_kvcache( # NOTE: this name is intentionally set as the old kernel in `sgl_kernel` -@maybe_wrap_jit_kernel_debug def apply_rope_with_cos_sin_cache_inplace( q: torch.Tensor, k: torch.Tensor, diff --git a/python/sglang/jit_kernel/timestep_embedding.py b/python/sglang/jit_kernel/timestep_embedding.py index c0213c5cd..c65c145fa 100644 --- a/python/sglang/jit_kernel/timestep_embedding.py +++ b/python/sglang/jit_kernel/timestep_embedding.py @@ -4,8 +4,8 @@ from typing import TYPE_CHECKING import torch -from sglang.jit_kernel.debug_utils import maybe_wrap_jit_kernel_debug from sglang.jit_kernel.utils import cache_once, load_jit, make_cpp_args +from sglang.kernel_api_logging import debug_kernel_api if TYPE_CHECKING: from tvm_ffi.module import Module @@ -22,7 +22,7 @@ def _jit_timestep_embedding_module(dtype: torch.dtype) -> Module: ) -@maybe_wrap_jit_kernel_debug +@debug_kernel_api def timestep_embedding( t: torch.Tensor, dim: int, diff --git a/python/sglang/kernel_api_logging.py b/python/sglang/kernel_api_logging.py index 661f1f02c..b2dae9c82 100644 --- a/python/sglang/kernel_api_logging.py +++ b/python/sglang/kernel_api_logging.py @@ -15,38 +15,51 @@ import os import sys from datetime import datetime from pathlib import Path -from typing import Any, Callable +from typing import Any, Callable, TypeVar, overload import torch +_logger = logging.getLogger("sglang.kernel_api") -def _substitute_process_id(path: str) -> str: +_T = TypeVar("_T") +_F = TypeVar("_F", bound=Callable[..., Any]) + + +def _str_with_pid(path: str) -> str: if "%i" in path: return path.replace("%i", str(os.getpid())) return path -_KERNEL_API_LOG_LEVEL = int(os.environ.get("SGLANG_KERNEL_API_LOGLEVEL", "0")) -_KERNEL_API_LOG_DEST = _substitute_process_id( - os.environ.get("SGLANG_KERNEL_API_LOGDEST", "stdout") -) +def _get_env(key: str, type: Callable[..., _T], default: _T) -> _T: + value_str = os.environ.get(key, None) + if value_str is None: + return default + try: + return type(value_str) + except Exception: + _logger.warning( + "Failed to parse environment variable %s=%r as %s, using default %r", + key, + value_str, + type.__name__, + default, + ) + return default + + +def _parse_pattern(value: str) -> list[str]: + return [p.strip() for p in value.split(",") if p.strip()] + + +_KERNEL_API_LOG_LEVEL = _get_env("SGLANG_KERNEL_API_LOGLEVEL", int, 0) +_KERNEL_API_LOG_DEST = _get_env("SGLANG_KERNEL_API_LOGDEST", _str_with_pid, "stdout") _DUMP_DIR = Path( - _substitute_process_id( - os.environ.get("SGLANG_KERNEL_API_DUMP_DIR", "sglang_kernel_api_dumps") - ) + _get_env("SGLANG_KERNEL_API_DUMP_DIR", _str_with_pid, "sglang_kernel_api_dumps") ) -_DUMP_INCLUDE_PATTERNS = [ - p.strip() - for p in os.environ.get("SGLANG_KERNEL_API_DUMP_INCLUDE", "").split(",") - if p.strip() -] -_DUMP_EXCLUDE_PATTERNS = [ - p.strip() - for p in os.environ.get("SGLANG_KERNEL_API_DUMP_EXCLUDE", "").split(",") - if p.strip() -] +_DUMP_INCLUDE_PATTERNS = _get_env("SGLANG_KERNEL_API_DUMP_INCLUDE", _parse_pattern, []) +_DUMP_EXCLUDE_PATTERNS = _get_env("SGLANG_KERNEL_API_DUMP_EXCLUDE", _parse_pattern, []) -_logger = logging.getLogger("sglang.kernel_api") _dump_call_counter: dict[str, int] = {} @@ -371,17 +384,36 @@ def _infer_func_name(func: Callable) -> str: return qualname +@overload +def debug_kernel_api( + func: _F, + *, + op_name: str | None = None, +) -> _F: ... + + +@overload +def debug_kernel_api( + *, + op_name: str | None = None, +) -> Callable[[_F], _F]: ... + + def debug_kernel_api( func: Callable | None = None, *, op_name: str | None = None, ) -> Callable: + # NOTE: avoid any overhead in the hot path when logging is disabled if _KERNEL_API_LOG_LEVEL == 0: if func is None: return lambda f: f return func def decorator(f: Callable) -> Callable: + if hasattr(f, "_debug_kernel_wrapped"): + return f + @functools.wraps(f) def wrapper(*args: Any, **kwargs: Any) -> Any: if _is_compiling(): @@ -434,18 +466,29 @@ def debug_kernel_api( _log_section("Output:", {"return": result}) return result + setattr(wrapper, "_debug_kernel_wrapped", True) return wrapper - if func is None: - return decorator - return decorator(func) + return decorator if func is None else decorator(func) -def debug_torch_op(op_name: str, *, namespace: str = "sglang") -> Callable: - def call(*args: Any, **kwargs: Any) -> Any: - return getattr(getattr(torch.ops, namespace), op_name)(*args, **kwargs) - - return debug_kernel_api(call, op_name=f"{namespace}.custom_op.{op_name}") +def debug_torch_op( + op_func: Callable, + op_name: str, + *, + namespace: str = "sglang", +) -> Callable: + """NOTE: For internal use. Prefer `debug_kernel_api` for general use cases.""" + # NOTE: avoid any overhead in the hot path when logging is disabled + impl = getattr(getattr(torch.ops, namespace), op_name) + if _KERNEL_API_LOG_LEVEL == 0: + return impl + # NOTE: propagate the marker to avoid double-wrapping + if hasattr(op_func, "_debug_kernel_wrapped"): + setattr(impl, "_debug_kernel_wrapped", True) + return impl + # NOTE: redirect the function name + return debug_kernel_api(impl, op_name=_infer_func_name(op_func)) def wrap_method_with_debug_kernel_once( @@ -455,6 +498,10 @@ def wrap_method_with_debug_kernel_once( op_name: str, marker_attr: str | None = None, ) -> Any: + # NOTE: avoid any overhead in the hot path when logging is disabled + if _KERNEL_API_LOG_LEVEL == 0: + return obj + if marker_attr is None: marker_attr = f"_debug_kernel_{method_name}_wrapped" diff --git a/python/sglang/multimodal_gen/runtime/layers/utils.py b/python/sglang/multimodal_gen/runtime/layers/utils.py index a6c7d3b28..1feeb3f36 100644 --- a/python/sglang/multimodal_gen/runtime/layers/utils.py +++ b/python/sglang/multimodal_gen/runtime/layers/utils.py @@ -156,7 +156,7 @@ class CustomOpWrapper: mutates_args=self.mutates_args, fake_impl=self.fake_impl, ) - self._impl = debug_torch_op(self.op_name) + self._impl = debug_torch_op(self.op_func, self.op_name) assert self._impl is not None return self._impl diff --git a/python/sglang/srt/layers/moe/moe_runner/flashinfer_trtllm.py b/python/sglang/srt/layers/moe/moe_runner/flashinfer_trtllm.py index 7914265dd..68decf875 100644 --- a/python/sglang/srt/layers/moe/moe_runner/flashinfer_trtllm.py +++ b/python/sglang/srt/layers/moe/moe_runner/flashinfer_trtllm.py @@ -8,13 +8,16 @@ from torch.nn import Module from torch.nn.parameter import Parameter # Import to register custom ops for torch.compile compatibility -import sglang.srt.layers.moe.flashinfer_trtllm_moe # noqa: F401 -from sglang.kernel_api_logging import debug_torch_op from sglang.srt.distributed import get_tp_group from sglang.srt.distributed.device_communicators.pynccl_allocator import ( use_symmetric_memory, ) from sglang.srt.layers.dp_attention import is_allocation_symmetric +from sglang.srt.layers.moe.flashinfer_trtllm_moe import ( + trtllm_fp8_block_scale_moe_wrapper, + trtllm_fp8_block_scale_routed_moe_wrapper, + trtllm_fp8_per_tensor_scale_moe_wrapper, +) from sglang.srt.layers.moe.moe_runner.base import ( MoeQuantInfo, MoeRunnerConfig, @@ -45,16 +48,6 @@ elif is_cuda_alike(): else: fp4_quantize = None -_trtllm_fp8_block_scale_routed_moe_wrapper = debug_torch_op( - "trtllm_fp8_block_scale_routed_moe_wrapper" -) -_trtllm_fp8_block_scale_moe_wrapper = debug_torch_op( - "trtllm_fp8_block_scale_moe_wrapper" -) -_trtllm_fp8_per_tensor_scale_moe = debug_torch_op( - "trtllm_fp8_per_tensor_scale_moe_wrapper" -) - def align_fp8_moe_weights_for_flashinfer_trtllm( layer: Module, swap_w13_halves: bool = False @@ -386,7 +379,7 @@ def fused_experts_none_to_flashinfer_trtllm_fp8( topk_weights=topk_output.topk_weights, ) - output = _trtllm_fp8_block_scale_routed_moe_wrapper( + output = trtllm_fp8_block_scale_routed_moe_wrapper( topk_ids=packed_topk_ids, routing_bias=None, hidden_states=a_q, @@ -419,7 +412,7 @@ def fused_experts_none_to_flashinfer_trtllm_fp8( else: assert TopKOutputChecker.format_is_bypassed(topk_output) - output = _trtllm_fp8_block_scale_moe_wrapper( + output = trtllm_fp8_block_scale_moe_wrapper( routing_logits=( router_logits.to(torch.float32) if routing_method_type == RoutingMethodType.DeepSeekV3 @@ -476,7 +469,7 @@ def fused_experts_none_to_flashinfer_trtllm_fp8( # Move kernel call outside context manager to avoid graph breaks # during torch.compile for piecewise cuda graph. # Use custom op wrapper for torch.compile compatibility. - output = _trtllm_fp8_per_tensor_scale_moe( + output = trtllm_fp8_per_tensor_scale_moe_wrapper( routing_logits=router_logits.to(torch.bfloat16), routing_bias=routing_bias_cast, hidden_states=a_q, diff --git a/python/sglang/srt/layers/quantization/bitsandbytes.py b/python/sglang/srt/layers/quantization/bitsandbytes.py index 3ee6da386..51b92624f 100644 --- a/python/sglang/srt/layers/quantization/bitsandbytes.py +++ b/python/sglang/srt/layers/quantization/bitsandbytes.py @@ -7,7 +7,6 @@ from typing import TYPE_CHECKING, Any, Optional import torch from packaging import version -from sglang.kernel_api_logging import debug_torch_op from sglang.srt.layers.linear import LinearBase from sglang.srt.layers.quantization.base_config import ( FusedMoEMethodBase, @@ -16,7 +15,8 @@ from sglang.srt.layers.quantization.base_config import ( QuantizeMethodBase, ) from sglang.srt.layers.quantization.unquant import UnquantizedLinearMethod -from sglang.srt.utils import direct_register_custom_op, set_weight_attrs +from sglang.srt.utils import set_weight_attrs +from sglang.srt.utils.custom_op import register_custom_op if TYPE_CHECKING: from sglang.srt.layers.moe.token_dispatcher import ( @@ -393,7 +393,8 @@ class BitsAndBytesLinearMethod(LinearMethodBase): return out -def _apply_bnb_4bit( +@register_custom_op(mutates_args=["out"]) +def apply_bnb_4bit( x: torch.Tensor, weight: torch.Tensor, offsets: torch.Tensor, @@ -416,28 +417,6 @@ def _apply_bnb_4bit( current_index += output_size -def _apply_bnb_4bit_fake( - x: torch.Tensor, - weight: torch.Tensor, - offsets: torch.Tensor, - out: torch.Tensor, -) -> None: - return - - -try: - direct_register_custom_op( - op_name="apply_bnb_4bit", - op_func=_apply_bnb_4bit, - mutates_args=["out"], - fake_impl=_apply_bnb_4bit_fake, - ) - apply_bnb_4bit = debug_torch_op("apply_bnb_4bit") - -except AttributeError as error: - raise error - - class BitsAndBytesMoEMethod(FusedMoEMethodBase): """MoE method for BitsAndBytes. diff --git a/python/sglang/srt/layers/quantization/fp8.py b/python/sglang/srt/layers/quantization/fp8.py index 987162922..234af35f5 100644 --- a/python/sglang/srt/layers/quantization/fp8.py +++ b/python/sglang/srt/layers/quantization/fp8.py @@ -10,7 +10,6 @@ import torch.nn.functional as F from torch.nn import Module from torch.nn.parameter import Parameter -from sglang.kernel_api_logging import debug_torch_op from sglang.srt.distributed import get_tensor_model_parallel_world_size, get_tp_group from sglang.srt.distributed.device_communicators.pynccl_allocator import ( use_symmetric_memory, @@ -58,7 +57,10 @@ from sglang.srt.layers.quantization.fp8_utils import ( requant_weight_ue8m0_inplace, ) from sglang.srt.layers.quantization.kv_cache import BaseKVCacheMethod -from sglang.srt.layers.quantization.marlin_utils_fp8 import prepare_fp8_layer_for_marlin +from sglang.srt.layers.quantization.marlin_utils_fp8 import ( + apply_fp8_marlin_linear, + prepare_fp8_layer_for_marlin, +) from sglang.srt.layers.quantization.unquant import ( UnquantizedFusedMoEMethod, UnquantizedLinearMethod, @@ -111,8 +113,6 @@ ACTIVATION_SCHEMES = ["static", "dynamic"] logger = logging.getLogger(__name__) -_apply_fp8_marlin_linear = debug_torch_op("apply_fp8_marlin_linear") - class Fp8Config(QuantizationConfig): """Config class for FP8.""" @@ -646,7 +646,7 @@ class Fp8LinearMethod(LinearMethodBase): bias: Optional[torch.Tensor] = None, ) -> torch.Tensor: if self.use_marlin: - return _apply_fp8_marlin_linear( + return apply_fp8_marlin_linear( input=x, weight=layer.weight, weight_scale=layer.weight_scale, diff --git a/python/sglang/srt/layers/quantization/marlin_utils_fp8.py b/python/sglang/srt/layers/quantization/marlin_utils_fp8.py index d5d008d5f..80afc12d4 100644 --- a/python/sglang/srt/layers/quantization/marlin_utils_fp8.py +++ b/python/sglang/srt/layers/quantization/marlin_utils_fp8.py @@ -13,7 +13,8 @@ from sglang.srt.layers.quantization.marlin_utils import ( should_use_atomic_add_reduce, ) from sglang.srt.layers.quantization.utils import get_scalar_types -from sglang.srt.utils import direct_register_custom_op, is_cuda +from sglang.srt.utils import is_cuda +from sglang.srt.utils.custom_op import register_custom_op _is_cuda = is_cuda() if _is_cuda: @@ -39,6 +40,22 @@ def fp8_fused_exponent_bias_into_scales(scales): return scales * s +def fake_apply_fp8_marlin_linear( + input: torch.Tensor, + weight: torch.Tensor, + weight_scale: torch.Tensor, + workspace: torch.Tensor, + size_n: int, + size_k: int, + bias: Optional[torch.Tensor], + use_fp32_reduce: bool = USE_FP32_REDUCE_DEFAULT, +) -> torch.Tensor: + out_shape = input.shape[:-1] + (size_n,) + fake_output = torch.empty(out_shape, dtype=input.dtype, device=input.device) + return fake_output + + +@register_custom_op(fake_impl=fake_apply_fp8_marlin_linear) def apply_fp8_marlin_linear( input: torch.Tensor, weight: torch.Tensor, @@ -83,30 +100,6 @@ def apply_fp8_marlin_linear( return output.reshape(out_shape) -def fake_apply_fp8_marlin_linear( - input: torch.Tensor, - weight: torch.Tensor, - weight_scale: torch.Tensor, - workspace: torch.Tensor, - size_n: int, - size_k: int, - bias: Optional[torch.Tensor], - use_fp32_reduce: bool = USE_FP32_REDUCE_DEFAULT, -) -> torch.Tensor: - - out_shape = input.shape[:-1] + (size_n,) - fake_output = torch.empty(out_shape, dtype=input.dtype, device=input.device) - return fake_output - - -direct_register_custom_op( - op_name="apply_fp8_marlin_linear", - op_func=apply_fp8_marlin_linear, - mutates_args=[], - fake_impl=fake_apply_fp8_marlin_linear, -) - - def prepare_fp8_layer_for_marlin( layer: torch.nn.Module, size_k_first: bool = True ) -> None: diff --git a/python/sglang/srt/mem_cache/memory_pool.py b/python/sglang/srt/mem_cache/memory_pool.py index 0e2f63cb4..881f3cad7 100644 --- a/python/sglang/srt/mem_cache/memory_pool.py +++ b/python/sglang/srt/mem_cache/memory_pool.py @@ -37,7 +37,6 @@ import triton import triton.language as tl from sglang.jit_kernel.kvcache import can_use_store_cache, store_cache -from sglang.kernel_api_logging import debug_kernel_api from sglang.srt.configs.mamba_utils import BaseLinearStateParams from sglang.srt.constants import GPU_MEMORY_TYPE_KV_CACHE from sglang.srt.environ import envs @@ -61,14 +60,8 @@ from sglang.srt.utils import ( is_npu, next_power_of_2, ) -from sglang.srt.utils.custom_op import register_custom_op from sglang.srt.utils.torch_memory_saver_adapter import TorchMemorySaverAdapter -store_cache = register_custom_op( - debug_kernel_api(store_cache, op_name="jit_kernel.kvcache.store_cache"), - mutates_args=["k_cache", "v_cache"], -) - if TYPE_CHECKING: from sglang.srt.managers.cache_controller import LayerDoneCounter from sglang.srt.managers.schedule_batch import Req diff --git a/python/sglang/srt/utils/common.py b/python/sglang/srt/utils/common.py index a4a90b5b1..57d85eb2f 100644 --- a/python/sglang/srt/utils/common.py +++ b/python/sglang/srt/utils/common.py @@ -1899,8 +1899,11 @@ def direct_register_custom_op( mutates_args: List[str], fake_impl: Optional[Callable] = None, target_lib: Optional[Library] = None, -): +) -> None: """ + NOTE: Please try to use `register_custom_op` instead of this function. + See `python/sglang/srt/utils/custom_op.py` for details. + `torch.library.custom_op` can have significant overhead because it needs to consider complicated dispatching logic. This function directly registers a custom op and dispatches it to the CUDA backend. diff --git a/python/sglang/srt/utils/custom_op.py b/python/sglang/srt/utils/custom_op.py index cce713c60..720776501 100644 --- a/python/sglang/srt/utils/custom_op.py +++ b/python/sglang/srt/utils/custom_op.py @@ -161,7 +161,7 @@ class CustomOpWrapper: mutates_args=self.mutates_args, fake_impl=self.fake_impl, ) - self._impl = debug_torch_op(self.op_name) + self._impl = debug_torch_op(self.op_func, self.op_name) assert self._impl is not None return self._impl @@ -290,7 +290,7 @@ def register_custom_op_from_extern( wrapper.__name__ = fn.__name__ wrapper.__qualname__ = fn.__qualname__ wrapper.__module__ = fn.__module__ - wrapper.__signature__ = new_sig + wrapper.__signature__ = new_sig # type: ignore[attr-defined] # Build annotations without computed args, preserving return type wrapper.__annotations__ = { k: v @@ -334,4 +334,4 @@ def register_custom_op_from_extern( fake_impl=fake_impl, ) - return debug_torch_op(name) + return debug_torch_op(fn, name)