flashinfer / wrapper96864e
flashinfer_wrapper_96864e · FlashInfer-Bench baselines · python · Apache-2.0
Kernel source · 66 lines ↓holds 20 records
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 66 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-flashinfer-wrapper-96864e?include=source"interfacepython
revisionda915083d4c7
symbolrun
pathmain.py
Compatibility
measured onNVIDIA B200
declared hardwareNVIDIA A100, NVIDIA B200, NVIDIA H100, NVIDIA H20, NVIDIA H200
architecturesunknown
dtypesbf16, fp32, int32
Benchmark evidence
20 measurements across 1 GPU, fastest first.
Operation / workload
Hardware
Latency
Rank
Observed
GQA paged decode h24 kv8 d128 ps1bf16 · [1, 24, 128] · num_pages=14 · num_kv_indices=13
NVIDIA B200
20.4µs
#1 of 1
2026-04-09
GQA paged decode h24 kv8 d128 ps1bf16 · [1, 24, 128] · num_pages=12 · num_kv_indices=11
NVIDIA B200
20.5µs
#1 of 1
2026-04-09
GQA paged decode h24 kv8 d128 ps1bf16 · [1, 24, 128] · num_pages=9 · num_kv_indices=8
NVIDIA B200
20.5µs
#1 of 1
2026-04-09
GQA paged decode h24 kv8 d128 ps1bf16 · [1, 24, 128] · num_pages=18 · num_kv_indices=2
NVIDIA B200
21.5µs
#1 of 1
2026-04-09
GQA paged decode h24 kv8 d128 ps1bf16 · [1, 24, 128] · num_pages=16 · num_kv_indices=15
NVIDIA B200
21.5µs
#1 of 1
2026-04-09
GQA paged decode h24 kv8 d128 ps1bf16 · [1, 24, 128] · num_pages=223 · num_kv_indices=205
NVIDIA B200
25.6µs
#1 of 1
2026-04-09
GQA paged decode h24 kv8 d128 ps1bf16 · [1, 24, 128] · num_pages=839 · num_kv_indices=616
NVIDIA B200
34.8µs
#1 of 1
2026-04-09
GQA paged decode h24 kv8 d128 ps1bf16 · [1, 24, 128] · num_pages=1143 · num_kv_indices=171
NVIDIA B200
37.8µs
#1 of 1
2026-04-09
GQA paged decode h24 kv8 d128 ps1bf16 · [1, 24, 128] · num_pages=2411 · num_kv_indices=648
NVIDIA B200
54.3µs
#1 of 1
2026-04-09
GQA paged decode h24 kv8 d128 ps1bf16 · [1, 24, 128] · num_pages=4799 · num_kv_indices=1220
NVIDIA B200
91.1µs
#1 of 1
2026-04-09
Show all 20 measurements ›Showing all 20 measurements ⌄
GQA paged decode h24 kv8 d128 ps1bf16 · [3, 24, 128] · num_pages=5560 · num_kv_indices=625
NVIDIA B200
102.4µs
#1 of 1
2026-04-09
GQA paged decode h24 kv8 d128 ps1bf16 · [3, 24, 128] · num_pages=7189 · num_kv_indices=1283
NVIDIA B200
127.6µs
#1 of 1
2026-04-09
GQA paged decode h24 kv8 d128 ps1bf16 · [3, 24, 128] · num_pages=11019 · num_kv_indices=2936
NVIDIA B200
188.4µs
#1 of 1
2026-04-09
GQA paged decode h24 kv8 d128 ps1bf16 · [7, 24, 128] · num_pages=12682 · num_kv_indices=1571
NVIDIA B200
208.9µs
#1 of 1
2026-04-09
GQA paged decode h24 kv8 d128 ps1bf16 · [7, 24, 128] · num_pages=16311 · num_kv_indices=3280
NVIDIA B200
268.1µs
#1 of 1
2026-04-09
GQA paged decode h24 kv8 d128 ps1bf16 · [7, 24, 128] · num_pages=26560 · num_kv_indices=9084
NVIDIA B200
428.0µs
#1 of 1
2026-04-09
GQA paged decode h24 kv8 d128 ps1bf16 · [15, 24, 128] · num_pages=32772 · num_kv_indices=6084
NVIDIA B200
504.8µs
#1 of 1
2026-04-09
GQA paged decode h24 kv8 d128 ps1bf16 · [15, 24, 128] · num_pages=42784 · num_kv_indices=9625
NVIDIA B200
659.5µs
#1 of 1
2026-04-09
GQA paged decode h24 kv8 d128 ps1bf16 · [1, 24, 128] · num_pages=61179 · num_kv_indices=1184
NVIDIA B200
897.0µs
#1 of 1
2026-04-09
Reproduction-ready · How evidence levels are derived →
Source and license
sourcehttps://huggingface.co/datasets/flashinfer-ai/flashinfer-trace
commitda915083d4c7c5e61aa3005e3d17ae488e0fc71c
revision digestsha256:33dbdb780e6716cd51f6c48859f9edba1b8da97fe4a128a02f1776982434e254
license declaredApache-2.0
license concludedApache-2.0
authorsbaseline
imported2026-08-16
Kernel source
main.py66 lines
import torch
import flashinfer
# GQA group_size=3 (24 qo_heads / 8 kv_heads) is not a power-of-2, so it is
# unsupported by the FlashInfer decode kernel. Work-around: expand KV heads
# from 8 to 24 (repeat_interleave x3) to make group_size=1 (MHA), which is
# mathematically equivalent. We also use BatchPrefillWithPagedKVCacheWrapper
# with causal=False, treating each decode step as a 1-token prefill per sequence.
_WORKSPACE_SIZE_BYTES = 128 * 1024 * 1024
_workspace_cache = {}
_wrapper_cache = {}
_plan_state = {}
def _get_workspace(device):
key = str(device)
buf = _workspace_cache.get(key)
if buf is None:
buf = torch.empty(_WORKSPACE_SIZE_BYTES, dtype=torch.uint8, device=device)
_workspace_cache[key] = buf
return buf
def _get_wrapper(key, device):
w = _wrapper_cache.get(key)
if w is None:
w = flashinfer.BatchPrefillWithPagedKVCacheWrapper(_get_workspace(device), kv_layout="NHD")
_wrapper_cache[key] = w
return w
def run(q, k_cache, v_cache, kv_indptr, kv_indices, sm_scale):
batch_size, num_qo_heads, head_dim = q.shape
_, page_size, num_kv_heads, _ = k_cache.shape
device = q.device
group_size = num_qo_heads // num_kv_heads # 3
# Expand KV heads: [num_pages, page_size, 8, 128] -> [num_pages, page_size, 24, 128]
k_exp = k_cache.repeat_interleave(group_size, dim=2)
v_exp = v_cache.repeat_interleave(group_size, dim=2)
paged_kv = torch.stack([k_exp, v_exp], dim=1) # [num_pages, 2, page_size, 24, 128]
expanded_heads = num_qo_heads # == 24
wkey = (str(device), num_qo_heads, expanded_heads, head_dim, page_size, q.dtype, k_cache.dtype)
wrapper = _get_wrapper(wkey, device)
state = _plan_state.get(wkey)
needs_plan = state is None or state["batch_size"] != batch_size or state["kv_ptr"] != kv_indptr.data_ptr()
if needs_plan:
qo_indptr = torch.arange(batch_size + 1, dtype=torch.int32, device=device)
wrapper.plan(
qo_indptr=qo_indptr,
paged_kv_indptr=kv_indptr,
paged_kv_indices=kv_indices,
paged_kv_last_page_len=torch.ones(batch_size, dtype=torch.int32, device=device),
num_qo_heads=num_qo_heads,
num_kv_heads=expanded_heads,
head_dim_qk=head_dim,
page_size=page_size,
causal=False,
sm_scale=float(sm_scale),
q_data_type=q.dtype,
kv_data_type=k_cache.dtype,
)
_plan_state[wkey] = {"batch_size": batch_size, "kv_ptr": kv_indptr.data_ptr()}
output, lse = wrapper.run(q, paged_kv, return_lse=True)
return output, lse
scrolls · 66 lines total
Source code from FlashInfer-Bench (flashinfer-ai/flashinfer-trace) · Apache-2.0
Best evidence level for this revision: reproducible
JSON