[Kernel] Reclassify kernel tests by ops group + move helpers out of the package (RFC #29630) (#32128)

Co-authored-by: Claude Opus 4.8 <noreply@anthropic.com>
This commit is contained in:
Xiaoyu Zhang
2026-07-23 12:18:27 +08:00
committed by GitHub
co-authored by Claude Opus 4.8
parent a2935ce329
commit 2d1a7be8c4
205 changed files with 204 additions and 173 deletions
@@ -0,0 +1,117 @@
import torch
import torch.nn.functional as F
from sgl_kernel import gelu_and_mul as gelu_and_mul_aot
from sgl_kernel import gelu_tanh_and_mul as gelu_tanh_and_mul_aot
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 (
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.test.ci.ci_register import register_amd_ci, register_cuda_ci
register_cuda_ci(
est_time=30, stage="base-b-kernel-benchmark", runner_config="1-gpu-large"
)
register_amd_ci(est_time=30, stage="jit-kernel-benchmark", runner_config="amd")
@torch.compile
def silu_and_mul(input: torch.Tensor) -> torch.Tensor:
lhs, rhs = input.split(input.shape[-1] // 2, dim=-1)
return F.silu(lhs) * rhs
@torch.compile
def gelu_and_mul(input: torch.Tensor) -> torch.Tensor:
lhs, rhs = input.split(input.shape[-1] // 2, dim=-1)
return F.gelu(lhs, approximate="none") * rhs
@torch.compile
def gelu_tanh_and_mul(input: torch.Tensor) -> torch.Tensor:
lhs, rhs = input.split(input.shape[-1] // 2, dim=-1)
return F.gelu(lhs, approximate="tanh") * rhs
OPS = {
"silu": (silu_and_mul_aot, silu_and_mul_jit, silu_and_mul),
"gelu": (gelu_and_mul_aot, gelu_and_mul_jit, gelu_and_mul),
"gelu_tanh": (gelu_tanh_and_mul_aot, gelu_tanh_and_mul_jit, gelu_tanh_and_mul),
}
@marker.parametrize("op_name", ["silu", "gelu", "gelu_tanh"])
@marker.parametrize("dim", [1024, 4096, 6144, 8192], [4096])
@marker.parametrize("batch_size", [2**x for x in range(0, 15)], [8, 512])
@marker.benchmark("impl", ["aot", "jit", "torch"])
def benchmark(op_name: str, dim: int, batch_size: int, impl: str):
x = create_random(batch_size, dim * 2)
aot_op, jit_op, torch_op = OPS[op_name]
fn = {"aot": aot_op, "jit": jit_op, "torch": torch_op}[impl]
return marker.do_bench(fn, input_args=(x,))
def _make_expert_ids(num_tokens: int, skip_ratio: float) -> torch.Tensor:
expert_ids = torch.randint(low=0, high=8, size=(num_tokens,), dtype=torch.int32)
if skip_ratio > 0:
skip = torch.rand(num_tokens) < skip_ratio
expert_ids[skip] = -1
return expert_ids
@marker.parametrize("op_name", ["silu", "gelu"])
@marker.parametrize("dim", [1024, 4096, 8192], [4096])
@marker.parametrize("batch_size", [64, 256, 1024, 4096, 16384], [1024])
@marker.parametrize("skip_ratio", [0.0, 0.25, 0.5], [0.25])
@marker.benchmark("impl", ["unfiltered", "filtered"])
def benchmark_filter(
op_name: str, dim: int, batch_size: int, skip_ratio: float, impl: str
):
torch.random.manual_seed(42)
x = create_random(batch_size, dim * 2)
jit_fn = silu_and_mul_jit if op_name == "silu" else gelu_and_mul_jit
extra_kwargs = {}
expert_ids = _make_expert_ids(batch_size, skip_ratio)
if impl == "filtered":
extra_kwargs = {"expert_ids": expert_ids.to(x.device), "expert_step": 1}
# NOTE: get the unmasked part from `experts_ids`
real_skip_ratio = (expert_ids == -1).sum().item() / batch_size
effective_bytes = int(x.nbytes * (1 - real_skip_ratio) * 1.5)
return marker.do_bench(
jit_fn,
input_args=(x,),
input_kwargs=extra_kwargs,
memory_args=None, # x is dynamic (counted in extra_memory_footprint)
memory_output=None, # same, output is dynamic
extra_memory_footprint=effective_bytes,
)
@torch.compile
def relu2_torch(input: torch.Tensor) -> torch.Tensor:
return F.relu(input).pow(2)
@marker.parametrize("dim", [1024, 4096, 6144, 8192], [4096])
@marker.parametrize("batch_size", [2**x for x in range(0, 15)], [8, 512])
@marker.benchmark("impl", ["jit", "torch"])
def benchmark_unary(dim: int, batch_size: int, impl: str):
x = create_random(batch_size, dim)
fn = {"jit": relu2_jit, "torch": relu2_torch}[impl]
return marker.do_bench(fn, input_args=(x,))
if __name__ == "__main__":
benchmark.run()
benchmark_filter.run()
benchmark_unary.run()
@@ -0,0 +1,65 @@
import torch
import triton
import triton.testing
from sglang.kernels.jit.benchmark.utils import (
DEFAULT_DEVICE,
get_benchmark_range,
run_benchmark_no_cudagraph,
)
from sglang.kernels.ops.attention.add_constant import (
_jit_add_constant_module,
add_constant,
)
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
register_cuda_ci(
est_time=15, stage="base-b-kernel-benchmark", runner_config="1-gpu-large"
)
register_amd_ci(est_time=15, stage="jit-kernel-benchmark", runner_config="amd")
CONSTANT = 7
SIZE_LIST = get_benchmark_range(
full_range=[128, 1024, 1025, 4096, 4097, 65536, 2**20, 2**22, 2**24],
ci_range=[4096, 2**20],
)
@triton.testing.perf_report(
triton.testing.Benchmark(
x_names=["size"],
x_vals=SIZE_LIST,
line_arg="provider",
line_vals=["jit_module", "jit_wrapper", "torch"],
line_names=["JIT module", "JIT wrapper", "PyTorch"],
styles=[("blue", "-"), ("orange", "-"), ("green", "--")],
ylabel="us",
plot_name="add-constant-performance",
args={},
)
)
def benchmark(size: int, provider: str):
src = torch.arange(size, dtype=torch.int32, device=DEFAULT_DEVICE)
if provider == "jit_module":
dst = torch.empty_like(src)
module = _jit_add_constant_module(CONSTANT)
def fn():
module.add_constant(dst, src)
elif provider == "jit_wrapper":
def fn():
add_constant(src, CONSTANT)
else:
def fn():
src + CONSTANT
return run_benchmark_no_cudagraph(fn)
if __name__ == "__main__":
benchmark.run(print_data=True)
@@ -0,0 +1,67 @@
import itertools
import torch
import triton
import triton.testing
from sglang.kernels.jit.benchmark.utils import (
DEFAULT_DEVICE,
get_benchmark_range,
run_benchmark,
)
from sglang.kernels.ops.attention.clamp_position import clamp_position_cuda
from sglang.srt.utils import get_compiler_backend
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
register_cuda_ci(
est_time=13, stage="base-b-kernel-benchmark", runner_config="1-gpu-large"
)
register_amd_ci(est_time=16, stage="jit-kernel-benchmark", runner_config="amd")
SIZE_LIST = get_benchmark_range(
full_range=[2**n for n in range(4, 16)],
ci_range=[256, 4096],
)
configs = list(itertools.product(SIZE_LIST))
def _torch_clamp_position(seq_lens):
return torch.clamp(seq_lens - 1, min=0).to(torch.int64)
_compiled_clamp_position = torch.compile(
_torch_clamp_position, dynamic=True, backend=get_compiler_backend()
)
@triton.testing.perf_report(
triton.testing.Benchmark(
x_names=["size"],
x_vals=configs,
line_arg="provider",
line_vals=["jit", "torch_compile", "torch"],
line_names=["SGL JIT Kernel", "torch.compile", "PyTorch"],
styles=[("blue", "-"), ("green", "-."), ("red", "--")],
ylabel="us",
plot_name="clamp-position-performance",
args={},
)
)
def benchmark(size: int, provider: str):
seq_lens = torch.randint(
0, 10000, (size,), dtype=torch.int64, device=DEFAULT_DEVICE
)
if provider == "jit":
fn = lambda: clamp_position_cuda(seq_lens)
elif provider == "torch_compile":
fn = lambda: _compiled_clamp_position(seq_lens)
else:
fn = lambda: _torch_clamp_position(seq_lens)
return run_benchmark(fn)
if __name__ == "__main__":
benchmark.run(print_data=True)
@@ -0,0 +1,164 @@
import itertools
import torch
import triton
import triton.testing
from sgl_kernel import concat_mla_absorb_q as aot_absorb_q
from sgl_kernel import concat_mla_k as aot_k
from sglang.kernels.jit.benchmark.utils import run_benchmark
from sglang.kernels.ops.attention.concat_mla import concat_mla_absorb_q as jit_absorb_q
from sglang.kernels.ops.attention.concat_mla import concat_mla_k as jit_k
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.utils import is_in_ci
register_cuda_ci(
est_time=6, stage="base-b-kernel-benchmark", runner_config="1-gpu-large"
)
IS_CI = is_in_ci()
NUM_LOCAL_HEADS = 128
QK_NOPE_HEAD_DIM = 128
QK_ROPE_HEAD_DIM = 64
K_HEAD_DIM = QK_NOPE_HEAD_DIM + QK_ROPE_HEAD_DIM
A_LAST_DIM = 512
B_LAST_DIM = 64
DTYPE = torch.bfloat16
DEVICE = "cuda"
def aot_concat_mla_k(k, k_nope, k_rope):
aot_k(k, k_nope, k_rope)
def jit_concat_mla_k(k, k_nope, k_rope):
jit_k(k, k_nope, k_rope)
def torch_concat_mla_k(k, k_nope, k_rope):
nope_head_dim = k_nope.shape[-1]
k[:, :, :nope_head_dim] = k_nope
k[:, :, nope_head_dim:] = k_rope.expand(-1, k.shape[1], -1)
def aot_concat_mla_absorb_q(a, b):
return aot_absorb_q(a, b)
def jit_concat_mla_absorb_q(a, b):
return jit_absorb_q(a, b)
def torch_concat_mla_absorb_q(a, b, out):
a_last_dim = a.shape[-1]
out[:, :, :a_last_dim] = a
out[:, :, a_last_dim:] = b
if IS_CI:
NUM_TOKENS_VALS = [256, 1024]
else:
NUM_TOKENS_VALS = [256, 512, 1024, 2048, 4096, 8192, 16384, 32768]
K_LINE_VALS = ["aot", "jit", "torch"]
K_LINE_NAMES = ["SGL AOT Kernel", "SGL JIT Kernel", "PyTorch"]
K_STYLES = [("orange", "-"), ("blue", "--"), ("green", "-.")]
def _create_concat_mla_k_data(num_tokens):
"""Allocate oversized containers and slice to produce non-contiguous tensors."""
k_nope_container = torch.randn(
(num_tokens, NUM_LOCAL_HEADS, QK_NOPE_HEAD_DIM + 128),
dtype=DTYPE,
device=DEVICE,
)
k_nope = k_nope_container[:, :, :QK_NOPE_HEAD_DIM]
k_rope_container = torch.randn(
(num_tokens, 1, 128 + QK_ROPE_HEAD_DIM),
dtype=DTYPE,
device=DEVICE,
)
k_rope = k_rope_container[:, :, -QK_ROPE_HEAD_DIM:]
k = torch.empty(
(num_tokens, NUM_LOCAL_HEADS, K_HEAD_DIM),
dtype=DTYPE,
device=DEVICE,
)
return k, k_nope, k_rope
@triton.testing.perf_report(
triton.testing.Benchmark(
x_names=["num_tokens"],
x_vals=NUM_TOKENS_VALS,
line_arg="provider",
line_vals=K_LINE_VALS,
line_names=K_LINE_NAMES,
styles=K_STYLES,
ylabel="us",
plot_name="concat-mla-k-performance",
args={},
)
)
def bench_concat_mla_k(num_tokens: int, provider: str):
k, k_nope, k_rope = _create_concat_mla_k_data(num_tokens)
FN_MAP = {
"aot": aot_concat_mla_k,
"jit": jit_concat_mla_k,
"torch": torch_concat_mla_k,
}
fn = lambda: FN_MAP[provider](k, k_nope, k_rope)
return run_benchmark(fn)
if IS_CI:
ABSORB_Q_VALS = list(itertools.product([4, 16], [16]))
else:
ABSORB_Q_VALS = list(itertools.product([1, 4, 8, 16, 32], [1, 8, 32, 128]))
Q_LINE_VALS = ["aot", "jit", "torch"]
Q_LINE_NAMES = ["SGL AOT Kernel", "SGL JIT Kernel", "PyTorch"]
Q_STYLES = [("orange", "-"), ("blue", "--"), ("green", "-.")]
@triton.testing.perf_report(
triton.testing.Benchmark(
x_names=["dim_0", "dim_1"],
x_vals=ABSORB_Q_VALS,
line_arg="provider",
line_vals=Q_LINE_VALS,
line_names=Q_LINE_NAMES,
styles=Q_STYLES,
ylabel="us",
plot_name="concat-mla-absorb-q-performance",
args={},
)
)
def bench_concat_mla_absorb_q(dim_0: int, dim_1: int, provider: str):
a = torch.randn(dim_0, dim_1, A_LAST_DIM, dtype=DTYPE, device=DEVICE)
b = torch.randn(dim_0, dim_1, B_LAST_DIM, dtype=DTYPE, device=DEVICE)
if provider == "torch":
out = torch.empty(
dim_0, dim_1, A_LAST_DIM + B_LAST_DIM, dtype=DTYPE, device=DEVICE
)
fn = lambda: torch_concat_mla_absorb_q(a, b, out)
else:
FN_MAP = {
"aot": aot_concat_mla_absorb_q,
"jit": jit_concat_mla_absorb_q,
}
fn = lambda: FN_MAP[provider](a, b)
return run_benchmark(fn)
if __name__ == "__main__":
bench_concat_mla_k.run(print_data=True)
bench_concat_mla_absorb_q.run(print_data=True)
@@ -0,0 +1,176 @@
from __future__ import annotations
import sys
import torch
import triton
from sglang.benchmark.bench_utils import run_bench
from sglang.kernels.jit.benchmark.utils import get_benchmark_range
from sglang.srt.utils import is_sm100_supported
from sglang.test.ci.ci_register import register_cuda_ci
register_cuda_ci(
est_time=5, stage="base-b-kernel-benchmark", runner_config="1-gpu-large"
)
try:
import deep_gemm
from deep_gemm.utils import per_token_cast_to_fp4
except Exception:
deep_gemm = None
per_token_cast_to_fp4 = None
HEAD_DIM = 128
NUM_HEADS = 64
BLOCK_KV = 64
NEXT_N = 1
shape_range = get_benchmark_range(
full_range=[(256, 8192), (256, 32768)],
ci_range=[(256, 8192)],
)
def _pack_fp8_cache(k: torch.Tensor, *, num_blocks: int) -> torch.Tensor:
k = k.view(num_blocks, BLOCK_KV, 1, HEAD_DIM)
scale = k.abs().float().amax(dim=3, keepdim=True).clamp(1.0e-4) / 448.0
k_fp8 = (k * (1.0 / scale)).to(torch.float8_e4m3fn)
buf = torch.empty(
(num_blocks, BLOCK_KV * (HEAD_DIM + 4)), dtype=torch.uint8, device="cuda"
)
buf[:, : BLOCK_KV * HEAD_DIM].copy_(
k_fp8.view(num_blocks, BLOCK_KV * HEAD_DIM).view(torch.uint8)
)
buf[:, BLOCK_KV * HEAD_DIM :].copy_(
scale.view(num_blocks, BLOCK_KV).view(torch.uint8)
)
return buf.view(num_blocks, BLOCK_KV, 1, HEAD_DIM + 4)
def _pack_fp4_cache(
k_fp4: torch.Tensor,
k_sf: torch.Tensor,
*,
num_blocks: int,
) -> torch.Tensor:
buf = torch.empty((num_blocks, BLOCK_KV * 68), dtype=torch.uint8, device="cuda")
buf[:, : BLOCK_KV * 64].view(num_blocks, BLOCK_KV, 64).copy_(
k_fp4.view(torch.uint8).view(num_blocks, BLOCK_KV, 64)
)
buf[:, BLOCK_KV * 64 :].view(num_blocks, BLOCK_KV, 4).copy_(
k_sf.contiguous().view(torch.uint8).view(num_blocks, BLOCK_KV, 4)
)
return buf.view(num_blocks, BLOCK_KV, 1, 68)
def _make_case(batch: int, seq_len_kv: int):
if deep_gemm is None or per_token_cast_to_fp4 is None:
raise RuntimeError("DeepGEMM is required for this benchmark.")
blocks_per_seq = triton.cdiv(seq_len_kv, BLOCK_KV)
padded_len = blocks_per_seq * BLOCK_KV
num_blocks = batch * blocks_per_seq
num_cache_tokens = num_blocks * BLOCK_KV
page_table = torch.arange(num_blocks, dtype=torch.int32, device="cuda").view(
batch, blocks_per_seq
)
context_lens = torch.full(
(batch, NEXT_N), seq_len_kv, dtype=torch.int32, device="cuda"
)
schedule = deep_gemm.get_paged_mqa_logits_metadata(
context_lens, BLOCK_KV, deep_gemm.get_num_sms(), indices=None
)
q = torch.randn(
batch, NEXT_N, NUM_HEADS, HEAD_DIM, device="cuda", dtype=torch.bfloat16
)
k = torch.randn(num_cache_tokens, HEAD_DIM, device="cuda", dtype=torch.bfloat16)
weights = torch.randn(batch * NEXT_N, NUM_HEADS, device="cuda", dtype=torch.float32)
q_scale = q.abs().float().amax(dim=-1, keepdim=True).clamp(1.0e-4) / 448.0
q_fp8 = (q.float() / q_scale).clamp(-448.0, 448.0).to(torch.float8_e4m3fn)
weights_fp8 = (
weights.view(batch, NEXT_N, NUM_HEADS)[:, :, :, None] * q_scale
).view(batch * NEXT_N, NUM_HEADS)
k_cache_fp8 = _pack_fp8_cache(k, num_blocks=num_blocks)
q_fp4_flat, q_sf_flat = per_token_cast_to_fp4(
q.view(-1, HEAD_DIM), use_ue8m0=True, gran_k=32, use_packed_ue8m0=True
)
q_fp4 = q_fp4_flat.view(batch, NEXT_N, NUM_HEADS, HEAD_DIM // 2)
q_sf = q_sf_flat.view(batch, NEXT_N, NUM_HEADS)
k_fp4, k_sf = per_token_cast_to_fp4(
k, use_ue8m0=True, gran_k=32, use_packed_ue8m0=True
)
k_cache_fp4 = _pack_fp4_cache(k_fp4, k_sf, num_blocks=num_blocks)
return {
"padded_len": padded_len,
"page_table": page_table,
"context_lens": context_lens,
"schedule": schedule,
"q_fp8": q_fp8,
"weights_fp8": weights_fp8,
"k_cache_fp8": k_cache_fp8,
"q_fp4": q_fp4,
"q_sf": q_sf,
"weights": weights,
"k_cache_fp4": k_cache_fp4,
}
@triton.testing.perf_report(
triton.testing.Benchmark(
x_names=["batch", "seq_len_kv"],
x_vals=shape_range,
x_log=False,
line_arg="provider",
line_vals=["fp8", "fp4"],
line_names=["Default FP8 indexer", "FP4 indexer"],
styles=[("blue", "-"), ("green", "-")],
ylabel="us",
plot_name="dsv4-fp4-indexer-performance",
args={},
)
)
def benchmark(batch: int, seq_len_kv: int, provider: str):
case = _make_case(batch, seq_len_kv)
if provider == "fp8":
fn = lambda: deep_gemm.fp8_paged_mqa_logits(
case["q_fp8"],
case["k_cache_fp8"],
case["weights_fp8"],
case["context_lens"],
case["page_table"],
case["schedule"],
case["padded_len"],
clean_logits=False,
indices=None,
)
elif provider == "fp4":
fn = lambda: deep_gemm.fp8_fp4_paged_mqa_logits(
(case["q_fp4"], case["q_sf"]),
case["k_cache_fp4"],
case["weights"],
case["context_lens"],
case["page_table"],
case["schedule"],
case["padded_len"],
clean_logits=False,
logits_dtype=torch.float32,
indices=None,
)
else:
raise ValueError(f"Unknown provider: {provider}")
return tuple(t * 1000 for t in run_bench(fn, use_cuda_graph=False))
if __name__ == "__main__":
if not is_sm100_supported():
print("[skip] DeepSeek V4 FP4 indexer benchmark requires SM100 CUDA.")
sys.exit(0)
if deep_gemm is None or per_token_cast_to_fp4 is None:
print("[skip] DeepGEMM is unavailable.")
sys.exit(0)
benchmark.run(print_data=True)
@@ -0,0 +1,269 @@
"""
Benchmark: fused_qknorm_rope JIT vs AOT (sgl_kernel)
Measures throughput (µs) for fused_qk_norm_rope across typical
LLM configurations (head_dim × num_heads × num_tokens).
Run:
python test/registered/jit/benchmark/bench_fused_qknorm_rope.py
"""
import itertools
import torch
import triton
import triton.testing
from sglang.kernels.jit.benchmark.utils import get_benchmark_range, run_benchmark
from sglang.kernels.ops.attention.fused_qknorm_rope import (
fused_qk_norm_rope as fused_qk_norm_rope_jit,
)
from sglang.test.ci.ci_register import register_cuda_ci
register_cuda_ci(
est_time=6, stage="base-b-kernel-benchmark", runner_config="1-gpu-large"
)
try:
from sgl_kernel import fused_qk_norm_rope as fused_qk_norm_rope_aot
AOT_AVAILABLE = True
except ImportError:
fused_qk_norm_rope_aot = None
AOT_AVAILABLE = False
# ---------------------------------------------------------------------------
# Benchmark configuration
# ---------------------------------------------------------------------------
NUM_TOKENS_RANGE = get_benchmark_range(
full_range=[1, 64, 256, 1024, 4096],
ci_range=[64, 512],
)
# (head_dim, num_heads_q, num_heads_k, num_heads_v) — typical MoE/dense configs
MODEL_CONFIGS = get_benchmark_range(
full_range=[
(64, 32, 8, 8), # small
(128, 32, 8, 8), # typical (e.g. Qwen3-8B)
(256, 16, 4, 4), # large head_dim
],
ci_range=[(128, 32, 8, 8)],
)
# Real production shapes (self-attention; num_heads_k == num_heads_v == num_heads_q).
# Format: (name, num_tokens, num_heads_q, num_heads_k, num_heads_v, head_dim, rotary_dim)
PRODUCTION_SHAPES = [
("flux_1024", 4096, 24, 24, 24, 128, 128),
("qwen_image_1024", 4096, 32, 32, 32, 128, 128),
("qwen_image_partial", 4096, 32, 32, 32, 128, 64),
("zimage_1024", 4096, 30, 30, 30, 128, 128),
("batch2_medium", 4096, 24, 24, 24, 128, 128), # B=2, T=2048
]
LINE_VALS = ["jit", "aot"] if AOT_AVAILABLE else ["jit"]
LINE_NAMES = ["JIT (new)", "AOT sgl_kernel"] if AOT_AVAILABLE else ["JIT (new)"]
STYLES = [("blue", "--"), ("orange", "-")] if AOT_AVAILABLE else [("blue", "--")]
# ---------------------------------------------------------------------------
# Benchmark: fused_qk_norm_rope (interleave style, no YaRN)
# ---------------------------------------------------------------------------
@triton.testing.perf_report(
triton.testing.Benchmark(
x_names=["num_tokens", "head_dim", "num_heads_q", "num_heads_k", "num_heads_v"],
x_vals=[
(nt, hd, nq, nk, nv)
for nt, (hd, nq, nk, nv) in itertools.product(
NUM_TOKENS_RANGE, MODEL_CONFIGS
)
],
line_arg="provider",
line_vals=LINE_VALS,
line_names=LINE_NAMES,
styles=STYLES,
ylabel="us",
plot_name="fused-qknorm-rope-performance",
args={},
)
)
def bench_fused_qknorm_rope(
num_tokens: int,
head_dim: int,
num_heads_q: int,
num_heads_k: int,
num_heads_v: int,
provider: str,
):
device = "cuda"
total_heads = num_heads_q + num_heads_k + num_heads_v
qkv = torch.randn(
(num_tokens, total_heads * head_dim), dtype=torch.bfloat16, device=device
)
q_weight = torch.ones(head_dim, dtype=torch.bfloat16, device=device)
k_weight = torch.ones(head_dim, dtype=torch.bfloat16, device=device)
position_ids = torch.arange(num_tokens, dtype=torch.int32, device=device)
common_kwargs = dict(
num_heads_q=num_heads_q,
num_heads_k=num_heads_k,
num_heads_v=num_heads_v,
head_dim=head_dim,
eps=1e-5,
q_weight=q_weight,
k_weight=k_weight,
base=10000.0,
is_neox=False,
position_ids=position_ids,
factor=1.0,
low=1.0,
high=32.0,
attention_factor=1.0,
rotary_dim=head_dim,
)
if provider == "jit":
fn = lambda: fused_qk_norm_rope_jit(qkv.clone(), **common_kwargs)
elif provider == "aot":
fn = lambda: fused_qk_norm_rope_aot(qkv.clone(), **common_kwargs)
else:
raise ValueError(f"Unknown provider: {provider}")
return run_benchmark(fn)
# ---------------------------------------------------------------------------
# Benchmark: fused_qk_norm_rope — real production shapes (with speedup column)
# ---------------------------------------------------------------------------
def bench_fused_qknorm_rope_production():
device = "cuda"
header = f"{'name':<22} {'tokens':>6} {'nq':>4} {'nk':>4} {'nv':>4} {'hd':>4} {'rdim':>5} {'JIT(us)':>9} {'AOT(us)':>9} {'speedup':>8}"
sep = "-" * len(header)
print("\nfused-qknorm-rope-production-shapes:")
print(sep)
print(header)
print(sep)
for (
name,
num_tokens,
num_heads_q,
num_heads_k,
num_heads_v,
head_dim,
rotary_dim,
) in PRODUCTION_SHAPES:
total_heads = num_heads_q + num_heads_k + num_heads_v
qkv = torch.randn(
(num_tokens, total_heads * head_dim), dtype=torch.bfloat16, device=device
)
q_weight = torch.ones(head_dim, dtype=torch.bfloat16, device=device)
k_weight = torch.ones(head_dim, dtype=torch.bfloat16, device=device)
position_ids = torch.arange(num_tokens, dtype=torch.int32, device=device)
common_kwargs = dict(
num_heads_q=num_heads_q,
num_heads_k=num_heads_k,
num_heads_v=num_heads_v,
head_dim=head_dim,
eps=1e-5,
q_weight=q_weight,
k_weight=k_weight,
base=10000.0,
is_neox=False,
position_ids=position_ids,
factor=1.0,
low=1.0,
high=32.0,
attention_factor=1.0,
rotary_dim=rotary_dim,
)
jit_us, _, _ = run_benchmark(
lambda: fused_qk_norm_rope_jit(qkv.clone(), **common_kwargs)
)
if AOT_AVAILABLE:
aot_us, _, _ = run_benchmark(
lambda: fused_qk_norm_rope_aot(qkv.clone(), **common_kwargs)
)
speedup = f"{aot_us / jit_us:.2f}x"
aot_str = f"{aot_us:9.3f}"
else:
aot_str = f"{'N/A':>9}"
speedup = "N/A"
print(
f"{name:<22} {num_tokens:>6} {num_heads_q:>4} {num_heads_k:>4} {num_heads_v:>4}"
f" {head_dim:>4} {rotary_dim:>5} {jit_us:9.3f} {aot_str} {speedup:>8}"
)
print(sep)
# ---------------------------------------------------------------------------
# Quick correctness diff
# ---------------------------------------------------------------------------
def calculate_diff():
if not AOT_AVAILABLE:
print("sgl_kernel not available — skipping AOT diff check")
return
device = "cuda"
print("Correctness diff (JIT vs AOT):")
for head_dim, is_neox in [(64, False), (128, False), (128, True), (256, False)]:
num_tokens = 32
num_heads_q, num_heads_k, num_heads_v = 4, 2, 2
total_heads = num_heads_q + num_heads_k + num_heads_v
qkv = torch.randn(
(num_tokens, total_heads * head_dim), dtype=torch.bfloat16, device=device
)
q_weight = torch.ones(head_dim, dtype=torch.bfloat16, device=device)
k_weight = torch.ones(head_dim, dtype=torch.bfloat16, device=device)
position_ids = torch.arange(num_tokens, dtype=torch.int32, device=device)
common = dict(
num_heads_q=num_heads_q,
num_heads_k=num_heads_k,
num_heads_v=num_heads_v,
head_dim=head_dim,
eps=1e-5,
q_weight=q_weight,
k_weight=k_weight,
base=10000.0,
is_neox=is_neox,
position_ids=position_ids,
factor=1.0,
low=1.0,
high=32.0,
attention_factor=1.0,
rotary_dim=head_dim,
)
qkv_jit = qkv.clone()
fused_qk_norm_rope_jit(qkv_jit, **common)
qkv_aot = qkv.clone()
fused_qk_norm_rope_aot(qkv_aot, **common)
match = torch.allclose(qkv_jit.float(), qkv_aot.float(), atol=1e-2, rtol=1e-2)
status = "OK" if match else "MISMATCH"
max_err = (qkv_jit.float() - qkv_aot.float()).abs().max().item()
print(
f" head_dim={head_dim:3d} is_neox={str(is_neox):5s} "
f"max_err={max_err:.2e} [{status}]"
)
if __name__ == "__main__":
calculate_diff()
print()
bench_fused_qknorm_rope.run(print_data=True)
print()
bench_fused_qknorm_rope_production()
@@ -0,0 +1,121 @@
import itertools
import math
from typing import Tuple
import torch
import torch.nn.functional as F
import triton
import triton.testing
from sglang.kernels.jit.benchmark.utils import (
DEFAULT_DEVICE,
DEFAULT_DTYPE,
get_benchmark_range,
run_benchmark,
)
from sglang.kernels.ops.attention.hadamard import hadamard_transform
from sglang.test.ci.ci_register import register_cuda_ci
register_cuda_ci(
est_time=5, stage="base-b-kernel-benchmark", runner_config="1-gpu-large"
)
# AOT kernel: might not be available in all environments.
# This is used for performance baseline comparison.
try:
from sgl_kernel import hadamard_transform as hadamard_transform_aot
AOT_AVAILABLE = True
except Exception:
AOT_AVAILABLE = False
# Naive reference implementation using scipy hadamard matrix.
try:
from scipy.linalg import hadamard
SCIPY_AVAILABLE = True
except ImportError:
SCIPY_AVAILABLE = False
# CI environment uses simplified parameters
batch_sizes = get_benchmark_range(
full_range=[1, 16, 64, 256],
ci_range=[16],
)
dim_range = get_benchmark_range(
full_range=[64, 256, 1024, 4096, 8192, 16384, 32768],
ci_range=[1024],
)
# Naive reference implementation using precomputed scipy hadamard matrix.
def torch_hadamard_transform(x, scale, H, dim, dim_padded):
flat = x.reshape(-1, dim)
if dim != dim_padded:
flat = F.pad(flat, (0, dim_padded - dim))
out = F.linear(flat, H) * scale
return out[..., :dim].reshape(x.shape)
available_providers = ["jit_kernel"]
available_names = ["JIT Kernel"]
available_styles = [("red", "-")]
if AOT_AVAILABLE:
available_providers.insert(0, "aot_kernel")
available_names.insert(0, "AOT Kernel")
available_styles.insert(0, ("green", "-"))
if SCIPY_AVAILABLE:
available_providers.append("naive")
available_names.append("Naive (scipy)")
available_styles.append(("blue", "-"))
configs = list(itertools.product(batch_sizes, dim_range))
@triton.testing.perf_report(
triton.testing.Benchmark(
x_names=["batch_size", "dim"],
x_vals=[list(c) for c in configs],
line_arg="provider",
line_vals=available_providers,
line_names=available_names,
styles=available_styles,
ylabel="us",
plot_name="hadamard-transform-performance",
args={},
)
)
def benchmark(batch_size: int, dim: int, provider: str) -> Tuple[float, float, float]:
scale = 1.0 / math.sqrt(dim)
x = torch.randn(batch_size, dim, device=DEFAULT_DEVICE, dtype=DEFAULT_DTYPE)
FN_MAP = {
"jit_kernel": lambda: hadamard_transform(x.clone(), scale=scale),
}
if AOT_AVAILABLE:
FN_MAP["aot_kernel"] = lambda: hadamard_transform_aot(x.clone(), scale=scale)
if SCIPY_AVAILABLE:
# Precompute Hadamard matrix on GPU to avoid CPU-GPU transfer
# during CUDA graph capture.
log_dim = math.ceil(math.log2(dim)) if dim > 0 else 0
dim_padded = 2**log_dim if dim > 0 else 1
H = torch.tensor(
hadamard(dim_padded, dtype=float),
dtype=DEFAULT_DTYPE,
device=DEFAULT_DEVICE,
)
FN_MAP["naive"] = lambda: torch_hadamard_transform(
x.clone(), scale, H, dim, dim_padded
)
fn = FN_MAP[provider]
return run_benchmark(fn)
if __name__ == "__main__":
print("=" * 80)
print("Benchmarking Fast Hadamard Transform")
print("=" * 80)
benchmark.run(print_data=True)
@@ -0,0 +1,124 @@
"""Benchmark: MiniMax-M3 single-stage radix-select decode topk (JIT CUDA) vs the
2-stage split-K Triton baseline (_topk_index_partial_kernel + _topk_index_merge_kernel).
Both consume the decode score tensor [num_heads, batch, max_seqblock] and produce
topk_idx [num_heads, batch, topk]. The JIT kernel is one launch with no
intermediate buffers; the baseline is two launches with split-K partials.
"""
import torch
import triton
from sglang.kernels.jit.benchmark import marker
from sglang.kernels.ops.attention.minimax_decode_topk import minimax_decode_topk
from sglang.kernels.ops.attention.minimax_sparse.decode.flash_with_topk_idx import (
_topk_index_merge_kernel,
_topk_index_partial_kernel,
)
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
register_cuda_ci(
est_time=8, stage="base-b-kernel-benchmark", runner_config="1-gpu-large"
)
register_amd_ci(est_time=8, stage="jit-kernel-benchmark", runner_config="amd")
BLOCK_SIZE = 128
TOPK = 16
NUM_HEADS = 1 # per-rank index heads at TP>=4
def _triton_2stage(score, seq_lens):
num_q_heads, batch_size, max_seqblock = score.shape
TOPK_TARGET_GRID = 64
MAX_NUM_TOPK_CHUNKS = 16
t = max(
1,
min(MAX_NUM_TOPK_CHUNKS, TOPK_TARGET_GRID // max(1, batch_size * num_q_heads)),
)
nchunks = 1 << (t.bit_length() - 1)
bt = triton.next_power_of_2(TOPK)
chunk_blocks = (max_seqblock + nchunks - 1) // nchunks
out = torch.empty(
(num_q_heads, batch_size, TOPK), device=score.device, dtype=torch.int32
)
tsp = torch.empty(
nchunks, num_q_heads, batch_size, bt, dtype=torch.float32, device=score.device
)
tip = torch.empty(
nchunks, num_q_heads, batch_size, bt, dtype=torch.int32, device=score.device
)
_topk_index_partial_kernel[(batch_size, num_q_heads, nchunks)](
score,
tsp,
tip,
seq_lens,
BLOCK_SIZE,
TOPK,
chunk_blocks,
score.stride(0),
score.stride(1),
score.stride(2),
tsp.stride(0),
tsp.stride(1),
tsp.stride(2),
tsp.stride(3),
tip.stride(0),
tip.stride(1),
tip.stride(2),
tip.stride(3),
)
_topk_index_merge_kernel[(batch_size, num_q_heads)](
tsp,
tip,
out,
seq_lens,
BLOCK_SIZE,
TOPK,
tsp.stride(0),
tsp.stride(1),
tsp.stride(2),
tsp.stride(3),
tip.stride(0),
tip.stride(1),
tip.stride(2),
tip.stride(3),
out.stride(0),
out.stride(1),
out.stride(2),
NUM_TOPK_CHUNKS=nchunks,
)
return out
def _jit(score, seq_lens):
return minimax_decode_topk(score, seq_lens, BLOCK_SIZE, TOPK)
FN_MAP = {"jit": _jit, "triton_2stage": _triton_2stage}
@marker.parametrize("ctx", [4096, 32768, 131072, 524288], [4096, 524288])
@marker.parametrize("batch", [1, 4, 16, 64, 256], [1, 64])
@marker.benchmark("impl", ["jit", "triton_2stage"])
def benchmark(ctx: int, batch: int, impl: str):
max_seqblock = (524288 + BLOCK_SIZE - 1) // BLOCK_SIZE
nb = min((ctx + BLOCK_SIZE - 1) // BLOCK_SIZE, max_seqblock)
score = torch.full(
(NUM_HEADS, batch, max_seqblock),
float("-inf"),
dtype=torch.float32,
device="cuda",
)
score[:, :, :nb] = torch.randn(NUM_HEADS, batch, nb, device="cuda") * 5.0
score[:, :, nb - 1] = 1e29 # forced local block
seq_lens = torch.full((batch,), ctx, device="cuda", dtype=torch.int32)
return marker.do_bench(
FN_MAP[impl],
input_args=(score, seq_lens),
graph_clone_args=(0, 1), # both read-only inputs
memory_args=(score,),
)
if __name__ == "__main__":
benchmark.run()
@@ -0,0 +1,148 @@
"""Benchmark: fused MiniMax-M3 Gemma-RMSNorm + partial NeoX RoPE (1 in-place
launch) vs the unfused path (GemmaRMSNorm(q) + GemmaRMSNorm(k) + rotary_emb,
3 launches + intermediates). Main attention branch, per-rank TP8 shape (nq=8, nk=1).
"""
import torch
from sglang.kernels.jit.benchmark import marker
from sglang.kernels.ops.attention.minimax_qknorm_rope import (
minimax_qknorm_rope,
minimax_qknorm_rope_grouped,
)
from sglang.test.ci.ci_register import register_cuda_ci
register_cuda_ci(
est_time=6, stage="base-b-kernel-benchmark", runner_config="1-gpu-large"
)
HEAD_DIM, ROTARY_DIM, BASE, EPS, MAXPOS = 128, 64, 5_000_000, 1e-6, 131072
NQ, NK = 64, 4
def _cache():
inv_freq = 1.0 / (
BASE
** (
torch.arange(0, ROTARY_DIM, 2, dtype=torch.float, device="cuda")
/ ROTARY_DIM
)
)
t = torch.arange(MAXPOS, dtype=torch.float, device="cuda")
freqs = torch.outer(t, inv_freq)
return torch.cat([freqs.cos(), freqs.sin()], dim=-1).contiguous()
def _unfused(qkv, wq, wk, cache, positions):
# GemmaRMSNorm (1+w) on q,k head-wise + partial neox rope, in plain torch
# (representative of the separate norm + rope launches).
T = qkv.shape[0]
q, k, v = qkv.split([NQ * HEAD_DIM, NK * HEAD_DIM, NK * HEAD_DIM], dim=-1)
def norm(x, w, nh):
x = x.reshape(T, nh, HEAD_DIM).float()
var = x.pow(2).mean(-1, keepdim=True)
return (x * torch.rsqrt(var + EPS) * (1.0 + w.float())).to(torch.bfloat16)
qn, kn = norm(q, wq, NQ), norm(k, wk, NK)
cs = cache.index_select(0, positions).float()
cos, sin = cs[:, None, :32], cs[:, None, 32:]
def rope(x):
x1, x2 = x[..., :32].float(), x[..., 32:64].float()
o1 = x1 * cos - x2 * sin
o2 = x2 * cos + x1 * sin
return torch.cat([o1, o2, x[..., 64:].float()], dim=-1).to(torch.bfloat16)
return rope(qn), rope(kn)
def _fused(qkv, wq, wk, cache, positions):
minimax_qknorm_rope(qkv, wq, wk, cache, positions, NQ, NK, NK, EPS)
return qkv
FN_MAP = {"fused": _fused, "unfused_torch": _unfused}
@marker.parametrize("T", [1, 16, 64, 256, 1024, 8192], [64, 1024])
@marker.benchmark("impl", ["fused", "unfused_torch"])
def benchmark(T: int, impl: str):
cache = _cache()
wq = (torch.randn(HEAD_DIM, device="cuda") * 0.1).to(torch.bfloat16)
wk = (torch.randn(HEAD_DIM, device="cuda") * 0.1).to(torch.bfloat16)
qkv = torch.randn(T, (NQ + 2 * NK) * HEAD_DIM, dtype=torch.bfloat16, device="cuda")
positions = torch.randint(0, MAXPOS, (T,), device="cuda", dtype=torch.int64)
return marker.do_bench(
FN_MAP[impl],
input_args=(qkv, wq, wk, cache, positions),
graph_clone_args=(0,),
memory_args=None,
)
# --- Combined main + index single launch (the fused qkv+index_qkv GEMM path) ---
# Per-rank TP8 sparse shape: main q=8/k=1/v=1 + index idx_q=1/idx_k=1 (value
# disabled). One grouped launch (q, k, idx_q, idx_k) vs two separate launches.
C_NQ, C_NKV, C_NIQ = 8, 1, 1
C_OFF_Q = 0
C_OFF_K = C_NQ
C_OFF_V = C_NQ + C_NKV
C_OFF_IQ = C_NQ + 2 * C_NKV
C_OFF_IK = C_OFF_IQ + C_NIQ
C_TOTAL_HEADS = C_OFF_IK + 1
def _combined_one(args):
qkv, wq, wk, wiq, wik, cache, positions = args
minimax_qknorm_rope_grouped(
qkv,
[
(wq, C_OFF_Q, C_NQ),
(wk, C_OFF_K, C_NKV),
(wiq, C_OFF_IQ, C_NIQ),
(wik, C_OFF_IK, 1),
],
cache,
positions,
EPS,
)
return qkv
def _combined_two(args):
# Two launches over the same buffer: main (q,k) then index (idx_q, idx_k).
qkv, wq, wk, wiq, wik, cache, positions = args
minimax_qknorm_rope_grouped(
qkv, [(wq, C_OFF_Q, C_NQ), (wk, C_OFF_K, C_NKV)], cache, positions, EPS
)
minimax_qknorm_rope_grouped(
qkv, [(wiq, C_OFF_IQ, C_NIQ), (wik, C_OFF_IK, 1)], cache, positions, EPS
)
return qkv
C_FN_MAP = {"combined_one_launch": _combined_one, "two_launches": _combined_two}
@marker.parametrize("T", [1, 16, 64, 256, 1024, 8192], [64, 1024])
@marker.benchmark("impl", ["combined_one_launch", "two_launches"])
def benchmark_combined(T: int, impl: str):
cache = _cache()
ws = [
(torch.randn(HEAD_DIM, device="cuda") * 0.1).to(torch.bfloat16)
for _ in range(4)
]
qkv = torch.randn(T, C_TOTAL_HEADS * HEAD_DIM, dtype=torch.bfloat16, device="cuda")
positions = torch.randint(0, MAXPOS, (T,), device="cuda", dtype=torch.int64)
return marker.do_bench(
C_FN_MAP[impl],
input_args=((qkv, *ws, cache, positions),),
graph_clone_args=(0,),
memory_args=None,
)
if __name__ == "__main__":
benchmark.run()
benchmark_combined.run()
@@ -0,0 +1,217 @@
"""Bench the hybrid ``mla_kv_pack_quantize_fp8`` against an inlined naive Triton baseline."""
import itertools
from typing import Tuple
import torch
import triton
import triton.language as tl
import triton.testing
from sglang.kernels.jit.benchmark.utils import (
DEFAULT_DEVICE,
DEFAULT_DTYPE,
DEFAULT_QUANTILES,
get_benchmark_range,
)
from sglang.kernels.jit.utils import is_arch_support_pdl
from sglang.kernels.ops.attention.mla_kv_pack_quantize_fp8 import (
mla_kv_pack_quantize_fp8 as hybrid_pack,
)
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
register_cuda_ci(
est_time=15, stage="base-b-kernel-benchmark", runner_config="1-gpu-large"
)
register_amd_ci(est_time=15, stage="jit-kernel-benchmark", runner_config="amd")
@triton.jit
def _triton_mla_kv_pack_quantize_fp8_kernel(
k_nope_ptr,
k_pe_ptr,
v_ptr,
k_out_ptr,
v_out_ptr,
k_scale_inv,
v_scale_inv,
s_total,
k_nope_stride_t,
k_nope_stride_h,
k_pe_stride_t,
v_stride_t,
v_stride_h,
k_out_stride_t,
k_out_stride_h,
v_out_stride_t,
v_out_stride_h,
QK_NOPE: tl.constexpr,
QK_ROPE: tl.constexpr,
V_HEAD: tl.constexpr,
FP8_DTYPE: tl.constexpr,
BLOCK_S: tl.constexpr,
ENABLE_PDL: tl.constexpr,
):
pid_s = tl.program_id(0)
pid_h = tl.program_id(1)
t_idx = pid_s * BLOCK_S + tl.arange(0, BLOCK_S)
t_mask = t_idx < s_total
nope_idx = tl.arange(0, QK_NOPE)
rope_idx = tl.arange(0, QK_ROPE)
v_idx = tl.arange(0, V_HEAD)
if ENABLE_PDL:
tl.extra.cuda.gdc_wait()
nope_off = (
t_idx[:, None] * k_nope_stride_t + pid_h * k_nope_stride_h + nope_idx[None, :]
)
k_nope = tl.load(k_nope_ptr + nope_off, mask=t_mask[:, None])
pe_off = t_idx[:, None] * k_pe_stride_t + rope_idx[None, :]
k_pe = tl.load(k_pe_ptr + pe_off, mask=t_mask[:, None])
v_off = t_idx[:, None] * v_stride_t + pid_h * v_stride_h + v_idx[None, :]
v = tl.load(v_ptr + v_off, mask=t_mask[:, None])
k_nope_fp8 = (k_nope.to(tl.float32) * k_scale_inv).to(FP8_DTYPE)
k_pe_fp8 = (k_pe.to(tl.float32) * k_scale_inv).to(FP8_DTYPE)
v_fp8 = (v.to(tl.float32) * v_scale_inv).to(FP8_DTYPE)
k_out_base = t_idx[:, None] * k_out_stride_t + pid_h * k_out_stride_h
tl.store(
k_out_ptr + k_out_base + nope_idx[None, :], k_nope_fp8, mask=t_mask[:, None]
)
tl.store(
k_out_ptr + k_out_base + QK_NOPE + rope_idx[None, :],
k_pe_fp8,
mask=t_mask[:, None],
)
v_out_off = (
t_idx[:, None] * v_out_stride_t + pid_h * v_out_stride_h + v_idx[None, :]
)
tl.store(v_out_ptr + v_out_off, v_fp8, mask=t_mask[:, None])
if ENABLE_PDL:
tl.extra.cuda.gdc_launch_dependents()
def _triton_pack(k_nope, k_pe, v, k_out, v_out):
s, num_heads, qk_nope = k_nope.shape
qk_rope = k_pe.shape[-1]
v_head = v.shape[-1]
k_pe_2d = k_pe.squeeze(1) if k_pe.dim() == 3 else k_pe
enable_pdl = is_arch_support_pdl()
if s < 512:
block_s, num_warps, num_stages = 1, 1, 2
elif s < 2048:
block_s, num_warps, num_stages = 4, 2, 3
else:
block_s, num_warps, num_stages = 16, 4, 3
extra = {"launch_pdl": True} if enable_pdl else {}
grid = (triton.cdiv(s, block_s), num_heads)
_triton_mla_kv_pack_quantize_fp8_kernel[grid](
k_nope,
k_pe_2d,
v,
k_out,
v_out,
1.0,
1.0,
s,
k_nope.stride(0),
k_nope.stride(1),
k_pe_2d.stride(0),
v.stride(0),
v.stride(1),
k_out.stride(0),
k_out.stride(1),
v_out.stride(0),
v_out.stride(1),
QK_NOPE=qk_nope,
QK_ROPE=qk_rope,
V_HEAD=v_head,
FP8_DTYPE=tl.float8e4nv,
BLOCK_S=block_s,
ENABLE_PDL=enable_pdl,
num_warps=num_warps,
num_stages=num_stages,
**extra,
)
QK_NOPE = 128
QK_ROPE = 64
V_HEAD = 128
NUM_HEADS = 32
NUM_LAYERS = 8
BS_RANGE = get_benchmark_range(
full_range=[1, 4, 16, 64, 256, 1024, 4096, 8192, 16384],
ci_range=[1, 64, 1024, 4096, 16384],
)
LINE_VALS = ["hybrid", "triton"]
LINE_NAMES = ["hybrid (v0+v1_flat)", "naive Triton"]
STYLES = [("green", "-"), ("red", "--")]
CONFIGS = list(itertools.product(BS_RANGE))
@triton.testing.perf_report(
triton.testing.Benchmark(
x_names=["batch_size"],
x_vals=CONFIGS,
line_arg="provider",
line_vals=LINE_VALS,
line_names=LINE_NAMES,
styles=STYLES,
ylabel="us",
plot_name="mla-kv-pack-quantize-fp8-performance",
args={},
)
)
def benchmark(batch_size: int, provider: str) -> Tuple[float, float, float]:
k_nope = torch.randn(
(NUM_LAYERS, batch_size, NUM_HEADS, QK_NOPE),
dtype=DEFAULT_DTYPE,
device=DEFAULT_DEVICE,
)
k_pe = torch.randn(
(NUM_LAYERS, batch_size, 1, QK_ROPE),
dtype=DEFAULT_DTYPE,
device=DEFAULT_DEVICE,
)
v = torch.randn(
(NUM_LAYERS, batch_size, NUM_HEADS, V_HEAD),
dtype=DEFAULT_DTYPE,
device=DEFAULT_DEVICE,
)
k_out = torch.empty(
(NUM_LAYERS, batch_size, NUM_HEADS, QK_NOPE + QK_ROPE),
dtype=torch.float8_e4m3fn,
device=DEFAULT_DEVICE,
)
v_out = torch.empty(
(NUM_LAYERS, batch_size, NUM_HEADS, V_HEAD),
dtype=torch.float8_e4m3fn,
device=DEFAULT_DEVICE,
)
torch.cuda.synchronize()
if provider == "hybrid":
def fn():
for i in range(NUM_LAYERS):
hybrid_pack(k_nope[i], k_pe[i], v[i], k_out=k_out[i], v_out=v_out[i])
else:
def fn():
for i in range(NUM_LAYERS):
_triton_pack(k_nope[i], k_pe[i], v[i], k_out[i], v_out[i])
ms, min_ms, max_ms = triton.testing.do_bench_cudagraph(
fn, quantiles=DEFAULT_QUANTILES
)
return (
1000 * ms / NUM_LAYERS,
1000 * max_ms / NUM_LAYERS,
1000 * min_ms / NUM_LAYERS,
)
if __name__ == "__main__":
benchmark.run(print_data=True)
@@ -0,0 +1,165 @@
"""Benchmark online c128 MTP write-prefix kernel."""
from __future__ import annotations
import itertools
from dataclasses import dataclass
import torch
import triton
import triton.testing
from sglang.kernels.jit.benchmark.utils import (
DEFAULT_DEVICE,
get_benchmark_range,
run_benchmark,
run_benchmark_no_cudagraph,
)
from sglang.kernels.ops.attention.dsv4.online_c128_mtp import (
_jit_online_c128_mtp_module,
)
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
register_cuda_ci(
est_time=10, stage="base-b-kernel-benchmark", runner_config="1-gpu-large"
)
register_amd_ci(est_time=10, stage="jit-kernel-benchmark", runner_config="amd")
HEAD_DIM = 512
STATE_DIM = HEAD_DIM * 3
SWA_PAGE_SIZE = 128
BATCH_SIZE_RANGE = get_benchmark_range(
full_range=[1, 2, 4, 8, 16, 32, 64, 128, 256],
ci_range=[8, 64],
)
NUM_VERIFY_TOKENS_RANGE = get_benchmark_range(
full_range=[1, 4, 8],
ci_range=[8],
)
BENCHMARK_CONFIGS = list(itertools.product(BATCH_SIZE_RANGE, NUM_VERIFY_TOKENS_RANGE))
@dataclass
class BenchmarkCase:
kv_score_input: torch.Tensor
seq_lens: torch.Tensor
req_pool_indices: torch.Tensor
req_to_token: torch.Tensor
ape: torch.Tensor
state: torch.Tensor
layer_bs: int
num_verify_tokens: int
state_slot_stride: int
def round_up_div(x: int, y: int) -> int:
return (x + y - 1) // y
def make_seq_lens(batch_size: int, num_verify_tokens: int) -> torch.Tensor:
# Cover chunk positions around the interesting boundaries. This exercises
# both the has-partial path and the final_seq % 128 == 0 skip-write path.
offsets = torch.tensor([0, 1, 2, 63, 120, 126, 127], dtype=torch.int64)
seq_offsets = offsets[torch.arange(batch_size, dtype=torch.int64) % offsets.numel()]
base = 8 * SWA_PAGE_SIZE
seq_lens = base + seq_offsets
assert int(seq_lens.max()) + num_verify_tokens < base + 2 * SWA_PAGE_SIZE
return seq_lens.to(device=DEFAULT_DEVICE)
def make_req_to_token(
batch_size: int, max_seq_len: int, num_chunks: int
) -> torch.Tensor:
chunk_ids = torch.arange(max_seq_len, dtype=torch.int32) // SWA_PAGE_SIZE
req_offsets = torch.arange(batch_size, dtype=torch.int32).unsqueeze(1) * num_chunks
req_to_token = req_offsets + chunk_ids.unsqueeze(0)
return req_to_token.contiguous().to(device=DEFAULT_DEVICE)
def make_case(batch_size: int, num_verify_tokens: int) -> BenchmarkCase:
seq_lens = make_seq_lens(batch_size, num_verify_tokens)
req_pool_indices = torch.arange(
batch_size, dtype=torch.int64, device=DEFAULT_DEVICE
)
max_seq_len = int(seq_lens.max().item()) + num_verify_tokens + SWA_PAGE_SIZE
num_chunks = round_up_div(max_seq_len, SWA_PAGE_SIZE)
req_to_token = make_req_to_token(batch_size, max_seq_len, num_chunks)
num_full_locs = batch_size * num_chunks
state_slot_stride = num_full_locs
state = torch.empty(
(state_slot_stride * (1 + num_verify_tokens), STATE_DIM),
dtype=torch.float32,
device=DEFAULT_DEVICE,
)
state.normal_(mean=0.0, std=0.01)
kv_score_input = torch.randn(
batch_size * num_verify_tokens,
HEAD_DIM * 2,
dtype=torch.float32,
device=DEFAULT_DEVICE,
)
ape = torch.randn(128, HEAD_DIM, dtype=torch.float32, device=DEFAULT_DEVICE)
return BenchmarkCase(
kv_score_input=kv_score_input,
seq_lens=seq_lens,
req_pool_indices=req_pool_indices,
req_to_token=req_to_token,
ape=ape,
state=state,
layer_bs=batch_size,
num_verify_tokens=num_verify_tokens,
state_slot_stride=state_slot_stride,
)
def call_write_prefix(module, case: BenchmarkCase) -> None:
module.write_prefix_states(
case.kv_score_input,
case.seq_lens,
case.req_pool_indices,
case.req_to_token,
case.ape,
case.state,
case.layer_bs,
case.num_verify_tokens,
case.state_slot_stride,
)
@triton.testing.perf_report(
triton.testing.Benchmark(
x_names=["batch_size", "num_verify_tokens"],
x_vals=BENCHMARK_CONFIGS,
line_arg="launch_mode",
line_vals=["cuda_graph", "eager"],
line_names=["CUDA graph", "Eager launch"],
styles=[("blue", "-"), ("orange", "--")],
ylabel="us",
plot_name="online-c128-mtp-write-prefix-performance",
args={},
)
)
def benchmark(
batch_size: int, num_verify_tokens: int, launch_mode: str
) -> tuple[float, float, float]:
case = make_case(batch_size, num_verify_tokens)
module = _jit_online_c128_mtp_module(
HEAD_DIM, case.seq_lens.dtype, case.req_pool_indices.dtype, case.state.dtype
)
fn = lambda: call_write_prefix(module, case)
if launch_mode == "cuda_graph":
return run_benchmark(fn)
if launch_mode == "eager":
return run_benchmark_no_cudagraph(fn)
raise ValueError(f"Unknown launch_mode: {launch_mode}")
if __name__ == "__main__":
benchmark.run(print_data=True)
@@ -0,0 +1,310 @@
import itertools
import torch
import triton
import triton.testing
from sglang.kernels.jit.benchmark.utils import (
DEFAULT_DEVICE,
DEFAULT_DTYPE,
get_benchmark_range,
run_benchmark,
)
from sglang.test.ci.ci_register import register_cuda_ci
register_cuda_ci(
est_time=6, stage="base-b-kernel-benchmark", runner_config="1-gpu-large"
)
MAX_SEQ_LEN = 131072
ROPE_BASE = 10000.0
ROPE_DIM = 128
CACHE_SIZE = 1024 * 1024
def create_cos_sin_cache(
rotary_dim: int = ROPE_DIM,
max_position: int = MAX_SEQ_LEN,
base: float = ROPE_BASE,
) -> torch.Tensor:
inv_freq = 1.0 / (
base
** (
torch.arange(0, rotary_dim, 2, dtype=torch.float32, device=DEFAULT_DEVICE)
/ rotary_dim
)
)
t = torch.arange(max_position, dtype=torch.float32, device=DEFAULT_DEVICE)
freqs = torch.einsum("i,j->ij", t, inv_freq)
cos = freqs.cos()
sin = freqs.sin()
return torch.cat((cos, sin), dim=-1)
# Pre-build the cache once
COS_SIN_CACHE = create_cos_sin_cache()
# ---------------------------------------------------------------------------
# RoPE-only provider implementations
# ---------------------------------------------------------------------------
def flashinfer_rope(
q: torch.Tensor,
k: torch.Tensor,
positions: torch.Tensor,
is_neox: bool,
) -> None:
from flashinfer.rope import apply_rope_with_cos_sin_cache_inplace
head_size = q.shape[-1]
apply_rope_with_cos_sin_cache_inplace(
positions=positions,
query=q.view(q.shape[0], -1),
key=k.view(k.shape[0], -1),
head_size=head_size,
cos_sin_cache=COS_SIN_CACHE,
is_neox=is_neox,
)
def sglang_pos_enc_rope(
q: torch.Tensor,
k: torch.Tensor,
positions: torch.Tensor,
is_neox: bool,
) -> None:
from sglang.kernels.ops.attention.rope import rotary_embedding_with_key
head_size = q.shape[-1]
rotary_embedding_with_key(
positions=positions,
query=q.view(q.shape[0], -1),
key=k.view(k.shape[0], -1),
head_size=head_size,
cos_sin_cache=COS_SIN_CACHE,
is_neox=is_neox,
)
def sglang_fused_rope(
q: torch.Tensor,
k: torch.Tensor,
positions: torch.Tensor,
is_neox: bool,
) -> None:
from sglang.kernels.ops.attention.rope import apply_rope_inplace
apply_rope_inplace(q, k, COS_SIN_CACHE, positions, is_neox=is_neox)
# ---------------------------------------------------------------------------
# RoPE + KV cache store provider implementations
# ---------------------------------------------------------------------------
def jit_rope_then_store(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
k_cache: torch.Tensor,
v_cache: torch.Tensor,
positions: torch.Tensor,
out_loc: torch.Tensor,
is_neox: bool,
) -> None:
from sglang.kernels.ops.attention.rope import apply_rope_inplace
from sglang.kernels.ops.kvcache.kvcache import store_cache
head_size = q.shape[-1]
row_dim = k.shape[-2] * head_size
apply_rope_inplace(
positions=positions,
q=q,
k=k,
rope_dim=head_size,
cos_sin_cache=COS_SIN_CACHE,
is_neox=is_neox,
)
store_cache(
k.view(-1, row_dim),
v.view(-1, row_dim),
k_cache,
v_cache,
out_loc,
)
def jit_fused_rope_store(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
k_cache: torch.Tensor,
v_cache: torch.Tensor,
positions: torch.Tensor,
out_loc: torch.Tensor,
is_neox: bool,
) -> None:
from sglang.kernels.ops.attention.rope import apply_rope_inplace_with_kvcache
apply_rope_inplace_with_kvcache(
q, k, v, k_cache, v_cache, COS_SIN_CACHE, positions, out_loc, is_neox=is_neox
)
# ---------------------------------------------------------------------------
# Benchmark configuration (shared)
# ---------------------------------------------------------------------------
BS_RANGE = get_benchmark_range(
full_range=[2**n for n in range(0, 16)],
ci_range=[16],
)
QK_HEAD_RANGE = get_benchmark_range(
full_range=[(8, 1), (16, 2), (32, 8)],
ci_range=[(16, 2)],
)
QK_HEAD_RANGE = [f"{q},{k}" for q, k in QK_HEAD_RANGE]
IS_NEOX_RANGE = get_benchmark_range(
full_range=[True, False],
ci_range=[True],
)
# ---------------------------------------------------------------------------
# Benchmark 1: RoPE only
# ---------------------------------------------------------------------------
ROPE_LINE_VALS = ["flashinfer", "jit_pos_enc", "jit_fused_rope"]
ROPE_LINE_NAMES = [
"FlashInfer",
"SGL JIT PosEnc",
"SGL JIT Fused RoPE",
]
ROPE_STYLES = [("green", "-."), ("red", "-"), ("blue", "--")]
rope_configs = list(itertools.product(QK_HEAD_RANGE, IS_NEOX_RANGE, BS_RANGE))
@triton.testing.perf_report(
triton.testing.Benchmark(
x_names=["num_q_k_heads", "is_neox", "batch_size"],
x_vals=rope_configs,
line_arg="provider",
line_vals=ROPE_LINE_VALS,
line_names=ROPE_LINE_NAMES,
styles=ROPE_STYLES,
ylabel="us",
plot_name="rope-performance",
args={},
)
)
def benchmark(batch_size: int, num_q_k_heads: str, is_neox: bool, provider: str):
qo, kv = num_q_k_heads.split(",")
num_qo_heads = int(qo)
num_kv_heads = int(kv)
q = torch.randn(
(batch_size, num_qo_heads, ROPE_DIM),
dtype=DEFAULT_DTYPE,
device=DEFAULT_DEVICE,
)
k = torch.randn(
(batch_size, num_kv_heads, ROPE_DIM),
dtype=DEFAULT_DTYPE,
device=DEFAULT_DEVICE,
)
seed = batch_size << 16 | num_qo_heads << 8 | num_kv_heads << 4 | is_neox
torch.random.manual_seed(seed)
positions = torch.randint(
MAX_SEQ_LEN, (batch_size,), device=DEFAULT_DEVICE, dtype=torch.int64
)
torch.cuda.synchronize()
FN_MAP = {
"flashinfer": flashinfer_rope,
"jit_pos_enc": sglang_pos_enc_rope,
"jit_fused_rope": sglang_fused_rope,
}
fn = lambda: FN_MAP[provider](q, k, positions, is_neox)
return run_benchmark(fn)
# ---------------------------------------------------------------------------
# Benchmark 2: RoPE + KV cache store
# ---------------------------------------------------------------------------
STORE_LINE_VALS = ["jit_rope_then_store", "jit_fused_store"]
STORE_LINE_NAMES = [
"SGL JIT RoPE + Store",
"SGL JIT Fused RoPE + Store",
]
STORE_STYLES = [("red", "-"), ("blue", "--")]
store_configs = list(itertools.product(QK_HEAD_RANGE, IS_NEOX_RANGE, BS_RANGE))
@triton.testing.perf_report(
triton.testing.Benchmark(
x_names=["num_q_k_heads", "is_neox", "batch_size"],
x_vals=store_configs,
line_arg="provider",
line_vals=STORE_LINE_VALS,
line_names=STORE_LINE_NAMES,
styles=STORE_STYLES,
ylabel="us",
plot_name="rope-store-performance",
args={},
)
)
def benchmark_store(batch_size: int, num_q_k_heads: str, is_neox: bool, provider: str):
qo, kv = num_q_k_heads.split(",")
num_qo_heads = int(qo)
num_kv_heads = int(kv)
q = torch.randn(
(batch_size, num_qo_heads, ROPE_DIM),
dtype=DEFAULT_DTYPE,
device=DEFAULT_DEVICE,
)
k = torch.randn(
(batch_size, num_kv_heads, ROPE_DIM),
dtype=DEFAULT_DTYPE,
device=DEFAULT_DEVICE,
)
v = torch.randn(
(batch_size, num_kv_heads, ROPE_DIM),
dtype=DEFAULT_DTYPE,
device=DEFAULT_DEVICE,
)
row_size = num_kv_heads * ROPE_DIM
k_cache = torch.zeros(
CACHE_SIZE, row_size, dtype=DEFAULT_DTYPE, device=DEFAULT_DEVICE
)
v_cache = torch.zeros(
CACHE_SIZE, row_size, dtype=DEFAULT_DTYPE, device=DEFAULT_DEVICE
)
out_loc = torch.randperm(CACHE_SIZE, device=DEFAULT_DEVICE, dtype=torch.int64)[
:batch_size
]
seed = batch_size << 16 | num_qo_heads << 8 | num_kv_heads << 4 | is_neox
torch.random.manual_seed(seed)
positions = torch.randint(
MAX_SEQ_LEN, (batch_size,), device=DEFAULT_DEVICE, dtype=torch.int64
)
torch.cuda.synchronize()
FN_MAP = {
"jit_rope_then_store": jit_rope_then_store,
"jit_fused_store": jit_fused_rope_store,
}
fn = lambda: FN_MAP[provider](
q, k, v, k_cache, v_cache, positions, out_loc, is_neox
)
return run_benchmark(fn)
if __name__ == "__main__":
print("Running RoPE performance benchmark...")
benchmark.run(print_data=True)
print("\nRunning RoPE + KV cache store performance benchmark...")
benchmark_store.run(print_data=True)
@@ -0,0 +1,154 @@
from __future__ import annotations
import math
import torch
import triton
import triton.testing
from sglang.kernels.jit.benchmark.utils import run_benchmark_no_cudagraph
from sglang.kernels.ops.attention.sparse_mla_q8kv8_prefill_sm90 import (
sparse_mla_q8kv8_prefill_fwd,
)
from sglang.srt.utils import is_sm90_supported
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.utils import is_in_ci
try:
from sgl_kernel.flash_mla import flash_mla_sparse_fwd
HAS_Q16_FLASHMLA = True
except ImportError:
flash_mla_sparse_fwd = None
HAS_Q16_FLASHMLA = False
register_cuda_ci(
est_time=120, stage="base-b-kernel-benchmark", runner_config="1-gpu-large"
)
IS_CI = is_in_ci()
DTYPE_FP8 = torch.float8_e4m3fn
D_V = 512
H_KV = 1
if IS_CI:
CASES = [
(2, 1024, 64, 512, 128),
(2, 1024, 64, 576, 128),
]
else:
CASES = [
(4096, 8192, 128, 576, 2048),
(4096, 32768, 128, 576, 2048),
(4096, 65536, 128, 576, 2048),
(4096, 8192, 64, 512, 512),
(4096, 32768, 64, 512, 512),
]
# This official benchmark intentionally measures the no-sink path. Current
# DeepSeek NSA E2E does not pass a per-head attention sink into sparse MLA, so
# sink-enabled timings are kernel feature coverage rather than E2E proxy data.
LINE_VALS = ["q8_fp8_jit"]
LINE_NAMES = ["Q8 FP8 JIT"]
STYLES = [("blue", "-")]
if HAS_Q16_FLASHMLA:
LINE_VALS.insert(0, "q16_bf16_flashmla")
LINE_NAMES.insert(0, "Q16 BF16 FlashMLA")
STYLES.insert(0, ("orange", "--"))
def _sm90_available() -> bool:
return is_sm90_supported()
def _make_indices(s_q: int, s_kv: int, topk: int, d_qk: int) -> torch.Tensor:
generator = torch.Generator(device="cuda")
generator.manual_seed(1000 + d_qk + topk)
return torch.randint(
0,
s_kv,
(s_q, H_KV, topk),
dtype=torch.int32,
device="cuda",
generator=generator,
)
def _make_q16_inputs(s_q: int, s_kv: int, h_q: int, d_qk: int, topk: int):
generator = torch.Generator(device="cuda")
generator.manual_seed(2000 + d_qk + s_kv)
q = torch.randn(
(s_q, h_q, d_qk), dtype=torch.bfloat16, device="cuda", generator=generator
)
kv = torch.randn(
(s_kv + 1, H_KV, d_qk), dtype=torch.bfloat16, device="cuda", generator=generator
)
indices = _make_indices(s_q, s_kv, topk, d_qk)
sm_scale = 1.0 / math.sqrt(d_qk)
return q, kv, indices, sm_scale
def _make_q8_inputs(s_q: int, s_kv: int, h_q: int, d_qk: int, topk: int):
generator = torch.Generator(device="cuda")
generator.manual_seed(3000 + d_qk + s_kv)
q = (torch.randn((s_q, h_q, d_qk), device="cuda", generator=generator) * 0.05).to(
DTYPE_FP8
)
kv = torch.zeros((s_kv + 1, H_KV, d_qk), dtype=DTYPE_FP8, device="cuda")
kv[:s_kv] = (
torch.randn((s_kv, H_KV, d_qk), device="cuda", generator=generator) * 0.05
).to(DTYPE_FP8)
indices = _make_indices(s_q, s_kv, topk, d_qk)
q_scale = torch.ones(1, dtype=torch.float32, device="cuda")
kv_scale = torch.ones(1, dtype=torch.float32, device="cuda")
sm_scale = 1.0 / math.sqrt(d_qk)
return q, kv, indices, sm_scale, q_scale, kv_scale
@triton.testing.perf_report(
triton.testing.Benchmark(
x_names=["s_q", "s_kv", "h_q", "d_qk", "topk"],
x_vals=CASES,
line_arg="provider",
line_vals=LINE_VALS,
line_names=LINE_NAMES,
styles=STYLES,
ylabel="us",
plot_name="sparse-mla-q8kv8-prefill-sm90-performance",
args={},
)
)
def bench_sparse_mla_q8kv8_prefill_sm90(
s_q: int, s_kv: int, h_q: int, d_qk: int, topk: int, provider: str
):
if provider == "q16_bf16_flashmla":
if not HAS_Q16_FLASHMLA:
raise RuntimeError(
"sgl_kernel.flash_mla.flash_mla_sparse_fwd is not available"
)
q, kv, indices, sm_scale = _make_q16_inputs(s_q, s_kv, h_q, d_qk, topk)
def fn():
return flash_mla_sparse_fwd(q, kv, indices, sm_scale, D_V)
elif provider == "q8_fp8_jit":
if not _sm90_available():
raise RuntimeError("Q8KV8 sparse prefill benchmark requires SM90 CUDA")
q, kv, indices, sm_scale, q_scale, kv_scale = _make_q8_inputs(
s_q, s_kv, h_q, d_qk, topk
)
def fn():
return sparse_mla_q8kv8_prefill_fwd(
q, kv, indices, sm_scale, q_scale, kv_scale, D_V
)
else:
raise ValueError(f"Unknown provider: {provider}")
return run_benchmark_no_cudagraph(fn)
if __name__ == "__main__":
bench_sparse_mla_q8kv8_prefill_sm90.run(print_data=True)
@@ -0,0 +1,90 @@
import torch
from sglang.kernels.jit.benchmark import marker
from sglang.kernels.ops.attention.dsv4.topk import (
plan_topk_v2,
topk_transform_512,
topk_transform_512_v2,
)
from sglang.test.ci.ci_register import register_cuda_ci
register_cuda_ci(
est_time=120, stage="base-b-kernel-benchmark", runner_config="1-gpu-large"
)
# Compressed page size used by the DSA indexer (real value is 256 // 4 = 64).
PAGE_SIZE = 64
def _make_inputs(batch_size: int, seq_len: int, k: int):
torch.random.manual_seed(42)
scores = torch.randn(batch_size, seq_len, dtype=torch.float32, device="cuda")
seq_lens = torch.full((batch_size,), seq_len, dtype=torch.int32, device="cuda")
num_pages = (seq_len + PAGE_SIZE - 1) // PAGE_SIZE
page_table = (
torch.arange(num_pages, dtype=torch.int32, device="cuda")
.unsqueeze(0)
.expand(batch_size, -1)
.contiguous()
)
out = torch.empty(batch_size, k, dtype=torch.int32, device="cuda")
return scores, seq_lens, page_table, out
def _make_p1_table(batch_size: int, seq_len: int):
# flashinfer / torch do a per-token (page_size=1) gather, so they need a
# (batch, seq) table (one entry per position) rather than the page-size-64 one.
src_page_table = (
torch.arange(seq_len, dtype=torch.int32, device="cuda")
.unsqueeze(0)
.expand(batch_size, -1)
.contiguous()
)
lengths = torch.full((batch_size,), seq_len, dtype=torch.int32, device="cuda")
return src_page_table, lengths
def _build_fn(provider: str, batch_size: int, seq_len: int, k: int):
scores, seq_lens, page_table, out = _make_inputs(batch_size, seq_len, k)
N = PAGE_SIZE
def fn(scores, seq_lens, page_table):
if provider == "jit_v1":
topk_transform_512(scores, seq_lens, page_table, out, N)
return out
elif provider == "jit_v2":
topk_transform_512_v2(scores, seq_lens, page_table, out, N, metadata)
return out
elif provider == "flashinfer":
from flashinfer import top_k_page_table_transform
return top_k_page_table_transform(scores, page_table, seq_lens, k)
elif provider == "torch":
idx = scores.topk(k, dim=-1).indices # (batch, k) int64
return torch.gather(page_table, 1, idx)
else:
raise ValueError(f"unknown provider {provider}")
if provider in ("flashinfer", "torch"):
page_table, seq_lens = _make_p1_table(batch_size, seq_len)
if provider == "jit_v2":
metadata = plan_topk_v2(seq_lens)
return fn, (scores, seq_lens, page_table)
@marker.parametrize("k", [512, 1024, 2048], [512])
@marker.parametrize("seq_len", [2**x for x in range(10, 19)], [4096, 65536])
@marker.parametrize("batch_size", [2**x for x in range(13)], [1, 128, 1024])
@marker.benchmark("provider", ["jit_v1", "jit_v2", "flashinfer", "torch"])
def benchmark(seq_len: int, batch_size: int, k: int, provider: str):
if k > seq_len:
marker.skip("k cannot be larger than seq_len")
if k == 2048 and provider == "jit_v1":
marker.skip("jit_v1 does not support k=2048")
fn, input_args = _build_fn(provider, batch_size, seq_len, k)
return marker.do_bench(fn, input_args=input_args, memory_args=input_args[:2])
if __name__ == "__main__":
benchmark.run()
@@ -0,0 +1,280 @@
from __future__ import annotations
import atexit
import contextlib
import logging
import os
from typing import Optional
import torch
import torch.distributed as dist
import sglang.srt.distributed.parallel_state as ps
from sglang.kernels.jit.benchmark import marker
from sglang.kernels.jit.benchmark.utils import get_benchmark_range, multigpu_bench_main
from sglang.kernels.jit.utils import cache_once, is_arch_support_pdl
from sglang.kernels.ops.communication.mp import register_comm_cleanup
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
register_cuda_ci(
est_time=120,
stage="base-b-kernel-benchmark",
runner_config="1-gpu-large",
disabled="requires multi-GPU, self-skips in CI",
)
register_amd_ci(est_time=120, stage="jit-kernel-benchmark", runner_config="amd")
# ---------------------------------------------------------------------------
# Sweep parameters
# ---------------------------------------------------------------------------
DTYPE = torch.bfloat16
DTYPE_ITEMSIZE = DTYPE.itemsize
MESSAGE_SIZES_KB = [2**x for x in range(2, 17)]
MESSAGE_SIZES_KB += [192, 384, 640, 768, 896, 1536, 3072]
MESSAGE_SIZES_KB.sort()
WORLD_SIZES = list(range(2, 9))
MAX_BYTES = max(MESSAGE_SIZES_KB) * 1024
# trtllm allreduce_fusion only supports these world sizes.
FI_SUPPORTED_WORLD_SIZES = (2, 4, 8)
# AOT custom_all_reduce (v1) only supports these world sizes.
AOT_SUPPORTED_WORLD_SIZES = (2, 4, 6, 8)
# jit-eager times the naive-loop dispatch (eager heuristics); jit-graph
# captures the calls in a CUDA graph (graph heuristics + pointer table).
PROVIDERS = ["nccl", "aot", "jit-eager", "jit-graph", "fi"]
WORLD_SIZES = get_benchmark_range(WORLD_SIZES, [2, 4, 8])
# ---------------------------------------------------------------------------
# Per-rank distributed init (run once per torchrun worker)
# ---------------------------------------------------------------------------
@cache_once
def _init_cpu_group() -> dist.ProcessGroup:
local_rank = int(os.environ["LOCAL_RANK"])
world_size = int(os.environ["WORLD_SIZE"])
torch.cuda.set_device(local_rank)
dist.init_process_group(backend="gloo")
ps._WORLD = coord = ps.init_world_group(
ranks=list(range(world_size)),
local_rank=local_rank,
backend="nccl",
)
atexit.register(dist.destroy_process_group)
torch.cuda.set_stream(torch.cuda.Stream())
return coord.cpu_group
@cache_once
def _init_nccl_group() -> dist.ProcessGroup:
_init_cpu_group()
local_rank = int(os.environ["LOCAL_RANK"])
device_group = torch.distributed.new_group(
backend="nccl",
device_id=torch.device(f"cuda:{local_rank}"),
)
assert isinstance(device_group, dist.ProcessGroup)
return device_group
# ---------------------------------------------------------------------------
# Backend wrappers - each exposes:
# .all_reduce(tensor) -> Tensor
# .graph_context() -> context manager wrapping cuda-graph capture
# (nullcontext when capture is not required)
# ---------------------------------------------------------------------------
class NCCLAllReduceBackend:
def __init__(self) -> None:
self.group = _init_nccl_group()
def graph_context(self):
return contextlib.nullcontext()
def all_reduce(self, tensor: torch.Tensor) -> torch.Tensor:
dist.all_reduce(tensor, group=self.group)
return tensor
class JITAllReduceBackend:
def __init__(self) -> None:
from sglang.srt.distributed.device_communicators.custom_all_reduce_v2 import (
CustomAllReduceV2,
)
device = torch.device(f"cuda:{int(os.environ['LOCAL_RANK'])}")
# tuned workspace sizes, capped at the sweep maximum
self.comm = CustomAllReduceV2(_init_cpu_group(), device, max_size=MAX_BYTES)
if self.comm.disabled:
raise RuntimeError("JIT CustomAllReduceV2 is disabled on this system")
# keep the whole sweep on the custom-AR path: the tuned config would
# otherwise send the largest sizes back to NCCL
self.comm.uncap_pull_thresholds()
register_comm_cleanup(self.comm)
def graph_context(self):
return self.comm.capture()
def all_reduce(self, tensor: torch.Tensor) -> Optional[torch.Tensor]:
assert self.comm.should_custom_ar(tensor), str(tensor.shape)
return self.comm.custom_all_reduce(tensor)
class AOTAllReduceBackend:
def __init__(self) -> None:
from sglang.srt.distributed.device_communicators.custom_all_reduce import (
CustomAllreduce,
)
device = torch.device(f"cuda:{int(os.environ['LOCAL_RANK'])}")
self.comm = CustomAllreduce(_init_cpu_group(), device, max_size=MAX_BYTES)
if self.comm.disabled:
raise RuntimeError("AOT CustomAllreduce is disabled on this system")
register_comm_cleanup(self.comm)
def graph_context(self):
return self.comm.capture()
def all_reduce(self, tensor: torch.Tensor) -> Optional[torch.Tensor]:
assert self.comm.should_custom_ar(tensor), str(tensor.shape)
return self.comm.custom_all_reduce(tensor)
class FlashInferAllReduceBackend:
def __init__(self) -> None:
import flashinfer.comm as comm
group = _init_cpu_group()
rank = dist.get_rank(group=group)
world_size = dist.get_world_size(group=group)
# Use the smallest message size as the inner hidden dim, so any
# message in the sweep is an integer multiple of it.
hidden_dim = 1024 * min(MESSAGE_SIZES_KB) // DTYPE_ITEMSIZE
num_tokens = MAX_BYTES // (hidden_dim * DTYPE_ITEMSIZE)
self._comm = comm
self._hidden_dim = hidden_dim
self._workspace = comm.create_allreduce_fusion_workspace(
backend="trtllm",
world_size=world_size,
rank=rank,
max_token_num=num_tokens,
hidden_dim=hidden_dim,
dtype=DTYPE,
)
def graph_context(self):
return contextlib.nullcontext()
def all_reduce(self, tensor: torch.Tensor) -> torch.Tensor:
return self._comm.allreduce_fusion(
input=tensor.view(-1, self._hidden_dim),
workspace=self._workspace,
pattern=self._comm.AllReduceFusionPattern.kAllReduce,
launch_with_pdl=is_arch_support_pdl(),
fp32_acc=True,
)
@cache_once
def _init_nccl_backend() -> NCCLAllReduceBackend:
return NCCLAllReduceBackend()
@cache_once
def _init_jit_backend() -> JITAllReduceBackend:
return JITAllReduceBackend()
@cache_once
def _init_aot_backend() -> AOTAllReduceBackend:
return AOTAllReduceBackend()
@cache_once
def _init_fi_backend() -> FlashInferAllReduceBackend:
return FlashInferAllReduceBackend()
BACKEND_FACTORY = {
"nccl": _init_nccl_backend,
"jit-eager": _init_jit_backend,
"jit-graph": _init_jit_backend,
"aot": _init_aot_backend,
"fi": _init_fi_backend,
}
@cache_once
def _init_all_backends() -> None:
"""Pre-build every supported backend before any timed iteration so JIT
compilation / IPC setup don't bleed into the first measured size.
"""
local_rank = int(os.environ["LOCAL_RANK"])
if local_rank == 0: # NOTE: log some verbose info on initialization
logging.basicConfig(level=logging.INFO)
world_size = dist.get_world_size(_init_cpu_group())
factories = dict(BACKEND_FACTORY)
if world_size not in AOT_SUPPORTED_WORLD_SIZES:
factories.pop("aot")
if world_size not in FI_SUPPORTED_WORLD_SIZES:
factories.pop("fi")
for fn in factories.values():
fn()
# reset level to warning
logging.getLogger().setLevel(logging.WARNING)
# ---------------------------------------------------------------------------
# Benchmark
# ---------------------------------------------------------------------------
@marker.parametrize("message_KB", MESSAGE_SIZES_KB)
@marker.benchmark("provider", PROVIDERS)
def benchmark(message_KB: int, provider: str):
cpu_group = _init_cpu_group()
gpu_group = _init_nccl_group()
world_size = dist.get_world_size(cpu_group)
if provider == "fi" and world_size not in FI_SUPPORTED_WORLD_SIZES:
marker.skip(
f"flashinfer trtllm allreduce_fusion needs world_size in "
f"{FI_SUPPORTED_WORLD_SIZES}"
)
if provider == "aot" and world_size not in AOT_SUPPORTED_WORLD_SIZES:
marker.skip(
f"AOT custom_all_reduce needs world_size in " f"{AOT_SUPPORTED_WORLD_SIZES}"
)
_init_all_backends()
backend = BACKEND_FACTORY[provider]()
message_bytes = message_KB * 1024
numel = message_bytes // DTYPE_ITEMSIZE
device_id = int(os.environ["LOCAL_RANK"])
device = torch.device(f"cuda:{device_id}")
x = torch.randn(numel, dtype=DTYPE, device=device)
ctx_fn = backend.graph_context if not provider.endswith("eager") else None
# Bandwidth-equivalent bytes moved by a ring all-reduce per rank.
effective_bytes = int(x.nbytes * 2 * (world_size - 1) / world_size)
return marker.do_bench(
backend.all_reduce,
input_args=(x,),
graph_context_fn=ctx_fn,
sync_multigpu_fn=lambda: dist.barrier(gpu_group),
# all-reduce is in-place w.r.t. its argument; explicit footprint
# captures the cross-GPU traffic instead.
memory_args=None,
memory_output=None,
extra_memory_footprint=effective_bytes,
)
if __name__ == "__main__":
multigpu_bench_main(
name=__name__,
file=__file__,
num_gpus=WORLD_SIZES,
main_fn=benchmark.run,
)
@@ -0,0 +1,171 @@
"""Benchmark the symmetric-memory multimem all-gather vs NCCL.
Providers:
- ``nccl`` : ``all_gather_into_tensor`` + concat-along-hidden reshape
(what ``tensor_model_parallel_all_gather(dim=-1)`` does)
- ``mm_safe`` : multimem kernel, ``safe=True`` (clones the buffer view)
- ``mm`` : multimem kernel, ``safe=False`` (fc gather config)
- ``mm_skipsync`` : multimem kernel, ``safe=False, skip_entry_sync=True``
(logits gather config)
Usage::
# Benchmark on the default world sizes (2, 4, 8 GPUs):
python test/registered/jit/benchmark/bench_symm_mem_all_gather.py
# Pick a specific world size (or comma-separated list):
python test/registered/jit/benchmark/bench_symm_mem_all_gather.py --num-gpu 8
"""
from __future__ import annotations
import atexit
import logging
import os
import torch
import torch.distributed as dist
import sglang.srt.distributed.parallel_state as ps
from sglang.kernels.jit.benchmark import marker
from sglang.kernels.jit.benchmark.utils import get_benchmark_range, multigpu_bench_main
from sglang.kernels.jit.utils import cache_once
from sglang.srt.distributed.device_communicators.triton_symm_mem_ag import (
all_gather_inner,
create_state,
)
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
register_cuda_ci(
est_time=120,
stage="base-b-kernel-benchmark",
runner_config="1-gpu-large",
disabled="requires multi-GPU, self-skips in CI",
)
register_amd_ci(est_time=120, stage="jit-kernel-benchmark", runner_config="amd")
# ---------------------------------------------------------------------------
# Sweep parameters
# ---------------------------------------------------------------------------
DTYPE = torch.bfloat16
PROVIDERS = ["nccl", "mm_safe", "mm", "mm_skipsync"]
# Full gathered hidden width H (per-rank shard is H / world_size).
HIDDENS = [7168, 16384, 163840]
NUM_TOKENS = [1, 8, 16, 32, 64, 128]
WORLD_SIZES = list(range(2, 9))
HIDDENS = get_benchmark_range(HIDDENS, [7168, 163840])
NUM_TOKENS = get_benchmark_range(NUM_TOKENS, [16, 64])
WORLD_SIZES = get_benchmark_range(WORLD_SIZES, [2, 4, 8])
MAX_HIDDEN = max(HIDDENS)
MAX_TOKENS = max(NUM_TOKENS)
# ---------------------------------------------------------------------------
# Per-rank distributed init (run once per torchrun worker)
# ---------------------------------------------------------------------------
@cache_once
def _init_cpu_group() -> dist.ProcessGroup:
local_rank = int(os.environ["LOCAL_RANK"])
world_size = int(os.environ["WORLD_SIZE"])
torch.cuda.set_device(local_rank)
dist.init_process_group(backend="gloo")
ps._WORLD = ps.init_world_group(
ranks=list(range(world_size)),
local_rank=local_rank,
backend="nccl",
)
atexit.register(dist.destroy_process_group)
logging.disable(logging.INFO)
torch.cuda.set_stream(torch.cuda.Stream())
return ps._WORLD.cpu_group
@cache_once
def _init_nccl_group() -> dist.ProcessGroup:
_init_cpu_group()
coord = ps._WORLD
assert coord is not None and coord.device_group is not None
return coord.device_group
@cache_once
def _init_state():
_init_cpu_group()
coord = ps._WORLD
return create_state(
group=coord.device_group,
rank_in_group=coord.rank_in_group,
max_tokens=MAX_TOKENS,
hidden_size=MAX_HIDDEN,
)
# ---------------------------------------------------------------------------
# Benchmark
# ---------------------------------------------------------------------------
@marker.parametrize("hidden", HIDDENS)
@marker.parametrize("num_tokens", NUM_TOKENS)
@marker.benchmark("provider", PROVIDERS)
def benchmark(num_tokens: int, hidden: int, provider: str):
gpu_group = _init_nccl_group()
state = _init_state()
world_size = state.world_size
local_hidden = hidden // world_size
if hidden % world_size != 0 or local_hidden % 8 != 0:
marker.skip(f"hidden={hidden} incompatible with world_size={world_size}")
if provider != "nccl" and state.symm_mem_hdl.multicast_ptr == 0:
marker.skip(f"multimem multicast unavailable for world_size={world_size}")
device = torch.device(f"cuda:{int(os.environ['LOCAL_RANK'])}")
x = torch.randn(num_tokens, local_hidden, dtype=DTYPE, device=device)
if provider == "nccl":
out_buf = torch.empty(
world_size * num_tokens, local_hidden, dtype=DTYPE, device=device
)
def fn(inp: torch.Tensor) -> torch.Tensor:
dist.all_gather_into_tensor(out_buf, inp, group=gpu_group)
return (
out_buf.reshape(world_size, num_tokens, local_hidden)
.movedim(0, 1)
.reshape(num_tokens, hidden)
)
else:
safe = provider == "mm_safe"
skip_entry_sync = provider == "mm_skipsync"
def fn(inp: torch.Tensor) -> torch.Tensor:
return all_gather_inner(
state,
inp,
tp_hidden_dim=hidden,
skip_entry_sync=skip_entry_sync,
safe=safe,
)
return marker.do_bench(
fn,
input_args=(x,),
graph_clone_args=(0,),
sync_multigpu_fn=lambda: dist.barrier(gpu_group),
# Footprint = the gathered output every rank ends up with.
memory_args=None,
memory_output=None,
extra_memory_footprint=num_tokens * hidden * x.element_size(),
)
if __name__ == "__main__":
multigpu_bench_main(
name=__name__,
file=__file__,
num_gpus=WORLD_SIZES,
main_fn=benchmark.run,
)
@@ -0,0 +1,233 @@
"""Benchmark fused TP QKNorm (push-mode custom-AR + RMSNorm) vs the serial
baseline (RMS sum-sq -> pull-mode all-reduce -> RMS apply).
Usage::
# Benchmark on every supported world size (2..8 GPUs):
python benchmark/bench_tp_qknorm.py
# Specific world sizes:
python benchmark/bench_tp_qknorm.py --num-gpu 4
python benchmark/bench_tp_qknorm.py --num-gpu 2,4,8
"""
from __future__ import annotations
import atexit
import logging
import multiprocessing
import os
from multiprocessing.context import SpawnProcess
from typing import List
import torch
import torch.distributed as dist
import sglang.srt.distributed.parallel_state as ps
from sglang.kernels.jit.benchmark import marker
from sglang.kernels.jit.benchmark.utils import multigpu_bench_main
from sglang.kernels.jit.utils import cache_once, get_ci_test_range
from sglang.kernels.ops.communication.all_reduce import (
fused_parallel_qknorm,
get_all_reduce_module,
get_fused_parallel_qknorm_max_occupancy,
get_fused_parallel_qknorm_module,
)
from sglang.kernels.ops.communication.mp import register_comm_cleanup
from sglang.srt.distributed.device_communicators.custom_all_reduce_v2 import (
CustomAllReduceV2,
)
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
register_cuda_ci(
est_time=120,
stage="base-b-kernel-benchmark",
runner_config="1-gpu-large",
disabled="requires multi-GPU, self-skips in CI",
)
register_amd_ci(est_time=120, stage="jit-kernel-benchmark", runner_config="amd")
# ---------------------------------------------------------------------------
# Sweep parameters
# ---------------------------------------------------------------------------
DTYPE = torch.bfloat16
EPS = 1e-6
Q_K_DIMS = [(6144, 1024)]
BATCH_SIZES = get_ci_test_range([2**i for i in range(15)], [1, 64, 1024])
MAX_PUSH_SIZE = 8 * max(BATCH_SIZES)
PROVIDERS = ["fused", "baseline"]
# ---------------------------------------------------------------------------
# Parallel JIT precompile (outer process, before any torchrun child starts)
# ---------------------------------------------------------------------------
def _compile_one(world_size: int) -> None:
"""Compile every kernel this bench touches for a single world_size.
Top-level so it survives ``spawn`` pickling. Compiled artifacts are
cached on disk by ``tvm_ffi``; torchrun children will reuse them.
"""
# baseline path: sum-sq -> all-reduce -> apply (also covers push mode)
get_all_reduce_module(DTYPE, world_size)
# fused path: fused QKNorm kernel (one per (dtype, world_size, q_dim, k_dim))
for q_dim, k_dim in Q_K_DIMS:
get_fused_parallel_qknorm_module(DTYPE, world_size, q_dim, k_dim)
def _precompile_kernels(num_gpus: List[int]) -> None:
ctx = multiprocessing.get_context("spawn")
procs: list[tuple[int, SpawnProcess]] = []
for world_size in num_gpus:
p = ctx.Process(target=_compile_one, args=(world_size,))
p.start()
procs.append((world_size, p))
for world_size, p in procs:
p.join()
if p.exitcode != 0:
raise RuntimeError(
f"TP QKNorm precompile failed for {world_size=} " f"(exit {p.exitcode})"
)
# ---------------------------------------------------------------------------
# Per-rank distributed init (run once per torchrun worker)
# ---------------------------------------------------------------------------
@cache_once
def _init_cpu_group() -> dist.ProcessGroup:
local_rank = int(os.environ["LOCAL_RANK"])
world_size = int(os.environ["WORLD_SIZE"])
torch.cuda.set_device(local_rank)
dist.init_process_group(backend="gloo")
ps._WORLD = coord = ps.init_world_group(
ranks=list(range(world_size)),
local_rank=local_rank,
backend="nccl",
)
atexit.register(dist.destroy_process_group)
logging.disable(logging.INFO)
torch.cuda.set_stream(torch.cuda.Stream())
return coord.cpu_group
@cache_once
def _init_gpu_group() -> dist.ProcessGroup:
_init_cpu_group()
device = torch.device(f"cuda:{int(os.environ['LOCAL_RANK'])}")
gpu_group = dist.new_group(backend="nccl", device_id=device)
assert isinstance(gpu_group, dist.ProcessGroup)
atexit.register(lambda: dist.destroy_process_group(gpu_group))
return gpu_group
@cache_once
def _init_fused_comm() -> CustomAllReduceV2:
"""Push-mode workspace sized for the fused-QKNorm bench."""
cpu_group = _init_cpu_group()
world_size = dist.get_world_size(cpu_group)
device = torch.device(f"cuda:{int(os.environ['LOCAL_RANK'])}")
q_dim, k_dim = Q_K_DIMS[0]
max_occupancy = get_fused_parallel_qknorm_max_occupancy(
DTYPE, world_size, q_dim, k_dim
)
if dist.get_rank(cpu_group) == 0:
print(f"Max occupancy for fused_parallel_qknorm: {max_occupancy} blocks/SM")
props = torch.cuda.get_device_properties(device)
comm = CustomAllReduceV2(
cpu_group,
device,
max_pull_size=0,
max_push_size=MAX_PUSH_SIZE,
max_push_blocks=props.multi_processor_count * max_occupancy,
)
if comm.disabled:
raise RuntimeError("JIT CustomAllReduceV2 is disabled on this system")
register_comm_cleanup(comm)
return comm
@cache_once
def _init_baseline_comm() -> CustomAllReduceV2:
"""Default (pull-mode) workspace for the serial baseline."""
cpu_group = _init_cpu_group()
device = torch.device(f"cuda:{int(os.environ['LOCAL_RANK'])}")
comm = CustomAllReduceV2(cpu_group, device)
if comm.disabled:
raise RuntimeError("JIT CustomAllReduceV2 is disabled on this system")
register_comm_cleanup(comm)
return comm
# ---------------------------------------------------------------------------
# Implementations
# ---------------------------------------------------------------------------
def _rmsnorm_baseline(
comm: CustomAllReduceV2,
q: torch.Tensor,
k: torch.Tensor,
q_weight: torch.Tensor,
k_weight: torch.Tensor,
world_size: int,
) -> None:
from sglang.srt.models.minimax_m2 import rms_apply_serial, rms_sumsq_serial
sum_sq = rms_sumsq_serial(q, k)
sum_sq = comm.custom_all_reduce(sum_sq)
rms_apply_serial(q, k, q_weight, k_weight, sum_sq, world_size, EPS)
# ---------------------------------------------------------------------------
# Benchmark
# ---------------------------------------------------------------------------
@marker.parametrize("q_dim,k_dim", Q_K_DIMS)
@marker.parametrize("batch_size", BATCH_SIZES)
@marker.benchmark("provider", PROVIDERS)
def benchmark(q_dim: int, k_dim: int, batch_size: int, provider: str):
cpu_group = _init_cpu_group()
gpu_group = _init_gpu_group()
world_size = dist.get_world_size(cpu_group)
device = torch.device(f"cuda:{int(os.environ['LOCAL_RANK'])}")
local_q_dim = q_dim // world_size
local_k_dim = k_dim // world_size
q = torch.randn(batch_size, local_q_dim, device=device, dtype=DTYPE)
k = torch.randn(batch_size, local_k_dim, device=device, dtype=DTYPE)
q_weight = torch.randn(local_q_dim, device=device, dtype=DTYPE)
k_weight = torch.randn(local_k_dim, device=device, dtype=DTYPE)
if provider == "fused":
comm = _init_fused_comm()
def fn(q, k, q_weight, k_weight):
fused_parallel_qknorm(comm.obj, q, k, q_weight, k_weight, EPS)
else:
comm = _init_baseline_comm()
def fn(q, k, q_weight, k_weight):
_rmsnorm_baseline(comm, q, k, q_weight, k_weight, world_size)
return marker.do_bench(
fn,
input_args=(q, k, q_weight, k_weight),
sync_multigpu_fn=lambda: dist.barrier(gpu_group),
memory_output=(q, k), # NOTE: In-place updates on q, k;
)
if __name__ == "__main__":
multigpu_bench_main(
name=__name__,
file=__file__,
num_gpus=[2, 4, 8], # NOTE: don't support other world size now
main_fn=benchmark.run,
pre_launch_fn=_precompile_kernels,
)
@@ -0,0 +1,95 @@
from dataclasses import dataclass
import torch
from sglang.kernels.jit.benchmark import marker
from sglang.kernels.ops.diffusion.causal_conv3d_cat_pad import (
fused_causal_conv3d_cat_pad_cuda,
)
from sglang.kernels.ops.diffusion.triton.causal_conv3d_pad import (
fused_causal_conv3d_cat_pad as fused_causal_conv3d_cat_pad_triton,
)
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
register_cuda_ci(
est_time=20,
stage="base-b-kernel-benchmark",
runner_config="1-gpu-large",
disabled="standalone benchmark",
)
register_amd_ci(est_time=20, stage="jit-kernel-benchmark", runner_config="amd")
DEVICE = "cuda"
DTYPE = torch.bfloat16
@dataclass(frozen=True)
class Case:
name: str
channels: int
t_size: int
h_size: int
w_size: int
cache_t: int
trace_count: int
CASES = [
Case("c1024_t1_h30_w52_cache1", 1024, 1, 30, 52, 1, 8),
Case("c1024_t1_h30_w52_cache2", 1024, 1, 30, 52, 2, 8),
Case("c1024_t2_h60_w104_cache1", 1024, 2, 60, 104, 1, 5),
Case("c1024_t2_h60_w104_cache2", 1024, 2, 60, 104, 2, 5),
Case("c512_t4_h120_w208_cache1", 512, 4, 120, 208, 1, 5),
Case("c512_t4_h120_w208_cache2", 512, 4, 120, 208, 2, 5),
Case("c256_t4_h240_w416_cache1", 256, 4, 240, 416, 1, 6),
Case("c256_t4_h240_w416_cache2", 256, 4, 240, 416, 2, 6),
]
CASE_BY_NAME = {case.name: case for case in CASES}
CASE_NAMES = [case.name for case in CASES]
def make_inputs(case: Case) -> tuple[torch.Tensor, torch.Tensor, tuple[int, ...]]:
generator = torch.Generator(device=DEVICE)
generator.manual_seed(case.channels * 1009 + case.t_size * 251 + case.cache_t)
x = torch.randn(
(1, case.channels, case.t_size, case.h_size, case.w_size),
device=DEVICE,
dtype=DTYPE,
generator=generator,
)
cache_x = torch.randn(
(1, case.channels, case.cache_t, case.h_size, case.w_size),
device=DEVICE,
dtype=DTYPE,
generator=generator,
)
padding = (1, 1, 1, 1, case.cache_t, 0)
return x, cache_x, padding
@marker.parametrize("case_name", CASE_NAMES, ci_vals=CASE_NAMES[:2])
@marker.benchmark("provider", ["triton", "cuda"])
def benchmark(case_name: str, provider: str) -> marker.BenchResult:
case = CASE_BY_NAME[case_name]
x, cache_x, padding = make_inputs(case)
fn = (
fused_causal_conv3d_cat_pad_triton
if provider == "triton"
else fused_causal_conv3d_cat_pad_cuda
)
actual = fused_causal_conv3d_cat_pad_cuda(x, cache_x, padding)
expected = fused_causal_conv3d_cat_pad_triton(x, cache_x, padding)
torch.testing.assert_close(actual, expected, atol=0, rtol=0)
return marker.do_bench(
fn,
input_args=(x, cache_x, padding),
use_cuda_graph=False,
replay_iters=200,
graph_clone_args=(0, 1),
memory_args=(x, cache_x),
memory_output="out",
)
if __name__ == "__main__":
benchmark.run()
@@ -0,0 +1,358 @@
import argparse
import csv
import json
import os
import re
import statistics
from pathlib import Path
from typing import Any, Callable
import flashinfer
import torch
from sglang.kernels.jit.benchmark.utils import DEFAULT_DTYPE
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
register_cuda_ci(
est_time=120,
stage="base-b-kernel-benchmark",
runner_config="1-gpu-large",
disabled="standalone diffusion NVFP4 benchmark",
)
SCRIPT_DIR = Path(__file__).resolve().parent
REPO_ROOT = (
Path(os.environ["SGLANG_NVFP4_REPO_ROOT"])
if os.environ.get("SGLANG_NVFP4_REPO_ROOT")
# Anchor on the installed jit_kernel package (python/sglang/kernels/jit) so
# this stays correct regardless of where the benchmark file lives.
else KERNEL_PATH.parents[2]
)
DEFAULT_OUTPUT_DIR = REPO_ROOT / "outputs" / "nvfp4_benchmarks"
DEFAULT_SHAPE_LIBRARY = SCRIPT_DIR / "diffusion_nvfp4_shapes.json"
DTYPE = DEFAULT_DTYPE
WARMUP = 8
ITERS = 20
FLOAT4_E2M1_MAX = 6.0
FLOAT8_E4M3_MAX = torch.finfo(torch.float8_e4m3fn).max
METHODS = ("flashinfer_auto", "flashinfer_cudnn")
def benchmark_provider(
fn: Callable[[], torch.Tensor],
warmup: int = WARMUP,
iters: int = ITERS,
) -> tuple[float, float, float]:
for _ in range(warmup):
y = fn()
del y
torch.cuda.synchronize()
times_ms: list[float] = []
for _ in range(iters):
start = torch.cuda.Event(enable_timing=True)
end = torch.cuda.Event(enable_timing=True)
start.record()
y = fn()
end.record()
end.synchronize()
times_ms.append(start.elapsed_time(end))
del y
return statistics.median(times_ms), max(times_ms), min(times_ms)
def make_global_scale(x: torch.Tensor) -> torch.Tensor:
max_abs = torch.amax(x.abs()).clamp_min_(1e-6)
return (FLOAT8_E4M3_MAX * FLOAT4_E2M1_MAX / max_abs).to(torch.float32)
def build_quantized_inputs(
m: int,
n: int,
k: int,
device: torch.device,
seed: int,
) -> dict[str, Any]:
assert k % 16 == 0, f"NVFP4 requires k % 16 == 0, got k={k}"
gen = torch.Generator(device=device)
gen.manual_seed(seed)
x = torch.randn((m, k), device=device, dtype=DTYPE, generator=gen)
w = torch.randn((n, k), device=device, dtype=DTYPE, generator=gen)
x_global_scale = make_global_scale(x)
w_global_scale = make_global_scale(w)
alpha = (1.0 / (x_global_scale * w_global_scale)).to(torch.float32)
x_fp4, x_sf = flashinfer.fp4_quantize(x, x_global_scale)
w_fp4, w_sf = flashinfer.fp4_quantize(w, w_global_scale)
if x_sf.dtype == torch.uint8:
x_sf = x_sf.view(torch.float8_e4m3fn)
if w_sf.dtype == torch.uint8:
w_sf = w_sf.view(torch.float8_e4m3fn)
return {
"x_fp4": x_fp4,
"w_fp4": w_fp4,
"x_sf": x_sf,
"w_sf": w_sf,
"alpha": alpha,
}
def make_shape_id(
model: str, shape_kind: str, prefix: str, m: int, n: int, k: int
) -> str:
prefix_slug = re.sub(r"[^a-zA-Z0-9]+", "_", prefix).strip("_")
return f"{model}_{shape_kind}_{prefix_slug}_{m}x{n}x{k}"
def load_shape_cases(shape_library: Path) -> list[dict[str, Any]]:
payload = json.loads(shape_library.read_text(encoding="utf-8"))
if not isinstance(payload, dict) or not payload:
raise RuntimeError(
f"Expected a non-empty model->shape list mapping in {shape_library}."
)
rows: list[dict[str, Any]] = []
for model, shapes in payload.items():
if not isinstance(shapes, list):
raise RuntimeError(
f"Expected {model} to map to a list of shapes in {shape_library}."
)
for shape in shapes:
m, n, k = (int(x) for x in shape["shape"])
count = int(shape["count"])
shape_kind = str(shape.get("kind", "actual_runtime_linear"))
prefix = str(shape.get("prefix", ""))
rows.append(
{
"shape_id": make_shape_id(model, shape_kind, prefix, m, n, k),
"source_model": model,
"shape_kind": shape_kind,
"runtime_prefix": prefix,
"m": m,
"n": n,
"k": k,
"count": count,
"approx_flops": 2 * m * n * k * count,
}
)
if not rows:
raise RuntimeError(f"No shapes found in {shape_library}.")
return rows
def split_csv_arg(text: str | None) -> set[str]:
if text is None or not text.strip():
return set()
return {item.strip() for item in text.split(",") if item.strip()}
def select_shape_cases(
rows: list[dict[str, Any]],
*,
models: set[str],
shape_kinds: set[str],
top_k: int,
rank_by: str,
) -> list[dict[str, Any]]:
filtered = [
row
for row in rows
if (not models or row["source_model"] in models)
and (not shape_kinds or row["shape_kind"] in shape_kinds)
]
key = "approx_flops" if rank_by == "flops" else "count"
return sorted(filtered, key=lambda row: int(row[key]), reverse=True)[:top_k]
def write_csv(rows: list[dict[str, Any]], output_path: Path) -> None:
with output_path.open("w", newline="", encoding="utf-8") as f:
writer = csv.DictWriter(
f,
fieldnames=[
"shape_id",
"source_model",
"shape_kind",
"runtime_prefix",
"m",
"n",
"k",
"count",
"approx_flops",
"method",
"median_ms",
"min_ms",
"max_ms",
"tflops",
],
)
writer.writeheader()
writer.writerows(rows)
def write_markdown(rows: list[dict[str, Any]], output_path: Path) -> None:
shape_rows = []
seen_shape_ids = set()
for row in rows:
if row["shape_id"] in seen_shape_ids:
continue
seen_shape_ids.add(row["shape_id"])
shape_rows.append(row)
lines: list[str] = []
lines.append("# Diffusion NVFP4 Scaled MM Benchmark")
lines.append("")
lines.append("## Shape Cases")
lines.append("")
lines.append("| Shape ID | Model | Shape Kind | Calls | Shape `(M,N,K)` | Prefix |")
lines.append("|---|---|---|---:|---|---|")
for row in shape_rows:
lines.append(
f"| {row['shape_id']} | {row['source_model']} | {row['shape_kind']} | {row['count']} | `({row['m']}, {row['n']}, {row['k']})` | `{row['runtime_prefix']}` |"
)
lines.append("")
for shape_row in shape_rows:
shape_id = shape_row["shape_id"]
scoped = [row for row in rows if row["shape_id"] == shape_id]
lines.append(f"## {shape_id}")
lines.append("")
lines.append("| Method | Median ms | TFLOPS |")
lines.append("|---|---:|---:|")
for row in sorted(scoped, key=lambda item: float(item["median_ms"])):
lines.append(
f"| {row['method']} | {float(row['median_ms']):.4f} | {float(row['tflops']):.1f} |"
)
lines.append("")
output_path.write_text("\n".join(lines) + "\n", encoding="utf-8")
def run_shape_suite(shape_cases: list[dict[str, Any]]) -> list[dict[str, Any]]:
device = torch.device("cuda")
rows: list[dict[str, Any]] = []
for idx, shape in enumerate(shape_cases):
m = int(shape["m"])
n = int(shape["n"])
k = int(shape["k"])
quantized = build_quantized_inputs(m, n, k, device, seed=idx)
metadata = {
"shape_id": str(shape["shape_id"]),
"source_model": str(shape["source_model"]),
"shape_kind": str(shape["shape_kind"]),
"runtime_prefix": str(shape["runtime_prefix"]),
"m": m,
"n": n,
"k": k,
"count": int(shape["count"]),
"approx_flops": int(shape["approx_flops"]),
}
providers: dict[str, Callable[[], torch.Tensor]] = {
"flashinfer_auto": lambda: flashinfer.mm_fp4(
quantized["x_fp4"],
quantized["w_fp4"].T,
quantized["x_sf"],
quantized["w_sf"].T,
quantized["alpha"],
DTYPE,
backend="auto",
),
"flashinfer_cudnn": lambda: flashinfer.mm_fp4(
quantized["x_fp4"],
quantized["w_fp4"].T,
quantized["x_sf"],
quantized["w_sf"].T,
quantized["alpha"],
DTYPE,
backend="cudnn",
),
}
for method in METHODS:
median_ms, max_ms, min_ms = benchmark_provider(providers[method])
rows.append(
{
**metadata,
"method": method,
"median_ms": median_ms,
"min_ms": min_ms,
"max_ms": max_ms,
"tflops": (2 * m * n * k) / (median_ms / 1e3) / 1e12,
}
)
return rows
def main() -> None:
parser = argparse.ArgumentParser(
description="Benchmark diffusion NVFP4 GEMM backends on the captured diffusion shape library."
)
parser.add_argument(
"--models",
help="Comma-separated source_model filter. Default: all models in the JSON shape library.",
)
parser.add_argument(
"--shape-kinds",
help="Comma-separated shape_kind filter. Default: benchmark every shape kind in the JSON shape library.",
)
parser.add_argument(
"--top-k",
type=int,
default=64,
help="Benchmark the top-k shapes after filtering and ranking.",
)
parser.add_argument(
"--rank-by",
choices=["flops", "count"],
default="flops",
help="How to rank shapes before selecting top-k.",
)
parser.add_argument(
"--output-dir",
default=str(DEFAULT_OUTPUT_DIR),
help="Directory for CSV/Markdown outputs.",
)
args = parser.parse_args()
if is_in_ci():
print("Skipping bench_diffusion_nvfp4_scaled_mm.py in CI")
return
if not torch.cuda.is_available():
raise RuntimeError("CUDA is required for NVFP4 scaled mm benchmarks.")
if not DEFAULT_SHAPE_LIBRARY.exists():
raise RuntimeError(
f"Shape library not found at {DEFAULT_SHAPE_LIBRARY}. "
"Commit or copy the generated diffusion_nvfp4_shapes.json first."
)
shape_cases = load_shape_cases(DEFAULT_SHAPE_LIBRARY)
selected_shapes = select_shape_cases(
shape_cases,
models=split_csv_arg(args.models),
shape_kinds=split_csv_arg(args.shape_kinds),
top_k=args.top_k,
rank_by=args.rank_by,
)
if not selected_shapes:
raise RuntimeError("No shapes matched the requested filters.")
rows = run_shape_suite(selected_shapes)
output_dir = Path(args.output_dir)
output_dir.mkdir(parents=True, exist_ok=True)
csv_path = output_dir / "diffusion_nvfp4_scaled_mm.csv"
md_path = output_dir / "diffusion_nvfp4_scaled_mm_summary.md"
write_csv(rows, csv_path)
write_markdown(rows, md_path)
print(f"Wrote {csv_path}")
print(f"Wrote {md_path}")
if __name__ == "__main__":
main()
@@ -0,0 +1,139 @@
# Benchmarks SGLang fused layernorm/rmsnorm scale shift kernels
# 1. fused_norm_scale_shift
# 2. fused_scale_residual_norm_scale_shift
import itertools
from typing import Tuple
import torch
import triton
import triton.testing
from sglang.kernels.jit.benchmark.utils import run_benchmark_no_cudagraph
from sglang.multimodal_gen.runtime.layers.layernorm import (
LayerNormScaleShift,
RMSNormScaleShift,
ScaleResidualLayerNormScaleShift,
ScaleResidualRMSNormScaleShift,
)
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.utils import is_in_ci
register_cuda_ci(
est_time=17,
stage="base-b-kernel-benchmark",
runner_config="1-gpu-large",
disabled="Temporarily skipped to unblock flashinfer upgrade. Ref: https://github.com/sgl-project/sglang/actions/runs/23735552939/job/69139238979?pr=21422",
)
if is_in_ci():
B_RANGE, S_RANGE, D_RANGE = [1], [128], [1024]
else:
B_RANGE, S_RANGE, D_RANGE = [1], [128, 1024, 4096], [1024, 3072, 4096]
NORM_TYPE_RANGE = ["layer", "rms"]
AFFINE_RANGE = [True, False]
DTYPE = torch.bfloat16
DEVICE = "cuda"
EPS = 1e-5
LINE_VALS = ["native", "cuda"]
LINE_NAMES = ["SGLang Native", "SGLang Fused"]
STYLES = [("red", "-"), ("blue", "--")]
config = list(
itertools.product(B_RANGE, S_RANGE, D_RANGE, NORM_TYPE_RANGE, AFFINE_RANGE)
)
def preprocess_layer(layer, affine: bool, D: int, DTYPE: torch.dtype):
if affine:
weight = torch.randn(D, dtype=DTYPE, device=DEVICE)
bias = torch.randn(D, dtype=DTYPE, device=DEVICE)
with torch.no_grad():
layer.norm.weight.copy_(weight)
if hasattr(layer.norm, "bias"):
layer.norm.bias.copy_(bias)
layer.requires_grad_(False)
return layer.to(DEVICE)
# ============================================================================
# Benchmark 1: fused_norm_scale_shift
# ============================================================================
@triton.testing.perf_report(
triton.testing.Benchmark(
x_names=["B", "S", "D", "norm_type", "affine"],
x_vals=config,
line_arg="provider",
line_vals=LINE_VALS,
line_names=LINE_NAMES,
styles=STYLES,
ylabel="us",
plot_name="fused_norm_scale_shift",
args={},
)
)
def bench_fused_norm_scale_shift(
B: int, S: int, D: int, norm_type, affine: bool, provider: str
) -> Tuple[float, float, float]:
x = torch.randn(B, S, D, dtype=DTYPE, device=DEVICE)
scale = torch.randn(B, S, D, dtype=DTYPE, device=DEVICE)
shift = torch.randn(B, S, D, dtype=DTYPE, device=DEVICE)
if norm_type == "layer":
layer = LayerNormScaleShift(D, EPS, affine, dtype=DTYPE)
else:
layer = RMSNormScaleShift(D, EPS, affine, dtype=DTYPE)
layer = preprocess_layer(layer, affine, D, DTYPE)
if provider == "native":
fn = lambda: layer.forward_native(x, shift, scale)
else:
fn = lambda: layer.forward_cuda(x, shift, scale)
return run_benchmark_no_cudagraph(fn)
# ============================================================================
# Benchmark 2: fused_scale_residual_norm_scale_shift
# ============================================================================
@triton.testing.perf_report(
triton.testing.Benchmark(
x_names=["B", "S", "D", "norm_type", "affine"],
x_vals=config,
line_arg="provider",
line_vals=LINE_VALS,
line_names=LINE_NAMES,
styles=STYLES,
ylabel="us",
plot_name="fused_scale_residual_norm_scale_shift",
args={},
)
)
def bench_fused_scale_residual_norm_scale_shift(
B: int, S: int, D: int, norm_type, affine: bool, provider: str
) -> Tuple[float, float, float]:
residual = torch.randn(B, S, D, dtype=DTYPE, device=DEVICE)
x = torch.randn(B, S, D, dtype=DTYPE, device=DEVICE)
scale = torch.randn(B, S, D, dtype=DTYPE, device=DEVICE)
shift = torch.randn(B, S, D, dtype=DTYPE, device=DEVICE)
gate = torch.randn(B, 1, D, dtype=DTYPE, device=DEVICE)
if norm_type == "layer":
layer = ScaleResidualLayerNormScaleShift(D, EPS, affine, dtype=DTYPE).to(DEVICE)
else:
layer = ScaleResidualRMSNormScaleShift(D, EPS, affine, dtype=DTYPE).to(DEVICE)
layer = preprocess_layer(layer, affine, D, DTYPE)
if provider == "native":
fn = lambda: layer.forward_native(residual, x, gate, shift, scale)
else:
fn = lambda: layer.forward_cuda(residual, x, gate, shift, scale)
return run_benchmark_no_cudagraph(fn)
if __name__ == "__main__":
print(f"\n{'='*80}")
print("Benchmark: fused_norm_scale_shift")
print(f"{'='*80}\n")
bench_fused_norm_scale_shift.run(print_data=True)
print(f"\n{'='*80}")
print("Benchmark: fused_scale_residual_norm_scale_shift")
print(f"{'='*80}\n")
bench_fused_scale_residual_norm_scale_shift.run(print_data=True)
@@ -0,0 +1,315 @@
import argparse
import csv
import statistics
import sys
from dataclasses import dataclass
from pathlib import Path
from typing import Callable
import torch
import torch.nn.functional as F
import triton.testing
from sglang.kernels.ops.diffusion.triton.group_norm_silu import triton_group_norm_silu
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
from sglang.utils import is_in_ci
register_cuda_ci(
est_time=45,
stage="base-b-kernel-benchmark",
runner_config="1-gpu-large",
disabled="standalone benchmark",
)
register_amd_ci(est_time=45, stage="jit-kernel-benchmark", runner_config="amd")
DEVICE = "cuda"
EPS = 1e-5
QUANTILES = [0.5, 0.2, 0.8]
@dataclass(frozen=True)
class Case:
name: str
shape: tuple[int, ...]
num_groups: int
CASES = [
Case("token_2d", (4, 128), 32),
Case("image_2d", (2, 64, 32, 32), 32),
Case("video_3d_small", (1, 64, 4, 16, 16), 32),
Case("threshold_3d", (1, 128, 1, 256, 256), 32),
Case("hunyuan_video_large", (1, 128, 20, 256, 256), 32),
# LTX-2 latent upsampler (`LatentUpsampler` + `ResBlock`) operates on
# `[B, mid_channels=512, F, H, W]` tensors with num_groups=32. The
# `small` and `pre_720p` cases stay in the default set; the larger
# `post_720p` case is opt-in via LARGE_CASES below.
Case("ltx2_upsampler_small", (1, 512, 8, 45, 80), 32),
Case("ltx2_upsampler_pre_720p", (1, 512, 16, 90, 160), 32),
]
# Cases too large to fit comfortably alongside the native-path intermediates
# on consumer GPUs (e.g. 24 GB L4). Opt in with `--cases large` (large only),
# `--cases all-large` (default + large), or by name.
#
# `ltx2_upsampler_post_720p` is ~471M bf16 elements (~940 MB tensor) and the
# eager `silu(group_norm(x))` reference materializes mean / variance /
# normalized / silu intermediates -- working set lands around 5 GB. On
# H100 / H200 this is fine and surfaces the asymptotic ~14x kernel speedup;
# on a 24 GB GPU it can OOM, so it's gated out of `--cases all`.
LARGE_CASES = [
Case("ltx2_upsampler_post_720p", (1, 512, 16, 180, 320), 32),
]
CASE_BY_NAME = {case.name: case for case in CASES + LARGE_CASES}
def dtype_from_name(name: str) -> torch.dtype:
mapping = {
"bf16": torch.bfloat16,
"bfloat16": torch.bfloat16,
"fp16": torch.float16,
"float16": torch.float16,
"fp32": torch.float32,
"float32": torch.float32,
}
return mapping[name]
def dtype_name(dtype: torch.dtype) -> str:
mapping = {
torch.bfloat16: "bf16",
torch.float16: "fp16",
torch.float32: "fp32",
}
return mapping[dtype]
def parse_dtypes(text: str) -> list[torch.dtype]:
return [dtype_from_name(item.strip()) for item in text.split(",") if item.strip()]
def parse_cases(text: str) -> list[Case]:
if text == "all":
return CASES
if text == "large":
return LARGE_CASES
if text == "all-large":
return CASES + LARGE_CASES
names = [item.strip() for item in text.split(",") if item.strip()]
missing = sorted(set(names) - CASE_BY_NAME.keys())
if missing:
raise ValueError(f"Unknown cases: {missing}")
return [CASE_BY_NAME[name] for name in names]
def tolerance(dtype: torch.dtype) -> tuple[float, float]:
if dtype == torch.float32:
return 1e-5, 1e-5
if dtype == torch.bfloat16:
return 7e-2, 2e-2
return 3e-3, 3e-3
def native_group_norm_silu(
x: torch.Tensor,
weight: torch.Tensor,
bias: torch.Tensor,
num_groups: int,
) -> torch.Tensor:
return F.silu(F.group_norm(x, num_groups, weight=weight, bias=bias, eps=EPS))
def make_inputs(case: Case, dtype: torch.dtype) -> tuple[torch.Tensor, ...]:
generator = torch.Generator(device=DEVICE)
generator.manual_seed(len(case.shape) * 1009 + case.shape[1] * 17 + case.num_groups)
x = torch.randn(case.shape, device=DEVICE, dtype=dtype, generator=generator)
weight = torch.randn(case.shape[1], device=DEVICE, dtype=dtype, generator=generator)
bias = torch.randn(case.shape[1], device=DEVICE, dtype=dtype, generator=generator)
return x, weight, bias
def do_bench_us(fn: Callable[[], object], warmup: int, rep: int) -> tuple[float, ...]:
median_ms, p20_ms, p80_ms = triton.testing.do_bench(
fn,
quantiles=QUANTILES,
warmup=warmup,
rep=rep,
)
return median_ms * 1000.0, p20_ms * 1000.0, p80_ms * 1000.0
def summarize(values: list[float]) -> float:
return statistics.median(values)
def run_case(
case: Case,
dtype: torch.dtype,
rounds: int,
warmup: int,
rep: int,
) -> dict[str, object]:
x, weight, bias = make_inputs(case, dtype)
with torch.inference_mode():
actual = triton_group_norm_silu(
x, weight, bias, num_groups=case.num_groups, eps=EPS
)
expected = native_group_norm_silu(x, weight, bias, case.num_groups)
atol, rtol = tolerance(dtype)
torch.testing.assert_close(actual, expected, atol=atol, rtol=rtol)
native_stats = []
fused_stats = []
for _ in range(rounds):
native_stats.append(
do_bench_us(
lambda: native_group_norm_silu(x, weight, bias, case.num_groups),
warmup=warmup,
rep=rep,
)
)
fused_stats.append(
do_bench_us(
lambda: triton_group_norm_silu(
x, weight, bias, num_groups=case.num_groups, eps=EPS
),
warmup=warmup,
rep=rep,
)
)
native_median_us = summarize([stats[0] for stats in native_stats])
fused_median_us = summarize([stats[0] for stats in fused_stats])
torch.cuda.empty_cache()
return {
"case": case.name,
"shape": "x".join(str(dim) for dim in case.shape),
"groups": case.num_groups,
"dtype": dtype_name(dtype),
"native_median_us": native_median_us,
"native_p20_us": summarize([stats[1] for stats in native_stats]),
"native_p80_us": summarize([stats[2] for stats in native_stats]),
"fused_median_us": fused_median_us,
"fused_p20_us": summarize([stats[1] for stats in fused_stats]),
"fused_p80_us": summarize([stats[2] for stats in fused_stats]),
"speedup": native_median_us / fused_median_us,
"rounds": rounds,
"warmup": warmup,
"rep": rep,
}
def run_profile(case: Case, dtype: torch.dtype, provider: str, iters: int) -> None:
x, weight, bias = make_inputs(case, dtype)
if provider == "native":
def fn() -> torch.Tensor:
return native_group_norm_silu(x, weight, bias, case.num_groups)
elif provider == "fused":
def fn() -> torch.Tensor:
return triton_group_norm_silu(
x, weight, bias, num_groups=case.num_groups, eps=EPS
)
else:
raise ValueError(f"Unknown provider: {provider}")
with torch.inference_mode():
for _ in range(5):
fn()
torch.cuda.synchronize()
for _ in range(iters):
fn()
torch.cuda.synchronize()
def write_csv(rows: list[dict[str, object]], output_path: Path) -> None:
output_path.parent.mkdir(parents=True, exist_ok=True)
fieldnames = list(rows[0].keys()) if rows else []
with output_path.open("w", newline="", encoding="utf-8") as f:
writer = csv.DictWriter(f, fieldnames=fieldnames)
writer.writeheader()
writer.writerows(rows)
def print_rows(rows: list[dict[str, object]]) -> None:
header = (
"case",
"dtype",
"shape",
"native_us",
"fused_us",
"speedup",
)
print("| " + " | ".join(header) + " |")
print("|---|---|---|---:|---:|---:|")
for row in rows:
print(
"| {case} | {dtype} | {shape} | {native:.2f} | {fused:.2f} | {speedup:.3f}x |".format(
case=row["case"],
dtype=row["dtype"],
shape=row["shape"],
native=row["native_median_us"],
fused=row["fused_median_us"],
speedup=row["speedup"],
)
)
def main() -> None:
parser = argparse.ArgumentParser(
description="Benchmark fused GroupNorm+SiLU against PyTorch GroupNorm+SiLU."
)
parser.add_argument(
"--cases",
default="all",
help=(
"Comma-separated case names, or one of: 'all' (default-sized "
"cases only), 'large' (high-memory cases only -- requires "
"H100/H200-class GPU), 'all-large' (both). See CASES + LARGE_CASES."
),
)
parser.add_argument("--dtypes", default="bf16,fp16")
parser.add_argument("--rounds", type=int, default=3)
parser.add_argument("--warmup", type=int, default=25)
parser.add_argument("--rep", type=int, default=100)
parser.add_argument("--output-csv", default="")
parser.add_argument("--profile-provider", choices=["native", "fused"], default="")
parser.add_argument("--profile-iters", type=int, default=20)
args = parser.parse_args()
if not torch.cuda.is_available():
raise RuntimeError("CUDA is required for this benchmark.")
cases = parse_cases(args.cases)
dtypes = parse_dtypes(args.dtypes)
if args.profile_provider:
if len(cases) != 1 or len(dtypes) != 1:
raise ValueError(
"--profile-provider requires exactly one case and one dtype"
)
run_profile(cases[0], dtypes[0], args.profile_provider, args.profile_iters)
return
rows = []
for case in cases:
for dtype in dtypes:
rows.append(run_case(case, dtype, args.rounds, args.warmup, args.rep))
print_rows(rows)
if args.output_csv:
write_csv(rows, Path(args.output_csv))
print(f"Wrote {args.output_csv}")
if __name__ == "__main__":
if is_in_ci():
print("Skipping bench_group_norm_silu.py in CI")
sys.exit(0)
main()
@@ -0,0 +1,209 @@
import random
import sys
from dataclasses import dataclass
import torch
from sglang.kernels.ops.diffusion.ltx2_qknorm_split_rope import (
ltx2_qknorm_split_rope_cuda,
)
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.utils import is_in_ci
register_cuda_ci(
est_time=30,
stage="base-b-kernel-benchmark",
runner_config="1-gpu-large",
disabled="standalone benchmark",
)
@dataclass(frozen=True)
class Workload:
name: str
batch: int
q_seq: int
k_seq: int
num_heads: int
head_dim: int
FULL_WORKLOADS = [
Workload("stage1_video_self_q1536_k1536_d4096", 2, 1536, 1536, 32, 128),
Workload("stage1_audio_self_q126_k126_d2048", 2, 126, 126, 32, 64),
Workload("stage1_audio_to_video_q1536_k126_d2048", 2, 1536, 126, 32, 64),
Workload("stage1_video_to_audio_q126_k1536_d2048", 2, 126, 1536, 32, 64),
Workload("stage2_video_self_q6144_k6144_d4096", 1, 6144, 6144, 32, 128),
Workload("stage2_audio_self_q126_k126_d2048", 1, 126, 126, 32, 64),
Workload("stage2_audio_to_video_q6144_k126_d2048", 1, 6144, 126, 32, 64),
Workload("stage2_video_to_audio_q126_k6144_d2048", 1, 126, 6144, 32, 64),
Workload("hq_stage1_video_self_q8160_k8160_d4096", 1, 8160, 8160, 32, 128),
Workload("hq_stage1_audio_to_video_q8160_k126_d2048", 1, 8160, 126, 32, 64),
Workload("hq_stage1_video_to_audio_q126_k8160_d2048", 1, 126, 8160, 32, 64),
Workload("hq_stage2_video_self_q32640_k32640_d4096", 1, 32640, 32640, 32, 128),
Workload("hq_stage2_audio_to_video_q32640_k126_d2048", 1, 32640, 126, 32, 64),
Workload("hq_stage2_video_to_audio_q126_k32640_d2048", 1, 126, 32640, 32, 64),
]
CI_WORKLOADS = [
Workload("stage1_video_self_q16_k16_d4096", 1, 16, 16, 32, 128),
Workload("stage1_audio_to_video_q16_k8_d2048", 1, 16, 8, 32, 64),
]
def _make_cos_sin(
batch: int, seq_len: int, num_heads: int, head_dim: int
) -> tuple[torch.Tensor, torch.Tensor]:
half_dim = head_dim // 2
cos = torch.randn(
batch, seq_len, num_heads, half_dim, device="cuda", dtype=torch.bfloat16
).transpose(1, 2)
sin = torch.randn(
batch, seq_len, num_heads, half_dim, device="cuda", dtype=torch.bfloat16
).transpose(1, 2)
return cos, sin
def _apply_split_rotary_ref(
x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor
) -> torch.Tensor:
x_dtype = x.dtype
batch = x.shape[0]
_, num_heads, seq_len, _ = cos.shape
x = x.reshape(batch, seq_len, num_heads, -1).swapaxes(1, 2)
last = x.shape[-1]
half = last // 2
split_x = x.reshape(*x.shape[:-1], 2, half)
first_x = split_x[..., :1, :]
second_x = split_x[..., 1:, :]
cos_u = cos.unsqueeze(-2)
sin_u = sin.unsqueeze(-2)
out = split_x * cos_u
out[..., :1, :].addcmul_(-sin_u, second_x)
out[..., 1:, :].addcmul_(sin_u, first_x)
out = out.reshape(*out.shape[:-2], last)
return out.swapaxes(1, 2).reshape(batch, seq_len, -1).to(dtype=x_dtype)
def _reference_pair(inputs):
q, k, q_cos, q_sin, k_cos, k_sin, q_norm, k_norm = inputs
with torch.autocast(device_type="cuda", dtype=torch.bfloat16, enabled=True):
q_out = _apply_split_rotary_ref(q_norm(q), q_cos, q_sin)
k_out = _apply_split_rotary_ref(k_norm(k), k_cos, k_sin)
return q_out.to(dtype=torch.bfloat16), k_out.to(dtype=torch.bfloat16)
def cuda_event_us(fn, warmups: int, repeats: int, rounds: int) -> float:
for _ in range(warmups):
fn()
torch.cuda.synchronize()
samples = []
for _ in range(rounds):
start = torch.cuda.Event(enable_timing=True)
end = torch.cuda.Event(enable_timing=True)
start.record()
for _ in range(repeats):
fn()
end.record()
end.synchronize()
samples.append(start.elapsed_time(end) * 1000.0 / repeats)
samples.sort()
return samples[len(samples) // 2]
def benchmark() -> None:
if not torch.cuda.is_available():
print("CUDA required")
return
torch.manual_seed(20260630)
random.seed(20260630)
torch.cuda.set_device(0)
workloads = CI_WORKLOADS if is_in_ci() else FULL_WORKLOADS
warmups = 3 if is_in_ci() else 10
repeats = 3 if is_in_ci() else 10
rounds = 3 if is_in_ci() else 7
print("| workload | torch us | cuda us | speedup |")
print("|---|---:|---:|---:|")
for workload in workloads:
hidden = workload.num_heads * workload.head_dim
q = torch.randn(
workload.batch,
workload.q_seq,
hidden,
device="cuda",
dtype=torch.bfloat16,
)
k = torch.randn(
workload.batch,
workload.k_seq,
hidden,
device="cuda",
dtype=torch.bfloat16,
)
q_cos, q_sin = _make_cos_sin(
workload.batch, workload.q_seq, workload.num_heads, workload.head_dim
)
k_cos, k_sin = _make_cos_sin(
workload.batch, workload.k_seq, workload.num_heads, workload.head_dim
)
q_norm = torch.nn.RMSNorm(hidden, eps=1e-6, device="cuda").to(
dtype=torch.bfloat16
)
k_norm = torch.nn.RMSNorm(hidden, eps=1e-6, device="cuda").to(
dtype=torch.bfloat16
)
inputs = (q, k, q_cos, q_sin, k_cos, k_sin, q_norm, k_norm)
q_ref, k_ref = _reference_pair(inputs)
q_out, k_out = ltx2_qknorm_split_rope_cuda(
q,
q_cos,
q_sin,
q_norm.weight,
k,
k_cos,
k_sin,
k_norm.weight,
eps=1e-6,
num_heads=workload.num_heads,
head_dim=workload.head_dim,
)
torch.cuda.synchronize()
assert torch.equal(q_ref, q_out)
assert torch.equal(k_ref, k_out)
fns = {
"torch": lambda: _reference_pair(inputs),
"cuda": lambda: ltx2_qknorm_split_rope_cuda(
q,
q_cos,
q_sin,
q_norm.weight,
k,
k_cos,
k_sin,
k_norm.weight,
eps=1e-6,
num_heads=workload.num_heads,
head_dim=workload.head_dim,
),
}
order = ["torch", "cuda"]
random.shuffle(order)
times = {
name: cuda_event_us(fns[name], warmups, repeats, rounds) for name in order
}
print(
f"| {workload.name} | {times['torch']:.2f} | "
f"{times['cuda']:.2f} | {times['torch'] / times['cuda']:.3f}x |"
)
torch.cuda.empty_cache()
if __name__ == "__main__":
benchmark()
sys.exit(0)
@@ -0,0 +1,757 @@
import argparse
import csv
import functools
import importlib
import math
import os
import statistics
import subprocess
import sys
from pathlib import Path
from typing import Callable
import torch
import torch.nn.functional as F
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.test.ci.ci_register import register_amd_ci, register_cuda_ci
from sglang.utils import is_in_ci
register_cuda_ci(
est_time=120,
stage="base-b-kernel-benchmark",
runner_config="1-gpu-large",
disabled="self-skips in CI, standalone tool",
)
register_amd_ci(est_time=120, stage="jit-kernel-benchmark", runner_config="amd")
os.environ.setdefault("FLASHINFER_DISABLE_VERSION_CHECK", "1")
REPO_ROOT = KERNEL_PATH.parents[2]
THIRD_PARTY_ROOT = REPO_ROOT / "third_party"
FLAGGEMS_REPO = "https://github.com/flagos-ai/FlagGems.git"
QUACK_REPO = "https://github.com/Dao-AILab/quack.git"
TORCH_LN = "torch.nn.LayerNorm"
SGL_RMS = "sglang.RMSNorm.forward_cuda"
SGL_FUSED = "sgl_kernel.fused_add_rmsnorm"
SGL_LN = "sglang.LayerNormScaleShift"
SGL_RES_LN = "sglang.ScaleResidualLayerNormScaleShift"
SGL_LN_PAIR = f"{SGL_LN} / {SGL_RES_LN}"
MOVA_LN_MIX = f"{TORCH_LN} / {SGL_LN_PAIR}"
ACTUAL_DIFFUSION_GROUPS: list[
tuple[str, str, list[tuple[str, str, tuple[int, ...], str]]]
] = [
(
"qwen",
"1 GPU",
[
("qwen_ln_4096x3072", "layernorm", (1, 4096, 3072), SGL_LN_PAIR),
("qwen_ln_26x3072", "layernorm", (1, 26, 3072), SGL_LN_PAIR),
("qwen_ln_6x3072", "layernorm", (1, 6, 3072), SGL_LN_PAIR),
("qwen_rms_26x3584", "rmsnorm", (1, 26, 3584), SGL_RMS),
("qwen_rms_6x3584", "rmsnorm", (1, 6, 3584), SGL_RMS),
],
),
(
"qwen-edit",
"1 GPU",
[
("qwen_edit_ln_200x3072", "layernorm", (1, 200, 3072), SGL_LN_PAIR),
("qwen_edit_ln_203x3072", "layernorm", (1, 203, 3072), SGL_LN_PAIR),
("qwen_edit_ln_8308x3072", "layernorm", (1, 8308, 3072), TORCH_LN),
("qwen_edit_rms_200x3584", "rmsnorm", (1, 200, 3584), SGL_RMS),
("qwen_edit_rms_203x3584", "rmsnorm", (1, 203, 3584), SGL_RMS),
],
),
(
"flux",
"1 GPU",
[
("flux_ln_77x768", "layernorm", (1, 77, 768), TORCH_LN),
("flux_ln_512x3072", "layernorm", (1, 512, 3072), TORCH_LN),
("flux_ln_4096x3072", "layernorm", (1, 4096, 3072), TORCH_LN),
("flux_ln_4608x3072", "layernorm", (1, 4608, 3072), TORCH_LN),
("flux_rms_512x4096", "rmsnorm", (1, 512, 4096), SGL_RMS),
],
),
(
"flux2",
"1 GPU",
[
("flux2_ln_512x6144", "layernorm", (1, 512, 6144), TORCH_LN),
("flux2_ln_4096x6144", "layernorm", (1, 4096, 6144), TORCH_LN),
("flux2_ln_4608x6144", "layernorm", (1, 4608, 6144), TORCH_LN),
("flux2_rms_4608x48x128", "rmsnorm", (1, 4608, 48, 128), SGL_RMS),
],
),
(
"zimage",
"1 GPU",
[
("zimage_ln_4128x3840", "layernorm", (1, 4128, 3840), TORCH_LN),
("zimage_rms_32x3840", "rmsnorm", (1, 32, 3840), SGL_RMS),
("zimage_rms_4096x3840", "rmsnorm", (1, 4096, 3840), SGL_RMS),
("zimage_rms_4128x3840", "rmsnorm", (1, 4128, 3840), SGL_RMS),
("zimage_rms_32x2560", "rmsnorm", (32, 2560), SGL_RMS),
],
),
(
"wan-ti2v",
"1 GPU",
[
("wan_ti2v_ln_17850x3072", "layernorm", (1, 17850, 3072), SGL_LN_PAIR),
("wan_ti2v_rms_17850x3072", "rmsnorm", (1, 17850, 3072), SGL_RMS),
("wan_ti2v_rms_512x3072", "rmsnorm", (1, 512, 3072), SGL_RMS),
("wan_ti2v_rms_512x4096", "rmsnorm", (1, 512, 4096), SGL_RMS),
],
),
(
"hunyuanvideo",
"1 GPU",
[
("hunyuan_ln_46x768", "layernorm", (1, 46, 768), TORCH_LN),
("hunyuan_ln_45x3072", "layernorm", (1, 45, 3072), SGL_LN_PAIR),
("hunyuan_ln_27030x3072", "layernorm", (1, 27030, 3072), SGL_LN_PAIR),
("hunyuan_ln_27075x3072", "layernorm", (1, 27075, 3072), SGL_LN),
("hunyuan_rms_140x4096", "rmsnorm", (1, 140, 4096), SGL_RMS),
("hunyuan_rms_45x24x128", "rmsnorm", (1, 45, 24, 128), SGL_RMS),
("hunyuan_rms_27030x24x128", "rmsnorm", (1, 27030, 24, 128), SGL_RMS),
("hunyuan_rms_27075x24x128", "rmsnorm", (1, 27075, 24, 128), SGL_RMS),
("hunyuan_fused_add_140x4096", "fused_add_rmsnorm", (140, 4096), SGL_FUSED),
],
),
(
"mova-720p",
"4 GPU, ulysses=4, ring=1",
[
("mova_ln_101x1536", "layernorm", (1, 101, 1536), MOVA_LN_MIX),
("mova_ln_403x1536", "layernorm", (1, 403, 1536), TORCH_LN),
("mova_ln_44100x5120", "layernorm", (1, 44100, 5120), MOVA_LN_MIX),
("mova_ln_176400x5120", "layernorm", (1, 176400, 5120), SGL_LN),
("mova_rms_101x1536", "rmsnorm", (1, 101, 1536), SGL_RMS),
("mova_rms_101x5120", "rmsnorm", (1, 101, 5120), SGL_RMS),
("mova_rms_44100x1536", "rmsnorm", (1, 44100, 1536), SGL_RMS),
("mova_rms_44100x5120", "rmsnorm", (1, 44100, 5120), SGL_RMS),
("mova_rms_512x1536", "rmsnorm", (1, 512, 1536), SGL_RMS),
("mova_rms_512x4096", "rmsnorm", (1, 512, 4096), SGL_RMS),
("mova_rms_512x5120", "rmsnorm", (1, 512, 5120), SGL_RMS),
],
),
]
ACTUAL_DIFFUSION_SHAPES: list[dict[str, object]] = [
{
"shape_id": shape_id,
"model": model,
"gpu_config": gpu_config,
"op": op,
"input_shape": list(input_shape),
"source_impl": source_impl,
}
for model, gpu_config, cases in ACTUAL_DIFFUSION_GROUPS
for shape_id, op, input_shape, source_impl in cases
]
def effective_rows_from_shape(input_shape: list[int]) -> int:
rows = 1
for dim in input_shape[:-1]:
rows *= dim
return rows
def ensure_repo(repo_name: str, repo_url: str) -> Path:
repo_path = THIRD_PARTY_ROOT / repo_name
if repo_path.exists():
return repo_path
repo_path.parent.mkdir(parents=True, exist_ok=True)
subprocess.run(
["git", "clone", "--depth", "1", repo_url, str(repo_path)],
check=True,
cwd=REPO_ROOT,
)
return repo_path
def ensure_python_dep(module_name: str, package_name: str | None = None) -> None:
package_name = package_name or module_name
try:
importlib.import_module(module_name)
except ModuleNotFoundError:
subprocess.run(
[sys.executable, "-m", "pip", "install", package_name],
check=True,
)
def dtype_from_name(name: str) -> torch.dtype:
mapping = {
"bf16": torch.bfloat16,
"bfloat16": torch.bfloat16,
"fp16": torch.float16,
"float16": torch.float16,
"fp32": torch.float32,
"float32": torch.float32,
}
return mapping[name]
def dtype_name(dtype: torch.dtype) -> str:
mapping = {
torch.bfloat16: "bf16",
torch.float16: "fp16",
torch.float32: "fp32",
}
return mapping[dtype]
def normalize_hidden_sizes(text: str) -> list[int]:
return [int(x) for x in text.split(",") if x]
def normalize_dtypes(text: str) -> list[torch.dtype]:
return [dtype_from_name(x.strip()) for x in text.split(",") if x.strip()]
def prewarm(fn: Callable[[], object], iters: int = 3) -> None:
for _ in range(iters):
fn()
torch.cuda.synchronize()
def benchmark_provider(
fn: Callable[[], object],
setup_fn: Callable[[], None] | None = None,
warmup: int = 10,
rep: int = 30,
) -> tuple[float, float, float]:
for _ in range(warmup):
if setup_fn is not None:
setup_fn()
fn()
torch.cuda.synchronize()
start_event = torch.cuda.Event(enable_timing=True)
end_event = torch.cuda.Event(enable_timing=True)
times_us: list[float] = []
for _ in range(rep):
if setup_fn is not None:
setup_fn()
start_event.record()
fn()
end_event.record()
end_event.synchronize()
times_us.append(start_event.elapsed_time(end_event) * 1000.0)
return statistics.median(times_us), max(times_us), min(times_us)
def geometric_mean(values: list[float]) -> float:
if not values:
return float("nan")
return math.exp(sum(math.log(v) for v in values) / len(values))
@functools.cache
def load_flaggems():
ensure_python_dep("sqlalchemy")
ensure_repo("FlagGems", FLAGGEMS_REPO)
src_root = THIRD_PARTY_ROOT / "FlagGems" / "src"
if str(src_root) not in sys.path:
sys.path.insert(0, str(src_root))
from flag_gems.fused.fused_add_rms_norm import fused_add_rms_norm
from flag_gems.ops.layernorm import layer_norm
from flag_gems.ops.rms_norm import rms_norm
return rms_norm, layer_norm, fused_add_rms_norm
@functools.cache
def load_quack():
repo_path = ensure_repo("quack", QUACK_REPO)
try:
quack_rmsnorm = importlib.import_module("quack.rmsnorm")
except ModuleNotFoundError:
subprocess.run(
[sys.executable, "-m", "pip", "install", "-e", str(repo_path)],
check=True,
)
quack_rmsnorm = importlib.import_module("quack.rmsnorm")
return quack_rmsnorm.rmsnorm_fwd, quack_rmsnorm.layernorm_fwd
def build_rmsnorm_providers(dtype: torch.dtype, batch_size: int, hidden_size: int):
import flashinfer.norm as flashinfer_norm
import sgl_kernel
x = torch.randn((batch_size, hidden_size), device=DEFAULT_DEVICE, dtype=dtype)
weight = torch.randn(hidden_size, device=DEFAULT_DEVICE, dtype=dtype)
jit_out = torch.empty_like(x)
sgl_out = torch.empty_like(x)
flashinfer_out = torch.empty_like(x)
flaggems_rms_norm, _, _ = load_flaggems()
quack_rmsnorm_fwd, _ = load_quack()
providers = {
"pytorch": lambda: F.rms_norm(x, (hidden_size,), weight, 1e-6),
"sgl_kernel": lambda: sgl_kernel.rmsnorm(x, weight, eps=1e-6, out=sgl_out),
"flashinfer": lambda: flashinfer_norm.rmsnorm(
x, weight, eps=1e-6, out=flashinfer_out
),
"jit_rmsnorm": lambda: jit_rmsnorm(x, weight, jit_out, 1e-6),
"quack": lambda: quack_rmsnorm_fwd(x, weight, eps=1e-6),
"triton_rms_norm_fn": lambda: rms_norm_fn(
x, weight, bias=None, residual=None, eps=1e-6
),
"flaggems": lambda: flaggems_rms_norm(x, (hidden_size,), weight, 1e-6),
}
if hidden_size <= 128:
providers["triton_one_pass"] = lambda: triton_one_pass_rms_norm(x, weight, 1e-6)
return providers
def build_fused_add_rmsnorm_providers(
dtype: torch.dtype, batch_size: int, hidden_size: int
):
import flashinfer.norm as flashinfer_norm
import sgl_kernel
base_x = torch.randn((batch_size, hidden_size), device=DEFAULT_DEVICE, dtype=dtype)
base_residual = torch.randn_like(base_x)
weight = torch.randn(hidden_size, device=DEFAULT_DEVICE, dtype=dtype)
x = base_x.clone()
residual = base_residual.clone()
def reset():
x.copy_(base_x)
residual.copy_(base_residual)
_, _, flaggems_fused_add_rms_norm = load_flaggems()
quack_rmsnorm_fwd, _ = load_quack()
def pytorch_impl():
out = x + residual
return F.rms_norm(out, (hidden_size,), weight, 1e-6)
providers = {
"pytorch": (pytorch_impl, reset),
"sgl_kernel": (
lambda: sgl_kernel.fused_add_rmsnorm(x, residual, weight, eps=1e-6),
reset,
),
"flashinfer": (
lambda: flashinfer_norm.fused_add_rmsnorm(x, residual, weight, eps=1e-6),
reset,
),
"jit_fused_add_rmsnorm": (
lambda: jit_fused_add_rmsnorm(x, residual, weight, 1e-6),
reset,
),
"quack": (
lambda: quack_rmsnorm_fwd(x, weight, residual=residual, eps=1e-6),
reset,
),
"flaggems": (
lambda: flaggems_fused_add_rms_norm(
x, residual, (hidden_size,), weight, 1e-6
),
reset,
),
}
return providers
def build_layernorm_providers(dtype: torch.dtype, batch_size: int, hidden_size: int):
import flashinfer.norm as flashinfer_norm
x = torch.randn((batch_size, hidden_size), device=DEFAULT_DEVICE, dtype=dtype)
weight = torch.randn(hidden_size, device=DEFAULT_DEVICE, dtype=dtype)
bias = torch.randn(hidden_size, device=DEFAULT_DEVICE, dtype=dtype)
flashinfer_weight = torch.randn(
hidden_size, device=DEFAULT_DEVICE, dtype=torch.float32
)
flashinfer_bias = torch.randn(
hidden_size, device=DEFAULT_DEVICE, dtype=torch.float32
)
triton_out = torch.empty_like(x)
_, flaggems_layer_norm, _ = load_flaggems()
_, quack_layernorm_fwd = load_quack()
providers = {
"pytorch": lambda: F.layer_norm(x, (hidden_size,), weight, bias, 1e-6),
"triton_norm_infer": lambda: norm_infer(
x, weight, bias, eps=1e-6, is_rms_norm=False, out=triton_out
),
"flashinfer": lambda: flashinfer_norm.layernorm(
x, flashinfer_weight, flashinfer_bias, 1e-6
),
"quack": lambda: quack_layernorm_fwd(
x, flashinfer_weight, flashinfer_bias, 1e-6
),
"flaggems": lambda: flaggems_layer_norm(x, (hidden_size,), weight, bias)[0],
}
return providers
def maybe_benchmark(
op_name: str,
provider_name: str,
fn: Callable[[], object],
rows: list[dict[str, object]],
dtype: torch.dtype,
batch_size: int,
hidden_size: int,
reset: Callable[[], None] | None = None,
metadata: dict[str, object] | None = None,
) -> None:
metadata = metadata or {}
try:
median_us, max_us, min_us = benchmark_provider(fn, reset)
rows.append(
{
"op": op_name,
"provider": provider_name,
"dtype": dtype_name(dtype),
"batch_size": batch_size,
"hidden_size": hidden_size,
"median_us": median_us,
"min_us": min_us,
"max_us": max_us,
"status": "ok",
"error": "",
**metadata,
}
)
except Exception as exc: # pragma: no cover - benchmark failures are data
rows.append(
{
"op": op_name,
"provider": provider_name,
"dtype": dtype_name(dtype),
"batch_size": batch_size,
"hidden_size": hidden_size,
"median_us": "",
"min_us": "",
"max_us": "",
"status": "unsupported",
"error": str(exc),
**metadata,
}
)
def write_csv(rows: list[dict[str, object]], output_path: Path) -> None:
output_path.parent.mkdir(parents=True, exist_ok=True)
with output_path.open("w", newline="", encoding="utf-8") as f:
writer = csv.DictWriter(
f,
fieldnames=[
"op",
"provider",
"dtype",
"batch_size",
"hidden_size",
"median_us",
"min_us",
"max_us",
"shape_id",
"source_model",
"source_gpu_config",
"source_input_shape",
"source_impl",
"status",
"error",
],
)
writer.writeheader()
writer.writerows(rows)
def write_markdown(rows: list[dict[str, object]], output_path: Path) -> None:
lines: list[str] = []
lines.append("# Norm Benchmark Summary")
lines.append("")
actual_shape_rows = [row for row in rows if row.get("shape_id")]
if actual_shape_rows:
seen: set[tuple[str, str, str, str, str, str]] = set()
lines.append("## Diffusion Shape Cases")
lines.append("")
lines.append(
"| Shape ID | Op | Model | GPU Config | Input Shape | Source Impl |"
)
lines.append("|---|---|---|---|---|---|")
for row in actual_shape_rows:
key = (
str(row.get("shape_id", "")),
str(row.get("op", "")),
str(row.get("source_model", "")),
str(row.get("source_gpu_config", "")),
str(row.get("source_input_shape", "")),
str(row.get("source_impl", "")),
)
if key in seen:
continue
seen.add(key)
lines.append(
f"| {key[0]} | {key[1]} | {key[2]} | {key[3]} | `{key[4]}` | {key[5]} |"
)
lines.append("")
for op_name in ("rmsnorm", "fused_add_rmsnorm", "layernorm"):
for dtype in sorted({row["dtype"] for row in rows}):
scoped = [
row
for row in rows
if row["op"] == op_name
and row["dtype"] == dtype
and row["status"] == "ok"
]
if not scoped:
continue
provider_to_values: dict[str, list[float]] = {}
provider_to_speedups: dict[str, list[float]] = {}
by_shape: dict[tuple[str, int, int], dict[str, float]] = {}
for row in scoped:
provider = str(row["provider"])
value = float(row["median_us"])
provider_to_values.setdefault(provider, []).append(value)
shape = (
str(row.get("shape_id", "")),
int(row["batch_size"]),
int(row["hidden_size"]),
)
by_shape.setdefault(shape, {})[provider] = value
for shape, perf in by_shape.items():
if "pytorch" not in perf:
continue
baseline = perf["pytorch"]
for provider, value in perf.items():
provider_to_speedups.setdefault(provider, []).append(
baseline / value
)
lines.append(f"## {op_name} ({dtype})")
lines.append("")
lines.append(
"| Provider | Geomean Speedup vs PyTorch | Median Latency (us) | Win Count |"
)
lines.append("|---|---:|---:|---:|")
wins: dict[str, int] = {}
for perf in by_shape.values():
best_provider = min(perf, key=perf.get)
wins[best_provider] = wins.get(best_provider, 0) + 1
for provider in sorted(provider_to_values):
geomean_speedup = geometric_mean(provider_to_speedups.get(provider, []))
median_latency = statistics.median(provider_to_values[provider])
win_count = wins.get(provider, 0)
lines.append(
f"| {provider} | {geomean_speedup:.3f}x | {median_latency:.2f} | {win_count} |"
)
lines.append("")
output_path.write_text("\n".join(lines) + "\n", encoding="utf-8")
def run_suite(
hidden_sizes: list[int],
batch_sizes: list[int],
dtypes: list[torch.dtype],
ops: list[str],
) -> list[dict[str, object]]:
rows: list[dict[str, object]] = []
for dtype in dtypes:
for batch_size in batch_sizes:
for hidden_size in hidden_sizes:
if "rmsnorm" in ops:
rms_providers = build_rmsnorm_providers(
dtype, batch_size, hidden_size
)
for provider_name, fn in rms_providers.items():
maybe_benchmark(
"rmsnorm",
provider_name,
fn,
rows,
dtype,
batch_size,
hidden_size,
)
if "fused_add_rmsnorm" in ops:
fused_providers = build_fused_add_rmsnorm_providers(
dtype, batch_size, hidden_size
)
for provider_name, provider in fused_providers.items():
fn, reset = provider
maybe_benchmark(
"fused_add_rmsnorm",
provider_name,
fn,
rows,
dtype,
batch_size,
hidden_size,
reset,
)
if "layernorm" in ops:
layernorm_providers = build_layernorm_providers(
dtype, batch_size, hidden_size
)
for provider_name, fn in layernorm_providers.items():
maybe_benchmark(
"layernorm",
provider_name,
fn,
rows,
dtype,
batch_size,
hidden_size,
)
return rows
def run_shape_suite(
shape_cases: list[dict[str, object]],
dtypes: list[torch.dtype],
) -> list[dict[str, object]]:
rows: list[dict[str, object]] = []
for case in shape_cases:
op_name = str(case["op"])
input_shape = [int(x) for x in case["input_shape"]]
batch_size = effective_rows_from_shape(input_shape)
hidden_size = input_shape[-1]
metadata = {
"shape_id": str(case["shape_id"]),
"source_model": str(case["model"]),
"source_gpu_config": str(case["gpu_config"]),
"source_input_shape": str(input_shape),
"source_impl": str(case["source_impl"]),
}
for dtype in dtypes:
if op_name == "rmsnorm":
providers = build_rmsnorm_providers(dtype, batch_size, hidden_size)
for provider_name, fn in providers.items():
maybe_benchmark(
op_name,
provider_name,
fn,
rows,
dtype,
batch_size,
hidden_size,
metadata=metadata,
)
elif op_name == "fused_add_rmsnorm":
providers = build_fused_add_rmsnorm_providers(
dtype, batch_size, hidden_size
)
for provider_name, provider in providers.items():
fn, reset = provider
maybe_benchmark(
op_name,
provider_name,
fn,
rows,
dtype,
batch_size,
hidden_size,
reset,
metadata=metadata,
)
elif op_name == "layernorm":
providers = build_layernorm_providers(dtype, batch_size, hidden_size)
for provider_name, fn in providers.items():
maybe_benchmark(
op_name,
provider_name,
fn,
rows,
dtype,
batch_size,
hidden_size,
metadata=metadata,
)
else:
raise ValueError(f"Unsupported op in shape preset: {op_name}")
return rows
def main() -> None:
parser = argparse.ArgumentParser(
description="Benchmark RMSNorm/LayerNorm implementations across providers."
)
parser.add_argument(
"--hidden-sizes",
default="64,128,256,512,1024,2048,4096,8192,16384",
help="Comma-separated hidden sizes.",
)
parser.add_argument(
"--batch-sizes",
default="1,16,128,1024",
help="Comma-separated batch sizes.",
)
parser.add_argument(
"--dtypes",
default="bf16,fp16",
help="Comma-separated dtypes: bf16, fp16, fp32.",
)
parser.add_argument(
"--output-dir",
default=str(REPO_ROOT / "outputs" / "norm_benchmarks"),
help="Directory for CSV/Markdown outputs.",
)
parser.add_argument(
"--ops",
default="rmsnorm,fused_add_rmsnorm,layernorm",
help="Comma-separated ops to benchmark.",
)
parser.add_argument(
"--shape-preset",
choices=["grid", "diffusion-actual"],
default="grid",
help="Use the default grid sweep or the captured diffusion workload shapes.",
)
args = parser.parse_args()
if not torch.cuda.is_available():
raise RuntimeError("CUDA is required for norm benchmarks.")
hidden_sizes = normalize_hidden_sizes(args.hidden_sizes)
batch_sizes = normalize_hidden_sizes(args.batch_sizes)
dtypes = normalize_dtypes(args.dtypes)
ops = [op.strip() for op in args.ops.split(",") if op.strip()]
if args.shape_preset == "diffusion-actual":
shape_cases = [case for case in ACTUAL_DIFFUSION_SHAPES if case["op"] in ops]
rows = run_shape_suite(shape_cases, dtypes)
else:
rows = run_suite(hidden_sizes, batch_sizes, dtypes, ops)
output_dir = Path(args.output_dir)
csv_path = output_dir / "norm_impls.csv"
md_path = output_dir / "norm_impls_summary.md"
write_csv(rows, csv_path)
write_markdown(rows, md_path)
print(f"Wrote {csv_path}")
print(f"Wrote {md_path}")
if __name__ == "__main__":
if is_in_ci():
print("Skipping bench_norm_impls.py in CI")
sys.exit(0)
main()
@@ -0,0 +1,192 @@
from dataclasses import dataclass
from typing import Tuple
import torch
import triton
import triton.testing
from sglang.kernels.jit.benchmark.utils import (
DEFAULT_DEVICE,
DEFAULT_DTYPE,
get_benchmark_range,
run_benchmark_no_cudagraph,
)
from sglang.test.ci.ci_register import register_cuda_ci
register_cuda_ci(
est_time=13, stage="base-b-kernel-benchmark", runner_config="1-gpu-large"
)
MAX_SEQ_LEN = 131072
ROPE_BASE = 10000.0
@dataclass(frozen=True)
class CaseSpec:
name: str
batch_size: int
num_tokens: int
num_heads: int
head_dim: int
rope_dim: int
is_neox: bool
BENCH_CASES = (
CaseSpec("flux_1024", 1, 4096, 24, 128, 128, False),
CaseSpec("qwen_image_1024", 1, 4096, 32, 128, 128, False),
CaseSpec("qwen_image_partial", 1, 4096, 32, 128, 64, False),
# Z-Image-Turbo default 1024x1024 config: dim=3840, num_heads=30 -> head_dim=128.
CaseSpec("zimage_1024", 1, 4096, 30, 128, 128, False),
CaseSpec("batch2_medium", 2, 2048, 24, 128, 128, False),
)
CASE_BY_NAME = {case.name: case for case in BENCH_CASES}
CASE_NAMES = get_benchmark_range(
full_range=[case.name for case in BENCH_CASES],
ci_range=[case.name for case in BENCH_CASES],
)
LINE_VALS = ["split", "fused"]
LINE_NAMES = ["JIT QKNorm + FlashInfer RoPE", "SGL JIT Fused QKNorm+RoPE"]
STYLES = [("red", "-"), ("blue", "--")]
def create_cos_sin_cache(
rotary_dim: int,
max_position: int = MAX_SEQ_LEN,
base: float = ROPE_BASE,
) -> torch.Tensor:
inv_freq = 1.0 / (
base
** (
torch.arange(0, rotary_dim, 2, dtype=torch.float32, device=DEFAULT_DEVICE)
/ rotary_dim
)
)
t = torch.arange(max_position, dtype=torch.float32, device=DEFAULT_DEVICE)
freqs = torch.einsum("i,j->ij", t, inv_freq)
return torch.cat((freqs.cos(), freqs.sin()), dim=-1)
def make_inputs(case: CaseSpec) -> dict[str, torch.Tensor | bool]:
seed = (
case.batch_size * 1_000_003
+ case.num_tokens * 8191
+ case.num_heads * 127
+ case.head_dim * 17
+ case.rope_dim
)
generator = torch.Generator(device=DEFAULT_DEVICE)
generator.manual_seed(seed)
return {
"q": torch.randn(
case.batch_size * case.num_tokens,
case.num_heads,
case.head_dim,
device=DEFAULT_DEVICE,
dtype=DEFAULT_DTYPE,
generator=generator,
),
"k": torch.randn(
case.batch_size * case.num_tokens,
case.num_heads,
case.head_dim,
device=DEFAULT_DEVICE,
dtype=DEFAULT_DTYPE,
generator=generator,
),
"q_weight": torch.randn(
case.head_dim,
device=DEFAULT_DEVICE,
dtype=DEFAULT_DTYPE,
generator=generator,
),
"k_weight": torch.randn(
case.head_dim,
device=DEFAULT_DEVICE,
dtype=DEFAULT_DTYPE,
generator=generator,
),
"positions": torch.randint(
0,
MAX_SEQ_LEN,
(case.batch_size * case.num_tokens,),
device=DEFAULT_DEVICE,
dtype=torch.int64,
generator=generator,
),
"cos_sin_cache": create_cos_sin_cache(case.rope_dim),
"is_neox": case.is_neox,
}
def clone_inputs(
inputs: dict[str, torch.Tensor | bool],
) -> dict[str, torch.Tensor | bool]:
out: dict[str, torch.Tensor | bool] = {}
for key, value in inputs.items():
out[key] = value.clone() if isinstance(value, torch.Tensor) else value
return out
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
q = inputs["q"]
k = inputs["k"]
q_weight = inputs["q_weight"]
k_weight = inputs["k_weight"]
positions = inputs["positions"]
cos_sin_cache = inputs["cos_sin_cache"]
is_neox = bool(inputs["is_neox"])
fused_inplace_qknorm(q, k, q_weight, k_weight)
apply_rope_with_cos_sin_cache_inplace(
positions=positions,
query=q.view(q.shape[0], -1),
key=k.view(k.shape[0], -1),
head_size=q.shape[-1],
cos_sin_cache=cos_sin_cache,
is_neox=is_neox,
)
def fused_qknorm_rope(inputs: dict[str, torch.Tensor | bool]) -> None:
from sglang.kernels.ops.diffusion.qknorm_rope import fused_inplace_qknorm_rope
fused_inplace_qknorm_rope(
inputs["q"],
inputs["k"],
inputs["q_weight"],
inputs["k_weight"],
inputs["cos_sin_cache"],
inputs["positions"],
is_neox=bool(inputs["is_neox"]),
rope_dim=inputs["cos_sin_cache"].shape[-1],
)
@triton.testing.perf_report(
triton.testing.Benchmark(
x_names=["case_name"],
x_vals=CASE_NAMES,
line_arg="provider",
line_vals=LINE_VALS,
line_names=LINE_NAMES,
styles=STYLES,
ylabel="us",
plot_name="diffusion-qknorm-rope-performance",
args={},
)
)
def benchmark(case_name: str, provider: str) -> Tuple[float, float, float]:
case = CASE_BY_NAME[case_name]
inputs = make_inputs(case)
fn = split_qknorm_rope if provider == "split" else fused_qknorm_rope
return run_benchmark_no_cudagraph(lambda: fn(inputs))
if __name__ == "__main__":
print("Running diffusion qknorm + rope performance benchmark...")
benchmark.run(print_data=True)
@@ -0,0 +1,187 @@
from typing import Tuple
import torch
import triton.testing
from sglang.kernels.jit.benchmark.utils import run_benchmark_no_cudagraph
from sglang.kernels.ops.diffusion.triton.norm import norm_infer
from sglang.kernels.ops.diffusion.triton.scale_shift import (
fuse_layernorm_scale_shift_gate_select01_kernel,
fuse_residual_layernorm_scale_shift_gate_select01_kernel,
)
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
from sglang.utils import is_in_ci
register_cuda_ci(
est_time=13, stage="base-b-kernel-benchmark", runner_config="1-gpu-large"
)
register_amd_ci(est_time=13, stage="jit-kernel-benchmark", runner_config="amd")
if is_in_ci():
B_RANGE, S_RANGE, D_RANGE = [1], [128], [3072]
else:
B_RANGE, S_RANGE, D_RANGE = [1, 2], [128, 512, 2048], [1024, 1536, 3072]
DTYPE = torch.bfloat16
DEVICE = "cuda"
EPS = 1e-6
LINE_VALS = ["split", "fused"]
LINE_NAMES = ["Triton Norm + Torch Select", "Fused Triton"]
STYLES = [("red", "-"), ("blue", "--")]
CONFIG = [(b, s, d) for b in B_RANGE for s in S_RANGE for d in D_RANGE]
def _make_common_inputs(batch_size: int, seq_len: int, hidden_size: int):
x = torch.randn(batch_size, seq_len, hidden_size, dtype=DTYPE, device=DEVICE)
weight = torch.randn(hidden_size, dtype=DTYPE, device=DEVICE)
bias = torch.randn(hidden_size, dtype=DTYPE, device=DEVICE)
index = torch.randint(0, 2, (batch_size, seq_len), dtype=torch.int32, device=DEVICE)
scale0 = torch.randn(batch_size, hidden_size, dtype=DTYPE, device=DEVICE)
shift0 = torch.randn(batch_size, hidden_size, dtype=DTYPE, device=DEVICE)
gate0 = torch.randn(batch_size, hidden_size, dtype=DTYPE, device=DEVICE)
scale1 = torch.randn(batch_size, hidden_size, dtype=DTYPE, device=DEVICE)
shift1 = torch.randn(batch_size, hidden_size, dtype=DTYPE, device=DEVICE)
gate1 = torch.randn(batch_size, hidden_size, dtype=DTYPE, device=DEVICE)
return x, weight, bias, index, scale0, shift0, gate0, scale1, shift1, gate1
def _apply_select01_modulation(
x: torch.Tensor,
scale0: torch.Tensor,
shift0: torch.Tensor,
gate0: torch.Tensor,
scale1: torch.Tensor,
shift1: torch.Tensor,
gate1: torch.Tensor,
index: torch.Tensor,
):
idx = index.bool().unsqueeze(-1)
scale = torch.where(idx, scale1.unsqueeze(1), scale0.unsqueeze(1))
shift = torch.where(idx, shift1.unsqueeze(1), shift0.unsqueeze(1))
gate = torch.where(idx, gate1.unsqueeze(1), gate0.unsqueeze(1))
return x * (1 + scale) + shift, gate
@triton.testing.perf_report(
triton.testing.Benchmark(
x_names=["B", "S", "D"],
x_vals=CONFIG,
line_arg="provider",
line_vals=LINE_VALS,
line_names=LINE_NAMES,
styles=STYLES,
ylabel="us",
plot_name="qwen_image_layernorm_scale_shift_gate_select01",
args={},
)
)
def bench_layernorm_scale_shift_gate_select01(
B: int, S: int, D: int, provider: str
) -> Tuple[float, float, float]:
x, weight, bias, index, scale0, shift0, gate0, scale1, shift1, gate1 = (
_make_common_inputs(B, S, D)
)
if provider == "split":
def fn():
normalized = norm_infer(
x.view(-1, x.shape[-1]),
weight,
bias,
eps=EPS,
is_rms_norm=False,
).view_as(x)
return _apply_select01_modulation(
normalized, scale0, shift0, gate0, scale1, shift1, gate1, index
)
else:
def fn():
return fuse_layernorm_scale_shift_gate_select01_kernel(
x,
weight=weight,
bias=bias,
scale0=scale0,
shift0=shift0,
gate0=gate0,
scale1=scale1,
shift1=shift1,
gate1=gate1,
index=index,
eps=EPS,
)
return run_benchmark_no_cudagraph(fn)
@triton.testing.perf_report(
triton.testing.Benchmark(
x_names=["B", "S", "D"],
x_vals=CONFIG,
line_arg="provider",
line_vals=LINE_VALS,
line_names=LINE_NAMES,
styles=STYLES,
ylabel="us",
plot_name="qwen_image_residual_layernorm_scale_shift_gate_select01",
args={},
)
)
def bench_residual_layernorm_scale_shift_gate_select01(
B: int, S: int, D: int, provider: str
) -> Tuple[float, float, float]:
x, weight, bias, index, scale0, shift0, gate0, scale1, shift1, gate1 = (
_make_common_inputs(B, S, D)
)
residual = torch.randn_like(x)
residual_gate = torch.randn_like(x)
if provider == "split":
def fn():
residual_out = residual + residual_gate * x
normalized = norm_infer(
residual_out.view(-1, residual_out.shape[-1]),
weight,
bias,
eps=EPS,
is_rms_norm=False,
).view_as(residual_out)
return _apply_select01_modulation(
normalized, scale0, shift0, gate0, scale1, shift1, gate1, index
)
else:
def fn():
return fuse_residual_layernorm_scale_shift_gate_select01_kernel(
x,
residual=residual,
residual_gate=residual_gate,
weight=weight,
bias=bias,
scale0=scale0,
shift0=shift0,
gate0=gate0,
scale1=scale1,
shift1=shift1,
gate1=gate1,
index=index,
eps=EPS,
)
return run_benchmark_no_cudagraph(fn)
if __name__ == "__main__":
print(f"\n{'=' * 80}")
print("Benchmark: qwen_image layernorm + scale_shift_gate_select01")
print(f"{'=' * 80}\n")
bench_layernorm_scale_shift_gate_select01.run(print_data=True)
print(f"\n{'=' * 80}")
print("Benchmark: qwen_image residual + layernorm + scale_shift_gate_select01")
print(f"{'=' * 80}\n")
bench_residual_layernorm_scale_shift_gate_select01.run(print_data=True)
@@ -0,0 +1,116 @@
import random
import sys
from dataclasses import dataclass
import torch
from sglang.kernels.ops.diffusion.residual_gate_add import residual_gate_add_cuda
from sglang.kernels.ops.diffusion.triton.scale_shift import fuse_scale_shift_kernel
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.utils import is_in_ci
register_cuda_ci(
est_time=30, stage="base-b-kernel-benchmark", runner_config="1-gpu-large"
)
@dataclass(frozen=True)
class Workload:
name: str
residual_shape: tuple[int, ...]
gate_shape: tuple[int, ...]
FULL_WORKLOADS = [
Workload("ltx2_bcast_s32640_c4096", (1, 32640, 4096), (1, 1, 4096)),
Workload("ltx2_full_s8160_c4096", (1, 8160, 4096), (1, 8160, 4096)),
Workload("ideogram4_bcast_s4096_c4608", (1, 4096, 4608), (1, 1, 4608)),
Workload("flux2_bcast_s4608_c3072", (1, 4608, 3072), (1, 1, 3072)),
Workload("flux2_bcast_s4096_c3072", (1, 4096, 3072), (1, 1, 3072)),
Workload("flux2_bcast_s512_c3072", (1, 512, 3072), (1, 1, 3072)),
Workload("ltx2_full_s126_c2048", (1, 126, 2048), (1, 126, 2048)),
]
CI_WORKLOADS = [
Workload("ltx2_bcast_s1024_c4096", (1, 1024, 4096), (1, 1, 4096)),
Workload("ltx2_full_s512_c4096", (1, 512, 4096), (1, 512, 4096)),
]
def cuda_event_us(fn, warmups: int, repeats: int, rounds: int) -> float:
for _ in range(warmups):
fn()
torch.cuda.synchronize()
samples = []
for _ in range(rounds):
start = torch.cuda.Event(enable_timing=True)
end = torch.cuda.Event(enable_timing=True)
start.record()
for _ in range(repeats):
fn()
end.record()
end.synchronize()
samples.append(start.elapsed_time(end) * 1000.0 / repeats)
samples.sort()
return samples[len(samples) // 2]
def benchmark() -> None:
if not torch.cuda.is_available():
print("CUDA required")
return
torch.manual_seed(20260625)
random.seed(20260625)
torch.cuda.set_device(0)
workloads = CI_WORKLOADS if is_in_ci() else FULL_WORKLOADS
warmups = 5 if is_in_ci() else 20
repeats = 5 if is_in_ci() else 20
rounds = 5 if is_in_ci() else 13
print("| workload | gate | torch us | triton us | cuda us | cuda/triton |")
print("|---|---|---:|---:|---:|---:|")
for workload in workloads:
residual = torch.randn(
workload.residual_shape, device="cuda", dtype=torch.bfloat16
)
update = torch.randn_like(residual)
gate = torch.randn(workload.gate_shape, device="cuda", dtype=torch.bfloat16)
ref = residual + update * gate
triton_out = fuse_scale_shift_kernel(update, gate, residual, scale_constant=0)
cuda_out = residual_gate_add_cuda(residual, update, gate)
torch.cuda.synchronize()
torch.testing.assert_close(triton_out, ref, atol=5e-2, rtol=5e-2)
torch.testing.assert_close(cuda_out, ref, atol=5e-2, rtol=5e-2)
fns = {
"torch": lambda: residual + update * gate,
"triton": lambda: fuse_scale_shift_kernel(
update, gate, residual, scale_constant=0
),
"cuda": lambda: residual_gate_add_cuda(residual, update, gate),
}
order = ["torch", "triton", "cuda"]
random.shuffle(order)
times = {
name: cuda_event_us(fns[name], warmups, repeats, rounds) for name in order
}
gate_kind = (
"bcast" if workload.gate_shape != workload.residual_shape else "full"
)
print(
f"| {workload.name} | {gate_kind} | {times['torch']:.2f} | "
f"{times['triton']:.2f} | {times['cuda']:.2f} | "
f"{times['triton'] / times['cuda']:.3f}x |"
)
torch.cuda.empty_cache()
if __name__ == "__main__":
benchmark()
sys.exit(0)
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,186 @@
import torch
import torch.nn.functional as F
import triton
import triton.testing
from sglang.kernels.jit.benchmark.utils import get_benchmark_range, run_benchmark
from sglang.kernels.ops.embeddings.vocab_parallel_embedding import (
vocab_parallel_embedding,
)
from sglang.srt.layers.vocab_parallel_embedding import get_masked_input_and_mask
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
register_cuda_ci(
est_time=10, stage="base-b-kernel-benchmark", runner_config="1-gpu-large"
)
register_amd_ci(est_time=10, stage="jit-kernel-benchmark", runner_config="amd")
# Key order must match the perf_report x_names.
DEFAULTS = dict(
batch_size=120,
hidden_size=6144,
vocab_size=128256,
tp_size=8,
token_pattern="uniform",
dtype="bf16",
)
# One-at-a-time star sweep around DEFAULTS (the full product would be 1170
# configs). Each entry overrides one field (or one coupled field group).
SWEEPS = [
(
"batch_size",
get_benchmark_range(
[1, 2, 4, 8, 16, 32, 64, 120, 256, 512, 1024, 2048, 4096],
ci_range=[1, 120],
),
),
("hidden_size", get_benchmark_range([4096, 6144, 7168], ci_range=[])),
(
("vocab_size", "tp_size"),
get_benchmark_range(
[(32000, 4), (32000, 8), (128256, 4), (128256, 8), (154880, 8)],
ci_range=[],
),
),
(
"token_pattern",
get_benchmark_range(["uniform", "all_local", "all_remote"], ci_range=[]),
),
("dtype", get_benchmark_range(["bf16", "fp16"], ci_range=[])),
]
def _make_benchmark_configs():
# Dict keying dedupes the all-defaults config each sweep re-produces.
configs = {}
for keys, values in SWEEPS:
for value in values:
override = (
dict(zip(keys, value)) if isinstance(keys, tuple) else {keys: value}
)
config = {**DEFAULTS, **override}
configs[tuple(config.values())] = None
return list(configs)
BENCHMARK_CONFIGS = _make_benchmark_configs()
def _dtype_from_name(dtype: str) -> torch.dtype:
if dtype == "bf16":
return torch.bfloat16
if dtype == "fp16":
return torch.float16
raise ValueError(f"Unknown dtype: {dtype}")
def _make_input_ids(
batch_size: int, vocab_size: int, tp_size: int, token_pattern: str
) -> torch.Tensor:
assert vocab_size % tp_size == 0
per_partition = vocab_size // tp_size
if token_pattern == "uniform":
return torch.randint(
0, vocab_size, (batch_size,), dtype=torch.int64, device="cuda"
)
if token_pattern == "all_local":
return torch.randint(
0,
per_partition,
(batch_size,),
dtype=torch.int64,
device="cuda",
)
if token_pattern == "all_remote":
return torch.randint(
per_partition,
vocab_size,
(batch_size,),
dtype=torch.int64,
device="cuda",
)
raise ValueError(f"Unknown token_pattern: {token_pattern}")
def _torch_vocab_parallel_embedding(
input_ids: torch.Tensor, weight: torch.Tensor, vocab_size: int
):
masked_input, input_mask = get_masked_input_and_mask(
input_ids,
0,
weight.shape[0],
0,
vocab_size,
vocab_size,
)
output = F.embedding(masked_input.long(), weight)
output.masked_fill_(input_mask.unsqueeze(-1), 0)
return output
def _triton_vocab_parallel_embedding(
input_ids: torch.Tensor, weight: torch.Tensor, vocab_size: int
):
return vocab_parallel_embedding(
input_ids,
weight,
0,
weight.shape[0],
0,
vocab_size,
vocab_size,
)
@triton.testing.perf_report(
triton.testing.Benchmark(
x_names=[
"batch_size",
"hidden_size",
"vocab_size",
"tp_size",
"token_pattern",
"dtype",
],
x_vals=BENCHMARK_CONFIGS,
line_arg="provider",
line_vals=["triton", "torch"],
line_names=["Fused Triton", "Compiled mask + embedding + masked_fill"],
styles=[("blue", "-"), ("red", "--")],
ylabel="us",
plot_name="vocab-parallel-embedding-performance",
args={},
)
)
def benchmark(
batch_size: int,
hidden_size: int,
vocab_size: int,
tp_size: int,
token_pattern: str,
dtype: str,
provider: str,
):
assert vocab_size % tp_size == 0
torch_dtype = _dtype_from_name(dtype)
per_partition = vocab_size // tp_size
input_ids = _make_input_ids(batch_size, vocab_size, tp_size, token_pattern)
weight = torch.randn((per_partition, hidden_size), dtype=torch_dtype, device="cuda")
expected = _torch_vocab_parallel_embedding(input_ids, weight, vocab_size)
actual = _triton_vocab_parallel_embedding(input_ids, weight, vocab_size)
torch.testing.assert_close(actual, expected, rtol=0, atol=0)
if provider == "triton":
fn = lambda: _triton_vocab_parallel_embedding(input_ids, weight, vocab_size)
elif provider == "torch":
fn = lambda: _torch_vocab_parallel_embedding(input_ids, weight, vocab_size)
else:
raise ValueError(f"Unknown provider: {provider}")
return run_benchmark(fn)
if __name__ == "__main__":
benchmark.run(print_data=True)
@@ -0,0 +1,86 @@
"""Benchmark for DeepSeek V3 fused QKV-A GEMM: CuTe DSL vs CUDA JIT vs torch.
Run on SM90+ (Hopper or later):
python test/registered/jit/benchmark/bench_dsv3_fused_a_gemm.py
"""
import torch
import torch.nn.functional as F
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.test.ci.ci_register import register_cuda_ci
from sglang.utils import is_in_ci
register_cuda_ci(
est_time=12, stage="base-b-kernel-benchmark", runner_config="1-gpu-large"
)
IS_CI = is_in_ci()
DTYPE = torch.bfloat16
DEVICE = "cuda"
HD_OUT = 2112
HD_IN_LIST = [6144, 7168]
NUM_TOKENS_LIST = [1, 8, 16] if IS_CI else list(range(1, 17))
LINE_VALS = ["cutedsl", "jit", "torch"]
LINE_NAMES = ["CuTe DSL", "CUDA JIT", "torch F.linear"]
STYLES = [("blue", "-"), ("orange", "--"), ("green", "-.")]
def _median_us(fn, *args) -> float:
result = marker.do_bench(
fn,
input_args=args,
use_cuda_graph=True,
metrics=(0.5,),
disable_log_bandwidth=True,
)
return result.times[0] * 1e6
def _bench(num_tokens, provider, hd_in):
mat_a = torch.randn((num_tokens, hd_in), dtype=DTYPE, device=DEVICE)
mat_b = torch.randn((HD_OUT, hd_in), dtype=DTYPE, device=DEVICE).transpose(0, 1)
fn_map = {
"cutedsl": cutedsl_dsv3_fused_a_gemm,
"jit": dsv3_fused_a_gemm,
"torch": lambda a, b: F.linear(a, b.T),
}
return _median_us(fn_map[provider], mat_a, mat_b)
@triton.testing.perf_report(
[
triton.testing.Benchmark(
x_names=["num_tokens"],
x_vals=NUM_TOKENS_LIST,
line_arg="provider",
line_vals=LINE_VALS,
line_names=LINE_NAMES,
styles=STYLES,
ylabel="us",
plot_name=f"dsv3-fused-a-gemm-bf16-K{hd_in}-N{HD_OUT}",
args={"hd_in": hd_in},
)
for hd_in in HD_IN_LIST
]
)
def benchmark(num_tokens, provider, hd_in):
return _bench(num_tokens, provider, hd_in)
if __name__ == "__main__":
if is_hip_runtime() or get_jit_cuda_arch().major < 9:
print(
"dsv3_fused_a_gemm JIT kernel requires SM90+ (Hopper). Skipping benchmark."
)
else:
benchmark.run(print_data=True)
@@ -0,0 +1,53 @@
"""Benchmark for DeepSeek V3 router GEMM (JIT kernel vs torch).
Run on a Hopper (SM90+) GPU:
python -m sglang.kernels.jit.benchmark.bench_dsv3_router_gemm
"""
import torch
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.test.ci.ci_register import register_amd_ci, register_cuda_ci
register_cuda_ci(
est_time=5, stage="base-b-kernel-benchmark", runner_config="1-gpu-large"
)
register_amd_ci(est_time=5, stage="jit-kernel-benchmark", runner_config="amd")
def _torch(mat_a, mat_b, out_dtype):
return F.linear(mat_a, mat_b).to(out_dtype)
FN_MAP = {
"jit": dsv3_router_gemm,
"torch": _torch,
}
@marker.parametrize("num_experts", [256, 384], [256])
@marker.parametrize("hidden_dim", [6144, 7168], [7168])
@marker.parametrize("num_tokens", list(range(1, 17)), [1, 8, 16])
@marker.parametrize("out_dtype", [torch.bfloat16, torch.float32])
@marker.benchmark("provider", ["jit", "torch"])
def benchmark(num_experts, hidden_dim, num_tokens, out_dtype, provider):
mat_a = create_random(num_tokens, hidden_dim)
mat_b = create_random(num_experts, hidden_dim)
return marker.do_bench(
FN_MAP[provider],
input_args=(mat_a, mat_b),
input_kwargs={"out_dtype": out_dtype},
)
if __name__ == "__main__":
if is_hip_runtime() or get_jit_cuda_arch().major < 9:
print(
"dsv3_router_gemm JIT kernel requires SM90+ (Hopper). Skipping benchmark."
)
else:
benchmark.run()
@@ -0,0 +1,104 @@
from __future__ import annotations
import sys
import torch
import triton
from sglang.kernels.jit.benchmark.utils import get_benchmark_range, run_benchmark
from sglang.kernels.ops.gemm.fp8_blockwise_gemm import fp8_blockwise_scaled_mm
from sglang.srt.utils import is_sm120_supported
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
register_cuda_ci(
est_time=5,
stage="base-b-kernel-benchmark",
runner_config="1-gpu-large",
)
register_amd_ci(est_time=5, stage="jit-kernel-benchmark", runner_config="amd")
def _make_inputs(m: int, n: int, k: int, device: str = "cuda"):
fp8_info = torch.finfo(torch.float8_e4m3fn)
fp8_max, fp8_min = fp8_info.max, fp8_info.min
a_fp32 = (torch.rand(m, k, dtype=torch.float32, device=device) - 0.5) * 2 * fp8_max
a_fp8 = a_fp32.clamp(min=fp8_min, max=fp8_max).to(torch.float8_e4m3fn)
b_fp32 = (torch.rand(n, k, dtype=torch.float32, device=device) - 0.5) * 2 * fp8_max
b_fp8 = b_fp32.clamp(min=fp8_min, max=fp8_max).to(torch.float8_e4m3fn).t()
scale_a = torch.randn((m, k // 128), device=device, dtype=torch.float32) * 0.001
scale_b = (
torch.randn((k // 128, n // 128), device=device, dtype=torch.float32) * 0.001
)
scale_a = scale_a.t().contiguous().t()
scale_b = scale_b.t().contiguous().t()
return a_fp8, b_fp8, scale_a, scale_b
def _torch_ref(a_fp8, b_fp8, scale_a, scale_b):
def group_broadcast(t, shape):
for i, s in enumerate(shape):
if t.shape[i] != s and t.shape[i] != 1:
assert s % t.shape[i] == 0
t = (
t.unsqueeze(i + 1)
.expand(*t.shape[: i + 1], s // t.shape[i], *t.shape[i + 1 :])
.flatten(i, i + 1)
)
return t
sa = group_broadcast(scale_a, a_fp8.shape)
sb = group_broadcast(scale_b, b_fp8.shape)
return torch.mm(sa * a_fp8.to(torch.float32), sb * b_fp8.to(torch.float32)).to(
torch.bfloat16
)
shape_range = get_benchmark_range(
full_range=[
(16, 4096, 4096), # swapAB tile N=32
(64, 4096, 4096), # swapAB tile N=64
(128, 4096, 4096), # non-swap 128
(512, 4096, 4096),
(1024, 8192, 4096),
],
ci_range=[(16, 4096, 4096), (128, 4096, 4096)],
)
@triton.testing.perf_report(
triton.testing.Benchmark(
x_names=["m", "n", "k"],
x_vals=shape_range,
x_log=False,
line_arg="provider",
line_vals=["jit", "torch_ref"],
line_names=["JIT FP8 Blockwise GEMM", "Torch Ref"],
styles=[("green", "-"), ("blue", "-")],
ylabel="us",
plot_name="fp8-blockwise-scaled-mm-performance",
args={},
)
)
def benchmark(m, n, k, provider):
a_fp8, b_fp8, scale_a, scale_b = _make_inputs(m, n, k)
if provider == "jit":
fn = lambda: fp8_blockwise_scaled_mm(
a_fp8, b_fp8, scale_a, scale_b, out_dtype=torch.bfloat16
)
elif provider == "torch_ref":
fn = lambda: _torch_ref(a_fp8, b_fp8, scale_a, scale_b)
else:
raise ValueError(f"Unknown provider: {provider}")
return run_benchmark(fn)
if __name__ == "__main__":
if not is_sm120_supported():
print(
"[skip] fp8_blockwise_scaled_mm benchmark requires SM120 with CUDA 12.8+."
)
sys.exit(0)
benchmark.run(print_data=True)
@@ -0,0 +1,347 @@
from __future__ import annotations
from dataclasses import dataclass
from typing import Optional, Tuple
import torch
import triton
import triton.testing
from sglang.kernels.jit.benchmark.kv_canary.utils import (
POOL_AXIS,
SWA_WINDOW,
BenchCase,
build_fast_matrix_cases,
build_full_matrix_cases,
naive_cumsum_fn,
)
from sglang.kernels.jit.benchmark.utils import (
DEFAULT_DEVICE,
get_benchmark_range,
run_benchmark,
)
from sglang.kernels.ops.kv_canary.plan import launch_canary_plan_kernels
from sglang.kernels.ops.kv_canary.verify import VerifyPlan
from sglang.kernels.ops.kv_canary.write import WritePlan
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
register_cuda_ci(est_time=900, suite="nightly-kernel-1-gpu", nightly=True)
# AMD mirrors the CUDA nightly registration (nightly-only, no per-PR suite).
register_amd_ci(est_time=900, suite="nightly-amd-kernel-1-gpu", nightly=True)
_TOTAL_TOKENS_AXIS: list[int] = [256, 4096, 65536, 262144]
_TOTAL_TOKENS_BS_AXIS: list[int] = [1, 32, 256]
@dataclass(frozen=True, slots=True, kw_only=True)
class _TotalTokensBenchCase:
bs: int
total_tokens: int
pool_kind: str
def _build_total_tokens_cases() -> list[_TotalTokensBenchCase]:
cases: list[_TotalTokensBenchCase] = []
for bs in _TOTAL_TOKENS_BS_AXIS:
for total_tokens in _TOTAL_TOKENS_AXIS:
if total_tokens < bs:
continue
for pool_kind in POOL_AXIS:
cases.append(
_TotalTokensBenchCase(
bs=bs, total_tokens=total_tokens, pool_kind=pool_kind
)
)
return cases
_POOL_CAPACITY_VERIFY_CAP_AXIS: list[int] = [16384, 262144, 1398028]
_POOL_CAPACITY_BS_AXIS: list[int] = [1, 4, 32]
_POOL_CAPACITY_PREFIX_LEN: int = 512
_POOL_CAPACITY_BS_PADDED_AXIS: list[Optional[int]] = [None, 4096]
@dataclass(frozen=True, slots=True, kw_only=True)
class _PoolCapacityBenchCase:
"""One pool-capacity bench point.
Attributes:
bs: Number of active (non-padding) requests in the launch.
bs_padded: Total request-axis size of the input tensors. ``None`` means
no padding (``bs_padded == bs``); a concrete value pads ``req_pool_indices``
with ``REQ_POOL_IDX_PADDING`` sentinels in rows ``[bs, bs_padded)``.
prefix_len: Per-active-request prefix length.
verify_capacity: Plan tensor row capacity.
pool_kind: "full".
"""
bs: int
bs_padded: Optional[int]
prefix_len: int
verify_capacity: int
pool_kind: str
def _build_pool_capacity_cases() -> list[_PoolCapacityBenchCase]:
cases: list[_PoolCapacityBenchCase] = []
for bs in _POOL_CAPACITY_BS_AXIS:
for verify_capacity in _POOL_CAPACITY_VERIFY_CAP_AXIS:
for bs_padded in _POOL_CAPACITY_BS_PADDED_AXIS:
if bs_padded is not None and bs_padded < bs:
continue
cases.append(
_PoolCapacityBenchCase(
bs=bs,
bs_padded=bs_padded,
prefix_len=_POOL_CAPACITY_PREFIX_LEN,
verify_capacity=verify_capacity,
pool_kind="full",
)
)
return cases
_X_NAMES_MATRIX = ["scenario", "bs", "prefix_len", "mode", "extend_len", "pool_kind"]
def _cases_to_matrix_x_vals(
cases: list[BenchCase],
) -> list[tuple[str, int, int, str, int, str]]:
return [
(c.scenario, c.bs, c.prefix_len, c.mode, c.extend_len, c.pool_kind)
for c in cases
]
_X_VALS_MATRIX = _cases_to_matrix_x_vals(
get_benchmark_range(
full_range=build_full_matrix_cases(),
ci_range=build_fast_matrix_cases(),
)
)
_X_NAMES_TT = ["bs", "total_tokens", "pool_kind"]
_X_VALS_TT = [(c.bs, c.total_tokens, c.pool_kind) for c in _build_total_tokens_cases()]
_X_NAMES_PC = ["bs", "bs_padded", "prefix_len", "verify_capacity", "pool_kind"]
_X_VALS_PC = [
(
c.bs,
c.bs_padded if c.bs_padded is not None else c.bs,
c.prefix_len,
c.verify_capacity,
c.pool_kind,
)
for c in _build_pool_capacity_cases()
]
def _build_plan_inputs(
*,
bs: int,
prefix_len: int,
extend_len: int,
pool_kind: str,
device: torch.device,
verify_capacity_override: Optional[int] = None,
bs_padded: Optional[int] = None,
) -> dict:
swa_window_size = SWA_WINDOW if pool_kind == "swa_window_128" else 0
verify_per_req = min(prefix_len, SWA_WINDOW) if swa_window_size > 0 else prefix_len
if verify_capacity_override is not None:
verify_capacity = max(1, verify_capacity_override)
else:
verify_capacity = max(1, bs * verify_per_req)
effective_bs = bs_padded if bs_padded is not None else bs
if effective_bs < bs:
raise ValueError(f"kv-canary bench: bs_padded={bs_padded} must be >= bs={bs}")
write_req_capacity = max(1, effective_bs)
verify_plan = VerifyPlan.allocate(verify_capacity=verify_capacity, device=device)
write_plan = WritePlan.allocate(
write_req_capacity=write_req_capacity, device=device
)
req_pool_indices = torch.zeros(effective_bs, dtype=torch.int64, device=device)
req_pool_indices[:bs] = torch.arange(1, bs + 1, dtype=torch.int64, device=device)
prefix_lens = torch.zeros(effective_bs, dtype=torch.int64, device=device)
prefix_lens[:bs] = prefix_len
extend_seq_lens = torch.zeros(effective_bs, dtype=torch.int64, device=device)
extend_seq_lens[:bs] = extend_len
max_seq_len = max(prefix_len + extend_len, 1)
req_to_token_rows = effective_bs + 1
req_to_token = torch.zeros(
req_to_token_rows,
max_seq_len,
dtype=torch.int32,
device=device,
)
if bs > 0:
row_idx = torch.arange(1, bs + 1, dtype=torch.int32, device=device).unsqueeze(1)
col_idx = torch.arange(max_seq_len, dtype=torch.int32, device=device).unsqueeze(
0
)
req_to_token[1 : bs + 1] = (row_idx - 1) * max_seq_len + col_idx
if swa_window_size > 0:
full_pool_size = effective_bs * max_seq_len + 1
full_to_swa: Optional[torch.Tensor] = torch.arange(
full_pool_size + 1, dtype=torch.int64, device=device
)
full_to_swa[-1] = -1
else:
full_to_swa = None
return dict(
verify_plan_out=verify_plan,
write_plan_out=write_plan,
req_pool_indices=req_pool_indices,
prefix_lens=prefix_lens,
extend_seq_lens=extend_seq_lens,
req_to_token=req_to_token,
swa_window_size=swa_window_size,
full_to_swa_index_mapping=full_to_swa,
verify_capacity=int(verify_plan.verify_slot_indices.shape[0]),
)
def _make_plan_callable(inputs: dict):
def fn() -> None:
launch_canary_plan_kernels(
verify_plan_out=inputs["verify_plan_out"],
write_plan_out=inputs["write_plan_out"],
req_pool_indices=inputs["req_pool_indices"],
prefix_lens=inputs["prefix_lens"],
extend_seq_lens=inputs["extend_seq_lens"],
req_to_token=inputs["req_to_token"],
swa_window_size=inputs["swa_window_size"],
full_to_swa_index_mapping=inputs["full_to_swa_index_mapping"],
verify_capacity=inputs["verify_capacity"],
req_to_verify_expected_tokens=None,
req_to_verify_expected_tokens_valid_lens=None,
kv_token_id_vs_position_offset=0,
)
return fn
@triton.testing.perf_report(
triton.testing.Benchmark(
x_names=_X_NAMES_MATRIX,
x_vals=_X_VALS_MATRIX,
line_arg="provider",
line_vals=["canary", "naive"],
line_names=["canary_plan_step", "naive torch.cumsum"],
styles=[("blue", "-"), ("red", "--")],
ylabel="us",
plot_name="kv-canary-plan-matrix-perf",
args={},
)
)
def benchmark_matrix(
scenario: str,
bs: int,
prefix_len: int,
mode: str,
extend_len: int,
pool_kind: str,
provider: str,
) -> Tuple[float, float, float]:
del scenario
del mode
device = torch.device(DEFAULT_DEVICE)
if provider == "canary":
inputs = _build_plan_inputs(
bs=bs,
prefix_len=prefix_len,
extend_len=extend_len,
pool_kind=pool_kind,
device=device,
)
fn = _make_plan_callable(inputs)
else:
fn = naive_cumsum_fn(bs=bs, device=device)
return run_benchmark(fn)
@triton.testing.perf_report(
triton.testing.Benchmark(
x_names=_X_NAMES_TT,
x_vals=_X_VALS_TT,
line_arg="provider",
line_vals=["canary", "naive"],
line_names=["canary_plan_step", "naive torch.cumsum"],
styles=[("blue", "-"), ("red", "--")],
ylabel="us",
plot_name="kv-canary-plan-total-tokens-perf",
args={},
)
)
def benchmark_total_tokens(
bs: int,
total_tokens: int,
pool_kind: str,
provider: str,
) -> Tuple[float, float, float]:
device = torch.device(DEFAULT_DEVICE)
per_req_prefix = max(1, total_tokens // max(bs, 1))
if provider == "canary":
inputs = _build_plan_inputs(
bs=bs,
prefix_len=per_req_prefix,
extend_len=1,
pool_kind=pool_kind,
device=device,
)
fn = _make_plan_callable(inputs)
else:
fn = naive_cumsum_fn(bs=bs, device=device)
return run_benchmark(fn)
@triton.testing.perf_report(
triton.testing.Benchmark(
x_names=_X_NAMES_PC,
x_vals=_X_VALS_PC,
line_arg="provider",
line_vals=["canary", "naive"],
line_names=["canary_plan_step", "naive torch.cumsum"],
styles=[("blue", "-"), ("red", "--")],
ylabel="us",
plot_name="kv-canary-plan-pool-capacity-perf",
args={},
)
)
def benchmark_pool_capacity(
bs: int,
bs_padded: int,
prefix_len: int,
verify_capacity: int,
pool_kind: str,
provider: str,
) -> Tuple[float, float, float]:
device = torch.device(DEFAULT_DEVICE)
if provider == "canary":
inputs = _build_plan_inputs(
bs=bs,
prefix_len=prefix_len,
extend_len=1,
pool_kind=pool_kind,
device=device,
verify_capacity_override=verify_capacity,
bs_padded=bs_padded,
)
fn = _make_plan_callable(inputs)
else:
fn = naive_cumsum_fn(bs=bs, device=device)
return run_benchmark(fn)
if __name__ == "__main__":
benchmark_matrix.run(print_data=True)
benchmark_total_tokens.run(print_data=True)
benchmark_pool_capacity.run(print_data=True)
@@ -0,0 +1,99 @@
from __future__ import annotations
from dataclasses import dataclass
import torch
import triton
import triton.testing
from sglang.kernels.jit.benchmark.utils import (
DEFAULT_DEVICE,
get_benchmark_range,
run_benchmark_no_cudagraph,
)
from sglang.kernels.ops.kv_canary.scatter_req_token_ids import (
launch_scatter_req_token_ids_kernel,
)
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
register_cuda_ci(est_time=180, suite="nightly-kernel-1-gpu", nightly=True)
# AMD mirrors the CUDA nightly registration (nightly-only, no per-PR suite).
# Note: amd_ci_exec.sh sets SGLANG_IS_IN_CI, so this runs the CI-reduced range
# (_BS_AXIS_CI/_SEQ_LEN_AXIS_CI via get_benchmark_range), same as CUDA nightly.
register_amd_ci(est_time=180, suite="nightly-amd-kernel-1-gpu", nightly=True)
_BS_AXIS_FULL: list[int] = [1, 8, 64, 256]
_SEQ_LEN_AXIS_FULL: list[int] = [128, 512, 2048, 8192]
_BS_AXIS_CI: list[int] = [1, 64]
_SEQ_LEN_AXIS_CI: list[int] = [512, 2048]
@dataclass(frozen=True, slots=True, kw_only=True)
class _BenchCase:
bs: int
seq_len: int
def _build_cases() -> list[_BenchCase]:
bs_axis = get_benchmark_range(full_range=_BS_AXIS_FULL, ci_range=_BS_AXIS_CI)
seq_axis = get_benchmark_range(
full_range=_SEQ_LEN_AXIS_FULL, ci_range=_SEQ_LEN_AXIS_CI
)
return [
_BenchCase(bs=bs, seq_len=seq_len) for bs in bs_axis for seq_len in seq_axis
]
_X_NAMES = ["bs", "seq_len"]
_X_VALS = [(c.bs, c.seq_len) for c in _build_cases()]
def _build_inputs(*, bs: int, seq_len: int, device: torch.device) -> dict:
max_reqs = max(bs + 1, 4)
max_context_len = max(seq_len + 1, 1)
total_tokens = bs * seq_len
flat = torch.randint(
low=0,
high=1 << 30,
size=(total_tokens,),
dtype=torch.int64,
device=device,
)
lens = torch.full((bs,), seq_len, dtype=torch.int64, device=device)
offsets = torch.zeros(bs + 1, dtype=torch.int64, device=device)
offsets[1:] = torch.cumsum(lens, dim=0)
req_pool_indices = torch.arange(1, bs + 1, dtype=torch.int64, device=device)
pool = torch.zeros((max_reqs, max_context_len), dtype=torch.int32, device=device)
return dict(
flat_in=flat,
offsets=offsets,
req_pool_indices=req_pool_indices,
pool_out=pool,
)
@triton.testing.perf_report(
triton.testing.Benchmark(
x_names=_X_NAMES,
x_vals=_X_VALS,
line_arg="provider",
line_vals=["triton"],
line_names=["Triton"],
styles=[("blue", "-")],
ylabel="time (us)",
plot_name="kv-canary-scatter-req-token-ids",
args={},
)
)
def benchmark(bs: int, seq_len: int, provider: str) -> tuple[float, float, float]:
inputs = _build_inputs(bs=bs, seq_len=seq_len, device=torch.device(DEFAULT_DEVICE))
return run_benchmark_no_cudagraph(
lambda: launch_scatter_req_token_ids_kernel(**inputs)
)
if __name__ == "__main__":
benchmark.run(print_data=True)
@@ -0,0 +1,318 @@
from __future__ import annotations
from typing import Tuple
import torch
import triton
import triton.testing
from sglang.kernels.jit.benchmark.kv_canary.utils import (
RING_CAPACITY,
SWA_WINDOW,
BenchCase,
build_fast_matrix_cases,
build_full_matrix_cases,
cases_to_x_vals,
make_real_kv_sources,
naive_slot_copy_fn,
)
from sglang.kernels.jit.benchmark.utils import (
DEFAULT_DEVICE,
get_benchmark_range,
run_benchmark,
)
from sglang.kernels.ops.kv_canary import consts
from sglang.kernels.ops.kv_canary.verify import (
CANARY_SLOT_BYTES,
CanaryLaunchTag,
RealKvSource,
VerifyOrWriteContext,
VerifyPlan,
launch_canary_verify_kernel,
)
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
register_cuda_ci(est_time=900, suite="nightly-kernel-1-gpu", nightly=True)
# AMD mirrors the CUDA nightly registration (nightly-only, no per-PR suite).
# Note: amd_ci_exec.sh sets SGLANG_IS_IN_CI, so this runs the CI-reduced range
# (build_fast_matrix_cases via get_benchmark_range), same as CUDA nightly.
register_amd_ci(est_time=900, suite="nightly-amd-kernel-1-gpu", nightly=True)
_X_NAMES = [
"scenario",
"bs",
"prefix_len",
"mode",
"extend_len",
"pool_kind",
"real_kv_kind",
"hash_mode",
]
_X_VALS = cases_to_x_vals(
get_benchmark_range(
full_range=build_full_matrix_cases(),
ci_range=build_fast_matrix_cases(),
)
)
_KERNEL_KIND_X_NAMES = ["kernel_kind_name"]
_KERNEL_KIND_X_VALS = [(tag.name,) for tag in CanaryLaunchTag]
def _verify_entry_count(case: BenchCase) -> int:
if case.pool_kind == "swa_window_128":
per_req = min(case.prefix_len, SWA_WINDOW)
else:
per_req = case.prefix_len
return case.bs * per_req
def _verify_num_slots(case: BenchCase) -> int:
if case.pool_kind == "swa_window_128":
per_req_slots = SWA_WINDOW
else:
per_req_slots = max(1, case.prefix_len)
return max(2, case.bs * per_req_slots + 1)
def _build_verify_inputs(case: BenchCase, *, device: torch.device) -> Tuple[
torch.Tensor,
VerifyPlan,
torch.Tensor,
torch.Tensor,
torch.Tensor,
torch.Tensor,
torch.Tensor,
tuple[RealKvSource, ...],
]:
total_entries = _verify_entry_count(case)
capacity = max(1, total_entries)
num_slots = _verify_num_slots(case)
canary_buf = torch.zeros(
num_slots, CANARY_SLOT_BYTES, dtype=torch.uint8, device=device
)
slot_indices = torch.empty(capacity, dtype=torch.int64, device=device)
positions = torch.empty(capacity, dtype=torch.int64, device=device)
prev_slots = torch.empty(capacity, dtype=torch.int64, device=device)
if total_entries > 0:
flat_idx = torch.arange(total_entries, device=device, dtype=torch.int64)
per_req = total_entries // case.bs if case.bs > 0 else 0
slot_indices[:total_entries] = (flat_idx % max(num_slots - 1, 1)).to(
torch.int64
)
positions[:total_entries] = (flat_idx % max(per_req, 1)).to(torch.int64)
is_head = (flat_idx % max(per_req, 1)) == 0
prev_seq = (flat_idx - 1) % max(num_slots - 1, 1)
prev_slots[:total_entries] = torch.where(
is_head, torch.full_like(flat_idx, -1), prev_seq
).to(torch.int64)
if capacity > total_entries:
slot_indices[total_entries:] = 0
positions[total_entries:] = 0
prev_slots[total_entries:] = -1
num_valid = torch.tensor([total_entries], dtype=torch.int32, device=device)
enable = torch.ones(1, dtype=torch.int32, device=device)
expected_input_ids = torch.full((capacity,), -1, dtype=torch.int64, device=device)
plan = VerifyPlan(
verify_slot_indices=slot_indices,
verify_expected_tokens=expected_input_ids,
verify_expected_positions=positions,
verify_prev_slot_indices=prev_slots,
verify_num_valid=num_valid,
enable=enable,
)
violation_ring = torch.zeros(
RING_CAPACITY, consts.VIOLATION_FIELDS, dtype=torch.int64, device=device
)
violation_write_index = torch.zeros(1, dtype=torch.int32, device=device)
slot_run_counter = torch.zeros(1, dtype=torch.int64, device=device)
kernel_run_counter = torch.zeros(1, dtype=torch.int64, device=device)
enable_chain_position_assert = torch.ones(1, dtype=torch.int32, device=device)
real_kv_sources = make_real_kv_sources(
kind=case.real_kv_kind, num_slots=num_slots, device=device
)
return (
canary_buf,
plan,
violation_ring,
violation_write_index,
slot_run_counter,
kernel_run_counter,
enable_chain_position_assert,
real_kv_sources,
)
def _build_context(
*,
canary_buf: torch.Tensor,
violation_ring: torch.Tensor,
violation_write_index: torch.Tensor,
slot_run_counter: torch.Tensor,
kernel_run_counter: torch.Tensor,
enable_chain_position_assert: torch.Tensor,
real_kv_sources: tuple[RealKvSource, ...],
kernel_kind: CanaryLaunchTag,
hash_mode: consts.RealKvHashMode,
) -> VerifyOrWriteContext:
return VerifyOrWriteContext(
canary_buf=canary_buf,
kernel_kind=kernel_kind,
violation_ring=violation_ring,
violation_write_index=violation_write_index,
slot_run_counter=slot_run_counter,
kernel_run_counter=kernel_run_counter,
real_kv_sources=real_kv_sources,
real_kv_hash_mode=hash_mode,
enable_chain_position_assert=enable_chain_position_assert,
)
@triton.testing.perf_report(
triton.testing.Benchmark(
x_names=_X_NAMES,
x_vals=_X_VALS,
line_arg="provider",
line_vals=["canary", "naive"],
line_names=["canary_verify_step", "naive index_copy_"],
styles=[("blue", "-"), ("red", "--")],
ylabel="us",
plot_name="kv-canary-verify-perf",
args={},
)
)
def benchmark(
scenario: str,
bs: int,
prefix_len: int,
mode: str,
extend_len: int,
pool_kind: str,
real_kv_kind: str,
hash_mode: str,
provider: str,
) -> Tuple[float, float, float]:
case = BenchCase(
scenario=scenario,
bs=bs,
prefix_len=prefix_len,
mode=mode,
extend_len=extend_len,
pool_kind=pool_kind,
real_kv_kind=real_kv_kind,
hash_mode=hash_mode,
)
device = torch.device(DEFAULT_DEVICE)
if provider == "canary":
(
canary_buf,
plan,
violation_ring,
violation_write_index,
slot_run_counter,
kernel_run_counter,
enable_chain_position_assert,
real_kv_sources,
) = _build_verify_inputs(case, device=device)
hash_mode_enum = consts.RealKvHashMode[case.hash_mode.upper()]
context = _build_context(
canary_buf=canary_buf,
violation_ring=violation_ring,
violation_write_index=violation_write_index,
slot_run_counter=slot_run_counter,
kernel_run_counter=kernel_run_counter,
enable_chain_position_assert=enable_chain_position_assert,
real_kv_sources=real_kv_sources,
kernel_kind=CanaryLaunchTag.HEAD_K_FULL,
hash_mode=hash_mode_enum,
)
def fn() -> None:
violation_write_index.zero_()
launch_canary_verify_kernel(
context=context,
plan=plan,
check_verify_expected_token=True,
)
else:
fn = naive_slot_copy_fn(total=_verify_entry_count(case), device=device)
return run_benchmark(fn)
@triton.testing.perf_report(
triton.testing.Benchmark(
x_names=_KERNEL_KIND_X_NAMES,
x_vals=_KERNEL_KIND_X_VALS,
line_arg="provider",
line_vals=["canary"],
line_names=["canary_verify_step"],
styles=[("blue", "-")],
ylabel="us",
plot_name="kv-canary-verify-kernel-kind-perf",
args={},
)
)
def benchmark_kernel_kind(
kernel_kind_name: str,
provider: str,
) -> Tuple[float, float, float]:
case = BenchCase(
scenario="kernel_kind",
bs=32,
prefix_len=4096,
mode="extend",
extend_len=128,
pool_kind="full",
real_kv_kind="none",
hash_mode="none",
)
device = torch.device(DEFAULT_DEVICE)
(
canary_buf,
plan,
violation_ring,
violation_write_index,
slot_run_counter,
kernel_run_counter,
enable_chain_position_assert,
real_kv_sources,
) = _build_verify_inputs(case, device=device)
kernel_kind = CanaryLaunchTag[kernel_kind_name]
hash_mode_enum = consts.RealKvHashMode[case.hash_mode.upper()]
context = _build_context(
canary_buf=canary_buf,
violation_ring=violation_ring,
violation_write_index=violation_write_index,
slot_run_counter=slot_run_counter,
kernel_run_counter=kernel_run_counter,
enable_chain_position_assert=enable_chain_position_assert,
real_kv_sources=real_kv_sources,
kernel_kind=kernel_kind,
hash_mode=hash_mode_enum,
)
def fn() -> None:
violation_write_index.zero_()
launch_canary_verify_kernel(
context=context,
plan=plan,
check_verify_expected_token=True,
)
return run_benchmark(fn)
if __name__ == "__main__":
benchmark.run(print_data=True)
benchmark_kernel_kind.run(print_data=True)
@@ -0,0 +1,314 @@
from __future__ import annotations
from typing import Tuple
import torch
import triton
import triton.testing
from sglang.kernels.jit.benchmark.kv_canary.utils import (
RING_CAPACITY,
SWA_WINDOW,
BenchCase,
build_fast_matrix_cases,
build_full_matrix_cases,
cases_to_x_vals,
make_real_kv_sources,
naive_slot_copy_fn,
)
from sglang.kernels.jit.benchmark.utils import (
DEFAULT_DEVICE,
get_benchmark_range,
run_benchmark,
)
from sglang.kernels.ops.kv_canary import consts
from sglang.kernels.ops.kv_canary.verify import (
CANARY_SLOT_BYTES,
CanaryLaunchTag,
VerifyOrWriteContext,
)
from sglang.kernels.ops.kv_canary.write import WritePlan, launch_canary_write_kernel
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
register_cuda_ci(est_time=900, suite="nightly-kernel-1-gpu", nightly=True)
# AMD mirrors the CUDA nightly registration (nightly-only, no per-PR suite).
# Note: amd_ci_exec.sh sets SGLANG_IS_IN_CI, so this runs the CI-reduced range
# (build_fast_matrix_cases via get_benchmark_range), same as CUDA nightly.
register_amd_ci(est_time=900, suite="nightly-amd-kernel-1-gpu", nightly=True)
_X_NAMES = [
"scenario",
"bs",
"prefix_len",
"mode",
"extend_len",
"pool_kind",
"real_kv_kind",
"hash_mode",
]
_X_VALS = cases_to_x_vals(
get_benchmark_range(
full_range=build_full_matrix_cases(),
ci_range=build_fast_matrix_cases(),
)
)
_KERNEL_KIND_X_NAMES = ["kernel_kind_name", "enable_write_verify_inputs_name"]
_KERNEL_KIND_X_VALS = [
(tag.name, str(enable)) for tag in CanaryLaunchTag for enable in (False, True)
]
def _write_entry_count(case: BenchCase) -> int:
return case.bs * case.extend_len
def _write_num_slots(case: BenchCase) -> int:
per_req_slots = max(
SWA_WINDOW if case.pool_kind == "swa_window_128" else 1,
case.prefix_len + case.extend_len,
)
return max(2, case.bs * per_req_slots + 1)
def _build_write_inputs(
case: BenchCase, *, device: torch.device, mirror_expected_inputs: bool = False
) -> dict:
total_entries = _write_entry_count(case)
num_tokens_padded = max(1, total_entries)
per_req_slots = max(
SWA_WINDOW if case.pool_kind == "swa_window_128" else 1,
case.prefix_len + case.extend_len,
)
num_slots = _write_num_slots(case)
canary_buf = torch.zeros(
num_slots, CANARY_SLOT_BYTES, dtype=torch.uint8, device=device
)
write_offsets = torch.zeros(case.bs + 1, dtype=torch.int64, device=device)
if case.bs > 0:
offsets_host = torch.arange(0, case.bs + 1, dtype=torch.int64) * case.extend_len
write_offsets.copy_(offsets_host.to(device))
write_seed_slots = torch.empty(case.bs, dtype=torch.int64, device=device)
if case.bs > 0:
if case.prefix_len == 0:
write_seed_slots.fill_(-1)
else:
per_req_stride = per_req_slots
seeds = (
torch.arange(case.bs, dtype=torch.int32, device=device) * per_req_stride
+ case.prefix_len
- 1
)
write_seed_slots.copy_(seeds.to(torch.int64))
write_num_valid_reqs = torch.tensor([case.bs], dtype=torch.int32, device=device)
plan = WritePlan(
write_offsets=write_offsets,
write_seed_slot_indices=write_seed_slots,
write_num_valid_reqs=write_num_valid_reqs,
)
input_ids = torch.zeros(num_tokens_padded, dtype=torch.int64, device=device)
positions = torch.zeros(num_tokens_padded, dtype=torch.int64, device=device)
out_cache_loc = torch.zeros(num_tokens_padded, dtype=torch.int64, device=device)
if total_entries > 0:
flat_idx = torch.arange(total_entries, device=device, dtype=torch.int64)
per_req_idx = flat_idx % max(case.extend_len, 1)
req_idx = flat_idx // max(case.extend_len, 1)
per_req_stride = per_req_slots
slots = (req_idx * per_req_stride + case.prefix_len + per_req_idx) % max(
num_slots, 1
)
input_ids[:total_entries] = (flat_idx % 32768).to(torch.int64)
positions[:total_entries] = (case.prefix_len + per_req_idx).to(torch.int64)
out_cache_loc[:total_entries] = slots.to(torch.int64)
if case.pool_kind == "swa_window_128":
full_to_swa = torch.arange(num_slots + 1, dtype=torch.int64, device=device)
full_to_swa[-1] = -1
out_cache_loc = full_to_swa[out_cache_loc]
if mirror_expected_inputs:
expected_input_tokens = input_ids.clone()
expected_input_positions = positions.clone()
else:
expected_input_tokens = None
expected_input_positions = None
violation_ring = torch.zeros(
RING_CAPACITY, consts.VIOLATION_FIELDS, dtype=torch.int64, device=device
)
violation_write_index = torch.zeros(1, dtype=torch.int32, device=device)
slot_run_counter = torch.zeros(1, dtype=torch.int64, device=device)
kernel_run_counter = torch.zeros(1, dtype=torch.int64, device=device)
enable_chain_position_assert = torch.ones(1, dtype=torch.int32, device=device)
real_kv_sources = make_real_kv_sources(
kind=case.real_kv_kind, num_slots=num_slots, device=device
)
return dict(
canary_buf=canary_buf,
plan=plan,
input_ids=input_ids,
positions=positions,
out_cache_loc=out_cache_loc,
expected_input_tokens=expected_input_tokens,
expected_input_positions=expected_input_positions,
violation_ring=violation_ring,
violation_write_index=violation_write_index,
slot_run_counter=slot_run_counter,
kernel_run_counter=kernel_run_counter,
enable_chain_position_assert=enable_chain_position_assert,
real_kv_sources=real_kv_sources,
)
def _build_context(
*,
inputs: dict,
kernel_kind: CanaryLaunchTag,
hash_mode: consts.RealKvHashMode,
) -> VerifyOrWriteContext:
return VerifyOrWriteContext(
canary_buf=inputs["canary_buf"],
kernel_kind=kernel_kind,
violation_ring=inputs["violation_ring"],
violation_write_index=inputs["violation_write_index"],
slot_run_counter=inputs["slot_run_counter"],
kernel_run_counter=inputs["kernel_run_counter"],
real_kv_sources=inputs["real_kv_sources"],
real_kv_hash_mode=hash_mode,
enable_chain_position_assert=inputs["enable_chain_position_assert"],
)
@triton.testing.perf_report(
triton.testing.Benchmark(
x_names=_X_NAMES,
x_vals=_X_VALS,
line_arg="provider",
line_vals=["canary", "naive"],
line_names=["canary_write_step", "naive index_copy_"],
styles=[("blue", "-"), ("red", "--")],
ylabel="us",
plot_name="kv-canary-write-perf",
args={},
)
)
def benchmark(
scenario: str,
bs: int,
prefix_len: int,
mode: str,
extend_len: int,
pool_kind: str,
real_kv_kind: str,
hash_mode: str,
provider: str,
) -> Tuple[float, float, float]:
case = BenchCase(
scenario=scenario,
bs=bs,
prefix_len=prefix_len,
mode=mode,
extend_len=extend_len,
pool_kind=pool_kind,
real_kv_kind=real_kv_kind,
hash_mode=hash_mode,
)
device = torch.device(DEFAULT_DEVICE)
if provider == "canary":
inputs = _build_write_inputs(case, device=device)
hash_mode_enum = consts.RealKvHashMode[case.hash_mode.upper()]
context = _build_context(
inputs=inputs,
kernel_kind=CanaryLaunchTag.HEAD_K_FULL,
hash_mode=hash_mode_enum,
)
def fn() -> None:
launch_canary_write_kernel(
context=context,
plan=inputs["plan"],
input_ids=inputs["input_ids"],
positions=inputs["positions"],
out_cache_loc=inputs["out_cache_loc"],
enable_write_input_assert=False,
expected_input_tokens=inputs["expected_input_tokens"],
expected_input_positions=inputs["expected_input_positions"],
)
else:
fn = naive_slot_copy_fn(total=_write_entry_count(case), device=device)
return run_benchmark(fn)
@triton.testing.perf_report(
triton.testing.Benchmark(
x_names=_KERNEL_KIND_X_NAMES,
x_vals=_KERNEL_KIND_X_VALS,
line_arg="provider",
line_vals=["canary"],
line_names=["canary_write_step"],
styles=[("blue", "-")],
ylabel="us",
plot_name="kv-canary-write-kernel-kind-perf",
args={},
)
)
def benchmark_kernel_kind(
kernel_kind_name: str,
enable_write_verify_inputs_name: str,
provider: str,
) -> Tuple[float, float, float]:
case = BenchCase(
scenario="kernel_kind",
bs=32,
prefix_len=4096,
mode="extend",
extend_len=128,
pool_kind="full",
real_kv_kind="none",
hash_mode="none",
)
device = torch.device(DEFAULT_DEVICE)
enable_write_verify_inputs = enable_write_verify_inputs_name == "True"
inputs = _build_write_inputs(
case, device=device, mirror_expected_inputs=enable_write_verify_inputs
)
kernel_kind = CanaryLaunchTag[kernel_kind_name]
hash_mode_enum = consts.RealKvHashMode[case.hash_mode.upper()]
context = _build_context(
inputs=inputs,
kernel_kind=kernel_kind,
hash_mode=hash_mode_enum,
)
def fn() -> None:
launch_canary_write_kernel(
context=context,
plan=inputs["plan"],
input_ids=inputs["input_ids"],
positions=inputs["positions"],
out_cache_loc=inputs["out_cache_loc"],
enable_write_input_assert=enable_write_verify_inputs,
expected_input_tokens=inputs["expected_input_tokens"],
expected_input_positions=inputs["expected_input_positions"],
)
return run_benchmark(fn)
if __name__ == "__main__":
benchmark.run(print_data=True)
benchmark_kernel_kind.run(print_data=True)
@@ -0,0 +1,53 @@
import torch
from sglang.kernels.jit.benchmark import marker
from sglang.kernels.ops.kvcache.fused_fp8_qkv_kv_cache import fused_fp8_qkv_kv_cache
from sglang.test.ci.ci_register import register_cuda_ci
register_cuda_ci(
est_time=6, stage="base-b-kernel-benchmark", runner_config="1-gpu-large"
)
FP8 = torch.float8_e4m3fn
D = 128
def fused_qkv(q, k, v, k_cache, v_cache, cache_loc, k_scale, v_scale):
return fused_fp8_qkv_kv_cache(
q, k, v, k_cache, v_cache, cache_loc, k_scale, v_scale
)
def fused_kv_only(q, k, v, k_cache, v_cache, cache_loc, k_scale, v_scale):
fused_fp8_qkv_kv_cache(None, k, v, k_cache, v_cache, cache_loc, k_scale, v_scale)
return q.to(FP8)
FN_MAP = {"fused_qkv": fused_qkv, "fused_kv_only": fused_kv_only}
@marker.parametrize("num_tokens", [8, 128, 2048, 4096, 8192, 16384], [8, 2048])
@marker.parametrize("hq,hkv", [(64, 2), (16, 1), (8, 1)])
@marker.benchmark("impl", ["fused_qkv", "fused_kv_only"])
def benchmark(num_tokens: int, hq: int, hkv: int, impl: str):
qd, kvd = hq * D, hkv * D
qkv = torch.randn(num_tokens, qd + 2 * kvd, dtype=torch.bfloat16, device="cuda")
q = qkv[:, :qd]
k = qkv[:, qd : qd + kvd].view(num_tokens, hkv, D)
v = qkv[:, qd + kvd :].view(num_tokens, hkv, D)
slots = num_tokens + 16
k_cache = torch.zeros(slots, hkv, D, dtype=FP8, device="cuda")
v_cache = torch.zeros(slots, hkv, D, dtype=FP8, device="cuda")
cache_loc = torch.arange(num_tokens, dtype=torch.int64, device="cuda")
k_scale = torch.tensor(0.5, dtype=torch.float32, device="cuda")
v_scale = torch.tensor(0.7, dtype=torch.float32, device="cuda")
return marker.do_bench(
FN_MAP[impl],
input_args=(q, k, v, k_cache, v_cache, cache_loc, k_scale, v_scale),
graph_clone_args=(0,),
memory_output=(k_cache, v_cache),
)
if __name__ == "__main__":
benchmark.run()
@@ -0,0 +1,424 @@
"""Benchmark for HiCache JIT kernel performance.
This benchmark tests the performance of KV cache transfer operations
between GPU and CPU (host pinned memory), comparing:
- SGL AOT Kernel: Pre-compiled transfer_kv kernels from sgl_kernel
- SGL JIT Kernel: JIT-compiled hicache kernels
- PyTorch Indexing: Plain PyTorch index copy
- PyTorch 2 Stream: PyTorch implementation using 2 CUDA streams
Tests cover:
- One Layer: CPU->GPU
- All Layer: GPU->CPU
Note: Uses do_bench instead of do_bench_cudagraph since CUDA graph
capture doesn't support CPU-GPU memory transfers.
"""
import itertools
import os
from dataclasses import dataclass
from typing import Tuple
import torch
import triton
import triton.testing
from sgl_kernel import transfer_kv_all_layer, transfer_kv_per_layer
from sglang.kernels.jit.benchmark.utils import DEFAULT_QUANTILES, get_benchmark_range
from sglang.kernels.ops.kvcache.hicache import (
can_use_hicache_jit_kernel,
transfer_hicache_all_layer,
transfer_hicache_one_layer,
)
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
register_cuda_ci(
est_time=29, stage="base-b-kernel-benchmark", runner_config="1-gpu-large"
)
register_amd_ci(est_time=29, stage="jit-kernel-benchmark", runner_config="amd")
DISABLE_TORCH = os.environ.get("DISABLE_TORCH", "0") == "1"
PAGE_SIZE = 1
ENABLE_SORT = True
GPU_CACHE_SIZE = 256 * 1024 # 256K tokens on GPU
HOST_CACHE_SIZE = 512 * 1024 # 512K tokens on CPU
NUM_LAYERS = 8
@dataclass(frozen=True)
class HiCacheCache:
k_cache_cuda: torch.Tensor
v_cache_cuda: torch.Tensor
k_cache_host: torch.Tensor
v_cache_host: torch.Tensor
def get_slice(self, num_layers: int, element_size: int) -> "HiCacheCache":
def slice_cuda(t: torch.Tensor) -> torch.Tensor:
needed_cuda = num_layers * GPU_CACHE_SIZE
return t.view(-1, element_size)[:needed_cuda].unflatten(0, (num_layers, -1))
def slice_host(t: torch.Tensor) -> torch.Tensor:
needed_host = num_layers * HOST_CACHE_SIZE
return t.view(-1, element_size)[:needed_host].unflatten(0, (num_layers, -1))
return HiCacheCache(
k_cache_cuda=slice_cuda(self.k_cache_cuda),
v_cache_cuda=slice_cuda(self.v_cache_cuda),
k_cache_host=slice_host(self.k_cache_host),
v_cache_host=slice_host(self.v_cache_host),
)
def gen_indices(
size: int, max_size: int, *, page_size: int = PAGE_SIZE
) -> torch.Tensor:
def align(x: int) -> int:
return (x + page_size - 1) // page_size
assert size <= max_size and max_size % page_size == 0
indices = torch.randperm(align(max_size))[: align(size)]
offsets = torch.arange(page_size)
return (indices[:, None] * page_size + offsets).flatten().cuda()[:size]
def sglang_aot_transfer_one(
k_cache_dst: torch.Tensor,
v_cache_dst: torch.Tensor,
indices_dst: torch.Tensor,
k_cache_src: torch.Tensor,
v_cache_src: torch.Tensor,
indices_src: torch.Tensor,
item_size: int,
) -> None:
"""SGL AOT Kernel for single layer transfer."""
transfer_kv_per_layer(
k_cache_src,
k_cache_dst,
v_cache_src,
v_cache_dst,
indices_src,
indices_dst,
item_size,
)
def sglang_jit_transfer_one(
k_cache_dst: torch.Tensor,
v_cache_dst: torch.Tensor,
indices_dst: torch.Tensor,
k_cache_src: torch.Tensor,
v_cache_src: torch.Tensor,
indices_src: torch.Tensor,
element_dim: int,
) -> None:
"""SGL JIT Kernel for single layer transfer."""
transfer_hicache_one_layer(
k_cache_dst,
v_cache_dst,
indices_dst,
k_cache_src,
v_cache_src,
indices_src,
element_dim=element_dim,
)
def sglang_aot_transfer_all(
k_ptrs_dst: torch.Tensor,
v_ptrs_dst: torch.Tensor,
indices_dst: torch.Tensor,
k_ptrs_src: torch.Tensor,
v_ptrs_src: torch.Tensor,
indices_src: torch.Tensor,
item_size: int,
num_layers: int,
) -> None:
"""SGL AOT Kernel for all layer transfer."""
transfer_kv_all_layer(
k_ptrs_src,
k_ptrs_dst,
v_ptrs_src,
v_ptrs_dst,
indices_src,
indices_dst,
item_size,
num_layers,
)
def sglang_jit_transfer_all(
k_ptrs_dst: torch.Tensor,
v_ptrs_dst: torch.Tensor,
indices_dst: torch.Tensor,
k_ptrs_src: torch.Tensor,
v_ptrs_src: torch.Tensor,
indices_src: torch.Tensor,
stride_bytes: int,
element_size: int,
) -> None:
"""SGL JIT Kernel for all layer transfer."""
transfer_hicache_all_layer(
k_ptrs_dst,
v_ptrs_dst,
indices_dst,
k_ptrs_src,
v_ptrs_src,
indices_src,
kv_cache_src_stride_bytes=stride_bytes,
kv_cache_dst_stride_bytes=stride_bytes,
element_size=element_size,
)
def pytorch_transfer(
k_cache_dst: torch.Tensor,
v_cache_dst: torch.Tensor,
indices_dst_on_dst: torch.Tensor,
k_cache_src: torch.Tensor,
v_cache_src: torch.Tensor,
indices_src_on_src: torch.Tensor,
) -> None:
"""PyTorch indexing baseline."""
dst_device = k_cache_dst.device
k_cache_dst[indices_dst_on_dst] = k_cache_src[indices_src_on_src].to(dst_device)
v_cache_dst[indices_dst_on_dst] = v_cache_src[indices_src_on_src].to(dst_device)
# Benchmark configuration
BS_RANGE = get_benchmark_range(
full_range=[2**n for n in range(0, 16)],
ci_range=[16],
)
ELEMENT_SIZE_RANGE = get_benchmark_range(
full_range=[64, 128, 256, 512, 1024],
ci_range=[1024],
)
LINE_VALS = ["aot", "jit", "torch"]
LINE_NAMES = ["SGL AOT Kernel", "SGL JIT Kernel", "PyTorch"]
STYLES = [("orange", "-"), ("blue", "--"), ("red", ":")]
CONFIGS = list(itertools.product(ELEMENT_SIZE_RANGE, BS_RANGE))
# =============================================================================
# One Layer Benchmarks
# =============================================================================
@triton.testing.perf_report(
triton.testing.Benchmark(
x_names=["element_size", "batch_size"],
x_vals=CONFIGS,
line_arg="provider",
line_vals=LINE_VALS,
line_names=LINE_NAMES,
styles=STYLES,
ylabel="us",
plot_name="hicache-one-layer-h2d",
args={},
)
)
def benchmark_one_layer_h2d(
element_size: int, batch_size: int, provider: str
) -> Tuple[float, float, float]:
"""One Layer: Host (CPU) -> Device (GPU)."""
global cache
cache_local = cache.get_slice(num_layers=NUM_LAYERS, element_size=element_size)
k_cache_src = cache_local.k_cache_host
v_cache_src = cache_local.v_cache_host
k_cache_dst = cache_local.k_cache_cuda
v_cache_dst = cache_local.v_cache_cuda
torch.manual_seed(batch_size * 65536 + element_size)
indices_src_gpu = gen_indices(batch_size, HOST_CACHE_SIZE)
indices_dst_gpu = gen_indices(batch_size, GPU_CACHE_SIZE)
if ENABLE_SORT:
indices_src_gpu, mapping = indices_src_gpu.sort()
indices_dst_gpu = indices_dst_gpu[mapping]
indices_src_cpu = indices_src_gpu.cpu()
torch.cuda.synchronize()
element_bytes = element_size * k_cache_src.element_size()
FN_MAP = {
"aot": lambda: [
sglang_aot_transfer_one(
k_cache_dst[i],
v_cache_dst[i],
indices_dst_gpu,
k_cache_src[i],
v_cache_src[i],
indices_src_gpu,
element_bytes,
)
for i in range(NUM_LAYERS)
],
"jit": lambda: [
sglang_jit_transfer_one(
k_cache_dst[i],
v_cache_dst[i],
indices_dst_gpu,
k_cache_src[i],
v_cache_src[i],
indices_src_gpu,
element_size,
)
for i in range(NUM_LAYERS)
],
"torch": lambda: [
pytorch_transfer(
k_cache_dst[i],
v_cache_dst[i],
indices_dst_gpu,
k_cache_src[i],
v_cache_src[i],
indices_src_cpu,
)
for i in range(NUM_LAYERS)
],
}
if provider == "jit" and not can_use_hicache_jit_kernel(element_size=element_bytes):
return (float("nan"), float("nan"), float("nan"))
if DISABLE_TORCH and provider in ["torch"]:
return (float("nan"), float("nan"), float("nan"))
ms, min_ms, max_ms = triton.testing.do_bench( # type: ignore
FN_MAP[provider], quantiles=DEFAULT_QUANTILES, warmup=5, rep=25
)
return (
1000 * ms / NUM_LAYERS,
1000 * max_ms / NUM_LAYERS,
1000 * min_ms / NUM_LAYERS,
)
# =============================================================================
# All Layer Benchmarks
# =============================================================================
def _create_ptr_tensor(tensors, device="cuda"):
"""Create a tensor of data pointers."""
return torch.tensor(
[t.data_ptr() for t in tensors],
dtype=torch.uint64,
device=device,
)
@triton.testing.perf_report(
triton.testing.Benchmark(
x_names=["element_size", "batch_size"],
x_vals=CONFIGS,
line_arg="provider",
line_vals=LINE_VALS,
line_names=LINE_NAMES,
styles=STYLES,
ylabel="us",
plot_name="hicache-all-layer-d2h",
args={},
)
)
def benchmark_all_layer_d2h(
element_size: int, batch_size: int, provider: str
) -> Tuple[float, float, float]:
"""All Layer: Device (GPU) -> Host (CPU)."""
global cache
cache_local = cache.get_slice(num_layers=NUM_LAYERS, element_size=element_size)
k_caches_src = cache_local.k_cache_cuda
v_caches_src = cache_local.v_cache_cuda
k_caches_dst = cache_local.k_cache_host
v_caches_dst = cache_local.v_cache_host
torch.manual_seed(batch_size * 65536 + element_size)
indices_src_gpu = gen_indices(batch_size, GPU_CACHE_SIZE)
indices_dst_gpu = gen_indices(batch_size, HOST_CACHE_SIZE)
if ENABLE_SORT:
indices_dst_gpu, mapping = indices_dst_gpu.sort()
indices_src_gpu = indices_src_gpu[mapping]
indices_dst_cpu = indices_dst_gpu.cpu()
torch.cuda.synchronize()
element_bytes = element_size * k_caches_src.element_size()
k_ptrs_src = _create_ptr_tensor([k_caches_src[i] for i in range(NUM_LAYERS)])
v_ptrs_src = _create_ptr_tensor([v_caches_src[i] for i in range(NUM_LAYERS)])
k_ptrs_dst = _create_ptr_tensor([k_caches_dst[i] for i in range(NUM_LAYERS)])
v_ptrs_dst = _create_ptr_tensor([v_caches_dst[i] for i in range(NUM_LAYERS)])
FN_MAP = {
"aot": lambda: sglang_aot_transfer_all(
k_ptrs_dst,
v_ptrs_dst,
indices_dst_gpu,
k_ptrs_src,
v_ptrs_src,
indices_src_gpu,
element_bytes,
NUM_LAYERS,
),
"jit": lambda: sglang_jit_transfer_all(
k_ptrs_dst,
v_ptrs_dst,
indices_dst_gpu,
k_ptrs_src,
v_ptrs_src,
indices_src_gpu,
element_bytes,
element_bytes,
),
"torch": lambda: [
pytorch_transfer(
k_caches_dst[i],
v_caches_dst[i],
indices_dst_cpu,
k_caches_src[i],
v_caches_src[i],
indices_src_gpu,
)
for i in range(NUM_LAYERS)
],
}
if provider == "jit" and not can_use_hicache_jit_kernel(element_size=element_bytes):
return (float("nan"), float("nan"), float("nan"))
if DISABLE_TORCH and provider in ["torch"]:
return (float("nan"), float("nan"), float("nan"))
ms, min_ms, max_ms = triton.testing.do_bench( # type: ignore
FN_MAP[provider], quantiles=DEFAULT_QUANTILES, warmup=5, rep=25
)
return (
1000 * ms / NUM_LAYERS,
1000 * max_ms / NUM_LAYERS,
1000 * min_ms / NUM_LAYERS,
)
if __name__ == "__main__":
MAX_SIZE = max(ELEMENT_SIZE_RANGE)
DEVICE_SHAPE = (NUM_LAYERS * GPU_CACHE_SIZE, MAX_SIZE)
HOST_SHAPE = (NUM_LAYERS * HOST_CACHE_SIZE, MAX_SIZE)
cache = HiCacheCache(
k_cache_cuda=torch.empty(DEVICE_SHAPE, dtype=torch.bfloat16, device="cuda"),
v_cache_cuda=torch.empty(DEVICE_SHAPE, dtype=torch.bfloat16, device="cuda"),
k_cache_host=torch.empty(HOST_SHAPE, dtype=torch.bfloat16, pin_memory=True),
v_cache_host=torch.empty(HOST_SHAPE, dtype=torch.bfloat16, pin_memory=True),
)
print("=" * 60)
print("One Layer: Host -> Device (CPU -> GPU)")
print("=" * 60)
benchmark_one_layer_h2d.run(print_data=True)
print("\n" + "=" * 60)
print("All Layer: Device -> Host (GPU -> CPU) [per-layer avg]")
print("=" * 60)
benchmark_all_layer_d2h.run(print_data=True)
@@ -0,0 +1,193 @@
import itertools
from typing import Dict, Tuple
import torch
import triton
import triton.testing
from sglang.kernels.jit.benchmark.utils import DEFAULT_DEVICE, DEFAULT_DTYPE
from sglang.kernels.ops.kvcache.hisparse import load_cache_to_device_buffer_mla
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
register_cuda_ci(
est_time=12, stage="base-b-kernel-benchmark", runner_config="1-gpu-large"
)
register_amd_ci(est_time=12, stage="jit-kernel-benchmark", runner_config="amd")
DEVICE = DEFAULT_DEVICE
DTYPE = DEFAULT_DTYPE
TOP_K = 2048
ITEM_SIZE_BYTES = 512
MISS_RATES = [0.2, 0.001]
ROUNDS = 5
WARMUP_ROUNDS = 5
BATCH_SIZES = [1, 10, 100]
HOT_BUFFER_SIZES = [4096, 8192]
CONFIGS = [
(
batch_size,
hot_buffer_size,
miss_rate,
batch_size * round(TOP_K * miss_rate),
)
for batch_size, hot_buffer_size, miss_rate in itertools.product(
BATCH_SIZES, HOT_BUFFER_SIZES, MISS_RATES
)
]
LINE_VALS = ["jit"]
LINE_NAMES = ["SGL JIT Kernel"]
STYLES = [("blue", "--")]
def _make_top_k_tokens(
num_hits: int, num_misses: int, hot_buffer_size: int
) -> torch.Tensor:
hit_tokens = torch.arange(num_hits, dtype=torch.int32, device=DEVICE)
miss_tokens = hot_buffer_size + torch.arange(
num_misses, dtype=torch.int32, device=DEVICE
)
return torch.cat([hit_tokens, miss_tokens])
def _miss_tokens_per_req(miss_rate: float) -> int:
return round(TOP_K * miss_rate)
def _build_inputs(
batch_size: int, hot_buffer_size: int, miss_rate: float
) -> Dict[str, torch.Tensor | int]:
dtype_bytes = torch.empty((), dtype=DTYPE).element_size()
kv_dim = ITEM_SIZE_BYTES // dtype_bytes
padded_buffer_size = hot_buffer_size + 1
seq_len = hot_buffer_size + TOP_K + 1
num_misses = _miss_tokens_per_req(miss_rate)
num_hits = TOP_K - num_misses
top_k_row = _make_top_k_tokens(num_hits, num_misses, hot_buffer_size)
top_k_tokens = top_k_row.view(1, -1).repeat(batch_size, 1).contiguous()
host_stride = seq_len
total_host_tokens = batch_size * host_stride
host_cache = torch.empty(
(total_host_tokens, 1, kv_dim), dtype=DTYPE, device="cpu", pin_memory=True
)
host_cache.copy_(torch.randn_like(host_cache))
total_device_tokens = batch_size * padded_buffer_size
device_buffer = torch.empty(
(total_device_tokens, 1, kv_dim), dtype=DTYPE, device=DEVICE
)
device_buffer.normal_()
device_buffer_locs = torch.arange(
total_device_tokens, dtype=torch.int32, device=DEVICE
).view(batch_size, padded_buffer_size)
device_buffer_tokens = torch.full(
(batch_size, padded_buffer_size), -1, dtype=torch.int32, device=DEVICE
)
device_buffer_tokens[:, :hot_buffer_size] = torch.arange(
hot_buffer_size, dtype=torch.int32, device=DEVICE
)
lru_slots = (
torch.arange(hot_buffer_size, dtype=torch.int16, device=DEVICE)
.view(1, -1)
.repeat(batch_size, 1)
)
return {
"top_k_tokens": top_k_tokens,
"device_buffer_tokens": device_buffer_tokens,
"initial_device_buffer_tokens": device_buffer_tokens.clone(),
"host_cache_locs": torch.arange(
total_host_tokens, dtype=torch.int64, device=DEVICE
).view(batch_size, host_stride),
"device_buffer_locs": device_buffer_locs,
"host_cache": host_cache,
"device_buffer": device_buffer,
"top_k_device_locs": torch.empty(
(batch_size, TOP_K), dtype=torch.int32, device=DEVICE
),
"req_pool_indices": torch.arange(batch_size, dtype=torch.int64, device=DEVICE),
"seq_lens": torch.full(
(batch_size,), seq_len, dtype=torch.int32, device=DEVICE
),
"lru_slots": lru_slots,
"initial_lru_slots": lru_slots.clone(),
"num_real_reqs": torch.tensor([batch_size], dtype=torch.int32, device=DEVICE),
}
def _time_kernel(batch_size: int, hot_buffer_size: int, miss_rate: float) -> float:
state = _build_inputs(batch_size, hot_buffer_size, miss_rate)
def run_once():
state["device_buffer_tokens"].copy_(state["initial_device_buffer_tokens"])
state["lru_slots"].copy_(state["initial_lru_slots"])
state["top_k_device_locs"].fill_(-1)
load_cache_to_device_buffer_mla(
top_k_tokens=state["top_k_tokens"],
device_buffer_tokens=state["device_buffer_tokens"],
host_cache_locs=state["host_cache_locs"],
device_buffer_locs=state["device_buffer_locs"],
host_cache=state["host_cache"],
device_buffer=state["device_buffer"],
top_k_device_locs=state["top_k_device_locs"],
req_pool_indices=state["req_pool_indices"],
seq_lens=state["seq_lens"],
lru_slots=state["lru_slots"],
item_size_bytes=ITEM_SIZE_BYTES,
num_top_k=TOP_K,
hot_buffer_size=hot_buffer_size,
block_size=1024,
num_real_reqs=state["num_real_reqs"],
)
run_once()
torch.cuda.synchronize()
for _ in range(WARMUP_ROUNDS):
run_once()
torch.cuda.synchronize()
start = torch.cuda.Event(enable_timing=True)
end = torch.cuda.Event(enable_timing=True)
start.record()
for _ in range(ROUNDS):
run_once()
end.record()
torch.cuda.synchronize()
return start.elapsed_time(end) * 1000.0 / ROUNDS
@triton.testing.perf_report(
triton.testing.Benchmark(
x_names=["batch_size", "hot_buffer_size", "miss_rate", "miss_tokens_cnt"],
x_vals=CONFIGS,
line_arg="provider",
line_vals=LINE_VALS,
line_names=LINE_NAMES,
styles=STYLES,
ylabel="us",
plot_name="hisparse-latency",
args={},
)
)
def benchmark_latency(
batch_size: int,
hot_buffer_size: int,
miss_rate: float,
miss_tokens_cnt: int,
provider: str,
) -> Tuple[float, float, float]:
assert provider == "jit"
batch_size = int(batch_size)
hot_buffer_size = int(hot_buffer_size)
miss_rate = float(miss_rate)
assert miss_tokens_cnt == batch_size * _miss_tokens_per_req(miss_rate)
avg_us = _time_kernel(batch_size, hot_buffer_size, miss_rate)
return avg_us, avg_us, avg_us
if __name__ == "__main__":
benchmark_latency.run(print_data=True)
@@ -0,0 +1,70 @@
"""Benchmark: fused MiniMax-M3 KV + index cache store (1 launch) vs the separate
per-buffer index_put_ stores (main K, main V, index K, optional index V)."""
import torch
from sglang.kernels.jit.benchmark import marker
from sglang.kernels.ops.kvcache.minimax_store_kv_index import store_kv_index
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
register_cuda_ci(
est_time=6, stage="base-b-kernel-benchmark", runner_config="1-gpu-large"
)
register_amd_ci(est_time=6, stage="jit-kernel-benchmark", runner_config="amd")
HEAD_DIM = 128
NUM_KV_HEADS = 1
HAS_V = False
N = 1 << 20
DTYPE = torch.bfloat16
def _separate(k, v, kc, vc, ik, ikc, loc, **_):
kc[loc] = k
vc[loc] = v
ikc[loc] = ik
def _fused(k, v, kc, vc, ik, ikc, loc, *, num_kv_heads, head_bytes):
store_kv_index(
k,
v,
kc,
vc,
ik,
ikc,
None,
None,
loc,
num_kv_heads=num_kv_heads,
head_bytes=head_bytes,
)
FN_MAP = {"fused": _fused, "separate": _separate}
@marker.parametrize("T", [16, 64, 256, 1024, 4096, 16384], [256, 4096])
@marker.benchmark("impl", ["fused", "separate"])
def benchmark(T: int, impl: str):
k = torch.randn(T, NUM_KV_HEADS * HEAD_DIM, dtype=DTYPE, device="cuda")
v = torch.randn_like(k)
ik = torch.randn(T, HEAD_DIM, dtype=DTYPE, device="cuda")
kc = torch.zeros(N, NUM_KV_HEADS * HEAD_DIM, dtype=DTYPE, device="cuda")
vc = torch.zeros_like(kc)
ikc = torch.zeros(N, HEAD_DIM, dtype=DTYPE, device="cuda")
loc = torch.randperm(N, device="cuda")[:T]
extra_kwargs = dict(num_kv_heads=NUM_KV_HEADS, head_bytes=HEAD_DIM * DTYPE.itemsize)
return marker.do_bench(
FN_MAP[impl],
input_args=(k, v, kc, vc, ik, ikc, loc),
input_kwargs=extra_kwargs if impl == "fused" else {},
# Read inputs cloned per iter; caches are write targets (kept hot).
graph_clone_args=(0, 1, 4, 6),
memory_args=(k, v, ik, loc),
memory_output=(k, v, ik),
)
if __name__ == "__main__":
benchmark.run()
@@ -0,0 +1,133 @@
"""Benchmark the set_mla_kv_buffer dispatcher.
Compares three providers across a batch-size sweep:
- ``wrapper``: the high-level wrapper exposed by ``set_mla_kv_buffer_triton``
(dispatches to TMA on SM90+, Triton fallback otherwise).
- ``jit_tma``: the JIT CUDA TMA bulk-store kernel directly.
- ``triton``: the BLOCK-tiled Triton kernel (SM<90 fallback path).
"""
import itertools
from typing import Tuple
import torch
import triton
import triton.testing
from sglang.kernels.jit.benchmark.utils import (
DEFAULT_DEVICE,
DEFAULT_DTYPE,
DEFAULT_QUANTILES,
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.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
register_cuda_ci(
est_time=9, stage="base-b-kernel-benchmark", runner_config="1-gpu-large"
)
def _triton_baseline(kv_buffer, loc, cache_k_nope, cache_k_rope):
nope_dim = cache_k_nope.shape[-1]
rope_dim = cache_k_rope.shape[-1]
total_dim = nope_dim + rope_dim
BLOCK = 128
n_loc = loc.numel()
grid = (n_loc, triton.cdiv(total_dim, BLOCK))
pdl_kwargs = {"USE_GDC": True, "launch_pdl": True} if is_arch_support_pdl() else {}
sglang_triton_kernel[grid](
kv_buffer,
cache_k_nope,
cache_k_rope,
loc,
kv_buffer.stride(0),
cache_k_nope.stride(0),
cache_k_rope.stride(0),
nope_dim,
rope_dim,
BLOCK=BLOCK,
DCP_RANK=0,
DCP_WORLD_SIZE=1,
**pdl_kwargs,
)
NUM_LAYERS = 8
CACHE_SIZE = (2 * 1024 * 1024) // NUM_LAYERS
NOPE_DIM = 512
ROPE_DIM = 64
BS_RANGE = get_benchmark_range(
full_range=[1, 8, 32, 128, 512, 1024, 2048, 4096, 8192, 16384],
ci_range=[1, 128, 2048, 4096, 8192],
)
LINE_VALS = ["wrapper", "jit_tma", "triton"]
LINE_NAMES = ["Wrapper (auto)", "JIT TMA bulk-store", "Triton (BLOCK=128 baseline)"]
STYLES = [("blue", "-"), ("green", "--"), ("red", "-.")]
X_NAMES = ["batch_size"]
CONFIGS = list(itertools.product(BS_RANGE))
@triton.testing.perf_report(
triton.testing.Benchmark(
x_names=X_NAMES,
x_vals=CONFIGS,
line_arg="provider",
line_vals=LINE_VALS,
line_names=LINE_NAMES,
styles=STYLES,
ylabel="us",
plot_name="set-mla-kv-buffer-performance",
args={},
)
)
def benchmark(batch_size: int, provider: str) -> Tuple[float, float, float]:
cache_k_nope = torch.randn(
(NUM_LAYERS, batch_size, 1, NOPE_DIM),
dtype=DEFAULT_DTYPE,
device=DEFAULT_DEVICE,
)
cache_k_rope = torch.randn(
(NUM_LAYERS, batch_size, 1, ROPE_DIM),
dtype=DEFAULT_DTYPE,
device=DEFAULT_DEVICE,
)
kv_buffer = torch.randn(
(NUM_LAYERS, CACHE_SIZE, 1, NOPE_DIM + ROPE_DIM),
dtype=DEFAULT_DTYPE,
device=DEFAULT_DEVICE,
)
loc = torch.randperm(CACHE_SIZE, device=DEFAULT_DEVICE)[:batch_size]
torch.cuda.synchronize()
FN_MAP = {
"wrapper": sglang_wrapper,
"jit_tma": lambda buf, loc, n, r: jit_set(buf, loc, n, r),
"triton": _triton_baseline,
}
def fn():
impl = FN_MAP[provider]
for i in range(NUM_LAYERS):
impl(kv_buffer[i], loc, cache_k_nope[i], cache_k_rope[i])
ms, min_ms, max_ms = triton.testing.do_bench_cudagraph(
fn, quantiles=DEFAULT_QUANTILES
)
return (
1000 * ms / NUM_LAYERS,
1000 * max_ms / NUM_LAYERS,
1000 * min_ms / NUM_LAYERS,
)
if __name__ == "__main__":
benchmark.run(print_data=True)
@@ -0,0 +1,76 @@
import torch
from sglang.kernels.jit.benchmark import marker
from sglang.kernels.jit.benchmark.utils import (
DEFAULT_DEVICE,
create_empty,
create_random,
)
from sglang.kernels.ops.kvcache.kvcache import store_cache
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
register_cuda_ci(
est_time=9, stage="base-b-kernel-benchmark", runner_config="1-gpu-large"
)
register_amd_ci(est_time=9, stage="jit-kernel-benchmark", runner_config="amd")
@torch.compile()
def torch_compile_store_cache(
k: torch.Tensor,
v: torch.Tensor,
k_cache: torch.Tensor,
v_cache: torch.Tensor,
indices: torch.Tensor,
) -> None:
k_cache[indices] = k
v_cache[indices] = v
alt_stream = torch.cuda.Stream()
def torch_streams_store_cache(
k: torch.Tensor,
v: torch.Tensor,
k_cache: torch.Tensor,
v_cache: torch.Tensor,
indices: torch.Tensor,
) -> None:
current_stream = torch.cuda.current_stream()
alt_stream.wait_stream(current_stream)
k_cache[indices] = k
with torch.cuda.stream(alt_stream):
v_cache[indices] = v
current_stream.wait_stream(alt_stream)
CACHE_SIZE = 2 * 1024 * 1024
FN_MAP = {
"jit": store_cache,
"torch_compile": torch_compile_store_cache,
"torch_streams": torch_streams_store_cache,
}
@marker.parametrize("item_size", [64, 128, 256, 512, 1024], [1024])
@marker.parametrize("batch_size", [2**n for n in range(0, 15)], [16])
@marker.benchmark("impl", ["jit", "torch_compile", "torch_streams"])
def benchmark(batch_size: int, item_size: int, impl: str):
torch.manual_seed(42)
k = create_random(batch_size, item_size)
k_cache = create_empty(CACHE_SIZE, item_size)
v = create_random(batch_size, item_size)
v_cache = create_empty(CACHE_SIZE, item_size)
indices = torch.randperm(CACHE_SIZE, device=DEFAULT_DEVICE)[:batch_size]
return marker.do_bench(
FN_MAP[impl],
input_args=(k, v, k_cache, v_cache, indices),
graph_clone_args=(0, 1, 4), # not need to clone cache, which is large
memory_args=(k, v, indices), # k_cache / v_cache excluded
memory_output=(k, v), # inplace write, size = k + v
)
if __name__ == "__main__":
benchmark.run()
@@ -0,0 +1,66 @@
from __future__ import annotations
import torch
from sglang.kernels.jit.benchmark import marker
from sglang.kernels.ops.layernorm.fused_eh_norm import fused_eh_norm
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
register_cuda_ci(
est_time=6, stage="base-b-kernel-benchmark", runner_config="1-gpu-large"
)
register_amd_ci(est_time=6, stage="jit-kernel-benchmark", runner_config="amd")
EPS = 1e-6
def reference(
x: torch.Tensor,
prev: torch.Tensor,
ew: torch.Tensor,
hw: torch.Tensor,
eps: float,
) -> torch.Tensor:
xf = x.float()
pf = prev.float()
x_var = xf.pow(2).mean(dim=-1, keepdim=True)
p_var = pf.pow(2).mean(dim=-1, keepdim=True)
return torch.cat(
(
(xf * torch.rsqrt(x_var + eps) * ew.float()).to(x.dtype),
(pf * torch.rsqrt(p_var + eps) * hw.float()).to(prev.dtype),
),
dim=-1,
)
FN_MAP = {
"jit": fused_eh_norm,
"torch": reference,
}
@marker.parametrize("dtype", [torch.bfloat16, torch.float16], [torch.bfloat16])
@marker.parametrize("hidden_size", [6144, 7168], [7168])
@marker.parametrize("num_tokens", [1, 4, 6, 8, 16, 32, 128, 512], [1, 6])
@marker.benchmark("impl", ["jit", "torch"])
def benchmark(num_tokens: int, hidden_size: int, dtype: torch.dtype, impl: str):
x = torch.randn(num_tokens, hidden_size, device="cuda", dtype=dtype)
prev = torch.randn_like(x)
ew = torch.randn(hidden_size, device="cuda", dtype=dtype)
hw = torch.randn(hidden_size, device="cuda", dtype=dtype)
expected = reference(x, prev, ew, hw, EPS)
actual = fused_eh_norm(x, prev, ew, hw, EPS)
torch.testing.assert_close(actual.float(), expected.float(), rtol=1e-2, atol=1e-2)
return marker.do_bench(
FN_MAP[impl],
input_args=(x, prev, ew, hw, EPS),
memory_args=(x, prev, ew, hw),
memory_output="out",
)
if __name__ == "__main__":
benchmark.run()
@@ -0,0 +1,104 @@
import itertools
import torch
import triton
import triton.testing
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.test.ci.ci_register import register_cuda_ci
register_cuda_ci(
est_time=30, stage="base-b-kernel-benchmark", runner_config="1-gpu-large"
)
DTYPE = torch.bfloat16
DEVICE = "cuda"
BS_LIST = get_benchmark_range(
full_range=[2**n for n in range(0, 14)],
ci_range=[16, 32],
)
HIDDEN_SIZE_LIST = get_benchmark_range(
full_range=sorted([1536, *range(1024, 8192 + 1, 1024)]),
ci_range=[512, 2048],
)
LINE_VALS = ["flashinfer", "jit"]
LINE_NAMES = ["FlashInfer", "SGL JIT Kernel"]
STYLES = [("blue", "--"), ("green", "-.")]
NUM_LAYERS = 4 # avoid L2 effect
configs_0 = list(itertools.product(HIDDEN_SIZE_LIST + [16384], BS_LIST))
configs_1 = list(itertools.product(HIDDEN_SIZE_LIST, BS_LIST))
@triton.testing.perf_report(
triton.testing.Benchmark(
x_names=["hidden_size", "batch_size"],
x_vals=configs_0,
line_arg="provider",
line_vals=LINE_VALS,
line_names=LINE_NAMES,
styles=STYLES,
ylabel="us",
plot_name="rmsnorm-performance",
args={},
)
)
def benchmark_rmsnorm(hidden_size: int, batch_size: int, provider: str):
input = torch.randn(
(NUM_LAYERS, batch_size, hidden_size), dtype=DTYPE, device=DEVICE
)
weight = torch.randn((NUM_LAYERS, hidden_size), dtype=DTYPE, device=DEVICE)
FN_MAP = {"jit": jit_rmsnorm, "flashinfer": fi_rmsnorm}
def f():
fn = FN_MAP[provider]
for i in range(NUM_LAYERS):
fn(input[i], weight[i], out=input[i])
return run_benchmark(f, scale=NUM_LAYERS)
@triton.testing.perf_report(
triton.testing.Benchmark(
x_names=["hidden_size", "batch_size"],
x_vals=configs_1,
line_arg="provider",
line_vals=LINE_VALS,
line_names=LINE_NAMES,
styles=STYLES,
ylabel="us",
plot_name="fused-add-rmsnorm-performance",
args={},
)
)
def benchmark_fused_add_rmsnorm(hidden_size: int, batch_size: int, provider: str):
input = torch.randn(
(NUM_LAYERS, batch_size, hidden_size), dtype=DTYPE, device=DEVICE
)
residual = torch.randn_like(input)
weight = torch.randn((NUM_LAYERS, hidden_size), dtype=DTYPE, device=DEVICE)
FN_MAP = {"jit": jit_fused_add_rmsnorm, "flashinfer": fi_fused_add_rmsnorm}
def f():
fn = FN_MAP[provider]
for i in range(NUM_LAYERS):
fn(input[i], residual[i], weight[i])
return run_benchmark(f, scale=NUM_LAYERS)
if __name__ == "__main__":
print("Benchmarking rmsnorm...")
benchmark_rmsnorm.run(print_data=True)
print("Benchmarking fused_add_rmsnorm...")
benchmark_fused_add_rmsnorm.run(print_data=True)
@@ -0,0 +1,77 @@
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.srt.utils import get_current_device_stream_fast
from sglang.test.ci.ci_register import register_cuda_ci
register_cuda_ci(
est_time=10, stage="base-b-kernel-benchmark", runner_config="1-gpu-large"
)
alt_stream = torch.cuda.Stream()
torch._dynamo.config.recompile_limit = 100
# NOTE: now aot fallback to flashinfer
def sglang_aot_qknorm(
q: torch.Tensor,
k: torch.Tensor,
q_weight: torch.Tensor,
k_weight: torch.Tensor,
) -> None:
from flashinfer import rmsnorm # lazy import to avoid crash
current_stream = get_current_device_stream_fast()
alt_stream.wait_stream(current_stream)
rmsnorm(q, q_weight, out=q)
with torch.cuda.stream(alt_stream):
rmsnorm(k, k_weight, out=k)
current_stream.wait_stream(alt_stream)
@torch.compile()
def torch_impl_qknorm(
q: torch.Tensor,
k: torch.Tensor,
q_weight: torch.Tensor,
k_weight: torch.Tensor,
eps: float = 1e-6,
) -> None:
q_mean = q.float().pow(2).mean(dim=-1, keepdim=True)
k_mean = k.float().pow(2).mean(dim=-1, keepdim=True)
q_norm = (q_mean + eps).rsqrt()
k_norm = (k_mean + eps).rsqrt()
q.copy_(q.float() * q_norm * q_weight.float())
k.copy_(k.float() * k_norm * k_weight.float())
FN_MAP = {
"aot": sglang_aot_qknorm,
"jit": fused_inplace_qknorm,
"torch": torch_impl_qknorm,
}
@marker.parametrize("head_dim", [128, 256, 512, 1024], [128])
@marker.parametrize("GQA", [4, 8], [4])
@marker.parametrize("num_kv_heads", [1, 2, 4, 8], [1])
@marker.parametrize("batch_size", [2**n for n in range(0, 14)], [16])
@marker.benchmark("impl", ["aot", "jit", "torch"])
def benchmark(head_dim: int, GQA: int, num_kv_heads: int, batch_size: int, impl: str):
num_qo_heads = GQA * num_kv_heads
q = create_random(batch_size, num_qo_heads, head_dim)
k = create_random(batch_size, num_kv_heads, head_dim)
q_weight = create_random(head_dim)
k_weight = create_random(head_dim)
return marker.do_bench(
FN_MAP[impl],
input_args=(q, k, q_weight, k_weight),
memory_output=(q, k), # inplace write to q, k
)
if __name__ == "__main__":
benchmark.run()
@@ -0,0 +1,125 @@
import itertools
from typing import Tuple
import torch
import triton
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.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
register_cuda_ci(
est_time=12, stage="base-b-kernel-benchmark", runner_config="1-gpu-large"
)
IS_CI = is_in_ci()
alt_stream = torch.cuda.Stream()
def sglang_jit_qknorm_across_heads(
q: torch.Tensor,
k: torch.Tensor,
q_weight: torch.Tensor,
k_weight: torch.Tensor,
) -> None:
fused_inplace_qknorm_across_heads(q, k, q_weight, k_weight)
def sglang_aot_qknorm_across_heads(
q: torch.Tensor,
k: torch.Tensor,
q_weight: torch.Tensor,
k_weight: torch.Tensor,
) -> None:
current_stream = get_current_device_stream_fast()
alt_stream.wait_stream(current_stream)
rmsnorm(q, q_weight, out=q)
with torch.cuda.stream(alt_stream):
rmsnorm(k, k_weight, out=k)
current_stream.wait_stream(alt_stream)
def flashinfer_qknorm_across_heads(
q: torch.Tensor,
k: torch.Tensor,
q_weight: torch.Tensor,
k_weight: torch.Tensor,
) -> None:
from flashinfer import rmsnorm
rmsnorm(q, q_weight, out=q)
rmsnorm(k, k_weight, out=k)
@torch.compile()
def torch_impl_qknorm_across_heads(
q: torch.Tensor,
k: torch.Tensor,
q_weight: torch.Tensor,
k_weight: torch.Tensor,
eps: float = 1e-6,
) -> None:
q_mean = q.float().pow(2).mean(dim=-1, keepdim=True)
k_mean = k.float().pow(2).mean(dim=-1, keepdim=True)
q_norm = (q_mean + eps).rsqrt()
k_norm = (k_mean + eps).rsqrt()
q.copy_(q.float() * q_norm * q_weight.float())
k.copy_(k.float() * k_norm * k_weight.float())
DTYPE = torch.bfloat16
DEVICE = "cuda"
if IS_CI:
BS_RANGE = [16]
HIDDEN_DIM_RANGE = [1024]
else:
BS_RANGE = [2**n for n in range(0, 14)]
HIDDEN_DIM_RANGE = [512, 1024, 2048, 4096, 8192]
LINE_VALS = ["jit", "aot", "flashinfer", "torch"]
LINE_NAMES = ["SGL JIT Kernel", "SGL AOT Kernel", "FlashInfer", "PyTorch"]
STYLES = [("blue", "-"), ("orange", "--"), ("green", "-."), ("red", ":")]
configs = list(itertools.product(BS_RANGE, HIDDEN_DIM_RANGE))
@triton.testing.perf_report(
triton.testing.Benchmark(
x_names=["batch_size", "hidden_dim"],
x_vals=configs,
line_arg="provider",
line_vals=LINE_VALS,
line_names=LINE_NAMES,
styles=STYLES,
ylabel="us",
plot_name="qknorm-across-heads-performance",
args={},
)
)
def benchmark(
batch_size: int, hidden_dim: int, provider: str
) -> Tuple[float, float, float]:
q = torch.randn((batch_size, hidden_dim), dtype=DTYPE, device=DEVICE)
k = torch.randn((batch_size, hidden_dim), dtype=DTYPE, device=DEVICE)
q_weight = torch.randn(hidden_dim, dtype=DTYPE, device=DEVICE)
k_weight = torch.randn(hidden_dim, dtype=DTYPE, device=DEVICE)
FN_MAP = {
"jit": sglang_jit_qknorm_across_heads,
"aot": sglang_aot_qknorm_across_heads,
"flashinfer": flashinfer_qknorm_across_heads,
"torch": torch_impl_qknorm_across_heads,
}
fn = lambda: FN_MAP[provider](q, k, q_weight, k_weight)
return run_benchmark(fn)
if __name__ == "__main__":
benchmark.run(print_data=True)
@@ -0,0 +1,63 @@
import torch
from sglang.kernels.jit.benchmark import marker
from sglang.kernels.jit.benchmark.utils import create_random
from sglang.kernels.ops.moe.moe_fused_gate import moe_fused_gate, moe_fused_gate_jit
from sglang.test.ci.ci_register import register_cuda_ci
register_cuda_ci(
est_time=20, stage="base-b-kernel-benchmark", runner_config="1-gpu-large"
)
TOPK = 8
SCALE = 2.5
@torch.compile
def torch_router(scores, bias, topk, scoring_func):
"""Reference PyTorch router: scoring + bias + top-k + renorm + scale."""
if scoring_func == "sigmoid":
activated = scores.sigmoid()
else:
activated = torch.nn.functional.softplus(scores).sqrt()
biased = activated + bias.unsqueeze(0)
_, ids = torch.topk(biased, k=topk, dim=-1)
weights = activated.gather(1, ids)
weights = weights / weights.sum(dim=-1, keepdim=True)
return weights * SCALE, ids.to(torch.int32)
@marker.parametrize("scoring_func", ["sigmoid", "sqrtsoftplus"])
@marker.parametrize("num_experts", [128, 256, 384, 512], [256, 384])
@marker.parametrize("num_tokens", [1, 4, 16, 64, 512, 1024, 8192], [16, 1024])
@marker.benchmark("provider", ["triton", "jit", "torch"])
def benchmark(num_tokens: int, num_experts: int, scoring_func: str, provider: str):
torch.manual_seed(0)
scores = create_random(num_tokens, num_experts, dtype=torch.float32)
bias = create_random(num_experts, dtype=torch.float32)
common = dict(
topk=TOPK,
scoring_func=scoring_func,
renormalize=True,
routed_scaling_factor=SCALE,
apply_routed_scaling_factor_on_output=True,
)
if provider == "triton":
return marker.do_bench(
moe_fused_gate, input_args=(scores, bias), input_kwargs=common
)
if provider == "jit":
return marker.do_bench(
moe_fused_gate_jit, input_args=(scores, bias), input_kwargs=common
)
if provider == "torch":
return marker.do_bench(
torch_router, input_args=(scores, bias, TOPK, scoring_func)
)
raise ValueError(f"unknown provider: {provider}")
if __name__ == "__main__":
benchmark.run()
@@ -0,0 +1,80 @@
import torch
from sglang.kernels.jit.benchmark import marker
from sglang.kernels.ops.moe.ep_moe_kernels import (
post_reorder_deepgemm,
post_reorder_triton_kernel,
)
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
register_cuda_ci(
est_time=8, stage="base-b-kernel-benchmark", runner_config="1-gpu-large"
)
register_amd_ci(est_time=8, stage="jit-kernel-benchmark", runner_config="amd")
HIDDEN = 6144
NUM_EXPERTS = 129
TOP_K = 5
RSF = 2.0
def _build(num_tokens):
m_max = (num_tokens // 256 + 1) * 256
down_output = torch.randn(
NUM_EXPERTS * m_max, HIDDEN, dtype=torch.bfloat16, device="cuda"
)
topk_ids = torch.randint(
0, NUM_EXPERTS, (num_tokens, TOP_K), dtype=torch.int32, device="cuda"
)
src2dst = (
topk_ids.long() * m_max
+ torch.randint(0, m_max, (num_tokens, TOP_K), device="cuda")
).to(torch.int32)
topk_weights = torch.rand(num_tokens, TOP_K, dtype=torch.float32, device="cuda")
return down_output, src2dst, topk_ids, topk_weights
def _new(down_output, src2dst, topk_ids, topk_weights):
num_tokens = topk_ids.shape[0]
out = torch.empty(num_tokens, HIDDEN, dtype=torch.bfloat16, device="cuda")
post_reorder_deepgemm(
down_output,
out,
src2dst,
topk_ids,
topk_weights,
TOP_K,
num_tokens,
HIDDEN,
RSF,
)
return out
def _old(down_output, src2dst, topk_ids, topk_weights):
num_tokens = topk_ids.shape[0]
out = torch.empty(num_tokens, HIDDEN, dtype=torch.bfloat16, device="cuda")
post_reorder_triton_kernel[(num_tokens,)](
down_output, out, src2dst, topk_ids, topk_weights, TOP_K, HIDDEN, BLOCK_SIZE=512
)
out *= RSF
return out
FN_MAP = {"fused": _new, "legacy": _old}
@marker.parametrize("num_tokens", [1, 8, 64, 256, 1024, 4096, 16384], [64, 4096])
@marker.benchmark("impl", ["fused", "legacy"])
def benchmark(num_tokens: int, impl: str):
down_output, src2dst, topk_ids, topk_weights = _build(num_tokens)
return marker.do_bench(
FN_MAP[impl],
input_args=(down_output, src2dst, topk_ids, topk_weights),
graph_clone_args=(0,),
memory_args=None,
)
if __name__ == "__main__":
benchmark.run()
@@ -0,0 +1,241 @@
import itertools
import sgl_kernel
import torch
import triton
import triton.testing
from sglang.kernels.jit.benchmark.utils import run_benchmark_no_cudagraph
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.utils import is_in_ci
register_cuda_ci(
est_time=5, stage="base-b-kernel-benchmark", runner_config="1-gpu-large"
)
def torch_top_k_renorm_probs(probs, top_k):
"""Vectorized PyTorch implementation of top-k renormalization."""
batch_size, vocab_size = probs.shape
# Handle scalar or tensor k
if isinstance(top_k, int):
k_val = min(max(top_k, 1), vocab_size)
# Get top-k indices for all batches at once
_, topk_indices = torch.topk(probs, k_val, dim=1, largest=True)
# Create mask: batch_size x vocab_size
mask = torch.zeros_like(probs)
mask.scatter_(1, topk_indices, 1.0)
# Vectorized renormalization
masked_probs = probs * mask
renorm_probs = masked_probs / (masked_probs.sum(dim=1, keepdim=True) + 1e-10)
return renorm_probs
else:
# Variable k per batch - need to handle separately
renorm_probs = torch.zeros_like(probs)
for i in range(batch_size):
k_val = min(max(top_k[i].item(), 1), vocab_size)
_, topk_indices = torch.topk(probs[i], k_val, largest=True)
mask = torch.zeros_like(probs[i])
mask[topk_indices] = 1.0
masked_probs = probs[i] * mask
renorm_probs[i] = masked_probs / (masked_probs.sum() + 1e-10)
return renorm_probs
def torch_top_p_renorm_probs(probs, top_p, eps=1e-5):
"""Vectorized PyTorch implementation of top-p renormalization."""
batch_size, vocab_size = probs.shape
# Handle scalar or tensor p
if isinstance(top_p, float):
p_val = top_p
# Vectorized implementation for uniform top_p
# Sort probs in descending order
sorted_probs, sorted_indices = torch.sort(probs, descending=True, dim=1)
cumsum_probs = torch.cumsum(sorted_probs, dim=1)
# Find cutoff: where cumsum exceeds top_p
cutoff_mask = cumsum_probs <= p_val
# Keep at least one token (the highest prob)
cutoff_mask[:, 0] = True
# Create mask in original order
mask = torch.zeros_like(probs)
mask.scatter_(1, sorted_indices, cutoff_mask.float())
# Vectorized renormalization
masked_probs = probs * mask
renorm_probs = masked_probs / (masked_probs.sum(dim=1, keepdim=True) + eps)
return renorm_probs
else:
# Variable p per batch - need to handle separately
renorm_probs = torch.zeros_like(probs)
for i in range(batch_size):
p_val = top_p[i].item()
sorted_prob, indices = torch.sort(probs[i], descending=False)
cdf = torch.cumsum(sorted_prob, dim=-1)
mask = torch.zeros(vocab_size, dtype=torch.float32, device=probs.device)
mask.scatter_(0, indices, (cdf >= (1 - p_val) - eps).float())
masked_probs = probs[i] * mask
renorm_probs[i] = masked_probs / (masked_probs.sum() + eps)
return renorm_probs
def calculate_diff_top_k_renorm(batch_size, vocab_size, k):
"""Compare Torch reference and SGLang kernel for top-k renorm correctness."""
torch.manual_seed(42)
device = torch.device("cuda")
pre_norm_prob = torch.rand(batch_size, vocab_size, device=device)
probs = pre_norm_prob / pre_norm_prob.sum(dim=-1, keepdim=True)
top_k_tensor = torch.full((batch_size,), k, device=device, dtype=torch.int32)
torch_output = torch_top_k_renorm_probs(probs, top_k_tensor)
sglang_output = sgl_kernel.top_k_renorm_prob(probs, top_k_tensor)
torch.testing.assert_close(torch_output, sglang_output, rtol=1e-3, atol=1e-3)
def calculate_diff_top_p_renorm(batch_size, vocab_size, p):
"""Compare Torch reference and SGLang kernel for top-p renorm correctness."""
torch.manual_seed(42)
device = torch.device("cuda")
pre_norm_prob = torch.rand(batch_size, vocab_size, device=device)
probs = pre_norm_prob / pre_norm_prob.sum(dim=-1, keepdim=True)
top_p_tensor = torch.full((batch_size,), p, device=device, dtype=torch.float32)
torch_output = torch_top_p_renorm_probs(probs, top_p_tensor)
sglang_output = sgl_kernel.top_p_renorm_prob(probs, top_p_tensor)
torch.testing.assert_close(torch_output, sglang_output, rtol=1e-3, atol=1e-3)
# Parameter space - simplified for CI
if is_in_ci():
batch_size_range = [16]
vocab_size_range = [111]
k_range = [10]
p_range = [0.5]
else:
batch_size_range = [16, 64, 128]
vocab_size_range = [111, 32000, 128256]
k_range = [10, 100, 500]
p_range = [0.1, 0.5, 0.9]
configs_k = list(itertools.product(batch_size_range, vocab_size_range, k_range))
configs_p = list(itertools.product(batch_size_range, vocab_size_range, p_range))
@triton.testing.perf_report(
triton.testing.Benchmark(
x_names=["batch_size", "vocab_size", "k"],
x_vals=configs_k,
line_arg="provider",
line_vals=["torch", "sglang"],
line_names=["Torch Reference", "SGL Kernel"],
styles=[("red", "-"), ("green", "-")],
ylabel="us",
plot_name="top-k-renorm-probs-performance",
args={},
)
)
def benchmark_top_k_renorm(batch_size, vocab_size, k, provider):
# Skip invalid configurations
if k >= vocab_size:
return float("nan"), float("nan"), float("nan")
torch.manual_seed(42)
device = torch.device("cuda")
pre_norm_prob = torch.rand(batch_size, vocab_size, device=device)
probs = pre_norm_prob / pre_norm_prob.sum(dim=-1, keepdim=True)
top_k_tensor = torch.full((batch_size,), k, device=device, dtype=torch.int32)
if provider == "torch":
fn = lambda: torch_top_k_renorm_probs(probs.clone(), top_k_tensor)
elif provider == "sglang":
fn = lambda: sgl_kernel.top_k_renorm_prob(probs.clone(), top_k_tensor)
return run_benchmark_no_cudagraph(fn)
@triton.testing.perf_report(
triton.testing.Benchmark(
x_names=["batch_size", "vocab_size", "p"],
x_vals=configs_p,
line_arg="provider",
line_vals=["torch", "sglang"],
line_names=["Torch Reference", "SGL Kernel"],
styles=[("red", "-"), ("blue", "-")],
ylabel="us",
plot_name="top-p-renorm-probs-performance",
args={},
)
)
def benchmark_top_p_renorm(batch_size, vocab_size, p, provider):
torch.manual_seed(42)
device = torch.device("cuda")
pre_norm_prob = torch.rand(batch_size, vocab_size, device=device)
probs = pre_norm_prob / pre_norm_prob.sum(dim=-1, keepdim=True)
top_p_tensor = torch.full((batch_size,), p, device=device, dtype=torch.float32)
if provider == "torch":
fn = lambda: torch_top_p_renorm_probs(probs.clone(), top_p_tensor)
elif provider == "sglang":
fn = lambda: sgl_kernel.top_p_renorm_prob(probs.clone(), top_p_tensor)
return run_benchmark_no_cudagraph(fn)
if __name__ == "__main__":
print("=" * 60)
print("Running correctness checks...")
print("=" * 60)
# Correctness checks - simplified for CI
if is_in_ci():
test_configs_k = [configs_k[0]] if configs_k else [(16, 111, 10)]
test_configs_p = [configs_p[0]] if configs_p else [(16, 111, 0.5)]
else:
test_configs_k = configs_k[:3] # Test first 3 configs
test_configs_p = configs_p[:3]
print("\n1. Testing top_k_renorm_probs...")
for cfg in test_configs_k:
batch_size, vocab_size, k = cfg
if k < vocab_size: # Skip invalid configs
calculate_diff_top_k_renorm(batch_size, vocab_size, k)
print(
f" ✓ Passed: batch_size={batch_size}, vocab_size={vocab_size}, k={k}"
)
print("\n2. Testing top_p_renorm_probs...")
for cfg in test_configs_p:
calculate_diff_top_p_renorm(*cfg)
batch_size, vocab_size, p = cfg
print(f" ✓ Passed: batch_size={batch_size}, vocab_size={vocab_size}, p={p}")
print("\n" + "=" * 60)
print("All correctness checks passed!")
print("=" * 60)
print("\n" + "=" * 60)
print("Starting performance benchmarks...")
print("=" * 60)
print("\n1. Benchmarking top_k_renorm_probs...")
benchmark_top_k_renorm.run(print_data=True)
print("\n2. Benchmarking top_p_renorm_probs...")
benchmark_top_p_renorm.run(print_data=True)
print("\n" + "=" * 60)
print("Benchmarking complete!")
print("=" * 60)
@@ -0,0 +1,126 @@
import itertools
import torch
import triton
import triton.testing
from sglang.kernels.jit.benchmark.utils import run_benchmark
from sglang.kernels.ops.quantization.awq_dequantize import (
awq_dequantize as jit_awq_dequantize,
)
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.utils import is_in_ci
register_cuda_ci(
est_time=5, stage="base-b-kernel-benchmark", runner_config="1-gpu-large"
)
try:
from sgl_kernel import awq_dequantize as aot_awq_dequantize
AOT_AVAILABLE = True
except ImportError:
AOT_AVAILABLE = False
IS_CI = is_in_ci()
if IS_CI:
qweight_row_range = [128]
qweight_cols_range = [16]
else:
qweight_row_range = [128, 256, 512, 1024, 3584]
qweight_cols_range = [16, 32, 64, 128, 448]
configs = list(itertools.product(qweight_row_range, qweight_cols_range))
def check_correctness():
if not AOT_AVAILABLE:
print("sgl_kernel AOT not available, skipping correctness check")
return
qweight_row, qweight_col = 128, 16
device = torch.device("cuda")
qweight = torch.randint(
0,
torch.iinfo(torch.int32).max,
(qweight_row, qweight_col),
dtype=torch.int32,
device=device,
)
group_size = qweight_row
scales_row = qweight_row // group_size
scales_col = qweight_col * 8
scales = torch.rand(scales_row, scales_col, dtype=torch.float16, device=device)
qzeros = torch.randint(
0,
torch.iinfo(torch.int32).max,
(scales_row, qweight_col),
dtype=torch.int32,
device=device,
)
jit_out = jit_awq_dequantize(qweight, scales, qzeros)
aot_out = aot_awq_dequantize(qweight, scales, qzeros)
torch.cuda.synchronize()
torch.testing.assert_close(jit_out, aot_out, rtol=0, atol=0)
print("Correctness check passed (JIT vs AOT)")
if AOT_AVAILABLE:
line_vals = ["jit", "aot"]
line_names = ["JIT Kernel", "AOT Kernel"]
styles = [("blue", "-"), ("green", "-")]
else:
line_vals = ["jit"]
line_names = ["JIT Kernel"]
styles = [("blue", "-")]
@triton.testing.perf_report(
triton.testing.Benchmark(
x_names=["qweight_row", "qweight_col"],
x_vals=configs,
line_arg="provider",
line_vals=line_vals,
line_names=line_names,
styles=styles,
ylabel="us",
plot_name="awq-dequantize-jit-vs-aot",
args={},
)
)
def benchmark(qweight_row, qweight_col, provider):
device = torch.device("cuda")
qweight = torch.randint(
0,
torch.iinfo(torch.int32).max,
(qweight_row, qweight_col),
dtype=torch.int32,
device=device,
)
group_size = qweight_row
scales_row = qweight_row // group_size
scales_col = qweight_col * 8
scales = torch.rand(scales_row, scales_col, dtype=torch.float16, device=device)
qzeros = torch.randint(
0,
torch.iinfo(torch.int32).max,
(scales_row, qweight_col),
dtype=torch.int32,
device=device,
)
if provider == "jit":
fn = lambda: jit_awq_dequantize(qweight, scales, qzeros)
elif provider == "aot":
fn = lambda: aot_awq_dequantize(qweight, scales, qzeros)
else:
raise ValueError(f"Unknown provider: {provider}")
return run_benchmark(fn)
if __name__ == "__main__":
check_correctness()
benchmark.run(print_data=True)
@@ -0,0 +1,292 @@
from __future__ import annotations
import sys
from typing import Any
import torch
import triton
from sglang.kernels.jit.benchmark.utils import get_benchmark_range, run_benchmark
from sglang.kernels.ops.quantization.mxfp8 import (
es_sm100_mxfp8_blockscaled_grouped_quant,
es_sm100_mxfp8_blockscaled_moe_grouped_gemm,
)
from sglang.test.ci.ci_register import register_cuda_ci
register_cuda_ci(
est_time=5, stage="base-b-kernel-benchmark", runner_config="1-gpu-large"
)
def is_sm100_supported(device=None) -> bool:
if not torch.cuda.is_available():
return False
return (torch.cuda.get_device_capability(device)[0] == 10) and (
torch.version.cuda >= "12.8"
)
_SM100_SUPPORTED = is_sm100_supported()
def _probe_sgl_kernel_group_mm() -> tuple[bool, str]:
if not _SM100_SUPPORTED:
return False, "MXFP8 MoE benchmark requires sm100+ with CUDA 12.8+."
try:
import sgl_kernel # noqa: F401
except Exception as e:
return False, f"import sgl_kernel failed: {e}"
if not hasattr(sgl_kernel, "es_sm100_mxfp8_blockscaled_grouped_mm"):
return False, "sgl_kernel.es_sm100_mxfp8_blockscaled_grouped_mm is missing."
try:
pass
# We assume if it's imported, it works
except Exception as e:
return False, f"calling sgl-kernel grouped_mm op failed: {e}"
return True, ""
_SGL_KERNEL_AVAILABLE, _SGL_KERNEL_REASON = _probe_sgl_kernel_group_mm()
def align(val: int, alignment: int = 128) -> int:
return int((val + alignment - 1) // alignment * alignment)
def _prepare_case(
total_tokens: int, n_g: int, k_g: int, num_experts: int, dtype: torch.dtype
) -> dict[str, Any]:
device = torch.device("cuda")
base = total_tokens // num_experts
rem = total_tokens % num_experts
m_per_expert = [base + (1 if i < rem else 0) for i in range(num_experts)]
expert_offset = 0
expert_offsets = []
aux_expert_offset = 0
aux_expert_offsets = []
a_blockscale_offset = 0
a_blockscale_offsets = []
b_blockscale_offset = 0
b_blockscale_offsets = []
tokens_per_expert_list = []
expert_ranges = []
problem_sizes = []
a_list = []
b_list = []
for g in range(num_experts):
m_g = m_per_expert[g]
tokens_per_expert_list.append(m_g)
expert_ranges.append((expert_offset, expert_offset + m_g))
expert_offsets.append(expert_offset)
expert_offset += m_g
aux_expert_offsets.append(aux_expert_offset)
aux_expert_offset += n_g
a_blockscale_offsets.append(a_blockscale_offset)
a_blockscale_offset += align(m_g, 128)
b_blockscale_offsets.append(b_blockscale_offset)
b_blockscale_offset += n_g # n_g already align to 128 in practice
problem_sizes.append([m_g, n_g, k_g])
a = torch.randn((m_g, k_g), device=device, dtype=dtype) * 0.1
b = torch.randn((n_g, k_g), device=device, dtype=dtype) * 0.1
a_list.append(a)
b_list.append(b)
a = torch.concat(a_list, dim=0)
b = torch.concat(b_list, dim=0)
_expert_offsets = torch.tensor(expert_offsets).to(device=device, dtype=torch.int32)
_aux_expert_offsets = torch.tensor(aux_expert_offsets).to(
device=device, dtype=torch.int32
)
_a_blockscale_offsets = torch.tensor(a_blockscale_offsets).to(
device=device, dtype=torch.int32
)
_b_blockscale_offsets = torch.tensor(b_blockscale_offsets).to(
device=device, dtype=torch.int32
)
_tokens_per_expert = torch.tensor(tokens_per_expert_list).to(
device=device, dtype=torch.int32
)
_problem_sizes = torch.tensor(problem_sizes).to(device=device, dtype=torch.int32)
a_quant = torch.zeros_like(a, dtype=torch.float8_e4m3fn, device=device)
a_scale_factor = torch.zeros(
(a_blockscale_offset, k_g // 32), dtype=torch.uint8, device=device
)
b_quant = torch.zeros_like(b, dtype=torch.float8_e4m3fn, device=device)
b_scale_factor = torch.zeros(
(num_experts * n_g, k_g // 32), dtype=torch.uint8, device=device
)
# Use a global workspace to avoid allocating 1GB every time
workspace = torch.empty((1024, 1024, 1024), dtype=torch.uint8, device=device)
es_sm100_mxfp8_blockscaled_grouped_quant(
a,
_tokens_per_expert,
_expert_offsets,
_a_blockscale_offsets,
a_quant,
a_scale_factor,
)
es_sm100_mxfp8_blockscaled_grouped_quant(
b,
torch.ones_like(_tokens_per_expert) * n_g,
_aux_expert_offsets,
_b_blockscale_offsets,
b_quant,
b_scale_factor,
)
b_quant = b_quant.view(num_experts, n_g, k_g)
b_scale_factor = b_scale_factor.view(num_experts, n_g, k_g // 32)
sgl_b_quant = b_quant.transpose(1, 2)
sgl_b_scale_factor = b_scale_factor.transpose(1, 2)
return {
"a": a,
"b": b.view(num_experts, n_g, k_g),
"b_quant": b_quant,
"a_quant": a_quant,
"b_scale_factor": b_scale_factor,
"a_scale_factor": a_scale_factor,
"expert_offsets": _expert_offsets,
"a_blockscale_offsets": _a_blockscale_offsets,
"tokens_per_expert": _tokens_per_expert,
"problem_sizes": _problem_sizes,
"sgl_b_quant": sgl_b_quant,
"sgl_b_scale_factor": sgl_b_scale_factor,
"workspace": workspace,
"expert_ranges": expert_ranges,
"dtype": dtype,
}
def _sgl_kernel_group_mm(case: dict[str, Any]) -> torch.Tensor:
from sgl_kernel import es_sm100_mxfp8_blockscaled_grouped_mm
a_quant = case["a_quant"]
sgl_b_quant = case["sgl_b_quant"]
a_scale_factor = case["a_scale_factor"]
sgl_b_scale_factor = case["sgl_b_scale_factor"]
problem_sizes = case["problem_sizes"]
expert_offsets = case["expert_offsets"]
a_blockscale_offsets = case["a_blockscale_offsets"]
dtype = case["dtype"]
total_tokens = a_quant.shape[0]
n_g = sgl_b_quant.shape[2]
# sgl-kernel takes output pre-allocated
d = torch.empty((total_tokens, n_g), device=a_quant.device, dtype=dtype)
es_sm100_mxfp8_blockscaled_grouped_mm(
d,
a_quant,
sgl_b_quant,
a_scale_factor,
sgl_b_scale_factor,
problem_sizes,
expert_offsets,
a_blockscale_offsets,
)
return d
shape_range = get_benchmark_range(
full_range=[
# (total_tokens, n_g, k_g, num_experts)
(1024, 4096, 4096, 64),
(2048, 4096, 4096, 64),
(4096, 4096, 4096, 64),
]
+ [
(total_tokens, n_g, k_g, num_experts)
for total_tokens in [32 * (2**i) for i in range(9)] # 32 to 8192
for n_g, k_g, num_experts in [
# DeepSeek-V3/R1, gateup, TP = 1, EP = 8
(4096, 7168, 32),
# DeepSeek-V3/R1, down, TP = 1, EP = 8
(7168, 2048, 32),
]
],
ci_range=[(1024, 2048, 2048, 8)],
)
line_vals = ["jit"]
line_names = ["JIT MXFP8 MoE GroupMM"]
styles = [("green", "-")]
if _SGL_KERNEL_AVAILABLE:
line_vals.append("sgl_kernel")
line_names.append("sgl-kernel MXFP8 MoE GroupMM")
styles.append(("orange", "-"))
@triton.testing.perf_report(
triton.testing.Benchmark(
x_names=["total_tokens", "n_g", "k_g", "num_experts"],
x_vals=shape_range,
x_log=False,
line_arg="provider",
line_vals=line_vals,
line_names=line_names,
styles=styles,
ylabel="us",
plot_name="mxfp8-moe-groupmm-performance",
args={},
)
)
def benchmark(total_tokens, n_g, k_g, num_experts, provider):
case = _prepare_case(total_tokens, n_g, k_g, num_experts, torch.bfloat16)
if provider == "jit":
fn = lambda: es_sm100_mxfp8_blockscaled_moe_grouped_gemm(
case["b_quant"],
case["a_quant"],
case["b_scale_factor"],
case["a_scale_factor"],
case["expert_offsets"],
case["a_blockscale_offsets"],
case["tokens_per_expert"],
case["workspace"],
case["dtype"],
)
elif provider == "sgl_kernel":
fn = lambda: _sgl_kernel_group_mm(case)
else:
raise ValueError(f"Unknown provider: {provider}")
# Warm up
fn()
# Profile
if provider == "jit":
torch.cuda.nvtx.range_push("jit")
fn()
torch.cuda.nvtx.range_pop()
elif provider == "sgl_kernel":
torch.cuda.nvtx.range_push("sgl_kernel")
fn()
torch.cuda.nvtx.range_pop()
return run_benchmark(fn)
if __name__ == "__main__":
if not _SM100_SUPPORTED:
print("[skip] MXFP8 MoE GroupMM benchmark requires sm100+ with CUDA 12.8+.")
sys.exit(0)
if not _SGL_KERNEL_AVAILABLE:
print(f"[info] sgl-kernel baseline unavailable: {_SGL_KERNEL_REASON}")
benchmark.run(print_data=True)
@@ -0,0 +1,124 @@
from typing import Optional, Tuple
import torch
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 (
per_tensor_quant_fp8,
)
from sglang.test.ci.ci_register import register_cuda_ci
register_cuda_ci(
est_time=5, stage="base-b-kernel-benchmark", runner_config="1-gpu-large"
)
try:
from vllm import _custom_ops as ops
VLLM_AVAILABLE = True
except ImportError:
ops = None
VLLM_AVAILABLE = False
try:
from sglang.srt.utils import is_hip
_is_hip = is_hip()
except ImportError:
_is_hip = False
fp8_type_ = torch.float8_e4m3fnuz if _is_hip else torch.float8_e4m3fn
def vllm_scaled_fp8_quant(
input: torch.Tensor,
scale: Optional[torch.Tensor] = None,
) -> Tuple[torch.Tensor, torch.Tensor]:
if not VLLM_AVAILABLE:
return sglang_scaled_fp8_quant(input, scale)
return ops.scaled_fp8_quant(input, scale)
def sglang_scaled_fp8_quant(
input: torch.Tensor,
scale: Optional[torch.Tensor] = None,
) -> Tuple[torch.Tensor, torch.Tensor]:
fp8_type_: torch.dtype = torch.float8_e4m3fn
output = torch.empty_like(input, device=input.device, dtype=fp8_type_)
is_static = True
if scale is None:
scale = torch.zeros(1, device=input.device, dtype=torch.float32)
is_static = False
per_tensor_quant_fp8(input, output, scale, is_static)
return output, scale
def calculate_diff(batch_size: int, seq_len: int):
device = torch.device("cuda")
x = torch.rand((batch_size, seq_len), dtype=torch.bfloat16, device=device)
if not VLLM_AVAILABLE:
print("vLLM not available, skipping comparison")
return
vllm_out, vllm_scale = vllm_scaled_fp8_quant(x)
sglang_out, sglang_scale = sglang_scaled_fp8_quant(x)
vllm_out = vllm_out.to(torch.float32)
sglang_out = sglang_out.to(torch.float32)
triton.testing.assert_close(vllm_out, sglang_out, rtol=1e-3, atol=1e-3)
triton.testing.assert_close(vllm_scale, sglang_scale, rtol=1e-3, atol=1e-3)
# Benchmark configuration
element_range = get_benchmark_range(
full_range=[2**n for n in range(10, 20)],
ci_range=[16384],
)
if VLLM_AVAILABLE:
line_vals = ["vllm", "sglang"]
line_names = ["VLLM", "SGL Kernel"]
styles = [("blue", "-"), ("green", "-")]
else:
line_vals = ["sglang"]
line_names = ["SGL Kernel"]
styles = [("green", "-")]
@triton.testing.perf_report(
triton.testing.Benchmark(
x_names=["element_count"],
x_vals=element_range,
line_arg="provider",
line_vals=line_vals,
line_names=line_names,
styles=styles,
ylabel="us",
plot_name="per-tensor-quant-fp8-performance",
args={},
)
)
def benchmark(element_count, provider):
dtype = torch.float16
device = torch.device("cuda")
x = torch.randn(element_count, 4096, device=device, dtype=dtype)
if provider == "vllm":
fn = lambda: vllm_scaled_fp8_quant(x.clone())
elif provider == "sglang":
fn = lambda: sglang_scaled_fp8_quant(x.clone())
else:
raise ValueError(f"Unknown provider: {provider}")
return run_benchmark(fn)
if __name__ == "__main__":
calculate_diff(batch_size=4, seq_len=4096)
benchmark.run(print_data=True)
@@ -0,0 +1,78 @@
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,
)
from sglang.test.ci.ci_register import register_cuda_ci
register_cuda_ci(
est_time=25, stage="base-b-kernel-benchmark", runner_config="1-gpu-large"
)
HIDDEN = 2048
LAYOUTS = {
"row_major_fp32": (False, False),
"col_major_fp32": (True, False),
"col_major_ue8m0": (True, True),
}
def _jit_v2(G, x, x_q, x_s, scale_ue8m0):
per_token_group_quant_8bit_v2(
x,
x_q,
x_s,
G,
1e-10,
float(fp8_min),
float(fp8_max),
scale_ue8m0=scale_ue8m0,
)
def _current(G, x, x_q, x_s, scale_ue8m0):
per_token_group_quant(x, x_q, x_s, G, scale_ue8m0=scale_ue8m0)
FN = {"jit_v2": _jit_v2, "current": _current}
@marker.parametrize("group_size", [32, 64, 128], ci_vals=[128])
@marker.parametrize("layout", list(LAYOUTS), ci_vals=["col_major_ue8m0"])
@marker.parametrize("num_tokens", [2**n for n in range(0, 14)], ci_vals=[1, 32, 2048])
@marker.benchmark("impl", ["jit_v2", "current"])
def benchmark(group_size: int, layout: str, num_tokens: int, impl: str):
column_major, scale_ue8m0 = LAYOUTS[layout]
x = create_random(num_tokens, HIDDEN)
x_q = create_empty(num_tokens, HIDDEN, dtype=fp8_dtype)
x_s = create_per_token_group_quant_fp8_output_scale(
x_shape=(num_tokens, HIDDEN),
device="cuda",
group_size=group_size,
column_major_scales=column_major,
scale_tma_aligned=column_major,
scale_ue8m0=scale_ue8m0,
)
return marker.do_bench(
FN[impl],
input_args=(group_size, x, x_q, x_s, scale_ue8m0),
graph_clone_args=(1,),
memory_args=(x,),
memory_output=(x_q, x_s),
)
if __name__ == "__main__":
benchmark.run()
@@ -0,0 +1,67 @@
import torch
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.test.ci.ci_register import register_cuda_ci
register_cuda_ci(
est_time=6, stage="base-b-kernel-benchmark", runner_config="1-gpu-large"
)
G = 128
HIDDEN = 4096
def _aot_v2(x, x_q, x_s):
# Low-level AOT op writing into the same preallocated x_q/x_s as the JIT
# path, so this is a kernel-vs-kernel comparison (no wrapper / no realloc).
sgl_per_token_group_quant_8bit(
x,
x_q,
x_s,
G,
1e-10,
float(fp8_min),
float(fp8_max),
False, # scale_ue8m0
False, # fuse_silu_and_mul
None, # masked_m
enable_v2=True,
)
def _jit_v2(x, x_q, x_s):
per_token_group_quant_8bit_v2(x, x_q, x_s, G, 1e-10, float(fp8_min), float(fp8_max))
FN = {"aot_v2": _aot_v2, "jit_v2": _jit_v2}
@marker.parametrize("num_tokens", [1, 8, 64, 512, 4096], ci_vals=[1, 512])
@marker.benchmark("impl", ["aot_v2", "jit_v2"])
def benchmark(num_tokens: int, impl: str):
x = create_random(num_tokens, HIDDEN)
x_q = torch.empty(num_tokens, HIDDEN, device="cuda", dtype=fp8_dtype)
x_s = create_per_token_group_quant_fp8_output_scale(
x_shape=(num_tokens, HIDDEN),
device="cuda",
group_size=G,
column_major_scales=True,
scale_tma_aligned=True,
scale_ue8m0=False,
)
return marker.do_bench(FN[impl], input_args=(x, x_q, x_s), graph_clone_args=(0,))
if __name__ == "__main__":
benchmark.run()
@@ -0,0 +1,123 @@
import math
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,
)
from sglang.test.ci.ci_register import register_cuda_ci
register_cuda_ci(
est_time=25, stage="base-b-kernel-benchmark", runner_config="1-gpu-large"
)
# name -> (moe_intermediate_size, topk, num_experts, group_size)
MODELS = {
"deepseek_v4": (3072, 6, 384, 32), # DeepSeek-V4 Pro
"deepseek_v3": (2048, 8, 256, 128), # DeepSeek-V3/R1
"qwen3_235b": (1536, 8, 128, 128), # Qwen3-235B-A22B
}
def _jit_v2(G, x, x_q, x_s, masked_m, expected_m, fuse):
per_token_group_quant_8bit_v2(
x,
x_q,
x_s,
G,
1e-10,
float(fp8_min),
float(fp8_max),
scale_ue8m0=True,
fuse_silu_and_mul=fuse,
masked_m=masked_m,
)
def _current(G, x, x_q, x_s, masked_m, expected_m, fuse):
per_token_group_quant(
x,
x_q,
x_s,
G,
scale_ue8m0=True,
fuse_silu_and_mul=fuse,
masked_m=masked_m,
expected_m=expected_m,
)
FN = {"jit_v2": _jit_v2, "current": _current}
@marker.parametrize("model", list(MODELS), ci_vals=["deepseek_v3"])
@marker.parametrize("num_gpus", [4, 8], ci_vals=[4])
@marker.parametrize("fuse_silu", [True, False], ci_vals=[False])
@marker.parametrize("balanced", [True, False], ci_vals=[True])
@marker.parametrize("num_tokens", [2**n for n in range(8)], ci_vals=[1, 128])
@marker.benchmark("impl", ["jit_v2", "current"], unit="us")
def benchmark(
model: str,
fuse_silu: bool,
num_gpus: int,
num_tokens: int,
balanced: bool,
impl: str,
) -> marker.BenchResult:
torch.cuda.random.manual_seed(42)
max_tokens = 128 # TODO: test other size
hidden_size, topk, num_experts, group_size = MODELS[model]
if num_experts % num_gpus != 0 or topk * num_gpus > num_experts:
marker.skip("Incompatible model configuration")
if impl == "jit_v2" and (hidden_size // group_size) % 16 != 0:
marker.skip("v2 masked requires num_groups % 16 == 0")
if num_tokens > max_tokens:
marker.skip("num_tokens exceeds max_tokens")
num_local_experts = num_experts // num_gpus
padded_tokens = max_tokens * num_gpus
expected_m = math.ceil(max_tokens * topk / num_local_experts)
in_hidden = hidden_size * (2 if fuse_silu else 1)
x = create_random(num_local_experts, padded_tokens, in_hidden)
x_q = create_empty(num_local_experts, padded_tokens, hidden_size, dtype=fp8_dtype)
x_s = create_per_token_group_quant_fp8_output_scale(
x_shape=(num_local_experts, padded_tokens, hidden_size),
device="cuda",
group_size=group_size,
column_major_scales=True,
scale_tma_aligned=True,
scale_ue8m0=True,
)
if balanced: # simulation
topk_ids = torch.randint(0, num_local_experts, (num_tokens * topk,))
masked_m = torch.bincount(topk_ids, minlength=num_local_experts)
masked_m = masked_m.cuda().int()
else: # only the last few experts receive all tokens
masked_m = create_empty(num_local_experts, dtype=torch.int32)
masked_m[:-topk].zero_()
masked_m[-topk:].fill_(num_tokens)
return marker.do_bench(
FN[impl],
input_args=(group_size, x, x_q, x_s, masked_m, expected_m, fuse_silu),
graph_clone_args=(0,),
memory_args=(x[:topk, :num_tokens], masked_m),
memory_output=(x_q[:topk, :num_tokens], x_s[:topk, :num_tokens]),
)
if __name__ == "__main__":
benchmark.run()
@@ -0,0 +1,127 @@
import torch
import triton
import triton.testing
from sglang.kernels.jit.benchmark.utils import (
DEFAULT_DEVICE,
get_benchmark_range,
run_benchmark_no_cudagraph,
)
from sglang.kernels.ops.speculative.ngram_embedding import (
compute_n_gram_ids,
compute_n_gram_ids_decode,
)
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
register_cuda_ci(
est_time=15, stage="base-b-kernel-benchmark", runner_config="1-gpu-large"
)
register_amd_ci(est_time=15, stage="jit-kernel-benchmark", runner_config="amd")
NE_N = 8
NE_K = 2
VOCAB_SIZE = 32000
EOS_TOKEN_ID = VOCAB_SIZE
MAX_CONTEXT_LEN = 1024
BATCH_SIZE_LIST = get_benchmark_range(
full_range=[1, 2, 8, 32, 128, 512, 1024, 2048, 4096],
ci_range=[32, 1024],
)
def _make_ngram_params():
ne_weights = torch.zeros([NE_N - 1, NE_K, NE_N], dtype=torch.int32)
ne_mods = torch.zeros([NE_N - 1, NE_K], dtype=torch.int32)
exclusive_sums = torch.zeros([(NE_N - 1) * NE_K + 1], dtype=torch.int32)
for n in range(2, NE_N + 1):
for k in range(NE_K):
config_id = (n - 2) * NE_K + k
mod = 65537 + 2 * config_id
ne_mods[n - 2][k] = mod
exclusive_sums[config_id + 1] = exclusive_sums[config_id] + mod
for delta in range(NE_N):
ne_weights[n - 2][k][delta] = pow(VOCAB_SIZE, delta, mod)
return (
ne_weights.to(DEFAULT_DEVICE),
ne_mods.to(DEFAULT_DEVICE),
exclusive_sums.to(DEFAULT_DEVICE),
)
@triton.testing.perf_report(
triton.testing.Benchmark(
x_names=["batch_size"],
x_vals=BATCH_SIZE_LIST,
line_arg="provider",
line_vals=["general", "decode"],
line_names=["general compute_n_gram_ids", "decode fast path"],
styles=[("blue", "-"), ("orange", "-")],
ylabel="us",
plot_name="ngram-compute-decode",
args={},
)
)
def benchmark(batch_size: int, provider: str):
num_configs = (NE_N - 1) * NE_K
max_running_reqs = batch_size + 8
ne_weights, ne_mods, exclusive_sums = _make_ngram_params()
ne_token_table = torch.randint(
0,
VOCAB_SIZE,
(max_running_reqs, MAX_CONTEXT_LEN),
dtype=torch.int32,
device=DEFAULT_DEVICE,
)
row_indices = torch.arange(batch_size, dtype=torch.int64, device=DEFAULT_DEVICE)
column_starts = torch.randint(
0, MAX_CONTEXT_LEN, (batch_size,), dtype=torch.int32, device=DEFAULT_DEVICE
)
n_gram_ids = torch.empty(
(batch_size, num_configs), dtype=torch.int32, device=DEFAULT_DEVICE
)
if provider == "general":
tokens = torch.empty(batch_size, dtype=torch.int32, device=DEFAULT_DEVICE)
exclusive_req_len_sums = torch.arange(
batch_size + 1, dtype=torch.int32, device=DEFAULT_DEVICE
)
def fn():
compute_n_gram_ids(
NE_N,
NE_K,
ne_weights,
ne_mods,
exclusive_sums,
tokens,
exclusive_req_len_sums,
ne_token_table,
row_indices,
column_starts,
n_gram_ids,
EOS_TOKEN_ID,
)
else:
def fn():
compute_n_gram_ids_decode(
NE_N,
NE_K,
ne_weights,
ne_mods,
exclusive_sums,
ne_token_table,
row_indices,
column_starts,
n_gram_ids,
EOS_TOKEN_ID,
)
return run_benchmark_no_cudagraph(fn)
if __name__ == "__main__":
benchmark.run(print_data=True)
@@ -0,0 +1,79 @@
import torch
import triton
import triton.testing
from sglang.kernels.jit.benchmark.utils import (
DEFAULT_DEVICE,
get_benchmark_range,
run_benchmark_no_cudagraph,
)
from sglang.kernels.ops.speculative.ngram_embedding import (
update_token_table,
update_token_table_decode,
)
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
register_cuda_ci(
est_time=15, stage="base-b-kernel-benchmark", runner_config="1-gpu-large"
)
register_amd_ci(est_time=15, stage="jit-kernel-benchmark", runner_config="amd")
MAX_CONTEXT_LEN = 4096
BATCH_SIZE_LIST = get_benchmark_range(
full_range=[1, 2, 8, 32, 128, 512, 1024, 2048, 4096],
ci_range=[32, 1024],
)
@triton.testing.perf_report(
triton.testing.Benchmark(
x_names=["batch_size"],
x_vals=BATCH_SIZE_LIST,
line_arg="provider",
line_vals=["general", "decode"],
line_names=["general update_token_table", "decode fast path"],
styles=[("blue", "-"), ("orange", "-")],
ylabel="us",
plot_name="ngram-update-token-table",
args={},
)
)
def benchmark(batch_size: int, provider: str):
max_running_reqs = batch_size + 8
tokens = torch.arange(batch_size, dtype=torch.int32, device=DEFAULT_DEVICE)
token_table = torch.empty(
(max_running_reqs, MAX_CONTEXT_LEN), dtype=torch.int32, device=DEFAULT_DEVICE
)
row_indices = torch.arange(batch_size, dtype=torch.int64, device=DEFAULT_DEVICE)
column_starts = torch.randint(
0, MAX_CONTEXT_LEN, (batch_size,), dtype=torch.int32, device=DEFAULT_DEVICE
)
req_lens = torch.ones(batch_size, dtype=torch.int32, device=DEFAULT_DEVICE)
if provider == "general":
def fn():
update_token_table(
tokens,
token_table,
row_indices,
column_starts,
req_lens,
None,
)
else:
def fn():
update_token_table_decode(
tokens,
token_table,
row_indices,
column_starts,
)
return run_benchmark_no_cudagraph(fn)
if __name__ == "__main__":
benchmark.run(print_data=True)
@@ -0,0 +1,77 @@
import itertools
import torch
import triton
import triton.testing
from sglang.kernels.jit.benchmark.utils import (
DEFAULT_DEVICE,
get_benchmark_range,
run_benchmark,
)
from sglang.kernels.ops.speculative.resolve_future_token_ids import (
resolve_future_token_ids_cuda,
)
from sglang.srt.utils import get_compiler_backend
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
register_cuda_ci(
est_time=10, stage="base-b-kernel-benchmark", runner_config="1-gpu-large"
)
register_amd_ci(est_time=10, stage="jit-kernel-benchmark", runner_config="amd")
SIZE_LIST = get_benchmark_range(
full_range=[2**n for n in range(4, 16)], # 16 … 32K elements
ci_range=[256, 4096],
)
configs = list(itertools.product(SIZE_LIST))
def _torch_resolve(input_ids, future_map):
input_ids[:] = torch.where(
input_ids < 0,
future_map[torch.clamp(-input_ids, min=0)],
input_ids,
)
_compiled_resolve = torch.compile(
_torch_resolve, dynamic=True, backend=get_compiler_backend()
)
@triton.testing.perf_report(
triton.testing.Benchmark(
x_names=["size"],
x_vals=configs,
line_arg="provider",
line_vals=["jit", "torch_compile", "torch"],
line_names=["SGL JIT Kernel", "torch.compile", "PyTorch"],
styles=[("blue", "-"), ("green", "-."), ("red", "--")],
ylabel="us",
plot_name="resolve-future-token-ids-performance",
args={},
)
)
def benchmark(size: int, provider: str):
map_size = 8192
future_map = torch.randint(
0, 50000, (map_size,), dtype=torch.int64, device=DEFAULT_DEVICE
)
input_ids = torch.randint(
-map_size + 1, 50000, (size,), dtype=torch.int64, device=DEFAULT_DEVICE
)
if provider == "jit":
fn = lambda: resolve_future_token_ids_cuda(input_ids.clone(), future_map)
elif provider == "torch_compile":
fn = lambda: _compiled_resolve(input_ids.clone(), future_map)
else:
fn = lambda: _torch_resolve(input_ids.clone(), future_map)
return run_benchmark(fn)
if __name__ == "__main__":
benchmark.run(print_data=True)
@@ -0,0 +1,166 @@
"""Benchmark CUDA topk=1 speculative decoding helpers."""
from __future__ import annotations
import torch
import triton
import triton.testing
from sglang.kernels.jit.benchmark.utils import (
DEFAULT_DEVICE,
get_benchmark_range,
run_benchmark,
)
from sglang.kernels.ops.speculative.topk1 import draft_topk1_postprocess
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
register_cuda_ci(
est_time=30, stage="base-b-kernel-benchmark", runner_config="1-gpu-large"
)
register_amd_ci(est_time=30, stage="jit-kernel-benchmark", runner_config="amd")
BATCH_SIZE_RANGE = get_benchmark_range(
full_range=[1, 2, 4, 8, 16, 32, 64, 128, 256, 512, 1024, 2048],
ci_range=[1, 16, 256, 2048],
)
VOCAB_SIZES = {
"dsv4": 129280,
"glm5_2": 151552,
}
VOCAB_SIZE_RANGE = get_benchmark_range(
full_range=list(VOCAB_SIZES.values()),
ci_range=list(VOCAB_SIZES.values()),
)
NUM_STEPS = 3
def make_logits(batch_size: int, vocab_size: int) -> torch.Tensor:
logits = torch.zeros(
(batch_size, vocab_size), dtype=torch.float32, device=DEFAULT_DEVICE
)
max_index = (
torch.arange(batch_size, dtype=torch.long, device=DEFAULT_DEVICE) * 9973 + 17
) % vocab_size
logits.scatter_(1, max_index[:, None], 1000.0)
return logits
def make_draft_case(batch_size: int, vocab_size: int):
logits = make_logits(batch_size, vocab_size)
positions = torch.zeros(batch_size, dtype=torch.long, device=DEFAULT_DEVICE)
return logits, positions
def make_chain_case(batch_size: int, vocab_size: int):
seed_topk_index = torch.randint(
0, vocab_size, (batch_size, 1), dtype=torch.long, device=DEFAULT_DEVICE
)
logits = [make_logits(batch_size, vocab_size) for _ in range(NUM_STEPS - 1)]
positions = torch.zeros(batch_size, dtype=torch.long, device=DEFAULT_DEVICE)
return seed_topk_index, logits, positions
def eager_draft_topk1_postprocess(logits: torch.Tensor, positions: torch.Tensor):
topk_index = torch.argmax(logits, dim=-1, keepdim=True)
topk_p = torch.ones_like(topk_index, dtype=torch.float32)
positions.add_(1)
return topk_p, topk_index
def fused_draft_topk1_postprocess(logits: torch.Tensor, positions: torch.Tensor):
return draft_topk1_postprocess(logits, positions)
def eager_chain_materialize(
seed_topk_index: torch.Tensor,
logits: list[torch.Tensor],
positions: torch.Tensor,
):
token_list = [seed_topk_index]
for step_logits in logits:
_, topk_index = eager_draft_topk1_postprocess(step_logits, positions)
token_list.append(topk_index)
return torch.cat(token_list, dim=1)
def fused_chain_materialize(
seed_topk_index: torch.Tensor,
logits: list[torch.Tensor],
positions: torch.Tensor,
):
draft_tokens = torch.empty(
(seed_topk_index.shape[0], NUM_STEPS),
dtype=torch.long,
device=DEFAULT_DEVICE,
)
draft_tokens[:, :1].copy_(seed_topk_index)
for i, step_logits in enumerate(logits, start=1):
draft_topk1_postprocess(
step_logits,
positions,
draft_tokens,
draft_token_column=i,
)
return draft_tokens
@triton.testing.perf_report(
triton.testing.Benchmark(
x_names=["batch_size", "vocab_size"],
x_vals=[(bs, vocab) for bs in BATCH_SIZE_RANGE for vocab in VOCAB_SIZE_RANGE],
line_arg="provider",
line_vals=["fused", "eager"],
line_names=["Fused Triton", "Eager torch"],
styles=[("blue", "-"), ("orange", "--")],
ylabel="us",
plot_name="spec-topk1-draft-postprocess",
args={},
)
)
def benchmark_draft_postprocess(
batch_size: int, vocab_size: int, provider: str
) -> tuple[float, float, float]:
logits, positions = make_draft_case(batch_size, vocab_size)
if provider == "fused":
fn = lambda: fused_draft_topk1_postprocess(logits, positions)
elif provider == "eager":
fn = lambda: eager_draft_topk1_postprocess(logits, positions)
else:
raise ValueError(f"Unknown provider: {provider}")
fn()
torch.cuda.synchronize()
return run_benchmark(fn)
@triton.testing.perf_report(
triton.testing.Benchmark(
x_names=["batch_size", "vocab_size"],
x_vals=[(bs, vocab) for bs in BATCH_SIZE_RANGE for vocab in VOCAB_SIZE_RANGE],
line_arg="provider",
line_vals=["fused", "eager"],
line_names=["Fused Triton", "Eager argmax + cat"],
styles=[("blue", "-"), ("orange", "--")],
ylabel="us",
plot_name="spec-topk1-chain-materialize",
args={},
)
)
def benchmark_chain_materialize(
batch_size: int, vocab_size: int, provider: str
) -> tuple[float, float, float]:
seed_topk_index, logits, positions = make_chain_case(batch_size, vocab_size)
if provider == "fused":
fn = lambda: fused_chain_materialize(seed_topk_index, logits, positions)
elif provider == "eager":
fn = lambda: eager_chain_materialize(seed_topk_index, logits, positions)
else:
raise ValueError(f"Unknown provider: {provider}")
fn()
torch.cuda.synchronize()
return run_benchmark(fn)
if __name__ == "__main__":
benchmark_draft_postprocess.run(print_data=True)
benchmark_chain_materialize.run(print_data=True)