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