diff --git a/sgl-kernel/benchmark/bench_per_token_group_quant_8bit.py b/sgl-kernel/benchmark/bench_per_token_group_quant_8bit.py index 58a3c643350c..9023718a4b42 100644 --- a/sgl-kernel/benchmark/bench_per_token_group_quant_8bit.py +++ b/sgl-kernel/benchmark/bench_per_token_group_quant_8bit.py @@ -1,16 +1,10 @@ import itertools import os -import time -from functools import partial -from pathlib import Path import torch import triton from sgl_kernel.test_utils import create_per_token_group_quant_test_data -from sglang.kernels.ops.quantization.fp8_kernel import ( - create_per_token_group_quant_fp8_output_scale, -) from sglang.kernels.ops.quantization.fp8_kernel import ( per_token_group_quant_8bit as triton_per_token_group_quant_8bit, ) @@ -223,17 +217,19 @@ def benchmark( "_per_token_group_quant_8bit|_silu_and_mul_post_quant_kernel", ), "sglang": ( - partial(sglang_per_token_group_quant_8bit, enable_v2=True), + sglang_per_token_group_quant_8bit, "per_token_group_quant_8bit_kernel", ), }[provider] - bench_fn = lambda: fn( - x=x, - masked_m=masked_m, - group_size=group_size, - dst_dtype=dst_dtype, - **{k: v for k, v in flags.items() if k not in ["masked_layout_mode"]}, - ) + + def bench_fn(): + return fn( + x=x, + masked_m=masked_m, + group_size=group_size, + dst_dtype=dst_dtype, + **{k: v for k, v in flags.items() if k not in ["masked_layout_mode"]}, + ) time_s = bench_kineto( bench_fn, kernel_names=kernel_names, num_tests=300 if mode_concentrated else 30 diff --git a/sgl-kernel/tests/test_per_token_group_quant_8bit.py b/sgl-kernel/tests/test_per_token_group_quant_8bit.py index 38699314a1d5..0432f4ecea52 100644 --- a/sgl-kernel/tests/test_per_token_group_quant_8bit.py +++ b/sgl-kernel/tests/test_per_token_group_quant_8bit.py @@ -1,8 +1,5 @@ import itertools -import os import sys -import time -from pathlib import Path import pytest import torch @@ -17,7 +14,7 @@ from sglang.kernels.ops.quantization.fp8_kernel import ( sglang_per_token_group_quant_8bit, ) -from sglang.srt.utils import get_bool_env_var, is_hip +from sglang.srt.utils import is_hip _is_hip = is_hip() fp8_type_ = torch.float8_e4m3fnuz if _is_hip else torch.float8_e4m3fn @@ -156,7 +153,7 @@ def _postprocess(x_q, x_s): *triton_per_token_group_quant_8bit(**execute_kwargs) ) x_q_sglang, x_s_sglang = _postprocess( - *sglang_per_token_group_quant_8bit(**execute_kwargs, enable_v2=True) + *sglang_per_token_group_quant_8bit(**execute_kwargs) ) try: