[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:
co-authored by
Claude Opus 4.8
parent
a2935ce329
commit
2d1a7be8c4
@@ -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)
|
||||
Reference in New Issue
Block a user