[Kernel] Classification cleanup: unify _jit_ naming, drop empty/model groups, add elementwise (RFC #29630) (#32148)

Co-authored-by: Claude Opus 4.8 <noreply@anthropic.com>
This commit is contained in:
Xiaoyu Zhang
2026-07-23 13:47:02 +08:00
committed by GitHub
co-authored by Claude Opus 4.8
parent 1b63155efe
commit 11b0e5c5ad
102 changed files with 173 additions and 220 deletions
@@ -3,7 +3,7 @@ import itertools
import pytest
import torch
from sglang.kernels.ops.layernorm.elementwise import fused_gate_sigmoid_mul_add
from sglang.kernels.ops.elementwise.elementwise import fused_gate_sigmoid_mul_add
DTYPES = [torch.float16, torch.bfloat16]
TOKEN_COUNTS = [1, 2, 4, 8, 16, 64, 512, 1024, 2048, 4096, 8192]
+1 -1
View File
@@ -3,7 +3,7 @@ import itertools
import pytest
import torch
from sglang.kernels.ops.layernorm.elementwise import fused_sigmoid_mul
from sglang.kernels.ops.elementwise.elementwise import fused_sigmoid_mul
DTYPES = [torch.float16, torch.bfloat16]
TOKEN_COUNTS = [1, 2, 4, 8, 16, 64, 512, 1024, 2048, 4096, 8192]
@@ -6,16 +6,12 @@ from sgl_kernel import silu_and_mul as silu_and_mul_aot
from sglang.kernels.jit.benchmark import marker
from sglang.kernels.jit.benchmark.utils import create_random
from sglang.kernels.ops.activation._jit_activation import (
gelu_and_mul as gelu_and_mul_jit,
)
from sglang.kernels.ops.activation._jit_activation import (
from sglang.kernels.ops.activation.activation import gelu_and_mul as gelu_and_mul_jit
from sglang.kernels.ops.activation.activation import (
gelu_tanh_and_mul as gelu_tanh_and_mul_jit,
)
from sglang.kernels.ops.activation._jit_activation import relu2 as relu2_jit
from sglang.kernels.ops.activation._jit_activation import (
silu_and_mul as silu_and_mul_jit,
)
from sglang.kernels.ops.activation.activation import relu2 as relu2_jit
from sglang.kernels.ops.activation.activation import silu_and_mul as silu_and_mul_jit
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
register_cuda_ci(
@@ -7,7 +7,7 @@ from sglang.kernels.jit.benchmark.utils import (
get_benchmark_range,
run_benchmark_no_cudagraph,
)
from sglang.kernels.ops.attention.add_constant import (
from sglang.kernels.ops.elementwise.add_constant import (
_jit_add_constant_module,
add_constant,
)
@@ -13,7 +13,7 @@ from sglang.kernels.jit.benchmark.utils import (
get_benchmark_range,
run_benchmark,
)
from sglang.kernels.ops.attention.hadamard import hadamard_transform
from sglang.kernels.ops.quantization.hadamard import hadamard_transform
from sglang.test.ci.ci_register import register_cuda_ci
register_cuda_ci(
@@ -17,10 +17,8 @@ from sglang.kernels.jit.benchmark.utils import DEFAULT_DEVICE
from sglang.kernels.jit.utils import KERNEL_PATH
from sglang.kernels.ops.diffusion.triton.norm import norm_infer, rms_norm_fn
from sglang.kernels.ops.diffusion.triton.rmsnorm_onepass import triton_one_pass_rms_norm
from sglang.kernels.ops.layernorm._jit_norm import (
fused_add_rmsnorm as jit_fused_add_rmsnorm,
)
from sglang.kernels.ops.layernorm._jit_norm import rmsnorm as jit_rmsnorm
from sglang.kernels.ops.layernorm.norm import fused_add_rmsnorm as jit_fused_add_rmsnorm
from sglang.kernels.ops.layernorm.norm import rmsnorm as jit_rmsnorm
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
from sglang.utils import is_in_ci
@@ -131,7 +131,7 @@ def clone_inputs(
def split_qknorm_rope(inputs: dict[str, torch.Tensor | bool]) -> None:
from flashinfer.rope import apply_rope_with_cos_sin_cache_inplace
from sglang.kernels.ops.layernorm._jit_norm import fused_inplace_qknorm
from sglang.kernels.ops.layernorm.norm import fused_inplace_qknorm
q = inputs["q"]
k = inputs["k"]
@@ -10,10 +10,10 @@ import triton.testing
from sglang.kernels.jit.benchmark import marker
from sglang.kernels.jit.utils import get_jit_cuda_arch, is_hip_runtime
from sglang.kernels.ops.gemm._jit_dsv3_fused_a_gemm import dsv3_fused_a_gemm
from sglang.kernels.ops.gemm.cutedsl_dsv3_fused_a_gemm import (
dsv3_fused_a_gemm as cutedsl_dsv3_fused_a_gemm,
)
from sglang.kernels.ops.gemm.dsv3_fused_a_gemm import dsv3_fused_a_gemm
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.utils import is_in_ci
@@ -10,7 +10,7 @@ import torch.nn.functional as F
from sglang.kernels.jit.benchmark import marker
from sglang.kernels.jit.benchmark.utils import create_random
from sglang.kernels.jit.utils import get_jit_cuda_arch, is_hip_runtime
from sglang.kernels.ops.gemm._jit_dsv3_router_gemm import dsv3_router_gemm
from sglang.kernels.ops.gemm.dsv3_router_gemm import dsv3_router_gemm
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
register_cuda_ci(
@@ -21,9 +21,7 @@ from sglang.kernels.jit.benchmark.utils import (
get_benchmark_range,
)
from sglang.kernels.jit.utils import is_arch_support_pdl
from sglang.kernels.ops.kvcache._jit_set_mla_kv_buffer import (
set_mla_kv_buffer as jit_set,
)
from sglang.kernels.ops.kvcache.set_mla_kv_buffer import set_mla_kv_buffer as jit_set
from sglang.srt.mem_cache.utils import set_mla_kv_buffer_kernel as sglang_triton_kernel
from sglang.srt.mem_cache.utils import set_mla_kv_buffer_triton as sglang_wrapper
from sglang.test.ci.ci_register import register_cuda_ci
@@ -7,10 +7,8 @@ from flashinfer.norm import fused_add_rmsnorm as fi_fused_add_rmsnorm
from flashinfer.norm import rmsnorm as fi_rmsnorm
from sglang.kernels.jit.benchmark.utils import get_benchmark_range, run_benchmark
from sglang.kernels.ops.layernorm._jit_norm import (
fused_add_rmsnorm as jit_fused_add_rmsnorm,
)
from sglang.kernels.ops.layernorm._jit_norm import rmsnorm as jit_rmsnorm
from sglang.kernels.ops.layernorm.norm import fused_add_rmsnorm as jit_fused_add_rmsnorm
from sglang.kernels.ops.layernorm.norm import rmsnorm as jit_rmsnorm
from sglang.test.ci.ci_register import register_cuda_ci
register_cuda_ci(
@@ -2,7 +2,7 @@ import torch
from sglang.kernels.jit.benchmark import marker
from sglang.kernels.jit.benchmark.utils import create_random
from sglang.kernels.ops.layernorm._jit_norm import fused_inplace_qknorm
from sglang.kernels.ops.layernorm.norm import fused_inplace_qknorm
from sglang.srt.utils import get_current_device_stream_fast
from sglang.test.ci.ci_register import register_cuda_ci
@@ -7,7 +7,7 @@ import triton.testing
from sgl_kernel import rmsnorm
from sglang.kernels.jit.benchmark.utils import run_benchmark
from sglang.kernels.ops.layernorm._jit_norm import fused_inplace_qknorm_across_heads
from sglang.kernels.ops.layernorm.norm import fused_inplace_qknorm_across_heads
from sglang.srt.utils import get_current_device_stream_fast
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.utils import is_in_ci
@@ -5,7 +5,7 @@ import triton
import triton.testing
from sglang.kernels.jit.benchmark.utils import get_benchmark_range, run_benchmark
from sglang.kernels.ops.quantization._jit_per_tensor_quant_fp8 import (
from sglang.kernels.ops.quantization.per_tensor_quant_fp8 import (
per_tensor_quant_fp8,
)
from sglang.test.ci.ci_register import register_cuda_ci
@@ -1,20 +1,20 @@
from sglang.kernels.jit.benchmark import marker
from sglang.kernels.jit.benchmark.utils import create_empty, create_random
# per_token_group_quant_8bit_v2 is DEPRECATED (no production call sites); the
# kernel is kept only as the perf baseline for this benchmark.
from sglang.kernels.ops.quantization._jit_per_token_group_quant import (
per_token_group_quant,
)
from sglang.kernels.ops.quantization._jit_per_token_group_quant_8bit_v2 import (
per_token_group_quant_8bit_v2,
)
from sglang.kernels.ops.quantization.fp8_kernel import (
create_per_token_group_quant_fp8_output_scale,
fp8_dtype,
fp8_max,
fp8_min,
)
# per_token_group_quant_8bit_v2 is DEPRECATED (no production call sites); the
# kernel is kept only as the perf baseline for this benchmark.
from sglang.kernels.ops.quantization.per_token_group_quant import (
per_token_group_quant,
)
from sglang.kernels.ops.quantization.per_token_group_quant_8bit_v2 import (
per_token_group_quant_8bit_v2,
)
from sglang.test.ci.ci_register import register_cuda_ci
register_cuda_ci(
@@ -3,15 +3,15 @@ from sgl_kernel import sgl_per_token_group_quant_8bit
from sglang.kernels.jit.benchmark import marker
from sglang.kernels.jit.benchmark.utils import create_random
from sglang.kernels.ops.quantization._jit_per_token_group_quant_8bit_v2 import (
per_token_group_quant_8bit_v2,
)
from sglang.kernels.ops.quantization.fp8_kernel import (
create_per_token_group_quant_fp8_output_scale,
fp8_dtype,
fp8_max,
fp8_min,
)
from sglang.kernels.ops.quantization.per_token_group_quant_8bit_v2 import (
per_token_group_quant_8bit_v2,
)
from sglang.test.ci.ci_register import register_cuda_ci
register_cuda_ci(
@@ -4,21 +4,21 @@ import torch
from sglang.kernels.jit.benchmark import marker
from sglang.kernels.jit.benchmark.utils import create_empty, create_random
# per_token_group_quant_8bit_v2 is DEPRECATED (no production call sites); the
# kernel is kept only as the perf baseline for this benchmark.
from sglang.kernels.ops.quantization._jit_per_token_group_quant import (
per_token_group_quant,
)
from sglang.kernels.ops.quantization._jit_per_token_group_quant_8bit_v2 import (
per_token_group_quant_8bit_v2,
)
from sglang.kernels.ops.quantization.fp8_kernel import (
create_per_token_group_quant_fp8_output_scale,
fp8_dtype,
fp8_max,
fp8_min,
)
# per_token_group_quant_8bit_v2 is DEPRECATED (no production call sites); the
# kernel is kept only as the perf baseline for this benchmark.
from sglang.kernels.ops.quantization.per_token_group_quant import (
per_token_group_quant,
)
from sglang.kernels.ops.quantization.per_token_group_quant_8bit_v2 import (
per_token_group_quant_8bit_v2,
)
from sglang.test.ci.ci_register import register_cuda_ci
register_cuda_ci(
@@ -5,7 +5,7 @@ import torch
import torch.nn.functional as F
from sglang.kernels.jit.utils import get_ci_test_range
from sglang.kernels.ops.activation._jit_activation import (
from sglang.kernels.ops.activation.activation import (
SUPPORTED_ACTIVATIONS,
relu2,
run_activation,
@@ -3,7 +3,7 @@ import sys
import pytest
import torch
from sglang.kernels.ops.attention.add_constant import add_constant
from sglang.kernels.ops.elementwise.add_constant import add_constant
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
register_cuda_ci(est_time=45, stage="base-b-kernel-unit", runner_config="1-gpu-large")
@@ -27,7 +27,7 @@ _is_xpu = is_xpu()
if _is_xpu:
from sgl_kernel import hadamard_transform
else:
from sglang.kernels.ops.attention.hadamard import hadamard_transform
from sglang.kernels.ops.quantization.hadamard import hadamard_transform
HEAD_DIM = 128
FP4_DIM = HEAD_DIM // 2
@@ -7,7 +7,7 @@ import torch
import torch.nn.functional as F
from scipy.linalg import hadamard
from sglang.kernels.ops.attention.hadamard import (
from sglang.kernels.ops.quantization.hadamard import (
hadamard_transform,
hadamard_transform_12n,
hadamard_transform_20n,
@@ -7,7 +7,7 @@ must be untouched by tau.
import pytest
import torch
from sglang.kernels.ops.model.inkling.inkling_attn_prologue import (
from sglang.kernels.ops.attention.inkling_attn_prologue import (
inkling_attn_prologue_decode,
)
from sglang.test.ci.ci_register import register_cuda_ci
@@ -5,10 +5,10 @@ including on the row-strided qkvr-slice layouts."""
import pytest
import torch
from sglang.kernels.ops.attention.inkling_row_scale import row_scale_bf16
from sglang.kernels.ops.attention.log_scaling_tau import (
_apply_log_scaling_tau_kernel,
)
from sglang.kernels.ops.model.inkling.inkling_row_scale import row_scale_bf16
from sglang.test.ci.ci_register import register_cuda_ci
register_cuda_ci(est_time=30, stage="base-b-kernel-unit", runner_config="1-gpu-large")
@@ -51,7 +51,7 @@ def test_row_compact_bitexact(rows, inner, strided):
"""The tau-less compaction flavor (kHasTau=false) must reproduce
.contiguous() exactly on the same strided layouts row_scale handles --
no other test exercises run_compact."""
from sglang.kernels.ops.model.inkling.inkling_row_scale import row_compact_bf16
from sglang.kernels.ops.attention.inkling_row_scale import row_compact_bf16
torch.manual_seed(rows + inner)
if strided:
@@ -48,7 +48,7 @@ def split_qknorm_rope(
) -> None:
from flashinfer.rope import apply_rope_with_cos_sin_cache_inplace
from sglang.kernels.ops.layernorm._jit_norm import fused_inplace_qknorm
from sglang.kernels.ops.layernorm.norm import fused_inplace_qknorm
fused_inplace_qknorm(q, k, q_weight, k_weight)
apply_rope_with_cos_sin_cache_inplace(
@@ -11,7 +11,7 @@ from sglang.kernels.jit.utils import (
get_jit_cuda_arch,
is_hip_runtime,
)
from sglang.kernels.ops.gemm._jit_dsv3_fused_a_gemm import dsv3_fused_a_gemm
from sglang.kernels.ops.gemm.dsv3_fused_a_gemm import dsv3_fused_a_gemm
from sglang.test.ci.ci_register import register_cuda_ci
register_cuda_ci(est_time=30, stage="base-b-kernel-unit", runner_config="1-gpu-large")
@@ -11,7 +11,7 @@ from sglang.kernels.jit.utils import (
get_jit_cuda_arch,
is_hip_runtime,
)
from sglang.kernels.ops.gemm._jit_dsv3_router_gemm import dsv3_router_gemm
from sglang.kernels.ops.gemm.dsv3_router_gemm import dsv3_router_gemm
from sglang.test.ci.ci_register import register_cuda_ci
register_cuda_ci(est_time=37, stage="base-b-kernel-unit", runner_config="1-gpu-large")
@@ -4,7 +4,7 @@ import pytest
import torch
from sglang.kernels.jit.utils import get_ci_test_range
from sglang.kernels.ops.kvcache._jit_set_mla_kv_buffer import (
from sglang.kernels.ops.kvcache.set_mla_kv_buffer import (
can_use_set_mla_kv_buffer,
set_mla_kv_buffer,
)
@@ -20,7 +20,7 @@ def sglang_jit_fused_add_rmsnorm(
*,
cast_x_before_out_mul: bool = False,
) -> None:
from sglang.kernels.ops.layernorm._jit_norm import fused_add_rmsnorm
from sglang.kernels.ops.layernorm.norm import fused_add_rmsnorm
fused_add_rmsnorm(
input, residual, weight, eps, cast_x_before_out_mul=cast_x_before_out_mul
@@ -21,6 +21,7 @@ GROUPS = [
"attention",
"communication",
"diffusion",
"elementwise",
"embeddings",
"gemm",
"grammar",
@@ -31,7 +32,6 @@ GROUPS = [
"moe",
"quantization",
"sampling",
"spatial",
"speculative",
]
@@ -34,7 +34,7 @@ def sglang_jit_qknorm(
q_weight: torch.Tensor,
k_weight: torch.Tensor,
) -> None:
from sglang.kernels.ops.layernorm._jit_norm import fused_inplace_qknorm
from sglang.kernels.ops.layernorm.norm import fused_inplace_qknorm
fused_inplace_qknorm(q, k, q_weight, k_weight)
@@ -19,7 +19,7 @@ def sglang_jit_qknorm_across_heads(
q_weight: torch.Tensor,
k_weight: torch.Tensor,
) -> None:
from sglang.kernels.ops.layernorm._jit_norm import fused_inplace_qknorm_across_heads
from sglang.kernels.ops.layernorm.norm import fused_inplace_qknorm_across_heads
fused_inplace_qknorm_across_heads(q, k, q_weight, k_weight)
@@ -26,7 +26,7 @@ def sglang_jit_rmsnorm(
output: torch.Tensor | None = None,
eps: float = EPS,
) -> None:
from sglang.kernels.ops.layernorm._jit_norm import rmsnorm
from sglang.kernels.ops.layernorm.norm import rmsnorm
rmsnorm(input, weight, out=output, eps=eps)
@@ -127,7 +127,7 @@ def test_rmsnorm(
@pytest.mark.parametrize("hidden_size", [64, 128, 256, 512, 8192, 8704, 16384])
def test_rmsnorm_hidden_size_support(hidden_size: int) -> None:
from sglang.kernels.ops.layernorm._jit_norm import _is_supported_rmsnorm_hidden_size
from sglang.kernels.ops.layernorm.norm import _is_supported_rmsnorm_hidden_size
assert _is_supported_rmsnorm_hidden_size(hidden_size)
@@ -148,7 +148,7 @@ def test_rmsnorm_hidden_size_support(hidden_size: int) -> None:
],
)
def test_rmsnorm_kernel_dispatch(hidden_size: int, expected: str) -> None:
from sglang.kernels.ops.layernorm._jit_norm import _rmsnorm_kernel_class
from sglang.kernels.ops.layernorm.norm import _rmsnorm_kernel_class
assert _rmsnorm_kernel_class(hidden_size) == expected
@@ -6,7 +6,7 @@ production strided-r layout and contiguous inputs."""
import pytest
import torch
from sglang.kernels.ops.model.inkling.inkling_rel_proj import rel_proj_small_t
from sglang.kernels.ops.attention.inkling_rel_proj import rel_proj_small_t
from sglang.test.ci.ci_register import register_cuda_ci
register_cuda_ci(est_time=30, stage="base-b-kernel-unit", runner_config="1-gpu-large")
@@ -6,7 +6,7 @@ import pytest
import torch
from sglang.kernels.jit.utils import get_ci_test_range
from sglang.kernels.ops.quantization._jit_per_tensor_quant_fp8 import (
from sglang.kernels.ops.quantization.per_tensor_quant_fp8 import (
per_tensor_quant_fp8,
)
from sglang.test.ci.ci_register import register_cuda_ci
@@ -22,14 +22,14 @@ import pytest
import torch
from sglang.kernels.jit.utils import get_ci_test_range
from sglang.kernels.ops.quantization._jit_per_token_group_quant import (
per_token_group_quant,
)
from sglang.kernels.ops.quantization.fp8_kernel import (
create_per_token_group_quant_fp8_output_scale,
fp8_dtype,
fp8_max,
)
from sglang.kernels.ops.quantization.per_token_group_quant import (
per_token_group_quant,
)
from sglang.test.ci.ci_register import register_cuda_ci
register_cuda_ci(est_time=90, stage="base-b-kernel-unit", runner_config="1-gpu-large")
@@ -4,7 +4,7 @@ import pytest
import torch
from sglang.kernels.jit.utils import get_ci_test_range
from sglang.kernels.ops.quantization._jit_per_token_group_quant_8bit_v2 import (
from sglang.kernels.ops.quantization.per_token_group_quant_8bit_v2 import (
per_token_group_quant_8bit_v2,
)
from sglang.test.ci.ci_register import register_cuda_ci