[misc] clean up kernel API (#21325)

This commit is contained in:
DarkSharpness
2026-03-25 09:10:23 +08:00
committed by GitHub
parent 281fe10b5e
commit dfc15b78b0
27 changed files with 160 additions and 198 deletions
+2
View File
@@ -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,
-45
View File
@@ -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,
+2 -2
View File
@@ -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 -3
View File
@@ -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,
+2
View File
@@ -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,
+2 -2
View File
@@ -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],
+3 -3
View File
@@ -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 -5
View File
@@ -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,
+7 -7
View File
@@ -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,
-2
View File
@@ -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,
+74 -27
View File
@@ -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.
+5 -5
View File
@@ -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
+4 -1
View File
@@ -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.
+3 -3
View File
@@ -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)