[misc] clean up kernel API (#21325)
This commit is contained in:
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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],
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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],
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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"],
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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"
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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.
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user