Add SGLang CUDA crash API logging inspired by FlashInfer (#20910)
This commit is contained in:
@@ -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,
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user