Add SGLang CUDA crash API logging inspired by FlashInfer (#20910)

This commit is contained in:
Xiaoyu Zhang
2026-03-22 16:39:40 +08:00
committed by GitHub
parent bb737d7a82
commit 766d225fcc
46 changed files with 1585 additions and 19 deletions
@@ -5,6 +5,8 @@ import triton # type: ignore
import triton.language as tl # type: ignore
from torch import Tensor
from sglang.jit_kernel.debug_utils import maybe_wrap_jit_kernel_debug
# RMSNorm-fp32
def maybe_contiguous_lastdim(x):
@@ -450,6 +452,7 @@ class LayerNormFn:
return y
@maybe_wrap_jit_kernel_debug
def layer_norm_fn(
x,
weight,
@@ -537,6 +540,7 @@ def _norm_infer_kernel(
tl.store(Y + cols, y, mask=cols < N)
@maybe_wrap_jit_kernel_debug
def norm_infer(
x: Tensor,
weight: Optional[Tensor],
@@ -579,6 +583,7 @@ def norm_infer(
return out
@maybe_wrap_jit_kernel_debug
def rms_norm_fn(
x,
weight,
@@ -625,5 +630,53 @@ from sglang.multimodal_gen.runtime.platforms import current_platform
if current_platform.is_mps():
from .mps_fallback import norm_infer_native, rms_norm_fn_native
norm_infer = norm_infer_native
rms_norm_fn = rms_norm_fn_native
@maybe_wrap_jit_kernel_debug
def norm_infer(
x: Tensor,
weight: Optional[Tensor],
bias: Optional[Tensor],
eps: float,
is_rms_norm: bool = False,
out: Optional[Tensor] = None,
):
return norm_infer_native(x, weight, bias, eps, is_rms_norm, out)
@maybe_wrap_jit_kernel_debug
def rms_norm_fn(
x,
weight,
bias,
residual=None,
x1=None,
weight1=None,
bias1=None,
eps=1e-6,
dropout_p=0.0,
rowscale=None,
prenorm=False,
residual_in_fp32=False,
zero_centered_weight=False,
return_dropout_mask=False,
out_dtype=None,
out=None,
residual_out=None,
):
return rms_norm_fn_native(
x,
weight,
bias,
residual,
x1,
weight1,
bias1,
eps,
dropout_p,
rowscale,
prenorm,
residual_in_fp32,
zero_centered_weight,
return_dropout_mask,
out_dtype,
out,
residual_out,
)
@@ -2,6 +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.srt.utils.custom_op import register_custom_op
@@ -35,6 +36,7 @@ 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
@@ -72,4 +74,6 @@ from sglang.multimodal_gen.runtime.platforms import current_platform
if current_platform.is_mps():
from .mps_fallback import triton_one_pass_rms_norm_native
triton_one_pass_rms_norm = triton_one_pass_rms_norm_native
@maybe_wrap_jit_kernel_debug
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)
@@ -2,6 +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.multimodal_gen.runtime.platforms import current_platform
@@ -64,6 +65,7 @@ def _rotary_embedding_kernel(
tl.store(output_row_ptr + offsets_x2, o2_vals.to(x2_vals.dtype), mask=mask)
@maybe_wrap_jit_kernel_debug
def apply_rotary_embedding(
x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor, interleaved: bool = False
) -> torch.Tensor:
@@ -110,9 +112,24 @@ def apply_rotary_embedding(
if current_platform.is_npu():
from .npu_fallback import apply_rotary_embedding_native
apply_rotary_embedding = apply_rotary_embedding_native
@maybe_wrap_jit_kernel_debug
def apply_rotary_embedding(
x: torch.Tensor,
cos: torch.Tensor,
sin: torch.Tensor,
interleaved: bool = False,
) -> torch.Tensor:
return apply_rotary_embedding_native(x, cos, sin, interleaved)
if current_platform.is_mps():
from .mps_fallback import apply_rotary_embedding_native
apply_rotary_embedding = apply_rotary_embedding_native
@maybe_wrap_jit_kernel_debug
def apply_rotary_embedding(
x: torch.Tensor,
cos: torch.Tensor,
sin: torch.Tensor,
interleaved: bool = False,
) -> torch.Tensor:
return apply_rotary_embedding_native(x, cos, sin, interleaved)
@@ -2,6 +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.multimodal_gen.runtime.platforms import current_platform
@@ -444,6 +445,7 @@ def fuse_scale_shift_gate_select01_kernel_blc_opt(
tl.store(gate_out_ptr + go_off, gate, mask=mask)
@maybe_wrap_jit_kernel_debug
def fuse_scale_shift_kernel(
x: torch.Tensor,
scale: torch.Tensor,
@@ -563,6 +565,7 @@ def fuse_scale_shift_kernel(
return output
@maybe_wrap_jit_kernel_debug
def fuse_scale_shift_gate_select01_kernel(
x: torch.Tensor,
scale0: torch.Tensor,
@@ -635,6 +638,7 @@ def fuse_scale_shift_gate_select01_kernel(
return output, gate_out
@maybe_wrap_jit_kernel_debug
def fuse_layernorm_scale_shift_gate_select01_kernel(
x: torch.Tensor,
weight: torch.Tensor | None,
@@ -724,6 +728,7 @@ def fuse_layernorm_scale_shift_gate_select01_kernel(
return output, gate_out
@maybe_wrap_jit_kernel_debug
def fuse_residual_layernorm_scale_shift_gate_select01_kernel(
x: torch.Tensor,
residual: torch.Tensor,
@@ -834,7 +839,19 @@ def fuse_residual_layernorm_scale_shift_gate_select01_kernel(
if current_platform.is_npu():
from .npu_fallback import fuse_scale_shift_native
fuse_scale_shift_kernel = fuse_scale_shift_native
@maybe_wrap_jit_kernel_debug
def fuse_scale_shift_kernel(
x: torch.Tensor,
scale: torch.Tensor,
shift: torch.Tensor,
scale_constant: float = 1.0,
block_l: int = 128,
block_c: int = 128,
):
return fuse_scale_shift_native(
x, scale, shift, scale_constant, block_l, block_c
)
if current_platform.is_mps():
from .mps_fallback import (
@@ -842,5 +859,41 @@ if current_platform.is_mps():
fuse_scale_shift_kernel_native,
)
fuse_scale_shift_kernel = fuse_scale_shift_kernel_native
fuse_scale_shift_gate_select01_kernel = fuse_scale_shift_gate_select01_kernel_native
@maybe_wrap_jit_kernel_debug
def fuse_scale_shift_kernel(
x: torch.Tensor,
scale: torch.Tensor,
shift: torch.Tensor,
scale_constant: float = 1.0,
block_l: int = 128,
block_c: int = 128,
):
return fuse_scale_shift_kernel_native(
x, scale, shift, scale_constant, block_l, block_c
)
@maybe_wrap_jit_kernel_debug
def fuse_scale_shift_gate_select01_kernel(
x: torch.Tensor,
scale0: torch.Tensor,
shift0: torch.Tensor,
gate0: torch.Tensor,
scale1: torch.Tensor,
shift1: torch.Tensor,
gate1: torch.Tensor,
index: torch.Tensor,
block_l: int = 128,
block_c: int = 128,
):
return fuse_scale_shift_gate_select01_kernel_native(
x,
scale0,
shift0,
gate0,
scale1,
shift1,
gate1,
index,
block_l,
block_c,
)