add hicache jit test (#17847)

Signed-off-by: Xuchun Shang <xuchun.shang@linux.alibaba.com>
This commit is contained in:
Xuchun Shang
2026-02-06 16:54:33 +08:00
committed by GitHub
parent f798ab9775
commit 3d68bd9d9b
6 changed files with 522 additions and 73 deletions

View File

@@ -1,17 +1,19 @@
import itertools
from typing import Tuple
import torch
import triton
import triton.testing
from sgl_kernel import rmsnorm
from sglang.jit_kernel.benchmark.utils import is_in_ci
from sglang.jit_kernel.benchmark.utils import (
DEFAULT_DEVICE,
DEFAULT_DTYPE,
get_benchmark_range,
run_benchmark,
)
from sglang.jit_kernel.norm import fused_inplace_qknorm
from sglang.srt.utils import get_current_device_stream_fast
IS_CI = is_in_ci()
alt_stream = torch.cuda.Stream()
@@ -73,17 +75,19 @@ def torch_impl_qknorm(
HEAD_DIM = 128
DTYPE = torch.bfloat16
DEVICE = "cuda"
if IS_CI:
BS_RANGE = [16]
GQA_RANGE = [4]
KV_HEAD_RANGE = [1]
else:
BS_RANGE = [2**n for n in range(0, 14)]
GQA_RANGE = [4, 8]
KV_HEAD_RANGE = [1, 2, 4, 8]
BS_RANGE = get_benchmark_range(
full_range=[2**n for n in range(0, 14)],
ci_range=[16],
)
GQA_RANGE = get_benchmark_range(
full_range=[4, 8],
ci_range=[4],
)
KV_HEAD_RANGE = get_benchmark_range(
full_range=[1, 2, 4, 8],
ci_range=[1],
)
LINE_VALS = ["aot", "jit", "fi", "torch"]
LINE_NAMES = ["SGL AOT Kernel", "SGL JIT Kernel", "FlashInfer", "PyTorch"]
@@ -105,14 +109,16 @@ configs = list(itertools.product(GQA_RANGE, KV_HEAD_RANGE, BS_RANGE))
args={},
)
)
def benchmark(
batch_size: int, GQA: int, num_kv_heads: int, provider: str
) -> Tuple[float, float, float]:
def benchmark(batch_size: int, GQA: int, num_kv_heads: int, provider: str):
num_qo_heads = GQA * num_kv_heads
q = torch.randn((batch_size, num_qo_heads, HEAD_DIM), dtype=DTYPE, device=DEVICE)
k = torch.randn((batch_size, num_kv_heads, HEAD_DIM), dtype=DTYPE, device=DEVICE)
q_weight = torch.randn(HEAD_DIM, dtype=DTYPE, device=DEVICE)
k_weight = torch.randn(HEAD_DIM, dtype=DTYPE, device=DEVICE)
q = torch.randn(
(batch_size, num_qo_heads, HEAD_DIM), dtype=DEFAULT_DTYPE, device=DEFAULT_DEVICE
)
k = torch.randn(
(batch_size, num_kv_heads, HEAD_DIM), dtype=DEFAULT_DTYPE, device=DEFAULT_DEVICE
)
q_weight = torch.randn(HEAD_DIM, dtype=DEFAULT_DTYPE, device=DEFAULT_DEVICE)
k_weight = torch.randn(HEAD_DIM, dtype=DEFAULT_DTYPE, device=DEFAULT_DEVICE)
FN_MAP = {
"aot": sglang_aot_qknorm,
"jit": sglang_jit_qknorm,
@@ -120,9 +126,7 @@ def benchmark(
"torch": torch_impl_qknorm,
}
fn = lambda: FN_MAP[provider](q, k, q_weight, k_weight)
quantiles = [0.5, 0.2, 0.8]
ms, min_ms, max_ms = triton.testing.do_bench_cudagraph(fn, quantiles=quantiles) # type: ignore
return 1000 * ms, 1000 * max_ms, 1000 * min_ms
return run_benchmark(fn)
if __name__ == "__main__":