diff --git a/python/sglang/jit_kernel/__main__.py b/python/sglang/jit_kernel/__main__.py index b9c0f9681..a7c52cf0d 100644 --- a/python/sglang/jit_kernel/__main__.py +++ b/python/sglang/jit_kernel/__main__.py @@ -7,10 +7,10 @@ import subprocess from tvm_ffi.libinfo import find_dlpack_include_path, find_include_path -from sglang.jit_kernel.utils import get_jit_cuda_arch, override_jit_cuda_arch -from sglang.jit_kernel.utils.arch import get_default_target_flags -from sglang.jit_kernel.utils.compile import DEFAULT_INCLUDE -from sglang.jit_kernel.utils.deps import REGISTERED_DEPENDENCIES +from sglang.kernels.jit.utils import get_jit_cuda_arch, override_jit_cuda_arch +from sglang.kernels.jit.utils.arch import get_default_target_flags +from sglang.kernels.jit.utils.compile import DEFAULT_INCLUDE +from sglang.kernels.jit.utils.deps import REGISTERED_DEPENDENCIES def _clangd_major_version() -> int | None: diff --git a/python/sglang/jit_kernel/activation.py b/python/sglang/jit_kernel/activation.py index e94aecb4c..520e78415 100644 --- a/python/sglang/jit_kernel/activation.py +++ b/python/sglang/jit_kernel/activation.py @@ -1,168 +1,5 @@ -from __future__ import annotations +"""Compatibility shim (RFC #29630 Phase 4) -> sglang.kernels.ops.activation._jit_activation.""" -from typing import TYPE_CHECKING, Optional +from sglang.kernels.ops.activation import _jit_activation as _impl -import torch - -from sglang.jit_kernel.utils import ( - cache_once, - get_jit_cuda_arch, - is_arch_support_pdl, - is_hip_runtime, - load_jit, - make_cpp_args, -) -from sglang.srt.utils.custom_op import register_custom_op - -if TYPE_CHECKING: - from tvm_ffi.module import Module - - -def _fast_math_flags() -> list[str]: - # Mirrors sgl-kernel's CMake policy: fast-math on SM90, precise on - # SM100+ (Blackwell needs bit-exact expf), off on HIP (clang rejects). - if is_hip_runtime(): - return [] - if get_jit_cuda_arch().major >= 10: - return [] - return ["--use_fast_math"] - - -@cache_once -def _jit_activation_module(dtype: torch.dtype) -> Module: - args = make_cpp_args(dtype, is_arch_support_pdl()) - return load_jit( - "activation", - *args, - cuda_files=["elementwise/activation.cuh"], - extra_cuda_cflags=_fast_math_flags(), - cuda_wrappers=[ - ("run_activation", f"ActivationKernel<{args}>::run_activation"), - ( - "run_activation_filtered", - f"ActivationKernel<{args}>::run_activation_filtered", - ), - ( - "run_unary_activation", - f"ActivationKernel<{args}>::run_unary_activation", - ), - ], - ) - - -SUPPORTED_ACTIVATIONS = {"silu", "gelu", "gelu_tanh"} -SUPPORTED_UNARY_ACTIVATIONS = {"relu2"} - - -@register_custom_op(mutates_args=["out"]) -def _run_activation_inplace( - op_name: str, input: torch.Tensor, out: torch.Tensor -) -> None: - hidden_size = input.shape[-1] // 2 - module = _jit_activation_module(input.dtype) - input_2d = input.view(-1, hidden_size * 2) - out_2d = out.view(-1, hidden_size) - module.run_activation(input_2d, out_2d, op_name) - - -@register_custom_op(mutates_args=["out"]) -def _run_activation_filtered_inplace( - op_name: str, - input: torch.Tensor, - out: torch.Tensor, - expert_ids: torch.Tensor, - expert_step: int, -) -> None: - hidden_size = input.shape[-1] // 2 - module = _jit_activation_module(input.dtype) - input_2d = input.view(-1, hidden_size * 2) - out_2d = out.view(-1, hidden_size) - module.run_activation_filtered(input_2d, out_2d, expert_ids, expert_step, op_name) - - -def run_activation( - op_name: str, - input: torch.Tensor, - out: Optional[torch.Tensor], - expert_ids: Optional[torch.Tensor] = None, - expert_step: int = 1, -) -> torch.Tensor: - """Apply ``op_name`` activation followed by element-wise multiplication. - - When ``expert_ids`` is provided, output rows are skipped for tokens whose - routed expert id is ``-1``. ``expert_step`` is 1 for per-token routing and - ``BLOCK_SIZE_M`` for sorted/TMA routing — i.e. ``expert_ids[token_id // - expert_step]`` is consulted before computing each row. - """ - assert op_name in SUPPORTED_ACTIVATIONS, f"Unsupported activation: {op_name}" - hidden_size = input.shape[-1] // 2 - if out is None: - out = input.new_empty(*input.shape[:-1], hidden_size) - if expert_ids is None: - _run_activation_inplace(op_name, input, out) - else: - _run_activation_filtered_inplace(op_name, input, out, expert_ids, expert_step) - return out - - -@register_custom_op(mutates_args=["out"]) -def _run_unary_activation_inplace( - op_name: str, input: torch.Tensor, out: torch.Tensor -) -> None: - last = input.shape[-1] - module = _jit_activation_module(input.dtype) - module.run_unary_activation(input.view(-1, last), out.view(-1, last), op_name) - - -def run_unary_activation( - op_name: str, - input: torch.Tensor, - out: Optional[torch.Tensor] = None, -) -> torch.Tensor: - """Apply a standalone (non-gated) element-wise activation: ``out = act(input)``. - - Unlike :func:`run_activation`, there is no gate/up split — ``input`` and - ``out`` share the same shape. - """ - assert ( - op_name in SUPPORTED_UNARY_ACTIVATIONS - ), f"Unsupported unary activation: {op_name}" - if out is None: - out = torch.empty_like(input) - _run_unary_activation_inplace(op_name, input, out) - return out - - -def relu2( - input: torch.Tensor, - out: Optional[torch.Tensor] = None, -) -> torch.Tensor: - """Squared ReLU: ``out = max(0, input) ** 2`` (element-wise).""" - return run_unary_activation("relu2", input, out) - - -def silu_and_mul( - input: torch.Tensor, - out: Optional[torch.Tensor] = None, - expert_ids: Optional[torch.Tensor] = None, - expert_step: int = 1, -) -> torch.Tensor: - return run_activation("silu", input, out, expert_ids, expert_step) - - -def gelu_and_mul( - input: torch.Tensor, - out: Optional[torch.Tensor] = None, - expert_ids: Optional[torch.Tensor] = None, - expert_step: int = 1, -) -> torch.Tensor: - return run_activation("gelu", input, out, expert_ids, expert_step) - - -def gelu_tanh_and_mul( - input: torch.Tensor, - out: Optional[torch.Tensor] = None, - expert_ids: Optional[torch.Tensor] = None, - expert_step: int = 1, -) -> torch.Tensor: - return run_activation("gelu_tanh", input, out, expert_ids, expert_step) +globals().update({k: getattr(_impl, k) for k in dir(_impl) if not k.startswith("__")}) diff --git a/python/sglang/jit_kernel/add_constant.py b/python/sglang/jit_kernel/add_constant.py index 228e0de60..e22fd2c40 100644 --- a/python/sglang/jit_kernel/add_constant.py +++ b/python/sglang/jit_kernel/add_constant.py @@ -4,7 +4,7 @@ from typing import TYPE_CHECKING import torch -from sglang.jit_kernel.utils import cache_once, load_jit, make_cpp_args +from sglang.kernels.jit.utils import cache_once, load_jit, make_cpp_args if TYPE_CHECKING: from tvm_ffi.module import Module diff --git a/python/sglang/jit_kernel/all_reduce.py b/python/sglang/jit_kernel/all_reduce.py index ec8dd12f4..54cd7757f 100644 --- a/python/sglang/jit_kernel/all_reduce.py +++ b/python/sglang/jit_kernel/all_reduce.py @@ -7,14 +7,14 @@ import torch import tvm_ffi from tvm_ffi import Module -from sglang.jit_kernel.utils import ( +from sglang.kernel_api_logging import debug_kernel_api +from sglang.kernels.jit.utils import ( cache_once, is_arch_support_pdl, lazy_register_class, load_jit, make_cpp_args, ) -from sglang.kernel_api_logging import debug_kernel_api class AllReduceAlgo(enum.Enum): diff --git a/python/sglang/jit_kernel/awq_dequantize.py b/python/sglang/jit_kernel/awq_dequantize.py index 4a188c02e..c416c9c72 100644 --- a/python/sglang/jit_kernel/awq_dequantize.py +++ b/python/sglang/jit_kernel/awq_dequantize.py @@ -4,7 +4,7 @@ from typing import TYPE_CHECKING import torch -from sglang.jit_kernel.utils import cache_once, load_jit, make_cpp_args +from sglang.kernels.jit.utils import cache_once, load_jit, make_cpp_args if TYPE_CHECKING: from tvm_ffi.module import Module diff --git a/python/sglang/jit_kernel/awq_marlin_repack.py b/python/sglang/jit_kernel/awq_marlin_repack.py index d51c1fd51..8e8707709 100644 --- a/python/sglang/jit_kernel/awq_marlin_repack.py +++ b/python/sglang/jit_kernel/awq_marlin_repack.py @@ -4,8 +4,8 @@ from typing import TYPE_CHECKING import torch -from sglang.jit_kernel.utils import cache_once, load_jit from sglang.kernel_api_logging import debug_kernel_api +from sglang.kernels.jit.utils import cache_once, load_jit if TYPE_CHECKING: from tvm_ffi.module import Module diff --git a/python/sglang/jit_kernel/benchmark/marker.py b/python/sglang/jit_kernel/benchmark/marker.py index 60e31783b..e34a00e31 100644 --- a/python/sglang/jit_kernel/benchmark/marker.py +++ b/python/sglang/jit_kernel/benchmark/marker.py @@ -21,7 +21,7 @@ from typing import ( import torch -from sglang.jit_kernel.utils import cache_once +from sglang.kernels.jit.utils import cache_once from sglang.utils import is_in_ci F = TypeVar("F", bound=Callable[..., "BenchResult"]) diff --git a/python/sglang/jit_kernel/clamp_position.py b/python/sglang/jit_kernel/clamp_position.py index ed57da776..d7156a66b 100644 --- a/python/sglang/jit_kernel/clamp_position.py +++ b/python/sglang/jit_kernel/clamp_position.py @@ -4,7 +4,7 @@ from typing import TYPE_CHECKING import torch -from sglang.jit_kernel.utils import cache_once, load_jit, make_cpp_args +from sglang.kernels.jit.utils import cache_once, load_jit, make_cpp_args if TYPE_CHECKING: from tvm_ffi.module import Module diff --git a/python/sglang/jit_kernel/concat_mla.py b/python/sglang/jit_kernel/concat_mla.py index 4945b73bc..3de6e9849 100644 --- a/python/sglang/jit_kernel/concat_mla.py +++ b/python/sglang/jit_kernel/concat_mla.py @@ -4,7 +4,7 @@ from typing import TYPE_CHECKING import torch -from sglang.jit_kernel.utils import cache_once, load_jit +from sglang.kernels.jit.utils import cache_once, load_jit if TYPE_CHECKING: from tvm_ffi.module import Module diff --git a/python/sglang/jit_kernel/diffusion/causal_conv3d_cat_pad.py b/python/sglang/jit_kernel/diffusion/causal_conv3d_cat_pad.py index 60ff5626d..4a07aaba9 100644 --- a/python/sglang/jit_kernel/diffusion/causal_conv3d_cat_pad.py +++ b/python/sglang/jit_kernel/diffusion/causal_conv3d_cat_pad.py @@ -4,7 +4,7 @@ from typing import TYPE_CHECKING import torch -from sglang.jit_kernel.utils import cache_once, load_jit, make_cpp_args +from sglang.kernels.jit.utils import cache_once, load_jit, make_cpp_args from sglang.srt.utils.custom_op import register_custom_op if TYPE_CHECKING: diff --git a/python/sglang/jit_kernel/diffusion/ltx2_qknorm_split_rope.py b/python/sglang/jit_kernel/diffusion/ltx2_qknorm_split_rope.py index 45277e2d2..d5598c888 100644 --- a/python/sglang/jit_kernel/diffusion/ltx2_qknorm_split_rope.py +++ b/python/sglang/jit_kernel/diffusion/ltx2_qknorm_split_rope.py @@ -4,7 +4,7 @@ from typing import TYPE_CHECKING import torch -from sglang.jit_kernel.utils import cache_once, load_jit +from sglang.kernels.jit.utils import cache_once, load_jit from sglang.srt.utils.custom_op import register_custom_op if TYPE_CHECKING: diff --git a/python/sglang/jit_kernel/diffusion/norm_scale_shift_native.py b/python/sglang/jit_kernel/diffusion/norm_scale_shift_native.py index cbc175918..fc093e0a7 100644 --- a/python/sglang/jit_kernel/diffusion/norm_scale_shift_native.py +++ b/python/sglang/jit_kernel/diffusion/norm_scale_shift_native.py @@ -4,7 +4,7 @@ from typing import TYPE_CHECKING import torch -from sglang.jit_kernel.utils import cache_once, load_jit +from sglang.kernels.jit.utils import cache_once, load_jit if TYPE_CHECKING: from tvm_ffi.module import Module diff --git a/python/sglang/jit_kernel/diffusion/qknorm_rope.py b/python/sglang/jit_kernel/diffusion/qknorm_rope.py index 8dfdf8d8d..80af7cdf7 100644 --- a/python/sglang/jit_kernel/diffusion/qknorm_rope.py +++ b/python/sglang/jit_kernel/diffusion/qknorm_rope.py @@ -5,7 +5,7 @@ from typing import TYPE_CHECKING import torch -from sglang.jit_kernel.utils import ( +from sglang.kernels.jit.utils import ( cache_once, is_arch_support_pdl, load_jit, diff --git a/python/sglang/jit_kernel/diffusion/residual_gate_add.py b/python/sglang/jit_kernel/diffusion/residual_gate_add.py index 9933e3d2a..9a4f32878 100644 --- a/python/sglang/jit_kernel/diffusion/residual_gate_add.py +++ b/python/sglang/jit_kernel/diffusion/residual_gate_add.py @@ -4,7 +4,7 @@ from typing import TYPE_CHECKING import torch -from sglang.jit_kernel.utils import cache_once, load_jit, make_cpp_args +from sglang.kernels.jit.utils import cache_once, load_jit, make_cpp_args from sglang.srt.utils.custom_op import register_custom_op if TYPE_CHECKING: diff --git a/python/sglang/jit_kernel/dsv32/elementwise.py b/python/sglang/jit_kernel/dsv32/elementwise.py index 4c07afe4b..e918679d5 100644 --- a/python/sglang/jit_kernel/dsv32/elementwise.py +++ b/python/sglang/jit_kernel/dsv32/elementwise.py @@ -2,7 +2,7 @@ import torch -from sglang.jit_kernel.utils import ( +from sglang.kernels.jit.utils import ( cache_once, is_arch_support_pdl, load_jit, diff --git a/python/sglang/jit_kernel/dsv3_fused_a_gemm.py b/python/sglang/jit_kernel/dsv3_fused_a_gemm.py index 3649beb20..56f770684 100644 --- a/python/sglang/jit_kernel/dsv3_fused_a_gemm.py +++ b/python/sglang/jit_kernel/dsv3_fused_a_gemm.py @@ -1,90 +1,5 @@ -""" -JIT kernel for DeepSeek V3 fused QKV-A GEMM (min-latency). +"""Compatibility shim (RFC #29630 Phase 4) -> sglang.kernels.ops.gemm._jit_dsv3_fused_a_gemm.""" -Runtime-compiled CUDA C++ kernel for SM90+ (Hopper) GPUs. -Shapes: hd_in a multiple of 256, hd_out a multiple of 16, num_tokens 1-16, bfloat16. -""" +from sglang.kernels.ops.gemm import _jit_dsv3_fused_a_gemm as _impl -from __future__ import annotations - -from typing import TYPE_CHECKING, Optional - -import torch - -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 -from sglang.srt.utils.common import direct_register_custom_op - -if TYPE_CHECKING: - from tvm_ffi.module import Module - - -@cache_once -def _jit_dsv3_fused_a_gemm_module(hd_in: int, hd_out: int, use_pdl: bool) -> Module: - args = make_cpp_args(hd_in, hd_out, use_pdl) - return load_jit( - "dsv3_fused_a_gemm", - *args, - cuda_files=["gemm/dsv3_fused_a_gemm.cuh"], - cuda_wrappers=[ - ("dsv3_fused_a_gemm", f"DSV3FusedAGemmKernel<{args}>::run"), - ], - ) - - -def _dsv3_fused_a_gemm_run(mat_a: torch.Tensor, mat_b: torch.Tensor) -> torch.Tensor: - assert mat_a.stride(1) == 1, "mat_a must be row-major [M, K]" - output = torch.empty( - (mat_a.shape[0], mat_b.shape[1]), - device=mat_a.device, - dtype=mat_a.dtype, - ) - module = _jit_dsv3_fused_a_gemm_module( - mat_a.shape[1], mat_b.shape[1], is_arch_support_pdl() - ) - module.dsv3_fused_a_gemm(mat_a, mat_b, output) - return output - - -def _dsv3_fused_a_gemm_fake(mat_a: torch.Tensor, mat_b: torch.Tensor) -> torch.Tensor: - return mat_a.new_empty((mat_a.shape[0], mat_b.shape[1]), dtype=torch.bfloat16) - - -direct_register_custom_op( - op_name="jit_dsv3_fused_a_gemm", - op_func=_dsv3_fused_a_gemm_run, - mutates_args=[], - fake_impl=_dsv3_fused_a_gemm_fake, -) - - -@debug_kernel_api -def dsv3_fused_a_gemm( - mat_a: torch.Tensor, - mat_b: torch.Tensor, - output: Optional[torch.Tensor] = None, -) -> torch.Tensor: - """ - DeepSeek V3 fused QKV-A GEMM kernel (JIT variant). - - Args: - mat_a: Input tensor of shape [num_tokens, hd_in], bfloat16, row-major. - hd_in must be a multiple of 256 and num_tokens in [1, 16]. - mat_b: Weight tensor of shape [hd_in, hd_out], bfloat16, column-major - (i.e. ``weight.T`` of a row-major [hd_out, hd_in] weight). - hd_out must be a multiple of 16. - output: Optional pre-allocated output tensor of shape [num_tokens, hd_out]. - - Returns: - Output tensor of shape [num_tokens, hd_out]. - """ - result = torch.ops.sglang.jit_dsv3_fused_a_gemm(mat_a, mat_b) - if output is not None: - output.copy_(result) - return output - return result +globals().update({k: getattr(_impl, k) for k in dir(_impl) if not k.startswith("__")}) diff --git a/python/sglang/jit_kernel/dsv3_router_gemm.py b/python/sglang/jit_kernel/dsv3_router_gemm.py index a7fadae52..fac174e38 100644 --- a/python/sglang/jit_kernel/dsv3_router_gemm.py +++ b/python/sglang/jit_kernel/dsv3_router_gemm.py @@ -1,92 +1,5 @@ -""" -JIT kernel for DeepSeek V3 router GEMM. +"""Compatibility shim (RFC #29630 Phase 4) -> sglang.kernels.ops.gemm._jit_dsv3_router_gemm.""" -Runtime-compiled CUDA C++ kernel for SM90+ (Hopper) GPUs. -Supports num_experts in {256, 384}, hidden_dim a multiple of 1024, num_tokens 1-16. -""" +from sglang.kernels.ops.gemm import _jit_dsv3_router_gemm as _impl -from __future__ import annotations - -from typing import TYPE_CHECKING, Optional - -import torch - -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 -from sglang.srt.utils.custom_op import register_custom_op - -if TYPE_CHECKING: - from tvm_ffi.module import Module - - -@cache_once -def _jit_dsv3_router_gemm_module( - num_experts: int, - hidden_dim: int, - use_pdl: bool, - out_float: bool, -) -> Module: - args = make_cpp_args(num_experts, hidden_dim, use_pdl, out_float) - return load_jit( - "dsv3_router_gemm", - *args, - cuda_files=["gemm/dsv3_router_gemm.cuh"], - cuda_wrappers=[ - ("dsv3_router_gemm", f"DSV3RouterGemmKernel<{args}>::run"), - ], - ) - - -@register_custom_op( - op_name="dsv3_router_gemm", - mutates_args=["output"], -) -def _dsv3_router_gemm_custom_op( - hidden_states: torch.Tensor, - router_weights: torch.Tensor, - output: torch.Tensor, -) -> None: - num_experts = router_weights.shape[0] - hidden_dim = hidden_states.shape[1] - out_float = output.dtype == torch.float32 - module = _jit_dsv3_router_gemm_module( - num_experts, hidden_dim, is_arch_support_pdl(), out_float - ) - module.dsv3_router_gemm(hidden_states, router_weights, output) - return None - - -@debug_kernel_api -def dsv3_router_gemm( - hidden_states: torch.Tensor, - router_weights: torch.Tensor, - out_dtype: torch.dtype = torch.bfloat16, - output: Optional[torch.Tensor] = None, -) -> torch.Tensor: - """ - DeepSeek V3 router GEMM kernel (JIT variant). - - Args: - hidden_states: Input tensor of shape [num_tokens, hidden_dim], bfloat16. - hidden_dim must be a multiple of 1024 and num_tokens in [1, 16]. - router_weights: Weight tensor of shape [num_experts, hidden_dim], bfloat16. - out_dtype: Output dtype, either torch.bfloat16 or torch.float32. - output: Optional pre-allocated output tensor. - - Returns: - Output tensor of shape [num_tokens, num_experts]. - """ - if output is None: - output = torch.empty( - hidden_states.shape[0], - router_weights.shape[0], - device=hidden_states.device, - dtype=out_dtype, - ) - _dsv3_router_gemm_custom_op(hidden_states, router_weights, output) - return output +globals().update({k: getattr(_impl, k) for k in dir(_impl) if not k.startswith("__")}) diff --git a/python/sglang/jit_kernel/dsv4/attn.py b/python/sglang/jit_kernel/dsv4/attn.py index 784711498..a436967a8 100644 --- a/python/sglang/jit_kernel/dsv4/attn.py +++ b/python/sglang/jit_kernel/dsv4/attn.py @@ -4,7 +4,7 @@ import torch import triton import triton.language as tl -from sglang.jit_kernel.utils import ( +from sglang.kernels.jit.utils import ( cache_once, is_arch_support_pdl, is_hip_runtime, diff --git a/python/sglang/jit_kernel/dsv4/compress.py b/python/sglang/jit_kernel/dsv4/compress.py index bf82bff04..148e23a10 100644 --- a/python/sglang/jit_kernel/dsv4/compress.py +++ b/python/sglang/jit_kernel/dsv4/compress.py @@ -4,7 +4,7 @@ from typing import TYPE_CHECKING, Literal, NamedTuple, Optional, Union import torch -from sglang.jit_kernel.utils import ( +from sglang.kernels.jit.utils import ( cache_once, is_arch_support_pdl, load_jit, diff --git a/python/sglang/jit_kernel/dsv4/compress_old.py b/python/sglang/jit_kernel/dsv4/compress_old.py index 9bf96a964..515f900b3 100644 --- a/python/sglang/jit_kernel/dsv4/compress_old.py +++ b/python/sglang/jit_kernel/dsv4/compress_old.py @@ -4,7 +4,7 @@ from typing import TYPE_CHECKING, Literal, NamedTuple, Optional, Union import torch -from sglang.jit_kernel.utils import ( +from sglang.kernels.jit.utils import ( cache_once, is_arch_support_pdl, load_jit, diff --git a/python/sglang/jit_kernel/dsv4/elementwise.py b/python/sglang/jit_kernel/dsv4/elementwise.py index b1395824d..309e3055e 100644 --- a/python/sglang/jit_kernel/dsv4/elementwise.py +++ b/python/sglang/jit_kernel/dsv4/elementwise.py @@ -2,7 +2,7 @@ from typing import Optional, Tuple import torch -from sglang.jit_kernel.utils import ( +from sglang.kernels.jit.utils import ( cache_once, is_arch_support_pdl, load_jit, diff --git a/python/sglang/jit_kernel/dsv4/fp8_wo_a.py b/python/sglang/jit_kernel/dsv4/fp8_wo_a.py index 907900bdb..df4709fee 100644 --- a/python/sglang/jit_kernel/dsv4/fp8_wo_a.py +++ b/python/sglang/jit_kernel/dsv4/fp8_wo_a.py @@ -4,13 +4,13 @@ from typing import TYPE_CHECKING, Tuple import torch -from sglang.jit_kernel.utils import ( +from sglang.kernel_api_logging import debug_kernel_api +from sglang.kernels.jit.utils import ( cache_once, is_arch_support_pdl, load_jit, make_cpp_args, ) -from sglang.kernel_api_logging import debug_kernel_api from sglang.srt.utils.custom_op import register_custom_op from .utils import make_name diff --git a/python/sglang/jit_kernel/dsv4/moe.py b/python/sglang/jit_kernel/dsv4/moe.py index 81a21c85c..7545c1f0d 100644 --- a/python/sglang/jit_kernel/dsv4/moe.py +++ b/python/sglang/jit_kernel/dsv4/moe.py @@ -2,7 +2,7 @@ from typing import Optional, Tuple import torch -from sglang.jit_kernel.utils import ( +from sglang.kernels.jit.utils import ( cache_once, is_arch_support_pdl, is_hip_runtime, diff --git a/python/sglang/jit_kernel/dsv4/online_c128_mtp.py b/python/sglang/jit_kernel/dsv4/online_c128_mtp.py index 5a901b4d9..f7f9cbbce 100644 --- a/python/sglang/jit_kernel/dsv4/online_c128_mtp.py +++ b/python/sglang/jit_kernel/dsv4/online_c128_mtp.py @@ -6,7 +6,7 @@ from typing import TYPE_CHECKING, Any, List, Optional import torch from sglang.jit_kernel.dsv4.utils import make_name -from sglang.jit_kernel.utils import cache_once, load_jit, make_cpp_args +from sglang.kernels.jit.utils import cache_once, load_jit, make_cpp_args from sglang.srt.environ import envs if TYPE_CHECKING: diff --git a/python/sglang/jit_kernel/dsv4/topk.py b/python/sglang/jit_kernel/dsv4/topk.py index fdaedfd75..53d7ebd52 100644 --- a/python/sglang/jit_kernel/dsv4/topk.py +++ b/python/sglang/jit_kernel/dsv4/topk.py @@ -4,7 +4,7 @@ from typing import Optional import torch -from sglang.jit_kernel.utils import ( +from sglang.kernels.jit.utils import ( cache_once, is_arch_support_pdl, is_hip_runtime, diff --git a/python/sglang/jit_kernel/fixup_zero_kv.py b/python/sglang/jit_kernel/fixup_zero_kv.py index 6175c0f37..6f42edf01 100644 --- a/python/sglang/jit_kernel/fixup_zero_kv.py +++ b/python/sglang/jit_kernel/fixup_zero_kv.py @@ -4,7 +4,7 @@ from typing import TYPE_CHECKING import torch -from sglang.jit_kernel.utils import cache_once, load_jit, make_cpp_args +from sglang.kernels.jit.utils import cache_once, load_jit, make_cpp_args if TYPE_CHECKING: from tvm_ffi.module import Module diff --git a/python/sglang/jit_kernel/flash_attention_v3.py b/python/sglang/jit_kernel/flash_attention_v3.py index fe7f42234..0d7cbf5ec 100644 --- a/python/sglang/jit_kernel/flash_attention_v3.py +++ b/python/sglang/jit_kernel/flash_attention_v3.py @@ -4,8 +4,8 @@ from typing import Optional, Union import torch -from sglang.jit_kernel.utils import cache_once from sglang.kernel_api_logging import debug_kernel_api +from sglang.kernels.jit.utils import cache_once from sglang.srt.environ import envs from sglang.srt.utils import get_device_capability, is_musa diff --git a/python/sglang/jit_kernel/flash_attn/cute/interface.py b/python/sglang/jit_kernel/flash_attn/cute/interface.py index cc8ae4e21..b5a31756a 100644 --- a/python/sglang/jit_kernel/flash_attn/cute/interface.py +++ b/python/sglang/jit_kernel/flash_attn/cute/interface.py @@ -15,7 +15,7 @@ from quack.compile_utils import make_fake_tensor as fake_tensor from sglang.jit_kernel.flash_attn.cute.cache_utils import get_jit_cache from sglang.jit_kernel.flash_attn.cute.testing import is_fake_mode -from sglang.jit_kernel.utils import is_arch_support_pdl +from sglang.kernels.jit.utils import is_arch_support_pdl if os.environ.get("CUTE_DSL_PTXAS_PATH", None) is not None: from sglang.jit_kernel.flash_attn.cute import cute_dsl_ptxas # noqa: F401 diff --git a/python/sglang/jit_kernel/fp8_blockwise_gemm.py b/python/sglang/jit_kernel/fp8_blockwise_gemm.py index 49b4c9606..2b883ee97 100644 --- a/python/sglang/jit_kernel/fp8_blockwise_gemm.py +++ b/python/sglang/jit_kernel/fp8_blockwise_gemm.py @@ -5,8 +5,8 @@ from typing import TYPE_CHECKING import torch -from sglang.jit_kernel.utils import cache_once, load_jit, override_jit_cuda_arch from sglang.kernel_api_logging import debug_kernel_api +from sglang.kernels.jit.utils import cache_once, load_jit, override_jit_cuda_arch from sglang.srt.utils.common import is_sm120_supported from sglang.srt.utils.custom_op import register_custom_op diff --git a/python/sglang/jit_kernel/fused_eh_norm.py b/python/sglang/jit_kernel/fused_eh_norm.py index d7b7c0817..43c8f26b0 100644 --- a/python/sglang/jit_kernel/fused_eh_norm.py +++ b/python/sglang/jit_kernel/fused_eh_norm.py @@ -4,7 +4,7 @@ from typing import TYPE_CHECKING import torch -from sglang.jit_kernel.utils import ( +from sglang.kernels.jit.utils import ( cache_once, is_arch_support_pdl, load_jit, diff --git a/python/sglang/jit_kernel/fused_fp8_qkv_kv_cache.py b/python/sglang/jit_kernel/fused_fp8_qkv_kv_cache.py index d412cae55..10a134e60 100644 --- a/python/sglang/jit_kernel/fused_fp8_qkv_kv_cache.py +++ b/python/sglang/jit_kernel/fused_fp8_qkv_kv_cache.py @@ -4,7 +4,7 @@ from typing import TYPE_CHECKING, Optional import torch -from sglang.jit_kernel.utils import ( +from sglang.kernels.jit.utils import ( cache_once, is_arch_support_pdl, load_jit, diff --git a/python/sglang/jit_kernel/fused_metadata_copy.py b/python/sglang/jit_kernel/fused_metadata_copy.py index 68d0f9227..569987ada 100644 --- a/python/sglang/jit_kernel/fused_metadata_copy.py +++ b/python/sglang/jit_kernel/fused_metadata_copy.py @@ -15,7 +15,7 @@ from typing import Optional import torch -from sglang.jit_kernel.utils import cache_once, load_jit, make_cpp_args +from sglang.kernels.jit.utils import cache_once, load_jit, make_cpp_args logger = logging.getLogger(__name__) diff --git a/python/sglang/jit_kernel/fused_qknorm_rope.py b/python/sglang/jit_kernel/fused_qknorm_rope.py index 00e872020..ca35fbd61 100644 --- a/python/sglang/jit_kernel/fused_qknorm_rope.py +++ b/python/sglang/jit_kernel/fused_qknorm_rope.py @@ -5,7 +5,7 @@ from typing import TYPE_CHECKING, Optional import torch -from sglang.jit_kernel.utils import cache_once, load_jit +from sglang.kernels.jit.utils import cache_once, load_jit from sglang.srt.utils.custom_op import register_custom_op if TYPE_CHECKING: diff --git a/python/sglang/jit_kernel/fused_store_index_cache.py b/python/sglang/jit_kernel/fused_store_index_cache.py index f8b3b1432..cb93b352c 100644 --- a/python/sglang/jit_kernel/fused_store_index_cache.py +++ b/python/sglang/jit_kernel/fused_store_index_cache.py @@ -13,13 +13,13 @@ from typing import TYPE_CHECKING import torch -from sglang.jit_kernel.utils import ( +from sglang.kernel_api_logging import debug_kernel_api +from sglang.kernels.jit.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 diff --git a/python/sglang/jit_kernel/gptq_marlin.py b/python/sglang/jit_kernel/gptq_marlin.py index d3bde5336..a980c1960 100644 --- a/python/sglang/jit_kernel/gptq_marlin.py +++ b/python/sglang/jit_kernel/gptq_marlin.py @@ -4,8 +4,8 @@ from typing import TYPE_CHECKING, Optional import torch -from sglang.jit_kernel.utils import cache_once, load_jit, make_cpp_args from sglang.kernel_api_logging import debug_kernel_api +from sglang.kernels.jit.utils import cache_once, load_jit, make_cpp_args if TYPE_CHECKING: from sgl_kernel.scalar_type import ScalarType diff --git a/python/sglang/jit_kernel/gptq_marlin_repack.py b/python/sglang/jit_kernel/gptq_marlin_repack.py index ea7fe9908..251f3b91f 100644 --- a/python/sglang/jit_kernel/gptq_marlin_repack.py +++ b/python/sglang/jit_kernel/gptq_marlin_repack.py @@ -4,8 +4,8 @@ from typing import TYPE_CHECKING import torch -from sglang.jit_kernel.utils import cache_once, load_jit from sglang.kernel_api_logging import debug_kernel_api +from sglang.kernels.jit.utils import cache_once, load_jit if TYPE_CHECKING: from tvm_ffi.module import Module diff --git a/python/sglang/jit_kernel/hadamard.py b/python/sglang/jit_kernel/hadamard.py index 6e8454749..f8a915bef 100644 --- a/python/sglang/jit_kernel/hadamard.py +++ b/python/sglang/jit_kernel/hadamard.py @@ -4,7 +4,7 @@ from typing import TYPE_CHECKING, Callable import torch -from sglang.jit_kernel.utils import KERNEL_PATH, cache_once, load_jit, make_cpp_args +from sglang.kernels.jit.utils import KERNEL_PATH, cache_once, load_jit, make_cpp_args from sglang.srt.utils.custom_op import register_custom_op if TYPE_CHECKING: diff --git a/python/sglang/jit_kernel/hicache.py b/python/sglang/jit_kernel/hicache.py index 268eabe18..1cce02fdb 100644 --- a/python/sglang/jit_kernel/hicache.py +++ b/python/sglang/jit_kernel/hicache.py @@ -3,8 +3,8 @@ from __future__ import annotations import logging from typing import TYPE_CHECKING -from sglang.jit_kernel.utils import cache_once, load_jit, make_cpp_args from sglang.kernel_api_logging import debug_kernel_api +from sglang.kernels.jit.utils import cache_once, load_jit, make_cpp_args if TYPE_CHECKING: import torch diff --git a/python/sglang/jit_kernel/hisparse.py b/python/sglang/jit_kernel/hisparse.py index b88143da1..0abdc438b 100644 --- a/python/sglang/jit_kernel/hisparse.py +++ b/python/sglang/jit_kernel/hisparse.py @@ -5,7 +5,7 @@ from typing import TYPE_CHECKING import torch -from sglang.jit_kernel.utils import load_jit, make_cpp_args +from sglang.kernels.jit.utils import load_jit, make_cpp_args if TYPE_CHECKING: from tvm_ffi.module import Module diff --git a/python/sglang/jit_kernel/inkling_all_reduce.py b/python/sglang/jit_kernel/inkling_all_reduce.py index b8dcbca3e..050d57ef4 100644 --- a/python/sglang/jit_kernel/inkling_all_reduce.py +++ b/python/sglang/jit_kernel/inkling_all_reduce.py @@ -10,7 +10,7 @@ from typing import TYPE_CHECKING import torch -from sglang.jit_kernel.utils import cache_once, empty_sentinel, load_jit, make_cpp_args +from sglang.kernels.jit.utils import cache_once, empty_sentinel, load_jit, make_cpp_args if TYPE_CHECKING: from tvm_ffi.module import Module diff --git a/python/sglang/jit_kernel/inkling_ar_fused.py b/python/sglang/jit_kernel/inkling_ar_fused.py index 455bf5379..a35515c05 100644 --- a/python/sglang/jit_kernel/inkling_ar_fused.py +++ b/python/sglang/jit_kernel/inkling_ar_fused.py @@ -9,7 +9,7 @@ from typing import TYPE_CHECKING import torch -from sglang.jit_kernel.utils import cache_once, empty_sentinel, load_jit, make_cpp_args +from sglang.kernels.jit.utils import cache_once, empty_sentinel, load_jit, make_cpp_args if TYPE_CHECKING: from tvm_ffi.module import Module diff --git a/python/sglang/jit_kernel/inkling_ar_scattered_sconv.py b/python/sglang/jit_kernel/inkling_ar_scattered_sconv.py index ed742e139..ecfad1489 100644 --- a/python/sglang/jit_kernel/inkling_ar_scattered_sconv.py +++ b/python/sglang/jit_kernel/inkling_ar_scattered_sconv.py @@ -10,7 +10,7 @@ from typing import TYPE_CHECKING import torch -from sglang.jit_kernel.utils import cache_once, load_jit, make_cpp_args +from sglang.kernels.jit.utils import cache_once, load_jit, make_cpp_args if TYPE_CHECKING: from tvm_ffi.module import Module diff --git a/python/sglang/jit_kernel/inkling_attn_prologue.py b/python/sglang/jit_kernel/inkling_attn_prologue.py index 50a00b3e5..252f9ab85 100644 --- a/python/sglang/jit_kernel/inkling_attn_prologue.py +++ b/python/sglang/jit_kernel/inkling_attn_prologue.py @@ -7,7 +7,7 @@ from typing import TYPE_CHECKING import torch -from sglang.jit_kernel.utils import ( +from sglang.kernels.jit.utils import ( cache_once, empty_sentinel, is_arch_support_pdl, diff --git a/python/sglang/jit_kernel/inkling_gate_topk_renorm.py b/python/sglang/jit_kernel/inkling_gate_topk_renorm.py index 46ccf054c..5601a2859 100644 --- a/python/sglang/jit_kernel/inkling_gate_topk_renorm.py +++ b/python/sglang/jit_kernel/inkling_gate_topk_renorm.py @@ -22,7 +22,7 @@ from typing import TYPE_CHECKING import torch -from sglang.jit_kernel.utils import cache_once, load_jit +from sglang.kernels.jit.utils import cache_once, load_jit if TYPE_CHECKING: from tvm_ffi.module import Module diff --git a/python/sglang/jit_kernel/inkling_rel_proj.py b/python/sglang/jit_kernel/inkling_rel_proj.py index ebefd9cfd..da7baf3fe 100644 --- a/python/sglang/jit_kernel/inkling_rel_proj.py +++ b/python/sglang/jit_kernel/inkling_rel_proj.py @@ -10,7 +10,7 @@ from typing import TYPE_CHECKING import torch -from sglang.jit_kernel.utils import ( +from sglang.kernels.jit.utils import ( cache_once, empty_sentinel, is_arch_support_pdl, diff --git a/python/sglang/jit_kernel/inkling_row_scale.py b/python/sglang/jit_kernel/inkling_row_scale.py index 4b831410a..48ca31388 100644 --- a/python/sglang/jit_kernel/inkling_row_scale.py +++ b/python/sglang/jit_kernel/inkling_row_scale.py @@ -9,7 +9,7 @@ from typing import TYPE_CHECKING import torch -from sglang.jit_kernel.utils import ( +from sglang.kernels.jit.utils import ( cache_once, is_arch_support_pdl, load_jit, diff --git a/python/sglang/jit_kernel/inkling_sconv.py b/python/sglang/jit_kernel/inkling_sconv.py index a3121dcea..871a9a12c 100644 --- a/python/sglang/jit_kernel/inkling_sconv.py +++ b/python/sglang/jit_kernel/inkling_sconv.py @@ -10,7 +10,7 @@ from typing import TYPE_CHECKING import torch -from sglang.jit_kernel.utils import cache_once, load_jit, make_cpp_args +from sglang.kernels.jit.utils import cache_once, load_jit, make_cpp_args if TYPE_CHECKING: from tvm_ffi.module import Module diff --git a/python/sglang/jit_kernel/kpool_topk_transform.py b/python/sglang/jit_kernel/kpool_topk_transform.py index 52dfdcf61..330293a68 100644 --- a/python/sglang/jit_kernel/kpool_topk_transform.py +++ b/python/sglang/jit_kernel/kpool_topk_transform.py @@ -4,7 +4,7 @@ from typing import TYPE_CHECKING, Optional import torch -from sglang.jit_kernel.utils import cache_once, load_jit +from sglang.kernels.jit.utils import cache_once, load_jit if TYPE_CHECKING: from tvm_ffi.module import Module diff --git a/python/sglang/jit_kernel/kv_canary/plan/entries_kernel.py b/python/sglang/jit_kernel/kv_canary/plan/entries_kernel.py index f63de2484..a5287f12f 100644 --- a/python/sglang/jit_kernel/kv_canary/plan/entries_kernel.py +++ b/python/sglang/jit_kernel/kv_canary/plan/entries_kernel.py @@ -4,7 +4,7 @@ from typing import TYPE_CHECKING, Optional import torch -from sglang.jit_kernel.utils import cache_once, load_jit, make_cpp_args +from sglang.kernels.jit.utils import cache_once, load_jit, make_cpp_args if TYPE_CHECKING: from tvm_ffi.module import Module diff --git a/python/sglang/jit_kernel/kv_canary/verify.py b/python/sglang/jit_kernel/kv_canary/verify.py index b7fd260bd..b17f484e9 100644 --- a/python/sglang/jit_kernel/kv_canary/verify.py +++ b/python/sglang/jit_kernel/kv_canary/verify.py @@ -7,7 +7,7 @@ from typing import TYPE_CHECKING, Final import torch from sglang.jit_kernel.kv_canary import consts -from sglang.jit_kernel.utils import cache_once, load_jit, make_cpp_args +from sglang.kernels.jit.utils import cache_once, load_jit, make_cpp_args if TYPE_CHECKING: from tvm_ffi.module import Module diff --git a/python/sglang/jit_kernel/kv_canary/write.py b/python/sglang/jit_kernel/kv_canary/write.py index 6926a2dc9..106e6118a 100644 --- a/python/sglang/jit_kernel/kv_canary/write.py +++ b/python/sglang/jit_kernel/kv_canary/write.py @@ -11,7 +11,7 @@ from sglang.jit_kernel.kv_canary.verify import ( _assert_contiguous, _build_real_kv_source_abi, ) -from sglang.jit_kernel.utils import cache_once, load_jit +from sglang.kernels.jit.utils import cache_once, load_jit if TYPE_CHECKING: from tvm_ffi.module import Module diff --git a/python/sglang/jit_kernel/kvcache.py b/python/sglang/jit_kernel/kvcache.py index b611e7657..8529a72c3 100644 --- a/python/sglang/jit_kernel/kvcache.py +++ b/python/sglang/jit_kernel/kvcache.py @@ -5,7 +5,7 @@ from typing import TYPE_CHECKING import torch -from sglang.jit_kernel.utils import ( +from sglang.kernels.jit.utils import ( cache_once, is_arch_support_pdl, load_jit, diff --git a/python/sglang/jit_kernel/lplb/cuda_solver.py b/python/sglang/jit_kernel/lplb/cuda_solver.py index 60a78c751..61ab5cfe6 100644 --- a/python/sglang/jit_kernel/lplb/cuda_solver.py +++ b/python/sglang/jit_kernel/lplb/cuda_solver.py @@ -17,7 +17,7 @@ from typing import TYPE_CHECKING, Optional import torch -from sglang.jit_kernel.utils import ( +from sglang.kernels.jit.utils import ( cache_once, get_jit_cuda_arch, load_jit, diff --git a/python/sglang/jit_kernel/minimax_decode_topk.py b/python/sglang/jit_kernel/minimax_decode_topk.py index 26940d66d..bc44e580f 100644 --- a/python/sglang/jit_kernel/minimax_decode_topk.py +++ b/python/sglang/jit_kernel/minimax_decode_topk.py @@ -17,7 +17,7 @@ from typing import TYPE_CHECKING, Tuple import torch -from sglang.jit_kernel.utils import cache_once, load_jit, make_cpp_args +from sglang.kernels.jit.utils import cache_once, load_jit, make_cpp_args if TYPE_CHECKING: from tvm_ffi.module import Module diff --git a/python/sglang/jit_kernel/minimax_qknorm_rope.py b/python/sglang/jit_kernel/minimax_qknorm_rope.py index 7815c398f..a67d25928 100644 --- a/python/sglang/jit_kernel/minimax_qknorm_rope.py +++ b/python/sglang/jit_kernel/minimax_qknorm_rope.py @@ -22,7 +22,7 @@ from typing import TYPE_CHECKING, List, Sequence, Tuple import torch -from sglang.jit_kernel.utils import ( +from sglang.kernels.jit.utils import ( cache_once, is_arch_support_pdl, load_jit, diff --git a/python/sglang/jit_kernel/minimax_quant_ue8m0.py b/python/sglang/jit_kernel/minimax_quant_ue8m0.py index 240e9cf75..d67786eca 100644 --- a/python/sglang/jit_kernel/minimax_quant_ue8m0.py +++ b/python/sglang/jit_kernel/minimax_quant_ue8m0.py @@ -4,7 +4,7 @@ from typing import TYPE_CHECKING, Tuple import torch -from sglang.jit_kernel.utils import ( +from sglang.kernels.jit.utils import ( cache_once, is_arch_support_pdl, load_jit, diff --git a/python/sglang/jit_kernel/minimax_store_kv_index.py b/python/sglang/jit_kernel/minimax_store_kv_index.py index 2c5fd4408..82c5a3a36 100644 --- a/python/sglang/jit_kernel/minimax_store_kv_index.py +++ b/python/sglang/jit_kernel/minimax_store_kv_index.py @@ -4,7 +4,7 @@ from typing import TYPE_CHECKING, Optional import torch -from sglang.jit_kernel.utils import ( +from sglang.kernels.jit.utils import ( cache_once, is_arch_support_pdl, load_jit, diff --git a/python/sglang/jit_kernel/mla_kv_pack_quantize_fp8.py b/python/sglang/jit_kernel/mla_kv_pack_quantize_fp8.py index bf1024eb0..2d221af4d 100644 --- a/python/sglang/jit_kernel/mla_kv_pack_quantize_fp8.py +++ b/python/sglang/jit_kernel/mla_kv_pack_quantize_fp8.py @@ -11,7 +11,7 @@ import torch import triton import triton.language as tl -from sglang.jit_kernel.utils import is_arch_support_pdl +from sglang.kernels.jit.utils import is_arch_support_pdl @triton.jit diff --git a/python/sglang/jit_kernel/moe_align.py b/python/sglang/jit_kernel/moe_align.py index ee136f1a8..0b071056a 100644 --- a/python/sglang/jit_kernel/moe_align.py +++ b/python/sglang/jit_kernel/moe_align.py @@ -4,7 +4,7 @@ from typing import TYPE_CHECKING import torch -from sglang.jit_kernel.utils import cache_once, load_jit, make_cpp_args +from sglang.kernels.jit.utils import cache_once, load_jit, make_cpp_args if TYPE_CHECKING: from tvm_ffi.module import Module diff --git a/python/sglang/jit_kernel/moe_finalize_fuse_shared.py b/python/sglang/jit_kernel/moe_finalize_fuse_shared.py index 743239732..9e862f2c7 100644 --- a/python/sglang/jit_kernel/moe_finalize_fuse_shared.py +++ b/python/sglang/jit_kernel/moe_finalize_fuse_shared.py @@ -4,7 +4,7 @@ from typing import Optional import torch -from sglang.jit_kernel.utils import cache_once, load_jit +from sglang.kernels.jit.utils import cache_once, load_jit @cache_once diff --git a/python/sglang/jit_kernel/moe_fused_gate.py b/python/sglang/jit_kernel/moe_fused_gate.py index 6cdfd9c3a..922d5e929 100644 --- a/python/sglang/jit_kernel/moe_fused_gate.py +++ b/python/sglang/jit_kernel/moe_fused_gate.py @@ -7,8 +7,8 @@ import torch import triton import triton.language as tl -from sglang.jit_kernel.utils import cache_once, is_arch_support_pdl, load_jit from sglang.kernel_api_logging import debug_kernel_api +from sglang.kernels.jit.utils import cache_once, is_arch_support_pdl, load_jit if TYPE_CHECKING: from tvm_ffi.module import Module diff --git a/python/sglang/jit_kernel/moe_lora_align.py b/python/sglang/jit_kernel/moe_lora_align.py index f18ad7ca0..da9486823 100644 --- a/python/sglang/jit_kernel/moe_lora_align.py +++ b/python/sglang/jit_kernel/moe_lora_align.py @@ -4,7 +4,7 @@ from typing import TYPE_CHECKING, Optional import torch -from sglang.jit_kernel.utils import cache_once, load_jit, make_cpp_args +from sglang.kernels.jit.utils import cache_once, load_jit, make_cpp_args if TYPE_CHECKING: from tvm_ffi.module import Module diff --git a/python/sglang/jit_kernel/moe_permute_prepare.py b/python/sglang/jit_kernel/moe_permute_prepare.py index 679c52332..1bbd40ed3 100644 --- a/python/sglang/jit_kernel/moe_permute_prepare.py +++ b/python/sglang/jit_kernel/moe_permute_prepare.py @@ -4,7 +4,7 @@ from typing import TYPE_CHECKING, Tuple import torch -from sglang.jit_kernel.utils import cache_once, load_jit +from sglang.kernels.jit.utils import cache_once, load_jit from sglang.srt.utils.custom_op import register_custom_op if TYPE_CHECKING: diff --git a/python/sglang/jit_kernel/moe_topk_sigmoid.py b/python/sglang/jit_kernel/moe_topk_sigmoid.py index 7fbdd22a6..b92d0ec0e 100644 --- a/python/sglang/jit_kernel/moe_topk_sigmoid.py +++ b/python/sglang/jit_kernel/moe_topk_sigmoid.py @@ -4,7 +4,7 @@ from typing import TYPE_CHECKING, Optional import torch -from sglang.jit_kernel.utils import cache_once, load_jit, make_cpp_args +from sglang.kernels.jit.utils import cache_once, load_jit, make_cpp_args from sglang.srt.utils.custom_op import register_custom_op if TYPE_CHECKING: diff --git a/python/sglang/jit_kernel/moe_wna16_marlin.py b/python/sglang/jit_kernel/moe_wna16_marlin.py index e9a8cd253..bbac46d10 100644 --- a/python/sglang/jit_kernel/moe_wna16_marlin.py +++ b/python/sglang/jit_kernel/moe_wna16_marlin.py @@ -4,8 +4,8 @@ from typing import TYPE_CHECKING, Optional import torch -from sglang.jit_kernel.utils import cache_once, load_jit, make_cpp_args from sglang.kernel_api_logging import debug_kernel_api +from sglang.kernels.jit.utils import cache_once, load_jit, make_cpp_args if TYPE_CHECKING: from sgl_kernel.scalar_type import ScalarType diff --git a/python/sglang/jit_kernel/mxfp8.py b/python/sglang/jit_kernel/mxfp8.py index 2f0a91f9f..8f65ba7db 100644 --- a/python/sglang/jit_kernel/mxfp8.py +++ b/python/sglang/jit_kernel/mxfp8.py @@ -4,7 +4,7 @@ from typing import TYPE_CHECKING import torch -from sglang.jit_kernel.utils import ( +from sglang.kernels.jit.utils import ( cache_once, load_jit, make_cpp_args, diff --git a/python/sglang/jit_kernel/ngram_corpus.py b/python/sglang/jit_kernel/ngram_corpus.py index d2121417c..90b6a2546 100644 --- a/python/sglang/jit_kernel/ngram_corpus.py +++ b/python/sglang/jit_kernel/ngram_corpus.py @@ -7,7 +7,7 @@ import numpy as np import torch import tvm_ffi -from sglang.jit_kernel.utils import cache_once, load_jit +from sglang.kernels.jit.utils import cache_once, load_jit _MATCH_TYPE_MAP = {"BFS": 0, "PROB": 1} diff --git a/python/sglang/jit_kernel/ngram_embedding.py b/python/sglang/jit_kernel/ngram_embedding.py index ea7da20ba..3b2de30c0 100644 --- a/python/sglang/jit_kernel/ngram_embedding.py +++ b/python/sglang/jit_kernel/ngram_embedding.py @@ -2,8 +2,8 @@ from __future__ import annotations from typing import TYPE_CHECKING -from sglang.jit_kernel.utils import cache_once, load_jit from sglang.kernel_api_logging import debug_kernel_api +from sglang.kernels.jit.utils import cache_once, load_jit if TYPE_CHECKING: import torch diff --git a/python/sglang/jit_kernel/norm.py b/python/sglang/jit_kernel/norm.py index 4fb3c9451..b1a334507 100644 --- a/python/sglang/jit_kernel/norm.py +++ b/python/sglang/jit_kernel/norm.py @@ -1,179 +1,5 @@ -from __future__ import annotations +"""Compatibility shim (RFC #29630 Phase 4) -> sglang.kernels.ops.layernorm._jit_norm.""" -import logging -from typing import TYPE_CHECKING, Optional +from sglang.kernels.ops.layernorm import _jit_norm as _impl -import torch - -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 - - -logger = logging.getLogger(__name__) - - -@cache_once -def _jit_qknorm_module(head_dim: int, dtype: torch.dtype) -> Module: - args = make_cpp_args(head_dim, is_arch_support_pdl(), dtype) - return load_jit( - "qknorm", - *args, - cuda_files=["elementwise/qknorm.cuh"], - cuda_wrappers=[("qknorm", f"QKNormKernel<{args}>::run")], - ) - - -_RMSNORM_WARP_SIZES = frozenset({64, 128, 256}) -_RMSNORM_MAX_HIDDEN_SIZE = 16384 -_RMSNORM_HALF_BLOCK_MIN_SIZE = 2048 - - -def _is_supported_rmsnorm_hidden_size(d: int) -> bool: - return d in _RMSNORM_WARP_SIZES or ( - (d > 256 and d % 256 == 0 and d <= 8192) - or (d >= 8192 and d % 512 == 0 and d <= 16384) - ) - - -def _rmsnorm_kernel_class(hidden_size: int) -> str: - if hidden_size in _RMSNORM_WARP_SIZES: - return "RMSNormWarpKernel" - if hidden_size == 512: - return "RMSNormHalfKernel" - if hidden_size >= _RMSNORM_HALF_BLOCK_MIN_SIZE: - if hidden_size % 512 == 0: - return "RMSNormHalfKernel" - return "RMSNormKernel" - - -@cache_once -def _jit_rmsnorm_module(hidden_size: int, dtype: torch.dtype) -> Module: - args = make_cpp_args(hidden_size, is_arch_support_pdl(), dtype) - kernel_class = f"{_rmsnorm_kernel_class(hidden_size)}<{args}>" - return load_jit( - "rmsnorm", - *args, - cuda_files=["elementwise/rmsnorm.cuh"], - cuda_wrappers=[("rmsnorm", f"{kernel_class}::run")], - ) - - -def is_supported_jit_fused_add_rmsnorm_hidden_size(hidden_size: int) -> bool: - return hidden_size > 0 and hidden_size % 16 == 0 and hidden_size <= 8192 - - -@cache_once -def _jit_fused_add_rmsnorm_module( - dtype: torch.dtype, cast_x_before_out_mul: bool -) -> Module: - args = make_cpp_args(cast_x_before_out_mul, dtype) - return load_jit( - "fused_add_rmsnorm", - *args, - cuda_files=["elementwise/fused_add_rmsnorm.cuh"], - cuda_wrappers=[("fused_add_rmsnorm", f"FusedAddRMSNormKernel<{args}>::run")], - ) - - -@cache_once -def _jit_qknorm_across_heads_module(dtype: torch.dtype) -> Module: - args = make_cpp_args(dtype) - return load_jit( - "qknorm_across_heads", - *args, - cuda_files=["elementwise/qknorm_across_heads.cuh"], - cuda_wrappers=[ - ("qknorm_across_heads", f"QKNormAcrossHeadsKernel<{args}>::run") - ], - ) - - -@torch.compiler.assume_constant_result -@cache_once -def can_use_fused_inplace_qknorm(head_dim: int, dtype: torch.dtype) -> bool: - if head_dim not in [64, 128, 256, 512, 1024]: - logger.warning(f"Unsupported head_dim={head_dim} for JIT QK-Norm kernel") - return False - try: - _jit_qknorm_module(head_dim, dtype) - return True - except Exception as e: - logger.warning(f"Failed to load JIT QK-Norm kernel: {e}") - return False - - -@debug_kernel_api -def fused_inplace_qknorm( - q: torch.Tensor, - k: torch.Tensor, - q_weight: torch.Tensor, - k_weight: torch.Tensor, - eps: float = 1e-6, - *, - head_dim: int = 0, -) -> None: - head_dim = head_dim or q.size(-1) - module = _jit_qknorm_module(head_dim, q.dtype) - module.qknorm(q, k, q_weight, k_weight, eps) - - -@debug_kernel_api -def rmsnorm( - input: torch.Tensor, - weight: torch.Tensor, - out: Optional[torch.Tensor] = None, - eps: float = 1e-6, -) -> None: - out = out if out is not None else input - hidden_size = input.size(-1) - if not _is_supported_rmsnorm_hidden_size(hidden_size): - raise RuntimeError( - f"jit rmsnorm: unsupported hidden_size={hidden_size}. " - f"Supported: {sorted(_RMSNORM_WARP_SIZES)}, and multiples of 256 in " - f"(256, {_RMSNORM_MAX_HIDDEN_SIZE}]." - ) - module = _jit_rmsnorm_module(hidden_size, input.dtype) - module.rmsnorm(input, weight, out, eps) - - -@debug_kernel_api -def fused_add_rmsnorm( - input: torch.Tensor, - residual: torch.Tensor, - weight: torch.Tensor, - eps: float = 1e-6, - *, - cast_x_before_out_mul: bool = False, -) -> None: - module = _jit_fused_add_rmsnorm_module(input.dtype, cast_x_before_out_mul) - module.fused_add_rmsnorm(input, residual, weight, eps) - - -@debug_kernel_api -def fused_inplace_qknorm_across_heads( - q: torch.Tensor, - k: torch.Tensor, - q_weight: torch.Tensor, - k_weight: torch.Tensor, - eps: float = 1e-6, -) -> None: - """ - Fused inplace QK normalization across all heads. - - Args: - q: Query tensor of shape [batch_size, num_heads * head_dim] - k: Key tensor of shape [batch_size, num_heads * head_dim] - q_weight: Query weight tensor of shape [num_heads * head_dim] - k_weight: Key weight tensor of shape [num_heads * head_dim] - eps: Epsilon for numerical stability - """ - module = _jit_qknorm_across_heads_module(q.dtype) - module.qknorm_across_heads(q, k, q_weight, k_weight, eps) +globals().update({k: getattr(_impl, k) for k in dir(_impl) if not k.startswith("__")}) diff --git a/python/sglang/jit_kernel/per_tensor_quant_fp8.py b/python/sglang/jit_kernel/per_tensor_quant_fp8.py index 40aee2607..b2b987ec0 100644 --- a/python/sglang/jit_kernel/per_tensor_quant_fp8.py +++ b/python/sglang/jit_kernel/per_tensor_quant_fp8.py @@ -1,77 +1,5 @@ -from __future__ import annotations +"""Compatibility shim (RFC #29630 Phase 4) -> sglang.kernels.ops.quantization._jit_per_tensor_quant_fp8.""" -from typing import TYPE_CHECKING +from sglang.kernels.ops.quantization import _jit_per_tensor_quant_fp8 as _impl -import torch - -from sglang.jit_kernel.utils import cache_once, load_jit, make_cpp_args -from sglang.srt.utils.custom_op import register_custom_op - -if TYPE_CHECKING: - from tvm_ffi.module import Module - - -@cache_once -def _jit_per_tensor_quant_fp8_module(is_static: bool, dtype: torch.dtype) -> Module: - args = make_cpp_args(is_static, dtype) - return load_jit( - "per_tensor_quant_fp8", - *args, - cuda_files=["gemm/per_tensor_quant_fp8.cuh"], - cuda_wrappers=[("per_tensor_quant_fp8", f"per_tensor_quant_fp8<{args}>")], - ) - - -@register_custom_op( - op_name="per_tensor_quant_fp8", - mutates_args=["output_q", "output_s"], -) -def per_tensor_quant_fp8( - input: torch.Tensor, - output_q: torch.Tensor, - output_s: torch.Tensor, - is_static: bool = False, -) -> None: - """ - Per-tensor quantization to FP8 format. - - Args: - input: Input tensor to quantize (float, half, or bfloat16) - output_q: Output quantized tensor (fp8_e4m3) - output_s: Output scale tensor (float scalar or 1D tensor with 1 element) - is_static: If True, assumes scale is pre-computed and skips absmax computation - """ - module = _jit_per_tensor_quant_fp8_module(is_static, input.dtype) - module.per_tensor_quant_fp8(input.view(-1), output_q.view(-1), output_s.view(-1)) - - -@cache_once -def _jit_per_tensor_absmax_fp8_module(dtype: torch.dtype) -> Module: - args = make_cpp_args(dtype) - return load_jit( - "per_tensor_absmax_fp8", - *args, - cuda_files=["gemm/per_tensor_quant_fp8.cuh"], - cuda_wrappers=[("per_tensor_absmax_fp8", f"per_tensor_absmax_fp8<{args}>")], - ) - - -@register_custom_op( - op_name="per_tensor_absmax_fp8", - mutates_args=["output_s"], -) -def per_tensor_absmax_fp8( - input: torch.Tensor, - output_s: torch.Tensor, -) -> None: - """Compute scale = max(abs(input)) / fp8_e4m3_max via atomic-max reduction. - - The caller must zero-initialise ``output_s`` before the call (the kernel - uses ``atomic_max`` across blocks, so starting from 0 is required). - - Args: - input: Input tensor (float16, bfloat16, or float32). Any shape. - output_s: Pre-allocated float32 tensor of shape (1,), zero-initialised. - """ - module = _jit_per_tensor_absmax_fp8_module(input.dtype) - module.per_tensor_absmax_fp8(input.view(-1), output_s.view(-1)) +globals().update({k: getattr(_impl, k) for k in dir(_impl) if not k.startswith("__")}) diff --git a/python/sglang/jit_kernel/per_token_group_quant.py b/python/sglang/jit_kernel/per_token_group_quant.py index a3b5eb67b..314c8c056 100644 --- a/python/sglang/jit_kernel/per_token_group_quant.py +++ b/python/sglang/jit_kernel/per_token_group_quant.py @@ -1,223 +1,5 @@ -from __future__ import annotations +"""Compatibility shim (RFC #29630 Phase 4) -> sglang.kernels.ops.quantization._jit_per_token_group_quant.""" -from typing import TYPE_CHECKING, Optional, Tuple +from sglang.kernels.ops.quantization import _jit_per_token_group_quant as _impl -import torch - -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 -from sglang.srt.utils.custom_op import register_custom_op - -if TYPE_CHECKING: - from tvm_ffi.module import Module - -_SUPPORTED_INPUT_DTYPES = (torch.bfloat16, torch.float16) -_SUPPORTED_OUTPUT_DTYPES = (torch.float8_e4m3fn, torch.int8) -_SUPPORTED_GROUP_SIZES = (16, 32, 64, 128, 256) - - -@cache_once -def _jit_module( - in_dtype: torch.dtype, - out_dtype: torch.dtype, - group_size: int, - scale_ue8m0: bool, - row_major: bool, - aligned: bool, - fuse_silu_and_mul: bool, - masked_layout: bool, - use_pdl: bool, -) -> Module: - assert in_dtype in _SUPPORTED_INPUT_DTYPES - assert out_dtype in _SUPPORTED_OUTPUT_DTYPES - assert group_size in _SUPPORTED_GROUP_SIZES - trait_args = make_cpp_args( - in_dtype, - out_dtype, - group_size, - scale_ue8m0, - row_major, - aligned, - fuse_silu_and_mul, - use_pdl, - ) - launcher = ( - "PerTokenGroupQuantMaskedKernel" - if masked_layout - else "PerTokenGroupQuantFlatKernel" - ) - return load_jit( - "per_token_group_quant", - *trait_args, - "masked" if masked_layout else "flat", - cuda_files=["gemm/per_token_group_quant.cuh"], - cuda_wrappers=[("per_token_group_quant", f"{launcher}<{trait_args}>::run")], - extra_cuda_cflags=["--use_fast_math"], - ) - - -def _infer_scale_layout( - output_s: torch.Tensor, scale_ue8m0: bool, num_groups: int -) -> Tuple[bool, bool]: - """Return ``(row_major, aligned)`` for ``output_s``. - - Column-major (transposed) scale buffers have token stride 1 and a larger - group stride; row-major buffers are contiguous. - """ - row_major = output_s.stride(-2) >= output_s.stride(-1) - if output_s.dtype == torch.int32: - if not scale_ue8m0: - raise ValueError("int32-packed scale buffers require scale_ue8m0=True") - aligned = num_groups % 4 == 0 - return row_major, aligned - if output_s.dtype == torch.float32: - if scale_ue8m0: - raise ValueError("scale_ue8m0=True requires an int32-packed output_s") - return row_major, True - raise ValueError(f"Unsupported output_s dtype {output_s.dtype}") - - -@register_custom_op( - op_name="per_token_group_quant", - mutates_args=["output_q", "output_s"], -) -def _per_token_group_quant_custom_op( - input: torch.Tensor, - output_q: torch.Tensor, - output_s: torch.Tensor, - group_size: int, - scale_ue8m0: bool = False, - fuse_silu_and_mul: bool = False, - masked_m: Optional[torch.Tensor] = None, - expected_m: Optional[int] = None, -) -> None: - num_groups = output_q.shape[-1] // group_size - row_major, aligned = _infer_scale_layout(output_s, scale_ue8m0, num_groups) - module = _jit_module( - input.dtype, - output_q.dtype, - int(group_size), - bool(scale_ue8m0), - row_major, - aligned, - bool(fuse_silu_and_mul), - masked_m is not None, - is_arch_support_pdl(), - ) - if masked_m is not None: - module.per_token_group_quant( - input, output_q, output_s, masked_m, int(expected_m or -1) - ) - else: - module.per_token_group_quant(input, output_q, output_s) - - -def _allocate_outputs( - input: torch.Tensor, - group_size: int, - out_dtype: torch.dtype, - scale_ue8m0: bool, - column_major_scales: bool, - fuse_silu_and_mul: bool, -) -> Tuple[torch.Tensor, torch.Tensor]: - """Allocate ``(output_q, output_s)`` in the requested major mode / scale - format, selected by ``(column_major_scales, scale_ue8m0)``.""" - hidden = input.shape[-1] // (2 if fuse_silu_and_mul else 1) - out_shape = (*input.shape[:-1], hidden) - output_q = torch.empty(out_shape, device=input.device, dtype=out_dtype) - - num_groups = hidden // group_size - if scale_ue8m0 and not column_major_scales: - # Row-major packed UE8M0: int32 [..., ceil(ng/4)] contiguous (an - # unaligned ng leaves a partially-used last int32 that the kernel zero- - # pads). The shared create_*_output_scale helper does not produce this - # layout. - output_s = torch.empty( - (*out_shape[:-1], (num_groups + 3) // 4), - device=input.device, - dtype=torch.int32, - ) - else: - from sglang.kernels.ops.quantization.fp8_kernel import ( - create_per_token_group_quant_fp8_output_scale, - ) - - output_s = create_per_token_group_quant_fp8_output_scale( - x_shape=out_shape, - device=input.device, - group_size=group_size, - column_major_scales=column_major_scales, - scale_tma_aligned=column_major_scales, - scale_ue8m0=scale_ue8m0, - ) - return output_q, output_s - - -@debug_kernel_api -def per_token_group_quant( - input: torch.Tensor, - output_q: Optional[torch.Tensor] = None, - output_s: Optional[torch.Tensor] = None, - group_size: int = 128, - scale_ue8m0: bool = False, - fuse_silu_and_mul: bool = False, - masked_m: Optional[torch.Tensor] = None, - expected_m: Optional[int] = None, - *, - out_dtype: Optional[torch.dtype] = None, - column_major_scales: bool = False, -) -> Tuple[torch.Tensor, torch.Tensor]: - """Per-token-group quantization. Returns ``(output_q, output_s)``. - - ``output_q`` / ``output_s`` are optional: pass them to quantize into - caller-owned buffers, or omit both to have them allocated per ``out_dtype`` - (default fp8_e4m3), ``scale_ue8m0`` and ``column_major_scales``. Either way - the two tensors are returned. - - Input / output shapes: - vanilla: input [T, hidden], output_q [T, hidden] - fuse_silu_and_mul: input [T, hidden*2], output_q [T, hidden] - masked (+ above): input [E, T_pad, ...], output_q [E, T_pad, hidden], - masked_m [E] int32 - ``output_s`` scale layouts (inferred from a supplied buffer's dtype/strides, - or allocated to match when omitted): - float32 contiguous -> row-major fp32 scales - float32 transposed -> col-major fp32 scales (TMA-aligned view) - int32 transposed -> col-major UE8M0 bytes packed 4-per-int32 - int32 contiguous -> row-major UE8M0 bytes packed 4-per-int32 - The packed layouts require ``scale_ue8m0=True``. - - ``expected_m`` (masked only) is an optional expected-tokens-per-expert hint. - - Inputs are bf16/fp16; group size is one of 16/32/64/128/256; the quant range - follows ``output_q.dtype`` (fp8_e4m3: +-448, int8: [-128, 127]). - """ - if output_q is None: - assert output_s is None - output_q, output_s = _allocate_outputs( - input, - group_size, - out_dtype or torch.float8_e4m3fn, - scale_ue8m0, - column_major_scales, - fuse_silu_and_mul, - ) - else: - assert output_s is not None - assert out_dtype is None or out_dtype == output_q.dtype - _per_token_group_quant_custom_op( - input=input, - output_q=output_q, - output_s=output_s, - group_size=group_size, - scale_ue8m0=scale_ue8m0, - fuse_silu_and_mul=fuse_silu_and_mul, - masked_m=masked_m, - expected_m=expected_m, - ) - return output_q, output_s +globals().update({k: getattr(_impl, k) for k in dir(_impl) if not k.startswith("__")}) diff --git a/python/sglang/jit_kernel/per_token_group_quant_8bit_v2.py b/python/sglang/jit_kernel/per_token_group_quant_8bit_v2.py index 84e6dc5fb..0b24d721f 100644 --- a/python/sglang/jit_kernel/per_token_group_quant_8bit_v2.py +++ b/python/sglang/jit_kernel/per_token_group_quant_8bit_v2.py @@ -1,136 +1,5 @@ -"""DEPRECATED: superseded by ``sglang.jit_kernel.per_token_group_quant`` (the -default CUDA path). No sglang runtime code may call this kernel; it is kept -only as the perf baseline for the per_token_group_quant benchmarks and its own -bit-parity tests, and will be deleted once those move to torch references. -""" +"""Compatibility shim (RFC #29630 Phase 4) -> sglang.kernels.ops.quantization._jit_per_token_group_quant_8bit_v2.""" -from __future__ import annotations +from sglang.kernels.ops.quantization import _jit_per_token_group_quant_8bit_v2 as _impl -from typing import TYPE_CHECKING, Optional - -import torch - -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 -from sglang.srt.utils.custom_op import register_custom_op - -if TYPE_CHECKING: - from tvm_ffi.module import Module - - -@cache_once -def _jit_module(in_dtype: torch.dtype, out_dtype: torch.dtype, use_pdl: bool) -> Module: - args = make_cpp_args(in_dtype, out_dtype, use_pdl) - return load_jit( - "per_token_group_quant_8bit_v2", - *args, - cuda_files=["gemm/per_token_group_quant_8bit_v2.cuh"], - cuda_wrappers=[ - ( - "per_token_group_quant_8bit_v2", - f"PerTokenGroupQuant8bitV2Kernel<{args}>::run", - ) - ], - # Match the AOT sgl-kernel build (-use_fast_math) so the FP8 scale - # division/rounding is bit-identical to sgl_per_token_group_quant_8bit_v2. - extra_cuda_cflags=["--use_fast_math"], - ) - - -@register_custom_op( - op_name="per_token_group_quant_8bit_v2", - mutates_args=["output_q", "output_s"], -) -def _per_token_group_quant_8bit_v2_custom_op( - input: torch.Tensor, - output_q: torch.Tensor, - output_s: torch.Tensor, - group_size: int, - eps: float, - min_8bit: float, - max_8bit: float, - scale_ue8m0: bool = False, - fuse_silu_and_mul: bool = False, - masked_m: Optional[torch.Tensor] = None, -) -> None: - """Opaque custom-op boundary around the JIT v2 kernel. - - Registering this as a custom op (instead of calling the tvm-ffi module - directly) keeps torch.compile / piecewise-CUDA-graph from tracing into the - tvm-ffi ``Function.__call__`` (which Dynamo cannot trace). All shape-derived - scalars are computed here and passed to the kernel. - - Layouts (matching the AOT v2): - vanilla: input (num_tokens, hidden), output_q (num_tokens, hidden) - fuse_silu_and_mul: input (num_tokens, hidden*2), output_q (num_tokens, hidden) - fuse_silu_and_mul+masked: input (num_experts, tokens_pad, hidden*2), - output_q (num_experts, tokens_pad, hidden), masked_m (num_experts,) - """ - masked_layout = masked_m is not None - numel = input.numel() - num_groups = numel // group_size // (2 if fuse_silu_and_mul else 1) - if num_groups == 0: # empty input -> grid 0 -> cudaErrorInvalidConfiguration - return - num_local_experts = input.shape[0] if masked_layout else 1 - last = output_q.dim() - 1 - is_column_major = output_s.stride(last - 1) < output_s.stride(last) - hidden_dim_num_groups = output_q.shape[last] // group_size - num_tokens_per_expert = output_q.shape[last - 1] - scale_expert_stride = output_s.stride(0) if masked_layout else 0 - scale_hidden_stride = output_s.stride(last) - - module = _jit_module(input.dtype, output_q.dtype, is_arch_support_pdl()) - module.per_token_group_quant_8bit_v2( - input, - output_q, - output_s, - masked_m if masked_layout else input, # unused (nullptr) when not masked - int(group_size), - bool(scale_ue8m0), - bool(fuse_silu_and_mul), - bool(masked_layout), - int(num_groups), - int(num_local_experts), - bool(is_column_major), - int(hidden_dim_num_groups), - int(num_tokens_per_expert), - int(scale_expert_stride), - int(scale_hidden_stride), - ) - - -@debug_kernel_api -def per_token_group_quant_8bit_v2( - input: torch.Tensor, - output_q: torch.Tensor, - output_s: torch.Tensor, - group_size: int, - eps: float, - min_8bit: float, - max_8bit: float, - scale_ue8m0: bool = False, - fuse_silu_and_mul: bool = False, - masked_m: Optional[torch.Tensor] = None, -) -> None: - """JIT port of sgl_per_token_group_quant_8bit_v2 (full feature parity). - - Wraps the registered custom op so torch.compile / piecewise CUDA graph treat - the tvm-ffi kernel call as an opaque boundary. - """ - _per_token_group_quant_8bit_v2_custom_op( - input=input, - output_q=output_q, - output_s=output_s, - group_size=group_size, - eps=eps, - min_8bit=min_8bit, - max_8bit=max_8bit, - scale_ue8m0=scale_ue8m0, - fuse_silu_and_mul=fuse_silu_and_mul, - masked_m=masked_m, - ) +globals().update({k: getattr(_impl, k) for k in dir(_impl) if not k.startswith("__")}) diff --git a/python/sglang/jit_kernel/resolve_future_token_ids.py b/python/sglang/jit_kernel/resolve_future_token_ids.py index d0068a4ad..43e786b10 100644 --- a/python/sglang/jit_kernel/resolve_future_token_ids.py +++ b/python/sglang/jit_kernel/resolve_future_token_ids.py @@ -4,7 +4,7 @@ from typing import TYPE_CHECKING import torch -from sglang.jit_kernel.utils import cache_once, load_jit, make_cpp_args +from sglang.kernels.jit.utils import cache_once, load_jit, make_cpp_args if TYPE_CHECKING: from tvm_ffi.module import Module diff --git a/python/sglang/jit_kernel/rmsnorm_hf.py b/python/sglang/jit_kernel/rmsnorm_hf.py index f4db56dac..572a097ab 100644 --- a/python/sglang/jit_kernel/rmsnorm_hf.py +++ b/python/sglang/jit_kernel/rmsnorm_hf.py @@ -6,7 +6,7 @@ from typing import TYPE_CHECKING, Optional import torch -from sglang.jit_kernel.utils import ( +from sglang.kernels.jit.utils import ( cache_once, is_arch_support_pdl, load_jit, diff --git a/python/sglang/jit_kernel/rope.py b/python/sglang/jit_kernel/rope.py index d9cbe0a8b..26b2f78e9 100644 --- a/python/sglang/jit_kernel/rope.py +++ b/python/sglang/jit_kernel/rope.py @@ -5,7 +5,7 @@ from typing import TYPE_CHECKING, Optional import torch -from sglang.jit_kernel.utils import ( +from sglang.kernels.jit.utils import ( cache_once, is_arch_support_pdl, load_jit, diff --git a/python/sglang/jit_kernel/set_mla_kv_buffer.py b/python/sglang/jit_kernel/set_mla_kv_buffer.py index 3624f5afc..cdda777b2 100644 --- a/python/sglang/jit_kernel/set_mla_kv_buffer.py +++ b/python/sglang/jit_kernel/set_mla_kv_buffer.py @@ -1,121 +1,5 @@ -"""JIT TMA bulk-store path for ``set_mla_kv_buffer``. +"""Compatibility shim (RFC #29630 Phase 4) -> sglang.kernels.ops.kvcache._jit_set_mla_kv_buffer.""" -Each warp scatter-writes one item's (nope, rope) row via a single -``cp.async.bulk.global.shared::cta`` store. Requires SM90+ (Hopper or later) -for the TMA bulk-store hardware. The host-side wrapper in -``sglang.srt.mem_cache.utils`` falls back to a Triton kernel for older arches. -""" +from sglang.kernels.ops.kvcache import _jit_set_mla_kv_buffer as _impl -from __future__ import annotations - -import logging -from typing import TYPE_CHECKING - -import torch - -from sglang.jit_kernel.utils import ( - cache_once, - is_arch_support_pdl, - load_jit, - make_cpp_args, -) - -if TYPE_CHECKING: - from tvm_ffi.module import Module - -logger = logging.getLogger(__name__) - - -@cache_once -def _jit_set_mla_kv_buffer_module( - nope_bytes: int, rope_bytes: int, use_pdl: bool -) -> Module: - args = make_cpp_args(nope_bytes, rope_bytes, use_pdl) - return load_jit( - f"set_mla_kv_buffer_{nope_bytes}_{rope_bytes}", - *args, - cuda_files=["elementwise/set_mla_kv_buffer.cuh"], - cuda_wrappers=[ - ("set_mla_kv_buffer", f"SetMlaKVBufferKernel<{args}>::run"), - ], - ) - - -@cache_once -def can_use_set_mla_kv_buffer(nope_bytes: int, rope_bytes: int) -> bool: - """Whether the TMA path can be used for these row byte widths. - - TMA bulk store requires ``(nope_bytes + rope_bytes)`` to be a multiple of - 16; both halves individually must also be a multiple of 4 (the warp-coop - smem load lower bound). - """ - if nope_bytes % 4 != 0 or rope_bytes % 4 != 0: - logger.warning( - "Unsupported nope_bytes=%d rope_bytes=%d for JIT set_mla_kv_buffer:" - " both must be multiples of 4", - nope_bytes, - rope_bytes, - ) - return False - if (nope_bytes + rope_bytes) % 16 != 0: - logger.warning( - "Unsupported nope_bytes=%d rope_bytes=%d for JIT set_mla_kv_buffer:" - " (nope_bytes + rope_bytes) must be a multiple of 16 for TMA bulk store", - nope_bytes, - rope_bytes, - ) - return False - try: - _jit_set_mla_kv_buffer_module(nope_bytes, rope_bytes, is_arch_support_pdl()) - return True - except Exception as e: # pragma: no cover - compile-time only - logger.warning( - "Failed to load JIT set_mla_kv_buffer kernel " - "with nope_bytes=%d rope_bytes=%d: %s", - nope_bytes, - rope_bytes, - e, - ) - return False - - -def _pick_num_warps(n_loc: int) -> int: - # Tuned on GB300: nw=4 wins below 1024 (more CTAs spread across SMs); - # nw=8 wins above (each CTA amortises the bulk-group commit better). - return 4 if n_loc <= 768 else 8 - - -def set_mla_kv_buffer( - kv_buffer: torch.Tensor, - loc: torch.Tensor, - cache_k_nope: torch.Tensor, - cache_k_rope: torch.Tensor, - num_warps: int = 0, -) -> None: - """Write packed [k_nope | k_rope] rows into ``kv_buffer`` at ``loc`` indices - via a TMA bulk-store. SM90+ only — the caller is expected to gate. - - Shapes (last dim is treated as the row payload; any leading singleton dims - on the source tensors are flattened away): - kv_buffer: [num_pages, total_dim] or [num_pages, 1, total_dim] - cache_k_nope: [n_loc, nope_dim] or [n_loc, 1, nope_dim] - cache_k_rope: [n_loc, rope_dim] or [n_loc, 1, rope_dim] - loc: [n_loc] - """ - n_loc = loc.shape[0] - if n_loc == 0: - return - - src_nope = cache_k_nope.view(n_loc, -1) if cache_k_nope.dim() != 2 else cache_k_nope - src_rope = cache_k_rope.view(n_loc, -1) if cache_k_rope.dim() != 2 else cache_k_rope - buf = kv_buffer.view(kv_buffer.shape[0], -1) if kv_buffer.dim() != 2 else kv_buffer - - nope_bytes = src_nope.shape[-1] * src_nope.element_size() - rope_bytes = src_rope.shape[-1] * src_rope.element_size() - if num_warps <= 0: - num_warps = _pick_num_warps(n_loc) - - module = _jit_set_mla_kv_buffer_module( - nope_bytes, rope_bytes, is_arch_support_pdl() - ) - module.set_mla_kv_buffer(buf, loc, src_nope, src_rope, num_warps) +globals().update({k: getattr(_impl, k) for k in dir(_impl) if not k.startswith("__")}) diff --git a/python/sglang/jit_kernel/sparse_mla_q8kv8_prefill_sm90.py b/python/sglang/jit_kernel/sparse_mla_q8kv8_prefill_sm90.py index 372f30c82..eb2315549 100644 --- a/python/sglang/jit_kernel/sparse_mla_q8kv8_prefill_sm90.py +++ b/python/sglang/jit_kernel/sparse_mla_q8kv8_prefill_sm90.py @@ -10,8 +10,8 @@ from typing import TYPE_CHECKING, Optional import torch -from sglang.jit_kernel.utils import cache_once, load_jit, override_jit_cuda_arch from sglang.kernel_api_logging import debug_kernel_api +from sglang.kernels.jit.utils import cache_once, load_jit, override_jit_cuda_arch from sglang.srt.utils.custom_op import register_custom_op if TYPE_CHECKING: diff --git a/python/sglang/jit_kernel/timestep_embedding.py b/python/sglang/jit_kernel/timestep_embedding.py index 08a941a6e..3269152e4 100644 --- a/python/sglang/jit_kernel/timestep_embedding.py +++ b/python/sglang/jit_kernel/timestep_embedding.py @@ -4,8 +4,8 @@ from typing import TYPE_CHECKING import torch -from sglang.jit_kernel.utils import cache_once, load_jit, make_cpp_args from sglang.kernel_api_logging import debug_kernel_api +from sglang.kernels.jit.utils import cache_once, load_jit, make_cpp_args if TYPE_CHECKING: from tvm_ffi.module import Module diff --git a/python/sglang/jit_kernel/transfer_mamba.py b/python/sglang/jit_kernel/transfer_mamba.py index a64afa060..2f7c42c23 100644 --- a/python/sglang/jit_kernel/transfer_mamba.py +++ b/python/sglang/jit_kernel/transfer_mamba.py @@ -4,7 +4,7 @@ Provides ``transfer_kv_mamba_pf_lf`` (load: page_first -> layer_first) and ``transfer_kv_mamba_lf_pf`` (backup: layer_first -> page_first). Uses the shared ``load_jit`` + ``cache_once`` infrastructure from -``sglang.jit_kernel.utils`` — the same mechanism used by ``hicache.py`` +``sglang.kernels.jit.utils`` — the same mechanism used by ``hicache.py`` for MHA/MLA staged write-back kernels. This ensures consistent content-addressed caching, CUDA arch detection, and multi-worker JIT compilation behavior across all JIT kernels. @@ -15,8 +15,8 @@ from __future__ import annotations import logging from typing import TYPE_CHECKING -from sglang.jit_kernel.utils import cache_once, load_jit from sglang.kernel_api_logging import debug_kernel_api +from sglang.kernels.jit.utils import cache_once, load_jit if TYPE_CHECKING: import torch diff --git a/python/sglang/jit_kernel/trtllm_lora_temp/kimi_k2_moe_fused_gate.py b/python/sglang/jit_kernel/trtllm_lora_temp/kimi_k2_moe_fused_gate.py index 0fd6231d3..a78bdd18d 100644 --- a/python/sglang/jit_kernel/trtllm_lora_temp/kimi_k2_moe_fused_gate.py +++ b/python/sglang/jit_kernel/trtllm_lora_temp/kimi_k2_moe_fused_gate.py @@ -4,7 +4,7 @@ from typing import TYPE_CHECKING, Tuple import torch -from sglang.jit_kernel.utils import cache_once, load_jit +from sglang.kernels.jit.utils import cache_once, load_jit if TYPE_CHECKING: from tvm_ffi.module import Module diff --git a/python/sglang/jit_kernel/trtllm_lora_temp/moe_lora_merged_align.py b/python/sglang/jit_kernel/trtllm_lora_temp/moe_lora_merged_align.py index 32b3fa8b5..c55e702ee 100644 --- a/python/sglang/jit_kernel/trtllm_lora_temp/moe_lora_merged_align.py +++ b/python/sglang/jit_kernel/trtllm_lora_temp/moe_lora_merged_align.py @@ -4,7 +4,7 @@ from typing import TYPE_CHECKING, Optional import torch -from sglang.jit_kernel.utils import cache_once, load_jit, make_cpp_args +from sglang.kernels.jit.utils import cache_once, load_jit, make_cpp_args if TYPE_CHECKING: from tvm_ffi.module import Module diff --git a/python/sglang/jit_kernel/trtllm_lora_temp/topk_softmax_pack.py b/python/sglang/jit_kernel/trtllm_lora_temp/topk_softmax_pack.py index 36a4cdef7..fb3455727 100644 --- a/python/sglang/jit_kernel/trtllm_lora_temp/topk_softmax_pack.py +++ b/python/sglang/jit_kernel/trtllm_lora_temp/topk_softmax_pack.py @@ -20,7 +20,7 @@ from typing import TYPE_CHECKING, Optional import torch -from sglang.jit_kernel.utils import cache_once, load_jit +from sglang.kernels.jit.utils import cache_once, load_jit from sglang.srt.utils.custom_op import register_custom_op if TYPE_CHECKING: diff --git a/python/sglang/kernels/jit/__init__.py b/python/sglang/kernels/jit/__init__.py new file mode 100644 index 000000000..3a456a260 --- /dev/null +++ b/python/sglang/kernels/jit/__init__.py @@ -0,0 +1,6 @@ +"""Internal JIT home under ``sglang.kernels`` (RFC #29630). + +Mirrors the legacy ``sglang.jit_kernel`` tree; shared build/runtime +infrastructure lives in :mod:`sglang.kernels.jit.utils`. csrc / include / +operators migrate here in later phases. +""" diff --git a/python/sglang/jit_kernel/utils/__init__.py b/python/sglang/kernels/jit/utils/__init__.py similarity index 71% rename from python/sglang/jit_kernel/utils/__init__.py rename to python/sglang/kernels/jit/utils/__init__.py index 452ee5441..41142291c 100644 --- a/python/sglang/jit_kernel/utils/__init__.py +++ b/python/sglang/kernels/jit/utils/__init__.py @@ -1,11 +1,11 @@ -"""Public interface of sglang.jit_kernel.utils.""" +"""Public interface of sglang.kernels.jit.utils.""" -from sglang.jit_kernel.utils.arch import ( +from sglang.kernels.jit.utils.arch import ( get_jit_cuda_arch, is_arch_support_pdl, override_jit_cuda_arch, ) -from sglang.jit_kernel.utils.common import ( +from sglang.kernels.jit.utils.common import ( cache_once, empty_sentinel, get_ci_test_range, @@ -14,7 +14,7 @@ from sglang.jit_kernel.utils.common import ( lazy_register_class, should_run_full_tests, ) -from sglang.jit_kernel.utils.compile import KERNEL_PATH, load_jit, make_cpp_args +from sglang.kernels.jit.utils.compile import KERNEL_PATH, load_jit, make_cpp_args __all__ = [ "empty_sentinel", diff --git a/python/sglang/jit_kernel/utils/arch.py b/python/sglang/kernels/jit/utils/arch.py similarity index 98% rename from python/sglang/jit_kernel/utils/arch.py rename to python/sglang/kernels/jit/utils/arch.py index 24c6492ee..107dc24f1 100644 --- a/python/sglang/jit_kernel/utils/arch.py +++ b/python/sglang/kernels/jit/utils/arch.py @@ -9,7 +9,7 @@ from typing import List import torch -from sglang.jit_kernel.utils.common import ( +from sglang.kernels.jit.utils.common import ( cache_once, is_hip_runtime, is_musa_runtime, diff --git a/python/sglang/jit_kernel/utils/common.py b/python/sglang/kernels/jit/utils/common.py similarity index 100% rename from python/sglang/jit_kernel/utils/common.py rename to python/sglang/kernels/jit/utils/common.py diff --git a/python/sglang/jit_kernel/utils/compile.py b/python/sglang/kernels/jit/utils/compile.py similarity index 98% rename from python/sglang/jit_kernel/utils/compile.py rename to python/sglang/kernels/jit/utils/compile.py index 1196d8559..f054eaa56 100644 --- a/python/sglang/jit_kernel/utils/compile.py +++ b/python/sglang/kernels/jit/utils/compile.py @@ -13,9 +13,9 @@ from typing import TYPE_CHECKING, List, Tuple, TypeAlias, Union import torch -from sglang.jit_kernel.utils.arch import get_default_target_flags, get_jit_cuda_arch -from sglang.jit_kernel.utils.common import cache_once, is_hip_runtime -from sglang.jit_kernel.utils.deps import REGISTERED_DEPENDENCIES +from sglang.kernels.jit.utils.arch import get_default_target_flags, get_jit_cuda_arch +from sglang.kernels.jit.utils.common import cache_once, is_hip_runtime +from sglang.kernels.jit.utils.deps import REGISTERED_DEPENDENCIES if TYPE_CHECKING: from tvm_ffi import Module diff --git a/python/sglang/jit_kernel/utils/deps.py b/python/sglang/kernels/jit/utils/deps.py similarity index 100% rename from python/sglang/jit_kernel/utils/deps.py rename to python/sglang/kernels/jit/utils/deps.py diff --git a/python/sglang/kernels/ops/activation/__init__.py b/python/sglang/kernels/ops/activation/__init__.py index eb0e2de96..4450dee8f 100644 --- a/python/sglang/kernels/ops/activation/__init__.py +++ b/python/sglang/kernels/ops/activation/__init__.py @@ -31,7 +31,7 @@ _HIP = frozenset({CapabilityRequirement.HIP}) # — the canonical OR-semantics case that a device-baked backend name couldn't. _CUDA_HIP = frozenset({CapabilityRequirement.CUDA, CapabilityRequirement.HIP}) # JIT before AOT to match the production path (srt/layers/activation.py imports -# from sglang.jit_kernel.activation on CUDA); auto-selection must not invert it. +# from sglang.kernels.ops.activation._jit_activation on CUDA); auto-selection must not invert it. _ACT_PRIORITY = ( KernelBackend.JIT, KernelBackend.AOT, @@ -82,7 +82,7 @@ class _GatedActivationOp(BaseFusedOp): expert_ids: Optional[torch.Tensor] = None, expert_step: int = 1, ) -> torch.Tensor: - import sglang.jit_kernel.activation as jit_activation + import sglang.kernels.ops.activation._jit_activation as jit_activation return getattr(jit_activation, self.kernel_attr)( input, out, expert_ids, expert_step @@ -183,7 +183,7 @@ class GeluTanhAndMulOp(_GatedActivationOp): class ReLU2Op(BaseFusedOp): """``out = relu(input) ** 2`` (single-input, not gated). - The real kernel is the CUDA JIT path (``sglang.jit_kernel.activation.relu2``, + The real kernel is the CUDA JIT path (``sglang.kernels.ops.activation._jit_activation.relu2``, used in production on CUDA); elsewhere the torch reference runs. """ @@ -214,7 +214,7 @@ class ReLU2Op(BaseFusedOp): def forward_jit( self, input: torch.Tensor, out: Optional[torch.Tensor] = None ) -> torch.Tensor: - from sglang.jit_kernel.activation import relu2 + from sglang.kernels.ops.activation._jit_activation import relu2 result = relu2(input) if out is None: diff --git a/python/sglang/kernels/ops/activation/_jit_activation.py b/python/sglang/kernels/ops/activation/_jit_activation.py new file mode 100644 index 000000000..7f39a89b0 --- /dev/null +++ b/python/sglang/kernels/ops/activation/_jit_activation.py @@ -0,0 +1,168 @@ +from __future__ import annotations + +from typing import TYPE_CHECKING, Optional + +import torch + +from sglang.kernels.jit.utils import ( + cache_once, + get_jit_cuda_arch, + is_arch_support_pdl, + is_hip_runtime, + load_jit, + make_cpp_args, +) +from sglang.srt.utils.custom_op import register_custom_op + +if TYPE_CHECKING: + from tvm_ffi.module import Module + + +def _fast_math_flags() -> list[str]: + # Mirrors sgl-kernel's CMake policy: fast-math on SM90, precise on + # SM100+ (Blackwell needs bit-exact expf), off on HIP (clang rejects). + if is_hip_runtime(): + return [] + if get_jit_cuda_arch().major >= 10: + return [] + return ["--use_fast_math"] + + +@cache_once +def _jit_activation_module(dtype: torch.dtype) -> Module: + args = make_cpp_args(dtype, is_arch_support_pdl()) + return load_jit( + "activation", + *args, + cuda_files=["elementwise/activation.cuh"], + extra_cuda_cflags=_fast_math_flags(), + cuda_wrappers=[ + ("run_activation", f"ActivationKernel<{args}>::run_activation"), + ( + "run_activation_filtered", + f"ActivationKernel<{args}>::run_activation_filtered", + ), + ( + "run_unary_activation", + f"ActivationKernel<{args}>::run_unary_activation", + ), + ], + ) + + +SUPPORTED_ACTIVATIONS = {"silu", "gelu", "gelu_tanh"} +SUPPORTED_UNARY_ACTIVATIONS = {"relu2"} + + +@register_custom_op(mutates_args=["out"]) +def _run_activation_inplace( + op_name: str, input: torch.Tensor, out: torch.Tensor +) -> None: + hidden_size = input.shape[-1] // 2 + module = _jit_activation_module(input.dtype) + input_2d = input.view(-1, hidden_size * 2) + out_2d = out.view(-1, hidden_size) + module.run_activation(input_2d, out_2d, op_name) + + +@register_custom_op(mutates_args=["out"]) +def _run_activation_filtered_inplace( + op_name: str, + input: torch.Tensor, + out: torch.Tensor, + expert_ids: torch.Tensor, + expert_step: int, +) -> None: + hidden_size = input.shape[-1] // 2 + module = _jit_activation_module(input.dtype) + input_2d = input.view(-1, hidden_size * 2) + out_2d = out.view(-1, hidden_size) + module.run_activation_filtered(input_2d, out_2d, expert_ids, expert_step, op_name) + + +def run_activation( + op_name: str, + input: torch.Tensor, + out: Optional[torch.Tensor], + expert_ids: Optional[torch.Tensor] = None, + expert_step: int = 1, +) -> torch.Tensor: + """Apply ``op_name`` activation followed by element-wise multiplication. + + When ``expert_ids`` is provided, output rows are skipped for tokens whose + routed expert id is ``-1``. ``expert_step`` is 1 for per-token routing and + ``BLOCK_SIZE_M`` for sorted/TMA routing — i.e. ``expert_ids[token_id // + expert_step]`` is consulted before computing each row. + """ + assert op_name in SUPPORTED_ACTIVATIONS, f"Unsupported activation: {op_name}" + hidden_size = input.shape[-1] // 2 + if out is None: + out = input.new_empty(*input.shape[:-1], hidden_size) + if expert_ids is None: + _run_activation_inplace(op_name, input, out) + else: + _run_activation_filtered_inplace(op_name, input, out, expert_ids, expert_step) + return out + + +@register_custom_op(mutates_args=["out"]) +def _run_unary_activation_inplace( + op_name: str, input: torch.Tensor, out: torch.Tensor +) -> None: + last = input.shape[-1] + module = _jit_activation_module(input.dtype) + module.run_unary_activation(input.view(-1, last), out.view(-1, last), op_name) + + +def run_unary_activation( + op_name: str, + input: torch.Tensor, + out: Optional[torch.Tensor] = None, +) -> torch.Tensor: + """Apply a standalone (non-gated) element-wise activation: ``out = act(input)``. + + Unlike :func:`run_activation`, there is no gate/up split — ``input`` and + ``out`` share the same shape. + """ + assert ( + op_name in SUPPORTED_UNARY_ACTIVATIONS + ), f"Unsupported unary activation: {op_name}" + if out is None: + out = torch.empty_like(input) + _run_unary_activation_inplace(op_name, input, out) + return out + + +def relu2( + input: torch.Tensor, + out: Optional[torch.Tensor] = None, +) -> torch.Tensor: + """Squared ReLU: ``out = max(0, input) ** 2`` (element-wise).""" + return run_unary_activation("relu2", input, out) + + +def silu_and_mul( + input: torch.Tensor, + out: Optional[torch.Tensor] = None, + expert_ids: Optional[torch.Tensor] = None, + expert_step: int = 1, +) -> torch.Tensor: + return run_activation("silu", input, out, expert_ids, expert_step) + + +def gelu_and_mul( + input: torch.Tensor, + out: Optional[torch.Tensor] = None, + expert_ids: Optional[torch.Tensor] = None, + expert_step: int = 1, +) -> torch.Tensor: + return run_activation("gelu", input, out, expert_ids, expert_step) + + +def gelu_tanh_and_mul( + input: torch.Tensor, + out: Optional[torch.Tensor] = None, + expert_ids: Optional[torch.Tensor] = None, + expert_step: int = 1, +) -> torch.Tensor: + return run_activation("gelu_tanh", input, out, expert_ids, expert_step) diff --git a/python/sglang/kernels/ops/attention/dsv4/compress_c128_hip.py b/python/sglang/kernels/ops/attention/dsv4/compress_c128_hip.py index 6838ebe8b..4cb67b8d9 100644 --- a/python/sglang/kernels/ops/attention/dsv4/compress_c128_hip.py +++ b/python/sglang/kernels/ops/attention/dsv4/compress_c128_hip.py @@ -14,7 +14,7 @@ from sglang.jit_kernel.dsv4 import ( CompressorDecodePlan, CompressorPrefillPlan, ) -from sglang.jit_kernel.utils import is_hip_runtime +from sglang.kernels.jit.utils import is_hip_runtime _is_hip = is_hip_runtime() diff --git a/python/sglang/kernels/ops/attention/fla/layernorm_gated.py b/python/sglang/kernels/ops/attention/fla/layernorm_gated.py index 0fa7e032f..a599e761f 100644 --- a/python/sglang/kernels/ops/attention/fla/layernorm_gated.py +++ b/python/sglang/kernels/ops/attention/fla/layernorm_gated.py @@ -15,7 +15,7 @@ import triton import triton.language as tl from einops import rearrange -from sglang.jit_kernel.utils import is_arch_support_pdl +from sglang.kernels.jit.utils import is_arch_support_pdl from sglang.srt.batch_invariant_ops import is_batch_invariant_mode_enabled from sglang.srt.model_executor.cuda_graph_config import ( Backend, diff --git a/python/sglang/kernels/ops/attention/utils.py b/python/sglang/kernels/ops/attention/utils.py index 74b130803..23c5d6abc 100644 --- a/python/sglang/kernels/ops/attention/utils.py +++ b/python/sglang/kernels/ops/attention/utils.py @@ -2,7 +2,7 @@ import torch import triton import triton.language as tl -from sglang.jit_kernel.utils import is_arch_support_pdl +from sglang.kernels.jit.utils import is_arch_support_pdl from sglang.kernels.ops.attention.pad import ( pad_sequence_with_mask as pad_sequence_with_mask, ) diff --git a/python/sglang/kernels/ops/gemm/__init__.py b/python/sglang/kernels/ops/gemm/__init__.py index 1851a9831..a07ae0671 100644 --- a/python/sglang/kernels/ops/gemm/__init__.py +++ b/python/sglang/kernels/ops/gemm/__init__.py @@ -59,7 +59,7 @@ register_kernel( KernelSpec( op="gemm.dsv3_fused_a_gemm", backend=KernelBackend.JIT, - target="sglang.jit_kernel.dsv3_fused_a_gemm:dsv3_fused_a_gemm", + target="sglang.kernels.ops.gemm._jit_dsv3_fused_a_gemm:dsv3_fused_a_gemm", capabilities=_CUDA, format_signature=FormatSignature( supported_dtypes=("bfloat16",), @@ -72,7 +72,7 @@ register_kernel( KernelSpec( op="gemm.dsv3_router_gemm", backend=KernelBackend.JIT, - target="sglang.jit_kernel.dsv3_router_gemm:dsv3_router_gemm", + target="sglang.kernels.ops.gemm._jit_dsv3_router_gemm:dsv3_router_gemm", capabilities=_CUDA, format_signature=FormatSignature( supported_dtypes=("bfloat16",), diff --git a/python/sglang/kernels/ops/gemm/_jit_dsv3_fused_a_gemm.py b/python/sglang/kernels/ops/gemm/_jit_dsv3_fused_a_gemm.py new file mode 100644 index 000000000..0cef3cbcf --- /dev/null +++ b/python/sglang/kernels/ops/gemm/_jit_dsv3_fused_a_gemm.py @@ -0,0 +1,90 @@ +""" +JIT kernel for DeepSeek V3 fused QKV-A GEMM (min-latency). + +Runtime-compiled CUDA C++ kernel for SM90+ (Hopper) GPUs. +Shapes: hd_in a multiple of 256, hd_out a multiple of 16, num_tokens 1-16, bfloat16. +""" + +from __future__ import annotations + +from typing import TYPE_CHECKING, Optional + +import torch + +from sglang.kernel_api_logging import debug_kernel_api +from sglang.kernels.jit.utils import ( + cache_once, + is_arch_support_pdl, + load_jit, + make_cpp_args, +) +from sglang.srt.utils.common import direct_register_custom_op + +if TYPE_CHECKING: + from tvm_ffi.module import Module + + +@cache_once +def _jit_dsv3_fused_a_gemm_module(hd_in: int, hd_out: int, use_pdl: bool) -> Module: + args = make_cpp_args(hd_in, hd_out, use_pdl) + return load_jit( + "dsv3_fused_a_gemm", + *args, + cuda_files=["gemm/dsv3_fused_a_gemm.cuh"], + cuda_wrappers=[ + ("dsv3_fused_a_gemm", f"DSV3FusedAGemmKernel<{args}>::run"), + ], + ) + + +def _dsv3_fused_a_gemm_run(mat_a: torch.Tensor, mat_b: torch.Tensor) -> torch.Tensor: + assert mat_a.stride(1) == 1, "mat_a must be row-major [M, K]" + output = torch.empty( + (mat_a.shape[0], mat_b.shape[1]), + device=mat_a.device, + dtype=mat_a.dtype, + ) + module = _jit_dsv3_fused_a_gemm_module( + mat_a.shape[1], mat_b.shape[1], is_arch_support_pdl() + ) + module.dsv3_fused_a_gemm(mat_a, mat_b, output) + return output + + +def _dsv3_fused_a_gemm_fake(mat_a: torch.Tensor, mat_b: torch.Tensor) -> torch.Tensor: + return mat_a.new_empty((mat_a.shape[0], mat_b.shape[1]), dtype=torch.bfloat16) + + +direct_register_custom_op( + op_name="jit_dsv3_fused_a_gemm", + op_func=_dsv3_fused_a_gemm_run, + mutates_args=[], + fake_impl=_dsv3_fused_a_gemm_fake, +) + + +@debug_kernel_api +def dsv3_fused_a_gemm( + mat_a: torch.Tensor, + mat_b: torch.Tensor, + output: Optional[torch.Tensor] = None, +) -> torch.Tensor: + """ + DeepSeek V3 fused QKV-A GEMM kernel (JIT variant). + + Args: + mat_a: Input tensor of shape [num_tokens, hd_in], bfloat16, row-major. + hd_in must be a multiple of 256 and num_tokens in [1, 16]. + mat_b: Weight tensor of shape [hd_in, hd_out], bfloat16, column-major + (i.e. ``weight.T`` of a row-major [hd_out, hd_in] weight). + hd_out must be a multiple of 16. + output: Optional pre-allocated output tensor of shape [num_tokens, hd_out]. + + Returns: + Output tensor of shape [num_tokens, hd_out]. + """ + result = torch.ops.sglang.jit_dsv3_fused_a_gemm(mat_a, mat_b) + if output is not None: + output.copy_(result) + return output + return result diff --git a/python/sglang/kernels/ops/gemm/_jit_dsv3_router_gemm.py b/python/sglang/kernels/ops/gemm/_jit_dsv3_router_gemm.py new file mode 100644 index 000000000..7279f1474 --- /dev/null +++ b/python/sglang/kernels/ops/gemm/_jit_dsv3_router_gemm.py @@ -0,0 +1,92 @@ +""" +JIT kernel for DeepSeek V3 router GEMM. + +Runtime-compiled CUDA C++ kernel for SM90+ (Hopper) GPUs. +Supports num_experts in {256, 384}, hidden_dim a multiple of 1024, num_tokens 1-16. +""" + +from __future__ import annotations + +from typing import TYPE_CHECKING, Optional + +import torch + +from sglang.kernel_api_logging import debug_kernel_api +from sglang.kernels.jit.utils import ( + cache_once, + is_arch_support_pdl, + load_jit, + make_cpp_args, +) +from sglang.srt.utils.custom_op import register_custom_op + +if TYPE_CHECKING: + from tvm_ffi.module import Module + + +@cache_once +def _jit_dsv3_router_gemm_module( + num_experts: int, + hidden_dim: int, + use_pdl: bool, + out_float: bool, +) -> Module: + args = make_cpp_args(num_experts, hidden_dim, use_pdl, out_float) + return load_jit( + "dsv3_router_gemm", + *args, + cuda_files=["gemm/dsv3_router_gemm.cuh"], + cuda_wrappers=[ + ("dsv3_router_gemm", f"DSV3RouterGemmKernel<{args}>::run"), + ], + ) + + +@register_custom_op( + op_name="dsv3_router_gemm", + mutates_args=["output"], +) +def _dsv3_router_gemm_custom_op( + hidden_states: torch.Tensor, + router_weights: torch.Tensor, + output: torch.Tensor, +) -> None: + num_experts = router_weights.shape[0] + hidden_dim = hidden_states.shape[1] + out_float = output.dtype == torch.float32 + module = _jit_dsv3_router_gemm_module( + num_experts, hidden_dim, is_arch_support_pdl(), out_float + ) + module.dsv3_router_gemm(hidden_states, router_weights, output) + return None + + +@debug_kernel_api +def dsv3_router_gemm( + hidden_states: torch.Tensor, + router_weights: torch.Tensor, + out_dtype: torch.dtype = torch.bfloat16, + output: Optional[torch.Tensor] = None, +) -> torch.Tensor: + """ + DeepSeek V3 router GEMM kernel (JIT variant). + + Args: + hidden_states: Input tensor of shape [num_tokens, hidden_dim], bfloat16. + hidden_dim must be a multiple of 1024 and num_tokens in [1, 16]. + router_weights: Weight tensor of shape [num_experts, hidden_dim], bfloat16. + out_dtype: Output dtype, either torch.bfloat16 or torch.float32. + output: Optional pre-allocated output tensor. + + Returns: + Output tensor of shape [num_tokens, num_experts]. + """ + if output is None: + output = torch.empty( + hidden_states.shape[0], + router_weights.shape[0], + device=hidden_states.device, + dtype=out_dtype, + ) + _dsv3_router_gemm_custom_op(hidden_states, router_weights, output) + return output diff --git a/python/sglang/kernels/ops/gemm/trtllm_lora_temp/kernel_utils.py b/python/sglang/kernels/ops/gemm/trtllm_lora_temp/kernel_utils.py index 541d1ccf6..a92fcd822 100644 --- a/python/sglang/kernels/ops/gemm/trtllm_lora_temp/kernel_utils.py +++ b/python/sglang/kernels/ops/gemm/trtllm_lora_temp/kernel_utils.py @@ -1,7 +1,7 @@ import triton import triton.language as tl -from sglang.jit_kernel.utils import is_arch_support_pdl +from sglang.kernels.jit.utils import is_arch_support_pdl def get_pdl_launch_metadata() -> tuple[bool, dict]: diff --git a/python/sglang/kernels/ops/kvcache/_jit_set_mla_kv_buffer.py b/python/sglang/kernels/ops/kvcache/_jit_set_mla_kv_buffer.py new file mode 100644 index 000000000..ad8cc879c --- /dev/null +++ b/python/sglang/kernels/ops/kvcache/_jit_set_mla_kv_buffer.py @@ -0,0 +1,121 @@ +"""JIT TMA bulk-store path for ``set_mla_kv_buffer``. + +Each warp scatter-writes one item's (nope, rope) row via a single +``cp.async.bulk.global.shared::cta`` store. Requires SM90+ (Hopper or later) +for the TMA bulk-store hardware. The host-side wrapper in +``sglang.srt.mem_cache.utils`` falls back to a Triton kernel for older arches. +""" + +from __future__ import annotations + +import logging +from typing import TYPE_CHECKING + +import torch + +from sglang.kernels.jit.utils import ( + cache_once, + is_arch_support_pdl, + load_jit, + make_cpp_args, +) + +if TYPE_CHECKING: + from tvm_ffi.module import Module + +logger = logging.getLogger(__name__) + + +@cache_once +def _jit_set_mla_kv_buffer_module( + nope_bytes: int, rope_bytes: int, use_pdl: bool +) -> Module: + args = make_cpp_args(nope_bytes, rope_bytes, use_pdl) + return load_jit( + f"set_mla_kv_buffer_{nope_bytes}_{rope_bytes}", + *args, + cuda_files=["elementwise/set_mla_kv_buffer.cuh"], + cuda_wrappers=[ + ("set_mla_kv_buffer", f"SetMlaKVBufferKernel<{args}>::run"), + ], + ) + + +@cache_once +def can_use_set_mla_kv_buffer(nope_bytes: int, rope_bytes: int) -> bool: + """Whether the TMA path can be used for these row byte widths. + + TMA bulk store requires ``(nope_bytes + rope_bytes)`` to be a multiple of + 16; both halves individually must also be a multiple of 4 (the warp-coop + smem load lower bound). + """ + if nope_bytes % 4 != 0 or rope_bytes % 4 != 0: + logger.warning( + "Unsupported nope_bytes=%d rope_bytes=%d for JIT set_mla_kv_buffer:" + " both must be multiples of 4", + nope_bytes, + rope_bytes, + ) + return False + if (nope_bytes + rope_bytes) % 16 != 0: + logger.warning( + "Unsupported nope_bytes=%d rope_bytes=%d for JIT set_mla_kv_buffer:" + " (nope_bytes + rope_bytes) must be a multiple of 16 for TMA bulk store", + nope_bytes, + rope_bytes, + ) + return False + try: + _jit_set_mla_kv_buffer_module(nope_bytes, rope_bytes, is_arch_support_pdl()) + return True + except Exception as e: # pragma: no cover - compile-time only + logger.warning( + "Failed to load JIT set_mla_kv_buffer kernel " + "with nope_bytes=%d rope_bytes=%d: %s", + nope_bytes, + rope_bytes, + e, + ) + return False + + +def _pick_num_warps(n_loc: int) -> int: + # Tuned on GB300: nw=4 wins below 1024 (more CTAs spread across SMs); + # nw=8 wins above (each CTA amortises the bulk-group commit better). + return 4 if n_loc <= 768 else 8 + + +def set_mla_kv_buffer( + kv_buffer: torch.Tensor, + loc: torch.Tensor, + cache_k_nope: torch.Tensor, + cache_k_rope: torch.Tensor, + num_warps: int = 0, +) -> None: + """Write packed [k_nope | k_rope] rows into ``kv_buffer`` at ``loc`` indices + via a TMA bulk-store. SM90+ only — the caller is expected to gate. + + Shapes (last dim is treated as the row payload; any leading singleton dims + on the source tensors are flattened away): + kv_buffer: [num_pages, total_dim] or [num_pages, 1, total_dim] + cache_k_nope: [n_loc, nope_dim] or [n_loc, 1, nope_dim] + cache_k_rope: [n_loc, rope_dim] or [n_loc, 1, rope_dim] + loc: [n_loc] + """ + n_loc = loc.shape[0] + if n_loc == 0: + return + + src_nope = cache_k_nope.view(n_loc, -1) if cache_k_nope.dim() != 2 else cache_k_nope + src_rope = cache_k_rope.view(n_loc, -1) if cache_k_rope.dim() != 2 else cache_k_rope + buf = kv_buffer.view(kv_buffer.shape[0], -1) if kv_buffer.dim() != 2 else kv_buffer + + nope_bytes = src_nope.shape[-1] * src_nope.element_size() + rope_bytes = src_rope.shape[-1] * src_rope.element_size() + if num_warps <= 0: + num_warps = _pick_num_warps(n_loc) + + module = _jit_set_mla_kv_buffer_module( + nope_bytes, rope_bytes, is_arch_support_pdl() + ) + module.set_mla_kv_buffer(buf, loc, src_nope, src_rope, num_warps) diff --git a/python/sglang/kernels/ops/kvcache/mla_buffer.py b/python/sglang/kernels/ops/kvcache/mla_buffer.py index 5bf285d97..7249e9d9b 100644 --- a/python/sglang/kernels/ops/kvcache/mla_buffer.py +++ b/python/sglang/kernels/ops/kvcache/mla_buffer.py @@ -4,7 +4,7 @@ import torch import triton import triton.language as tl -from sglang.jit_kernel.utils import is_arch_support_pdl +from sglang.kernels.jit.utils import is_arch_support_pdl from sglang.srt.runtime_context import get_parallel @@ -116,10 +116,10 @@ def set_mla_kv_buffer_triton( Name retained for caller compatibility; the implementation is no longer Triton-only. """ - from sglang.jit_kernel.set_mla_kv_buffer import ( + from sglang.kernels.ops.kvcache._jit_set_mla_kv_buffer import ( can_use_set_mla_kv_buffer, ) - from sglang.jit_kernel.set_mla_kv_buffer import ( + from sglang.kernels.ops.kvcache._jit_set_mla_kv_buffer import ( set_mla_kv_buffer as jit_set_mla_kv_buffer, ) diff --git a/python/sglang/kernels/ops/layernorm/__init__.py b/python/sglang/kernels/ops/layernorm/__init__.py index 9f9953114..d42dd8560 100644 --- a/python/sglang/kernels/ops/layernorm/__init__.py +++ b/python/sglang/kernels/ops/layernorm/__init__.py @@ -113,7 +113,7 @@ class RMSNormOp(BaseFusedOp): ) -> torch.Tensor: import torch - from sglang.jit_kernel.norm import rmsnorm as jit_rmsnorm + from sglang.kernels.ops.layernorm._jit_norm import rmsnorm as jit_rmsnorm if out is None: out = torch.empty_like(input) @@ -227,7 +227,9 @@ class FusedAddRMSNormOp(BaseFusedOp): eps: float = 1e-6, enable_pdl: Optional[bool] = None, ) -> None: - from sglang.jit_kernel.norm import fused_add_rmsnorm as jit_fused_add_rmsnorm + from sglang.kernels.ops.layernorm._jit_norm import ( + fused_add_rmsnorm as jit_fused_add_rmsnorm, + ) return jit_fused_add_rmsnorm(input, residual, weight, eps) diff --git a/python/sglang/kernels/ops/layernorm/_jit_norm.py b/python/sglang/kernels/ops/layernorm/_jit_norm.py new file mode 100644 index 000000000..046afbe48 --- /dev/null +++ b/python/sglang/kernels/ops/layernorm/_jit_norm.py @@ -0,0 +1,179 @@ +from __future__ import annotations + +import logging +from typing import TYPE_CHECKING, Optional + +import torch + +from sglang.kernel_api_logging import debug_kernel_api +from sglang.kernels.jit.utils import ( + cache_once, + is_arch_support_pdl, + load_jit, + make_cpp_args, +) + +if TYPE_CHECKING: + from tvm_ffi.module import Module + + +logger = logging.getLogger(__name__) + + +@cache_once +def _jit_qknorm_module(head_dim: int, dtype: torch.dtype) -> Module: + args = make_cpp_args(head_dim, is_arch_support_pdl(), dtype) + return load_jit( + "qknorm", + *args, + cuda_files=["elementwise/qknorm.cuh"], + cuda_wrappers=[("qknorm", f"QKNormKernel<{args}>::run")], + ) + + +_RMSNORM_WARP_SIZES = frozenset({64, 128, 256}) +_RMSNORM_MAX_HIDDEN_SIZE = 16384 +_RMSNORM_HALF_BLOCK_MIN_SIZE = 2048 + + +def _is_supported_rmsnorm_hidden_size(d: int) -> bool: + return d in _RMSNORM_WARP_SIZES or ( + (d > 256 and d % 256 == 0 and d <= 8192) + or (d >= 8192 and d % 512 == 0 and d <= 16384) + ) + + +def _rmsnorm_kernel_class(hidden_size: int) -> str: + if hidden_size in _RMSNORM_WARP_SIZES: + return "RMSNormWarpKernel" + if hidden_size == 512: + return "RMSNormHalfKernel" + if hidden_size >= _RMSNORM_HALF_BLOCK_MIN_SIZE: + if hidden_size % 512 == 0: + return "RMSNormHalfKernel" + return "RMSNormKernel" + + +@cache_once +def _jit_rmsnorm_module(hidden_size: int, dtype: torch.dtype) -> Module: + args = make_cpp_args(hidden_size, is_arch_support_pdl(), dtype) + kernel_class = f"{_rmsnorm_kernel_class(hidden_size)}<{args}>" + return load_jit( + "rmsnorm", + *args, + cuda_files=["elementwise/rmsnorm.cuh"], + cuda_wrappers=[("rmsnorm", f"{kernel_class}::run")], + ) + + +def is_supported_jit_fused_add_rmsnorm_hidden_size(hidden_size: int) -> bool: + return hidden_size > 0 and hidden_size % 16 == 0 and hidden_size <= 8192 + + +@cache_once +def _jit_fused_add_rmsnorm_module( + dtype: torch.dtype, cast_x_before_out_mul: bool +) -> Module: + args = make_cpp_args(cast_x_before_out_mul, dtype) + return load_jit( + "fused_add_rmsnorm", + *args, + cuda_files=["elementwise/fused_add_rmsnorm.cuh"], + cuda_wrappers=[("fused_add_rmsnorm", f"FusedAddRMSNormKernel<{args}>::run")], + ) + + +@cache_once +def _jit_qknorm_across_heads_module(dtype: torch.dtype) -> Module: + args = make_cpp_args(dtype) + return load_jit( + "qknorm_across_heads", + *args, + cuda_files=["elementwise/qknorm_across_heads.cuh"], + cuda_wrappers=[ + ("qknorm_across_heads", f"QKNormAcrossHeadsKernel<{args}>::run") + ], + ) + + +@torch.compiler.assume_constant_result +@cache_once +def can_use_fused_inplace_qknorm(head_dim: int, dtype: torch.dtype) -> bool: + if head_dim not in [64, 128, 256, 512, 1024]: + logger.warning(f"Unsupported head_dim={head_dim} for JIT QK-Norm kernel") + return False + try: + _jit_qknorm_module(head_dim, dtype) + return True + except Exception as e: + logger.warning(f"Failed to load JIT QK-Norm kernel: {e}") + return False + + +@debug_kernel_api +def fused_inplace_qknorm( + q: torch.Tensor, + k: torch.Tensor, + q_weight: torch.Tensor, + k_weight: torch.Tensor, + eps: float = 1e-6, + *, + head_dim: int = 0, +) -> None: + head_dim = head_dim or q.size(-1) + module = _jit_qknorm_module(head_dim, q.dtype) + module.qknorm(q, k, q_weight, k_weight, eps) + + +@debug_kernel_api +def rmsnorm( + input: torch.Tensor, + weight: torch.Tensor, + out: Optional[torch.Tensor] = None, + eps: float = 1e-6, +) -> None: + out = out if out is not None else input + hidden_size = input.size(-1) + if not _is_supported_rmsnorm_hidden_size(hidden_size): + raise RuntimeError( + f"jit rmsnorm: unsupported hidden_size={hidden_size}. " + f"Supported: {sorted(_RMSNORM_WARP_SIZES)}, and multiples of 256 in " + f"(256, {_RMSNORM_MAX_HIDDEN_SIZE}]." + ) + module = _jit_rmsnorm_module(hidden_size, input.dtype) + module.rmsnorm(input, weight, out, eps) + + +@debug_kernel_api +def fused_add_rmsnorm( + input: torch.Tensor, + residual: torch.Tensor, + weight: torch.Tensor, + eps: float = 1e-6, + *, + cast_x_before_out_mul: bool = False, +) -> None: + module = _jit_fused_add_rmsnorm_module(input.dtype, cast_x_before_out_mul) + module.fused_add_rmsnorm(input, residual, weight, eps) + + +@debug_kernel_api +def fused_inplace_qknorm_across_heads( + q: torch.Tensor, + k: torch.Tensor, + q_weight: torch.Tensor, + k_weight: torch.Tensor, + eps: float = 1e-6, +) -> None: + """ + Fused inplace QK normalization across all heads. + + Args: + q: Query tensor of shape [batch_size, num_heads * head_dim] + k: Key tensor of shape [batch_size, num_heads * head_dim] + q_weight: Query weight tensor of shape [num_heads * head_dim] + k_weight: Key weight tensor of shape [num_heads * head_dim] + eps: Epsilon for numerical stability + """ + module = _jit_qknorm_across_heads_module(q.dtype) + module.qknorm_across_heads(q, k, q_weight, k_weight, eps) diff --git a/python/sglang/kernels/ops/layernorm/elementwise.py b/python/sglang/kernels/ops/layernorm/elementwise.py index 1414e0038..c7a65cd6f 100644 --- a/python/sglang/kernels/ops/layernorm/elementwise.py +++ b/python/sglang/kernels/ops/layernorm/elementwise.py @@ -4,7 +4,7 @@ import torch import triton import triton.language as tl -from sglang.jit_kernel.utils import is_arch_support_pdl +from sglang.kernels.jit.utils import is_arch_support_pdl from sglang.kernels.ops.activation.softcap import softcap_out as fused_softcap from sglang.srt.utils import is_hip from sglang.srt.utils.custom_op import register_custom_op diff --git a/python/sglang/kernels/ops/layernorm/mhc.py b/python/sglang/kernels/ops/layernorm/mhc.py index 190d40874..114afac1a 100644 --- a/python/sglang/kernels/ops/layernorm/mhc.py +++ b/python/sglang/kernels/ops/layernorm/mhc.py @@ -7,7 +7,7 @@ from typing import Tuple import torch -from sglang.jit_kernel.utils import is_arch_support_pdl +from sglang.kernels.jit.utils import is_arch_support_pdl from sglang.srt.distributed.device_communicators.pynccl_allocator import ( use_symmetric_memory, ) diff --git a/python/sglang/kernels/ops/mamba/causal_conv1d_triton.py b/python/sglang/kernels/ops/mamba/causal_conv1d_triton.py index d70f176e8..6080d4001 100644 --- a/python/sglang/kernels/ops/mamba/causal_conv1d_triton.py +++ b/python/sglang/kernels/ops/mamba/causal_conv1d_triton.py @@ -10,7 +10,7 @@ import torch import triton import triton.language as tl -from sglang.jit_kernel.utils import is_arch_support_pdl +from sglang.kernels.jit.utils import is_arch_support_pdl PAD_SLOT_ID = -1 diff --git a/python/sglang/kernels/ops/mamba/triton_ops/mamba_ssm.py b/python/sglang/kernels/ops/mamba/triton_ops/mamba_ssm.py index 09dbd73db..3ed64d131 100644 --- a/python/sglang/kernels/ops/mamba/triton_ops/mamba_ssm.py +++ b/python/sglang/kernels/ops/mamba/triton_ops/mamba_ssm.py @@ -11,7 +11,7 @@ import triton import triton.language as tl from packaging import version -from sglang.jit_kernel.utils import is_arch_support_pdl +from sglang.kernels.jit.utils import is_arch_support_pdl PAD_SLOT_ID = -1 diff --git a/python/sglang/kernels/ops/quantization/__init__.py b/python/sglang/kernels/ops/quantization/__init__.py index 4187053a4..a744b97ba 100644 --- a/python/sglang/kernels/ops/quantization/__init__.py +++ b/python/sglang/kernels/ops/quantization/__init__.py @@ -56,7 +56,7 @@ register_kernel( KernelSpec( op="quantization.per_token_group_quant", backend=KernelBackend.JIT, - target="sglang.jit_kernel.per_token_group_quant:per_token_group_quant", + target="sglang.kernels.ops.quantization._jit_per_token_group_quant:per_token_group_quant", capabilities=_CUDA, format_signature=FormatSignature( supported_dtypes=("float8_e4m3fn", "int8"), diff --git a/python/sglang/kernels/ops/quantization/_jit_per_tensor_quant_fp8.py b/python/sglang/kernels/ops/quantization/_jit_per_tensor_quant_fp8.py new file mode 100644 index 000000000..4a45eb61d --- /dev/null +++ b/python/sglang/kernels/ops/quantization/_jit_per_tensor_quant_fp8.py @@ -0,0 +1,77 @@ +from __future__ import annotations + +from typing import TYPE_CHECKING + +import torch + +from sglang.kernels.jit.utils import cache_once, load_jit, make_cpp_args +from sglang.srt.utils.custom_op import register_custom_op + +if TYPE_CHECKING: + from tvm_ffi.module import Module + + +@cache_once +def _jit_per_tensor_quant_fp8_module(is_static: bool, dtype: torch.dtype) -> Module: + args = make_cpp_args(is_static, dtype) + return load_jit( + "per_tensor_quant_fp8", + *args, + cuda_files=["gemm/per_tensor_quant_fp8.cuh"], + cuda_wrappers=[("per_tensor_quant_fp8", f"per_tensor_quant_fp8<{args}>")], + ) + + +@register_custom_op( + op_name="per_tensor_quant_fp8", + mutates_args=["output_q", "output_s"], +) +def per_tensor_quant_fp8( + input: torch.Tensor, + output_q: torch.Tensor, + output_s: torch.Tensor, + is_static: bool = False, +) -> None: + """ + Per-tensor quantization to FP8 format. + + Args: + input: Input tensor to quantize (float, half, or bfloat16) + output_q: Output quantized tensor (fp8_e4m3) + output_s: Output scale tensor (float scalar or 1D tensor with 1 element) + is_static: If True, assumes scale is pre-computed and skips absmax computation + """ + module = _jit_per_tensor_quant_fp8_module(is_static, input.dtype) + module.per_tensor_quant_fp8(input.view(-1), output_q.view(-1), output_s.view(-1)) + + +@cache_once +def _jit_per_tensor_absmax_fp8_module(dtype: torch.dtype) -> Module: + args = make_cpp_args(dtype) + return load_jit( + "per_tensor_absmax_fp8", + *args, + cuda_files=["gemm/per_tensor_quant_fp8.cuh"], + cuda_wrappers=[("per_tensor_absmax_fp8", f"per_tensor_absmax_fp8<{args}>")], + ) + + +@register_custom_op( + op_name="per_tensor_absmax_fp8", + mutates_args=["output_s"], +) +def per_tensor_absmax_fp8( + input: torch.Tensor, + output_s: torch.Tensor, +) -> None: + """Compute scale = max(abs(input)) / fp8_e4m3_max via atomic-max reduction. + + The caller must zero-initialise ``output_s`` before the call (the kernel + uses ``atomic_max`` across blocks, so starting from 0 is required). + + Args: + input: Input tensor (float16, bfloat16, or float32). Any shape. + output_s: Pre-allocated float32 tensor of shape (1,), zero-initialised. + """ + module = _jit_per_tensor_absmax_fp8_module(input.dtype) + module.per_tensor_absmax_fp8(input.view(-1), output_s.view(-1)) diff --git a/python/sglang/kernels/ops/quantization/_jit_per_token_group_quant.py b/python/sglang/kernels/ops/quantization/_jit_per_token_group_quant.py new file mode 100644 index 000000000..d485a5816 --- /dev/null +++ b/python/sglang/kernels/ops/quantization/_jit_per_token_group_quant.py @@ -0,0 +1,223 @@ +from __future__ import annotations + +from typing import TYPE_CHECKING, Optional, Tuple + +import torch + +from sglang.kernel_api_logging import debug_kernel_api +from sglang.kernels.jit.utils import ( + cache_once, + is_arch_support_pdl, + load_jit, + make_cpp_args, +) +from sglang.srt.utils.custom_op import register_custom_op + +if TYPE_CHECKING: + from tvm_ffi.module import Module + +_SUPPORTED_INPUT_DTYPES = (torch.bfloat16, torch.float16) +_SUPPORTED_OUTPUT_DTYPES = (torch.float8_e4m3fn, torch.int8) +_SUPPORTED_GROUP_SIZES = (16, 32, 64, 128, 256) + + +@cache_once +def _jit_module( + in_dtype: torch.dtype, + out_dtype: torch.dtype, + group_size: int, + scale_ue8m0: bool, + row_major: bool, + aligned: bool, + fuse_silu_and_mul: bool, + masked_layout: bool, + use_pdl: bool, +) -> Module: + assert in_dtype in _SUPPORTED_INPUT_DTYPES + assert out_dtype in _SUPPORTED_OUTPUT_DTYPES + assert group_size in _SUPPORTED_GROUP_SIZES + trait_args = make_cpp_args( + in_dtype, + out_dtype, + group_size, + scale_ue8m0, + row_major, + aligned, + fuse_silu_and_mul, + use_pdl, + ) + launcher = ( + "PerTokenGroupQuantMaskedKernel" + if masked_layout + else "PerTokenGroupQuantFlatKernel" + ) + return load_jit( + "per_token_group_quant", + *trait_args, + "masked" if masked_layout else "flat", + cuda_files=["gemm/per_token_group_quant.cuh"], + cuda_wrappers=[("per_token_group_quant", f"{launcher}<{trait_args}>::run")], + extra_cuda_cflags=["--use_fast_math"], + ) + + +def _infer_scale_layout( + output_s: torch.Tensor, scale_ue8m0: bool, num_groups: int +) -> Tuple[bool, bool]: + """Return ``(row_major, aligned)`` for ``output_s``. + + Column-major (transposed) scale buffers have token stride 1 and a larger + group stride; row-major buffers are contiguous. + """ + row_major = output_s.stride(-2) >= output_s.stride(-1) + if output_s.dtype == torch.int32: + if not scale_ue8m0: + raise ValueError("int32-packed scale buffers require scale_ue8m0=True") + aligned = num_groups % 4 == 0 + return row_major, aligned + if output_s.dtype == torch.float32: + if scale_ue8m0: + raise ValueError("scale_ue8m0=True requires an int32-packed output_s") + return row_major, True + raise ValueError(f"Unsupported output_s dtype {output_s.dtype}") + + +@register_custom_op( + op_name="per_token_group_quant", + mutates_args=["output_q", "output_s"], +) +def _per_token_group_quant_custom_op( + input: torch.Tensor, + output_q: torch.Tensor, + output_s: torch.Tensor, + group_size: int, + scale_ue8m0: bool = False, + fuse_silu_and_mul: bool = False, + masked_m: Optional[torch.Tensor] = None, + expected_m: Optional[int] = None, +) -> None: + num_groups = output_q.shape[-1] // group_size + row_major, aligned = _infer_scale_layout(output_s, scale_ue8m0, num_groups) + module = _jit_module( + input.dtype, + output_q.dtype, + int(group_size), + bool(scale_ue8m0), + row_major, + aligned, + bool(fuse_silu_and_mul), + masked_m is not None, + is_arch_support_pdl(), + ) + if masked_m is not None: + module.per_token_group_quant( + input, output_q, output_s, masked_m, int(expected_m or -1) + ) + else: + module.per_token_group_quant(input, output_q, output_s) + + +def _allocate_outputs( + input: torch.Tensor, + group_size: int, + out_dtype: torch.dtype, + scale_ue8m0: bool, + column_major_scales: bool, + fuse_silu_and_mul: bool, +) -> Tuple[torch.Tensor, torch.Tensor]: + """Allocate ``(output_q, output_s)`` in the requested major mode / scale + format, selected by ``(column_major_scales, scale_ue8m0)``.""" + hidden = input.shape[-1] // (2 if fuse_silu_and_mul else 1) + out_shape = (*input.shape[:-1], hidden) + output_q = torch.empty(out_shape, device=input.device, dtype=out_dtype) + + num_groups = hidden // group_size + if scale_ue8m0 and not column_major_scales: + # Row-major packed UE8M0: int32 [..., ceil(ng/4)] contiguous (an + # unaligned ng leaves a partially-used last int32 that the kernel zero- + # pads). The shared create_*_output_scale helper does not produce this + # layout. + output_s = torch.empty( + (*out_shape[:-1], (num_groups + 3) // 4), + device=input.device, + dtype=torch.int32, + ) + else: + from sglang.kernels.ops.quantization.fp8_kernel import ( + create_per_token_group_quant_fp8_output_scale, + ) + + output_s = create_per_token_group_quant_fp8_output_scale( + x_shape=out_shape, + device=input.device, + group_size=group_size, + column_major_scales=column_major_scales, + scale_tma_aligned=column_major_scales, + scale_ue8m0=scale_ue8m0, + ) + return output_q, output_s + + +@debug_kernel_api +def per_token_group_quant( + input: torch.Tensor, + output_q: Optional[torch.Tensor] = None, + output_s: Optional[torch.Tensor] = None, + group_size: int = 128, + scale_ue8m0: bool = False, + fuse_silu_and_mul: bool = False, + masked_m: Optional[torch.Tensor] = None, + expected_m: Optional[int] = None, + *, + out_dtype: Optional[torch.dtype] = None, + column_major_scales: bool = False, +) -> Tuple[torch.Tensor, torch.Tensor]: + """Per-token-group quantization. Returns ``(output_q, output_s)``. + + ``output_q`` / ``output_s`` are optional: pass them to quantize into + caller-owned buffers, or omit both to have them allocated per ``out_dtype`` + (default fp8_e4m3), ``scale_ue8m0`` and ``column_major_scales``. Either way + the two tensors are returned. + + Input / output shapes: + vanilla: input [T, hidden], output_q [T, hidden] + fuse_silu_and_mul: input [T, hidden*2], output_q [T, hidden] + masked (+ above): input [E, T_pad, ...], output_q [E, T_pad, hidden], + masked_m [E] int32 + ``output_s`` scale layouts (inferred from a supplied buffer's dtype/strides, + or allocated to match when omitted): + float32 contiguous -> row-major fp32 scales + float32 transposed -> col-major fp32 scales (TMA-aligned view) + int32 transposed -> col-major UE8M0 bytes packed 4-per-int32 + int32 contiguous -> row-major UE8M0 bytes packed 4-per-int32 + The packed layouts require ``scale_ue8m0=True``. + + ``expected_m`` (masked only) is an optional expected-tokens-per-expert hint. + + Inputs are bf16/fp16; group size is one of 16/32/64/128/256; the quant range + follows ``output_q.dtype`` (fp8_e4m3: +-448, int8: [-128, 127]). + """ + if output_q is None: + assert output_s is None + output_q, output_s = _allocate_outputs( + input, + group_size, + out_dtype or torch.float8_e4m3fn, + scale_ue8m0, + column_major_scales, + fuse_silu_and_mul, + ) + else: + assert output_s is not None + assert out_dtype is None or out_dtype == output_q.dtype + _per_token_group_quant_custom_op( + input=input, + output_q=output_q, + output_s=output_s, + group_size=group_size, + scale_ue8m0=scale_ue8m0, + fuse_silu_and_mul=fuse_silu_and_mul, + masked_m=masked_m, + expected_m=expected_m, + ) + return output_q, output_s diff --git a/python/sglang/kernels/ops/quantization/_jit_per_token_group_quant_8bit_v2.py b/python/sglang/kernels/ops/quantization/_jit_per_token_group_quant_8bit_v2.py new file mode 100644 index 000000000..54eb87739 --- /dev/null +++ b/python/sglang/kernels/ops/quantization/_jit_per_token_group_quant_8bit_v2.py @@ -0,0 +1,136 @@ +"""DEPRECATED: superseded by ``sglang.kernels.ops.quantization._jit_per_token_group_quant`` (the +default CUDA path). No sglang runtime code may call this kernel; it is kept +only as the perf baseline for the per_token_group_quant benchmarks and its own +bit-parity tests, and will be deleted once those move to torch references. +""" + +from __future__ import annotations + +from typing import TYPE_CHECKING, Optional + +import torch + +from sglang.kernel_api_logging import debug_kernel_api +from sglang.kernels.jit.utils import ( + cache_once, + is_arch_support_pdl, + load_jit, + make_cpp_args, +) +from sglang.srt.utils.custom_op import register_custom_op + +if TYPE_CHECKING: + from tvm_ffi.module import Module + + +@cache_once +def _jit_module(in_dtype: torch.dtype, out_dtype: torch.dtype, use_pdl: bool) -> Module: + args = make_cpp_args(in_dtype, out_dtype, use_pdl) + return load_jit( + "per_token_group_quant_8bit_v2", + *args, + cuda_files=["gemm/per_token_group_quant_8bit_v2.cuh"], + cuda_wrappers=[ + ( + "per_token_group_quant_8bit_v2", + f"PerTokenGroupQuant8bitV2Kernel<{args}>::run", + ) + ], + # Match the AOT sgl-kernel build (-use_fast_math) so the FP8 scale + # division/rounding is bit-identical to sgl_per_token_group_quant_8bit_v2. + extra_cuda_cflags=["--use_fast_math"], + ) + + +@register_custom_op( + op_name="per_token_group_quant_8bit_v2", + mutates_args=["output_q", "output_s"], +) +def _per_token_group_quant_8bit_v2_custom_op( + input: torch.Tensor, + output_q: torch.Tensor, + output_s: torch.Tensor, + group_size: int, + eps: float, + min_8bit: float, + max_8bit: float, + scale_ue8m0: bool = False, + fuse_silu_and_mul: bool = False, + masked_m: Optional[torch.Tensor] = None, +) -> None: + """Opaque custom-op boundary around the JIT v2 kernel. + + Registering this as a custom op (instead of calling the tvm-ffi module + directly) keeps torch.compile / piecewise-CUDA-graph from tracing into the + tvm-ffi ``Function.__call__`` (which Dynamo cannot trace). All shape-derived + scalars are computed here and passed to the kernel. + + Layouts (matching the AOT v2): + vanilla: input (num_tokens, hidden), output_q (num_tokens, hidden) + fuse_silu_and_mul: input (num_tokens, hidden*2), output_q (num_tokens, hidden) + fuse_silu_and_mul+masked: input (num_experts, tokens_pad, hidden*2), + output_q (num_experts, tokens_pad, hidden), masked_m (num_experts,) + """ + masked_layout = masked_m is not None + numel = input.numel() + num_groups = numel // group_size // (2 if fuse_silu_and_mul else 1) + if num_groups == 0: # empty input -> grid 0 -> cudaErrorInvalidConfiguration + return + num_local_experts = input.shape[0] if masked_layout else 1 + last = output_q.dim() - 1 + is_column_major = output_s.stride(last - 1) < output_s.stride(last) + hidden_dim_num_groups = output_q.shape[last] // group_size + num_tokens_per_expert = output_q.shape[last - 1] + scale_expert_stride = output_s.stride(0) if masked_layout else 0 + scale_hidden_stride = output_s.stride(last) + + module = _jit_module(input.dtype, output_q.dtype, is_arch_support_pdl()) + module.per_token_group_quant_8bit_v2( + input, + output_q, + output_s, + masked_m if masked_layout else input, # unused (nullptr) when not masked + int(group_size), + bool(scale_ue8m0), + bool(fuse_silu_and_mul), + bool(masked_layout), + int(num_groups), + int(num_local_experts), + bool(is_column_major), + int(hidden_dim_num_groups), + int(num_tokens_per_expert), + int(scale_expert_stride), + int(scale_hidden_stride), + ) + + +@debug_kernel_api +def per_token_group_quant_8bit_v2( + input: torch.Tensor, + output_q: torch.Tensor, + output_s: torch.Tensor, + group_size: int, + eps: float, + min_8bit: float, + max_8bit: float, + scale_ue8m0: bool = False, + fuse_silu_and_mul: bool = False, + masked_m: Optional[torch.Tensor] = None, +) -> None: + """JIT port of sgl_per_token_group_quant_8bit_v2 (full feature parity). + + Wraps the registered custom op so torch.compile / piecewise CUDA graph treat + the tvm-ffi kernel call as an opaque boundary. + """ + _per_token_group_quant_8bit_v2_custom_op( + input=input, + output_q=output_q, + output_s=output_s, + group_size=group_size, + eps=eps, + min_8bit=min_8bit, + max_8bit=max_8bit, + scale_ue8m0=scale_ue8m0, + fuse_silu_and_mul=fuse_silu_and_mul, + masked_m=masked_m, + ) diff --git a/python/sglang/kernels/ops/quantization/fp8_kernel.py b/python/sglang/kernels/ops/quantization/fp8_kernel.py index a9f500241..cf4c87d9d 100644 --- a/python/sglang/kernels/ops/quantization/fp8_kernel.py +++ b/python/sglang/kernels/ops/quantization/fp8_kernel.py @@ -28,7 +28,7 @@ try: except: pass -from sglang.jit_kernel.utils import is_arch_support_pdl +from sglang.kernels.jit.utils import is_arch_support_pdl from sglang.srt.layers import deep_gemm_wrapper from sglang.srt.utils import ( ceil_align, @@ -55,13 +55,13 @@ _is_sm120_supported = is_sm120_supported() _use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip if _is_cuda or _is_musa: - from sglang.jit_kernel.per_tensor_quant_fp8 import ( - per_tensor_quant_fp8 as sgl_per_tensor_quant_fp8, - ) from sglang.kernels.ops.quantization import ( per_token_group_quant, sgl_per_token_quant_fp8, ) + from sglang.kernels.ops.quantization._jit_per_tensor_quant_fp8 import ( + per_tensor_quant_fp8 as sgl_per_tensor_quant_fp8, + ) if _is_musa: # per_token_group_quant is CUDA-only JIT; MUSA keeps the AOT v2 group-quant op. @@ -543,7 +543,7 @@ def _run_per_token_group_quant_8bit_kernel( ``sglang_per_token_quant_fp8``. """ if scale_ue8m0 and x_s.dtype == torch.float32 and not _is_musa: - from sglang.jit_kernel.per_token_group_quant_8bit_v2 import ( + from sglang.kernels.ops.quantization._jit_per_token_group_quant_8bit_v2 import ( per_token_group_quant_8bit_v2, ) diff --git a/python/sglang/srt/layers/attention/dsv4/compressor_v2.py b/python/sglang/srt/layers/attention/dsv4/compressor_v2.py index 12f5a8bf2..1a0efee8f 100644 --- a/python/sglang/srt/layers/attention/dsv4/compressor_v2.py +++ b/python/sglang/srt/layers/attention/dsv4/compressor_v2.py @@ -10,7 +10,7 @@ from sglang.jit_kernel.dsv4 import ( compress_forward, compress_norm_rope_store, ) -from sglang.jit_kernel.utils import is_hip_runtime +from sglang.kernels.jit.utils import is_hip_runtime from sglang.srt.environ import envs if TYPE_CHECKING: diff --git a/python/sglang/srt/layers/attention/tokenspeed_mla_backend.py b/python/sglang/srt/layers/attention/tokenspeed_mla_backend.py index 75b4ae06a..30f68f407 100644 --- a/python/sglang/srt/layers/attention/tokenspeed_mla_backend.py +++ b/python/sglang/srt/layers/attention/tokenspeed_mla_backend.py @@ -35,7 +35,7 @@ import torch from sglang.jit_kernel.fp8_quantize import fp8_quantize from sglang.jit_kernel.mla_kv_pack_quantize_fp8 import mla_kv_pack_quantize_fp8 -from sglang.jit_kernel.utils import is_arch_support_pdl +from sglang.kernels.jit.utils import is_arch_support_pdl from sglang.srt.layers.attention.trtllm_mla_backend import ( TRTLLMMLABackend, TRTLLMMLAMultiStepDraftBackend, diff --git a/python/sglang/srt/layers/moe/moe_runner/flashinfer_trtllm.py b/python/sglang/srt/layers/moe/moe_runner/flashinfer_trtllm.py index 72d44a474..866d44ff9 100644 --- a/python/sglang/srt/layers/moe/moe_runner/flashinfer_trtllm.py +++ b/python/sglang/srt/layers/moe/moe_runner/flashinfer_trtllm.py @@ -75,7 +75,7 @@ def finalize_flashinfer_trtllm_deferred_output( shared_output: torch.Tensor, ) -> torch.Tensor: from sglang.jit_kernel.moe_finalize_fuse_shared import moe_finalize_fuse_shared - from sglang.jit_kernel.utils import is_arch_support_pdl + from sglang.kernels.jit.utils import is_arch_support_pdl return moe_finalize_fuse_shared( deferred_output.gemm2_out, diff --git a/python/sglang/srt/layers/moe/moe_runner/triton_utils/inkling_moe.py b/python/sglang/srt/layers/moe/moe_runner/triton_utils/inkling_moe.py index e7e3c2f28..2ec710989 100644 --- a/python/sglang/srt/layers/moe/moe_runner/triton_utils/inkling_moe.py +++ b/python/sglang/srt/layers/moe/moe_runner/triton_utils/inkling_moe.py @@ -6,7 +6,7 @@ import torch import triton import triton.language as tl -from sglang.jit_kernel.utils import is_arch_support_pdl +from sglang.kernels.jit.utils import is_arch_support_pdl from sglang.srt.layers.moe.moe_runner.triton_utils.helion_utils import ( get_model_depths, helion_aot_autotune, diff --git a/python/sglang/srt/layers/moe/moe_runner/triton_utils/sigmoid_gate_topk_renorm.py b/python/sglang/srt/layers/moe/moe_runner/triton_utils/sigmoid_gate_topk_renorm.py index bd3c3c7af..9ec5e1648 100644 --- a/python/sglang/srt/layers/moe/moe_runner/triton_utils/sigmoid_gate_topk_renorm.py +++ b/python/sglang/srt/layers/moe/moe_runner/triton_utils/sigmoid_gate_topk_renorm.py @@ -13,7 +13,7 @@ import triton import triton.language as tl from sglang.jit_kernel.inkling_gate_topk_renorm import inkling_gate_topk_renorm_v2 -from sglang.jit_kernel.utils import is_arch_support_pdl +from sglang.kernels.jit.utils import is_arch_support_pdl from sglang.srt.environ import envs from sglang.srt.layers.moe.moe_runner.triton_utils.gate_topk import ( fpval_to_key, diff --git a/python/sglang/srt/models/gpt_oss.py b/python/sglang/srt/models/gpt_oss.py index eae718dd1..cd5731671 100644 --- a/python/sglang/srt/models/gpt_oss.py +++ b/python/sglang/srt/models/gpt_oss.py @@ -26,7 +26,7 @@ import torch from torch import nn from transformers import PretrainedConfig -from sglang.jit_kernel.utils import is_arch_support_pdl +from sglang.kernels.jit.utils import is_arch_support_pdl from sglang.srt.distributed import ( get_pp_group, tensor_model_parallel_all_reduce, diff --git a/python/sglang/srt/models/inkling_common/moe.py b/python/sglang/srt/models/inkling_common/moe.py index 9746d9a5c..0c5e8c357 100644 --- a/python/sglang/srt/models/inkling_common/moe.py +++ b/python/sglang/srt/models/inkling_common/moe.py @@ -14,7 +14,7 @@ from sglang.jit_kernel.inkling_gate_topk_renorm import ( inkling_gate_gemv, inkling_gate_gemv_fused, ) -from sglang.jit_kernel.utils import is_arch_support_pdl +from sglang.kernels.jit.utils import is_arch_support_pdl from sglang.srt.configs.inkling import InklingModelConfig from sglang.srt.distributed import ( get_tensor_model_parallel_group, diff --git a/test/registered/jit/benchmark/bench_custom_all_reduce.py b/test/registered/jit/benchmark/bench_custom_all_reduce.py index 4bc187693..b2ea3e9d6 100644 --- a/test/registered/jit/benchmark/bench_custom_all_reduce.py +++ b/test/registered/jit/benchmark/bench_custom_all_reduce.py @@ -13,7 +13,7 @@ import sglang.srt.distributed.parallel_state as ps from sglang.jit_kernel.benchmark import marker from sglang.jit_kernel.benchmark.utils import get_benchmark_range, multigpu_bench_main from sglang.jit_kernel.mp import register_comm_cleanup -from sglang.jit_kernel.utils import cache_once, is_arch_support_pdl +from sglang.kernels.jit.utils import cache_once, is_arch_support_pdl from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci register_cuda_ci( diff --git a/test/registered/jit/benchmark/bench_dsv3_fused_a_gemm.py b/test/registered/jit/benchmark/bench_dsv3_fused_a_gemm.py index 9fa7462ec..f66b89f27 100644 --- a/test/registered/jit/benchmark/bench_dsv3_fused_a_gemm.py +++ b/test/registered/jit/benchmark/bench_dsv3_fused_a_gemm.py @@ -13,7 +13,7 @@ from sglang.jit_kernel.cutedsl_dsv3_fused_a_gemm import ( dsv3_fused_a_gemm as cutedsl_dsv3_fused_a_gemm, ) from sglang.jit_kernel.dsv3_fused_a_gemm import dsv3_fused_a_gemm -from sglang.jit_kernel.utils import get_jit_cuda_arch, is_hip_runtime +from sglang.kernels.jit.utils import get_jit_cuda_arch, is_hip_runtime from sglang.test.ci.ci_register import register_cuda_ci from sglang.utils import is_in_ci diff --git a/test/registered/jit/benchmark/bench_dsv3_router_gemm.py b/test/registered/jit/benchmark/bench_dsv3_router_gemm.py index dd3d0adf0..d5f04a9de 100644 --- a/test/registered/jit/benchmark/bench_dsv3_router_gemm.py +++ b/test/registered/jit/benchmark/bench_dsv3_router_gemm.py @@ -10,7 +10,7 @@ import torch.nn.functional as F from sglang.jit_kernel.benchmark import marker from sglang.jit_kernel.benchmark.utils import create_random from sglang.jit_kernel.dsv3_router_gemm import dsv3_router_gemm -from sglang.jit_kernel.utils import get_jit_cuda_arch, is_hip_runtime +from sglang.kernels.jit.utils import get_jit_cuda_arch, is_hip_runtime from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci register_cuda_ci( diff --git a/test/registered/jit/benchmark/bench_mla_kv_pack_quantize_fp8.py b/test/registered/jit/benchmark/bench_mla_kv_pack_quantize_fp8.py index aa3f4080d..436923209 100644 --- a/test/registered/jit/benchmark/bench_mla_kv_pack_quantize_fp8.py +++ b/test/registered/jit/benchmark/bench_mla_kv_pack_quantize_fp8.py @@ -17,7 +17,7 @@ from sglang.jit_kernel.benchmark.utils import ( from sglang.jit_kernel.mla_kv_pack_quantize_fp8 import ( mla_kv_pack_quantize_fp8 as hybrid_pack, ) -from sglang.jit_kernel.utils import is_arch_support_pdl +from sglang.kernels.jit.utils import is_arch_support_pdl from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci register_cuda_ci( diff --git a/test/registered/jit/benchmark/bench_set_mla_kv_buffer.py b/test/registered/jit/benchmark/bench_set_mla_kv_buffer.py index 28f05e964..42da45253 100644 --- a/test/registered/jit/benchmark/bench_set_mla_kv_buffer.py +++ b/test/registered/jit/benchmark/bench_set_mla_kv_buffer.py @@ -21,7 +21,7 @@ from sglang.jit_kernel.benchmark.utils import ( get_benchmark_range, ) from sglang.jit_kernel.set_mla_kv_buffer import set_mla_kv_buffer as jit_set -from sglang.jit_kernel.utils import is_arch_support_pdl +from sglang.kernels.jit.utils import is_arch_support_pdl 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 diff --git a/test/registered/jit/benchmark/bench_symm_mem_all_gather.py b/test/registered/jit/benchmark/bench_symm_mem_all_gather.py index 86c31ff84..a06d72c53 100644 --- a/test/registered/jit/benchmark/bench_symm_mem_all_gather.py +++ b/test/registered/jit/benchmark/bench_symm_mem_all_gather.py @@ -28,7 +28,7 @@ import torch.distributed as dist import sglang.srt.distributed.parallel_state as ps from sglang.jit_kernel.benchmark import marker from sglang.jit_kernel.benchmark.utils import get_benchmark_range, multigpu_bench_main -from sglang.jit_kernel.utils import cache_once +from sglang.kernels.jit.utils import cache_once from sglang.srt.distributed.device_communicators.triton_symm_mem_ag import ( all_gather_inner, create_state, diff --git a/test/registered/jit/benchmark/bench_tp_qknorm.py b/test/registered/jit/benchmark/bench_tp_qknorm.py index 290f5b71c..1f32defbe 100644 --- a/test/registered/jit/benchmark/bench_tp_qknorm.py +++ b/test/registered/jit/benchmark/bench_tp_qknorm.py @@ -32,7 +32,7 @@ from sglang.jit_kernel.all_reduce import ( from sglang.jit_kernel.benchmark import marker from sglang.jit_kernel.benchmark.utils import multigpu_bench_main from sglang.jit_kernel.mp import register_comm_cleanup -from sglang.jit_kernel.utils import cache_once, get_ci_test_range +from sglang.kernels.jit.utils import cache_once, get_ci_test_range from sglang.srt.distributed.device_communicators.custom_all_reduce_v2 import ( CustomAllReduceV2, ) diff --git a/test/registered/jit/benchmark/diffusion/bench_diffusion_nvfp4_scaled_mm.py b/test/registered/jit/benchmark/diffusion/bench_diffusion_nvfp4_scaled_mm.py index cfabb8ff5..495ff93ed 100644 --- a/test/registered/jit/benchmark/diffusion/bench_diffusion_nvfp4_scaled_mm.py +++ b/test/registered/jit/benchmark/diffusion/bench_diffusion_nvfp4_scaled_mm.py @@ -11,7 +11,7 @@ import flashinfer import torch from sglang.jit_kernel.benchmark.utils import DEFAULT_DTYPE -from sglang.jit_kernel.utils import KERNEL_PATH +from sglang.kernels.jit.utils import KERNEL_PATH from sglang.test.ci.ci_register import register_cuda_ci from sglang.utils import is_in_ci diff --git a/test/registered/jit/benchmark/diffusion/bench_norm_impls.py b/test/registered/jit/benchmark/diffusion/bench_norm_impls.py index 489f10693..f55c0fe3d 100644 --- a/test/registered/jit/benchmark/diffusion/bench_norm_impls.py +++ b/test/registered/jit/benchmark/diffusion/bench_norm_impls.py @@ -18,7 +18,7 @@ from sglang.jit_kernel.diffusion.triton.norm import norm_infer, rms_norm_fn from sglang.jit_kernel.diffusion.triton.rmsnorm_onepass import triton_one_pass_rms_norm from sglang.jit_kernel.norm import fused_add_rmsnorm as jit_fused_add_rmsnorm from sglang.jit_kernel.norm import rmsnorm as jit_rmsnorm -from sglang.jit_kernel.utils import KERNEL_PATH +from sglang.kernels.jit.utils import KERNEL_PATH from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci from sglang.utils import is_in_ci diff --git a/test/registered/jit/diffusion/test_causal_conv3d_cat_pad.py b/test/registered/jit/diffusion/test_causal_conv3d_cat_pad.py index 25c2a3b0b..907c59c06 100644 --- a/test/registered/jit/diffusion/test_causal_conv3d_cat_pad.py +++ b/test/registered/jit/diffusion/test_causal_conv3d_cat_pad.py @@ -9,7 +9,7 @@ from sglang.jit_kernel.diffusion.causal_conv3d_cat_pad import ( from sglang.jit_kernel.diffusion.triton.causal_conv3d_pad import ( fused_causal_conv3d_cat_pad as fused_causal_conv3d_cat_pad_triton, ) -from sglang.jit_kernel.utils import get_ci_test_range +from sglang.kernels.jit.utils import get_ci_test_range 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") diff --git a/test/registered/jit/diffusion/test_qknorm_rope.py b/test/registered/jit/diffusion/test_qknorm_rope.py index 8d775d7b8..717b8024e 100644 --- a/test/registered/jit/diffusion/test_qknorm_rope.py +++ b/test/registered/jit/diffusion/test_qknorm_rope.py @@ -5,7 +5,7 @@ import pytest import torch import triton -from sglang.jit_kernel.utils import get_ci_test_range +from sglang.kernels.jit.utils import get_ci_test_range from sglang.test.ci.ci_register import register_cuda_ci register_cuda_ci(est_time=44, stage="base-b-kernel-unit", runner_config="1-gpu-large") diff --git a/test/registered/jit/diffusion/test_qwen_image_modulation.py b/test/registered/jit/diffusion/test_qwen_image_modulation.py index 596f6bcc7..8730605e7 100644 --- a/test/registered/jit/diffusion/test_qwen_image_modulation.py +++ b/test/registered/jit/diffusion/test_qwen_image_modulation.py @@ -9,7 +9,7 @@ from sglang.jit_kernel.diffusion.triton.scale_shift import ( fuse_layernorm_scale_shift_gate_select01_kernel, fuse_residual_layernorm_scale_shift_gate_select01_kernel, ) -from sglang.jit_kernel.utils import get_ci_test_range +from sglang.kernels.jit.utils import get_ci_test_range from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci register_cuda_ci(est_time=15, stage="base-b-kernel-unit", runner_config="1-gpu-large") diff --git a/test/registered/jit/diffusion/test_varlen_pack_pad.py b/test/registered/jit/diffusion/test_varlen_pack_pad.py index 614c4b711..e3e8e7e86 100644 --- a/test/registered/jit/diffusion/test_varlen_pack_pad.py +++ b/test/registered/jit/diffusion/test_varlen_pack_pad.py @@ -12,7 +12,7 @@ from sglang.jit_kernel.diffusion.triton.varlen_pack_pad import ( fused_pack_qkv, fused_scatter_to_padded, ) -from sglang.jit_kernel.utils import get_ci_test_range +from sglang.kernels.jit.utils import get_ci_test_range from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci register_cuda_ci(est_time=10, stage="base-b-kernel-unit", runner_config="1-gpu-large") diff --git a/test/registered/jit/diffusion/test_varlen_uspattn_equivalence.py b/test/registered/jit/diffusion/test_varlen_uspattn_equivalence.py index 51837e17d..13e6a7ec7 100644 --- a/test/registered/jit/diffusion/test_varlen_uspattn_equivalence.py +++ b/test/registered/jit/diffusion/test_varlen_uspattn_equivalence.py @@ -20,7 +20,7 @@ from sglang.jit_kernel.diffusion.triton.varlen_pack_pad import ( fused_scatter_to_padded, ) from sglang.jit_kernel.flash_attention import flash_attn_varlen_func -from sglang.jit_kernel.utils import get_ci_test_range +from sglang.kernels.jit.utils import get_ci_test_range from sglang.multimodal_gen.runtime.layers.attention.backends import ( flash_attn as _fa_backend, ) diff --git a/test/registered/jit/test_activation.py b/test/registered/jit/test_activation.py index 62c2becba..dd9975f7b 100644 --- a/test/registered/jit/test_activation.py +++ b/test/registered/jit/test_activation.py @@ -9,7 +9,7 @@ from sglang.jit_kernel.activation import ( relu2, run_activation, ) -from sglang.jit_kernel.utils import get_ci_test_range +from sglang.kernels.jit.utils import get_ci_test_range from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci register_cuda_ci(est_time=20, stage="base-b-kernel-unit", runner_config="1-gpu-large") diff --git a/test/registered/jit/test_custom_all_reduce.py b/test/registered/jit/test_custom_all_reduce.py index 8fd238828..84571d81c 100644 --- a/test/registered/jit/test_custom_all_reduce.py +++ b/test/registered/jit/test_custom_all_reduce.py @@ -31,7 +31,7 @@ import sglang.srt.distributed.parallel_state as ps from sglang.jit_kernel.all_reduce import AllReduceAlgo, get_all_reduce_module from sglang.jit_kernel.mp import register_comm_cleanup from sglang.jit_kernel.tests.utils import multigpu_pytest_main -from sglang.jit_kernel.utils import cache_once, get_ci_test_range +from sglang.kernels.jit.utils import cache_once, get_ci_test_range from sglang.srt.distributed.device_communicators.custom_all_reduce_v2 import ( CustomAllReduceV2, ) diff --git a/test/registered/jit/test_cutedsl_bf16_gemm.py b/test/registered/jit/test_cutedsl_bf16_gemm.py index 9d293d529..efb555cc6 100644 --- a/test/registered/jit/test_cutedsl_bf16_gemm.py +++ b/test/registered/jit/test_cutedsl_bf16_gemm.py @@ -5,7 +5,11 @@ import sys import pytest import torch -from sglang.jit_kernel.utils import get_ci_test_range, get_jit_cuda_arch, is_hip_runtime +from sglang.kernels.jit.utils import ( + get_ci_test_range, + get_jit_cuda_arch, + is_hip_runtime, +) from sglang.test.ci.ci_register import register_cuda_ci register_cuda_ci(est_time=30, stage="base-b-kernel-unit", runner_config="4-gpu-b200") diff --git a/test/registered/jit/test_cutedsl_dsv3_fused_a_gemm.py b/test/registered/jit/test_cutedsl_dsv3_fused_a_gemm.py index 96993fc01..6fa9c224d 100644 --- a/test/registered/jit/test_cutedsl_dsv3_fused_a_gemm.py +++ b/test/registered/jit/test_cutedsl_dsv3_fused_a_gemm.py @@ -5,7 +5,11 @@ import sys import pytest import torch -from sglang.jit_kernel.utils import get_ci_test_range, get_jit_cuda_arch, is_hip_runtime +from sglang.kernels.jit.utils import ( + get_ci_test_range, + get_jit_cuda_arch, + is_hip_runtime, +) 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") diff --git a/test/registered/jit/test_dsv3_fused_a_gemm.py b/test/registered/jit/test_dsv3_fused_a_gemm.py index 83287f0e0..cc6489ebf 100644 --- a/test/registered/jit/test_dsv3_fused_a_gemm.py +++ b/test/registered/jit/test_dsv3_fused_a_gemm.py @@ -7,7 +7,11 @@ import torch import torch.nn.functional as F from sglang.jit_kernel.dsv3_fused_a_gemm import dsv3_fused_a_gemm -from sglang.jit_kernel.utils import get_ci_test_range, get_jit_cuda_arch, is_hip_runtime +from sglang.kernels.jit.utils import ( + get_ci_test_range, + get_jit_cuda_arch, + is_hip_runtime, +) 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") diff --git a/test/registered/jit/test_dsv3_router_gemm.py b/test/registered/jit/test_dsv3_router_gemm.py index b934fe903..928318ac3 100644 --- a/test/registered/jit/test_dsv3_router_gemm.py +++ b/test/registered/jit/test_dsv3_router_gemm.py @@ -7,7 +7,11 @@ import pytest import torch from sglang.jit_kernel.dsv3_router_gemm import dsv3_router_gemm -from sglang.jit_kernel.utils import get_ci_test_range, get_jit_cuda_arch, is_hip_runtime +from sglang.kernels.jit.utils import ( + get_ci_test_range, + get_jit_cuda_arch, + is_hip_runtime, +) 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") diff --git a/test/registered/jit/test_fused_add_rmsnorm.py b/test/registered/jit/test_fused_add_rmsnorm.py index 8edb56345..afdabb241 100644 --- a/test/registered/jit/test_fused_add_rmsnorm.py +++ b/test/registered/jit/test_fused_add_rmsnorm.py @@ -4,7 +4,7 @@ import sys import pytest import torch -from sglang.jit_kernel.utils import get_ci_test_range +from sglang.kernels.jit.utils import get_ci_test_range from sglang.test.ci.ci_register import register_cuda_ci register_cuda_ci(est_time=10, stage="base-b-kernel-unit", runner_config="1-gpu-large") diff --git a/test/registered/jit/test_mla_kv_pack_quantize_fp8.py b/test/registered/jit/test_mla_kv_pack_quantize_fp8.py index 3266a4715..3b75e4a37 100644 --- a/test/registered/jit/test_mla_kv_pack_quantize_fp8.py +++ b/test/registered/jit/test_mla_kv_pack_quantize_fp8.py @@ -4,7 +4,7 @@ import pytest import torch from sglang.jit_kernel.mla_kv_pack_quantize_fp8 import mla_kv_pack_quantize_fp8 -from sglang.jit_kernel.utils import get_ci_test_range +from sglang.kernels.jit.utils import get_ci_test_range from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci register_cuda_ci(est_time=60, stage="base-b-kernel-unit", runner_config="1-gpu-large") diff --git a/test/registered/jit/test_moe_align_block_size.py b/test/registered/jit/test_moe_align_block_size.py index 78bf6e121..461791d50 100644 --- a/test/registered/jit/test_moe_align_block_size.py +++ b/test/registered/jit/test_moe_align_block_size.py @@ -7,7 +7,7 @@ import triton import triton.language as tl from sglang.jit_kernel.moe_align import moe_align_block_size -from sglang.jit_kernel.utils import get_ci_test_range +from sglang.kernels.jit.utils import get_ci_test_range from sglang.test.ci.ci_register import register_cuda_ci register_cuda_ci(est_time=28, stage="base-b-kernel-unit", runner_config="1-gpu-large") diff --git a/test/registered/jit/test_moe_fused_gate.py b/test/registered/jit/test_moe_fused_gate.py index 07ac21f66..58a06104a 100644 --- a/test/registered/jit/test_moe_fused_gate.py +++ b/test/registered/jit/test_moe_fused_gate.py @@ -22,7 +22,7 @@ import pytest import torch from sglang.jit_kernel.moe_fused_gate import moe_fused_gate, moe_fused_gate_jit -from sglang.jit_kernel.utils import get_ci_test_range +from sglang.kernels.jit.utils import get_ci_test_range from sglang.test.ci.ci_register import register_cuda_ci register_cuda_ci(est_time=8, stage="base-b-kernel-unit", runner_config="1-gpu-large") diff --git a/test/registered/jit/test_per_tensor_quant_fp8.py b/test/registered/jit/test_per_tensor_quant_fp8.py index 75a389453..589b0ae69 100644 --- a/test/registered/jit/test_per_tensor_quant_fp8.py +++ b/test/registered/jit/test_per_tensor_quant_fp8.py @@ -6,7 +6,7 @@ import pytest import torch from sglang.jit_kernel.per_tensor_quant_fp8 import per_tensor_quant_fp8 -from sglang.jit_kernel.utils import get_ci_test_range +from sglang.kernels.jit.utils import get_ci_test_range from sglang.test.ci.ci_register import register_cuda_ci register_cuda_ci(est_time=16, stage="base-b-kernel-unit", runner_config="1-gpu-large") diff --git a/test/registered/jit/test_per_token_group_quant.py b/test/registered/jit/test_per_token_group_quant.py index 3ce482480..485ab432f 100644 --- a/test/registered/jit/test_per_token_group_quant.py +++ b/test/registered/jit/test_per_token_group_quant.py @@ -22,7 +22,7 @@ import pytest import torch from sglang.jit_kernel.per_token_group_quant import per_token_group_quant -from sglang.jit_kernel.utils import get_ci_test_range +from sglang.kernels.jit.utils import get_ci_test_range from sglang.kernels.ops.quantization.fp8_kernel import ( create_per_token_group_quant_fp8_output_scale, fp8_dtype, diff --git a/test/registered/jit/test_per_token_group_quant_8bit_v2.py b/test/registered/jit/test_per_token_group_quant_8bit_v2.py index 075dcec0c..3fcbca018 100644 --- a/test/registered/jit/test_per_token_group_quant_8bit_v2.py +++ b/test/registered/jit/test_per_token_group_quant_8bit_v2.py @@ -6,7 +6,7 @@ import torch from sglang.jit_kernel.per_token_group_quant_8bit_v2 import ( per_token_group_quant_8bit_v2, ) -from sglang.jit_kernel.utils import get_ci_test_range +from sglang.kernels.jit.utils import get_ci_test_range 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") diff --git a/test/registered/jit/test_qknorm.py b/test/registered/jit/test_qknorm.py index a540ae3a1..136cd189e 100644 --- a/test/registered/jit/test_qknorm.py +++ b/test/registered/jit/test_qknorm.py @@ -5,7 +5,7 @@ import pytest import torch import triton -from sglang.jit_kernel.utils import get_ci_test_range +from sglang.kernels.jit.utils import get_ci_test_range 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") diff --git a/test/registered/jit/test_qknorm_across_heads.py b/test/registered/jit/test_qknorm_across_heads.py index e21258b55..8ab741d0d 100644 --- a/test/registered/jit/test_qknorm_across_heads.py +++ b/test/registered/jit/test_qknorm_across_heads.py @@ -5,7 +5,7 @@ import pytest import torch import triton -from sglang.jit_kernel.utils import get_ci_test_range +from sglang.kernels.jit.utils import get_ci_test_range from sglang.test.ci.ci_register import register_cuda_ci register_cuda_ci(est_time=15, stage="base-b-kernel-unit", runner_config="1-gpu-large") diff --git a/test/registered/jit/test_rmsnorm.py b/test/registered/jit/test_rmsnorm.py index 824d77704..206942240 100644 --- a/test/registered/jit/test_rmsnorm.py +++ b/test/registered/jit/test_rmsnorm.py @@ -4,7 +4,7 @@ import sys import pytest import torch -from sglang.jit_kernel.utils import get_ci_test_range +from sglang.kernels.jit.utils import get_ci_test_range from sglang.srt.utils import is_hip from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci diff --git a/test/registered/jit/test_rmsnorm_hf.py b/test/registered/jit/test_rmsnorm_hf.py index 8152e9456..800b31671 100644 --- a/test/registered/jit/test_rmsnorm_hf.py +++ b/test/registered/jit/test_rmsnorm_hf.py @@ -10,7 +10,7 @@ from sglang.jit_kernel.rmsnorm_hf import ( is_supported_rmsnorm_hf_hidden_size, rmsnorm_hf, ) -from sglang.jit_kernel.utils import get_ci_test_range +from sglang.kernels.jit.utils import get_ci_test_range from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci register_cuda_ci(est_time=30, stage="base-b-kernel-unit", runner_config="1-gpu-large") diff --git a/test/registered/jit/test_rope.py b/test/registered/jit/test_rope.py index e1ceb7287..7e779908f 100644 --- a/test/registered/jit/test_rope.py +++ b/test/registered/jit/test_rope.py @@ -4,7 +4,7 @@ import pytest import torch import triton -from sglang.jit_kernel.utils import get_ci_test_range +from sglang.kernels.jit.utils import get_ci_test_range from sglang.srt.utils import is_hip from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci diff --git a/test/registered/jit/test_set_mla_kv_buffer.py b/test/registered/jit/test_set_mla_kv_buffer.py index 4062355ef..2015a04eb 100644 --- a/test/registered/jit/test_set_mla_kv_buffer.py +++ b/test/registered/jit/test_set_mla_kv_buffer.py @@ -7,7 +7,7 @@ from sglang.jit_kernel.set_mla_kv_buffer import ( can_use_set_mla_kv_buffer, set_mla_kv_buffer, ) -from sglang.jit_kernel.utils import get_ci_test_range +from sglang.kernels.jit.utils import get_ci_test_range 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") diff --git a/test/registered/jit/test_store_cache.py b/test/registered/jit/test_store_cache.py index c731da48c..fdb620207 100644 --- a/test/registered/jit/test_store_cache.py +++ b/test/registered/jit/test_store_cache.py @@ -5,7 +5,7 @@ import pytest import torch from sglang.jit_kernel.kvcache import can_use_store_cache, store_cache -from sglang.jit_kernel.utils import get_ci_test_range +from sglang.kernels.jit.utils import get_ci_test_range from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci register_cuda_ci(est_time=28, stage="base-b-kernel-unit", runner_config="1-gpu-large") diff --git a/test/registered/jit/test_symm_mem_all_gather.py b/test/registered/jit/test_symm_mem_all_gather.py index fd3341a07..ea6b7e650 100644 --- a/test/registered/jit/test_symm_mem_all_gather.py +++ b/test/registered/jit/test_symm_mem_all_gather.py @@ -27,7 +27,7 @@ import torch.distributed as dist import sglang.srt.distributed.parallel_state as ps from sglang.jit_kernel.tests.utils import multigpu_pytest_main -from sglang.jit_kernel.utils import cache_once, get_ci_test_range +from sglang.kernels.jit.utils import cache_once, get_ci_test_range from sglang.srt.distributed.device_communicators.triton_symm_mem_ag import ( all_gather_inner, create_state, diff --git a/test/registered/jit/test_timestep_embedding.py b/test/registered/jit/test_timestep_embedding.py index bb92a8e7d..3c0834e82 100644 --- a/test/registered/jit/test_timestep_embedding.py +++ b/test/registered/jit/test_timestep_embedding.py @@ -13,7 +13,7 @@ except Exception: from sglang.jit_kernel.timestep_embedding import ( timestep_embedding as timestep_embedding_cuda, ) -from sglang.jit_kernel.utils import get_ci_test_range +from sglang.kernels.jit.utils import get_ci_test_range from sglang.test.ci.ci_register import register_cuda_ci register_cuda_ci(est_time=16, stage="base-b-kernel-unit", runner_config="1-gpu-large") diff --git a/test/registered/jit/test_tp_qknorm.py b/test/registered/jit/test_tp_qknorm.py index 6eccd8416..ab0730f27 100644 --- a/test/registered/jit/test_tp_qknorm.py +++ b/test/registered/jit/test_tp_qknorm.py @@ -21,7 +21,7 @@ from sglang.jit_kernel.all_reduce import ( ) from sglang.jit_kernel.mp import register_comm_cleanup from sglang.jit_kernel.tests.utils import multigpu_pytest_main -from sglang.jit_kernel.utils import cache_once +from sglang.kernels.jit.utils import cache_once from sglang.srt.distributed.device_communicators.custom_all_reduce_v2 import ( CustomAllReduceV2, ) diff --git a/test/registered/unit/distributed/test_vmm_utils.py b/test/registered/unit/distributed/test_vmm_utils.py index 8637969cf..f1cd47551 100644 --- a/test/registered/unit/distributed/test_vmm_utils.py +++ b/test/registered/unit/distributed/test_vmm_utils.py @@ -21,7 +21,7 @@ import torch.distributed as dist from cuda.bindings import driver as drv from sglang.jit_kernel.tests.utils import multigpu_pytest_main -from sglang.jit_kernel.utils import cache_once +from sglang.kernels.jit.utils import cache_once from sglang.srt.distributed.device_communicators.vmm_utils import ( check_drv, exchange_posix_fds,