Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
24 changes: 10 additions & 14 deletions sgl-kernel/benchmark/bench_per_token_group_quant_8bit.py
Original file line number Diff line number Diff line change
@@ -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,
)
Expand Down Expand Up @@ -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
Expand Down
7 changes: 2 additions & 5 deletions sgl-kernel/tests/test_per_token_group_quant_8bit.py
Original file line number Diff line number Diff line change
@@ -1,8 +1,5 @@
import itertools
import os
import sys
import time
from pathlib import Path

import pytest
import torch
Expand All @@ -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
Expand Down Expand Up @@ -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:
Expand Down
Loading