flashinfer / wrapperda7954
flashinfer_wrapper_da7954 · FlashInfer-Bench baselines · python · Apache-2.0
Kernel source · 94 lines ↓holds 20 records
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 94 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-flashinfer-wrapper-da7954?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 h32 kv4 d128 ps64bf16 · [1, 32, 128] · num_pages=122 · num_kv_indices=58
NVIDIA B200
19.5µs
#1 of 1
2026-04-09
GQA paged decode h32 kv4 d128 ps64bf16 · [1, 32, 128] · num_pages=115 · num_kv_indices=51
NVIDIA B200
19.5µs
#1 of 1
2026-04-09
GQA paged decode h32 kv4 d128 ps64bf16 · [2, 32, 128] · num_pages=1200 · num_kv_indices=100
NVIDIA B200
23.6µs
#1 of 1
2026-04-09
GQA paged decode h32 kv4 d128 ps64bf16 · [4, 32, 128] · num_pages=4013 · num_kv_indices=201
NVIDIA B200
27.7µs
#1 of 1
2026-04-09
GQA paged decode h32 kv4 d128 ps64bf16 · [1, 32, 128] · num_pages=353 · num_kv_indices=289
NVIDIA B200
35.9µs
#1 of 1
2026-04-09
GQA paged decode h32 kv4 d128 ps64bf16 · [8, 32, 128] · num_pages=7795 · num_kv_indices=419
NVIDIA B200
44.1µs
#1 of 1
2026-04-09
GQA paged decode h32 kv4 d128 ps64bf16 · [2, 32, 128] · num_pages=1638 · num_kv_indices=500
NVIDIA B200
52.2µs
#1 of 1
2026-04-09
GQA paged decode h32 kv4 d128 ps64bf16 · [16, 32, 128] · num_pages=11058 · num_kv_indices=808
NVIDIA B200
74.8µs
#1 of 1
2026-04-09
GQA paged decode h32 kv4 d128 ps64bf16 · [1, 32, 128] · num_pages=1136 · num_kv_indices=816
NVIDIA B200
78.9µs
#1 of 1
2026-04-09
Show all 20 measurements ›Showing all 20 measurements ⌄
GQA paged decode h32 kv4 d128 ps64bf16 · [4, 32, 128] · num_pages=5265 · num_kv_indices=1280
NVIDIA B200
105.5µs
#1 of 1
2026-04-09
GQA paged decode h32 kv4 d128 ps64bf16 · [32, 32, 128] · num_pages=17008 · num_kv_indices=1645
NVIDIA B200
136.2µs
#1 of 1
2026-04-09
GQA paged decode h32 kv4 d128 ps64bf16 · [2, 32, 128] · num_pages=3848 · num_kv_indices=2292
NVIDIA B200
179.3µs
#1 of 1
2026-04-09
GQA paged decode h32 kv4 d128 ps64bf16 · [8, 32, 128] · num_pages=8227 · num_kv_indices=2426
NVIDIA B200
195.6µs
#1 of 1
2026-04-09
GQA paged decode h32 kv4 d128 ps64bf16 · [4, 32, 128] · num_pages=7489 · num_kv_indices=3334
NVIDIA B200
259.1µs
#1 of 1
2026-04-09
GQA paged decode h32 kv4 d128 ps64bf16 · [16, 32, 128] · num_pages=11709 · num_kv_indices=5106
NVIDIA B200
378.0µs
#1 of 1
2026-04-09
GQA paged decode h32 kv4 d128 ps64bf16 · [8, 32, 128] · num_pages=10568 · num_kv_indices=6728
NVIDIA B200
500.6µs
#1 of 1
2026-04-09
GQA paged decode h32 kv4 d128 ps64bf16 · [32, 32, 128] · num_pages=17092 · num_kv_indices=10156
NVIDIA B200
764.9µs
#1 of 1
2026-04-09
GQA paged decode h32 kv4 d128 ps64bf16 · [16, 32, 128] · num_pages=16065 · num_kv_indices=13578
NVIDIA B200
965.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:6146d881a59895ae79e81b5a755b9317d404c6b04a3505418af27f4e3fe496e1
license declaredApache-2.0
license concludedApache-2.0
authorsbaseline
imported2026-08-16
Kernel source
main.py94 lines
import torch
import flashinfer
_WORKSPACE_SIZE_BYTES = 128 * 1024 * 1024
_workspace_cache = {}
_wrapper_cache = {}
_plan_state = {}
def _get_workspace(device):
key = str(device)
buffer = _workspace_cache.get(key)
if buffer is None or buffer.device != device or buffer.numel() < _WORKSPACE_SIZE_BYTES:
buffer = torch.empty(_WORKSPACE_SIZE_BYTES, dtype=torch.uint8, device=device)
_workspace_cache[key] = buffer
return buffer
def _get_wrapper(key, device):
wrapper = _wrapper_cache.get(key)
if wrapper is None:
workspace = _get_workspace(device)
wrapper = flashinfer.BatchDecodeWithPagedKVCacheWrapper(
workspace,
kv_layout="NHD",
)
_wrapper_cache[key] = wrapper
return wrapper
def run(q, k_cache, v_cache, kv_indptr, kv_indices, kv_last_page_len, sm_scale):
batch_size, num_qo_heads, head_dim = q.shape
_, page_size, num_kv_heads, _ = k_cache.shape
len_indptr = kv_indptr.shape[0]
num_kv_indices = kv_indices.shape[0]
device = q.device
wrapper_key = (
str(device),
num_qo_heads,
num_kv_heads,
head_dim,
page_size,
q.dtype,
k_cache.dtype,
)
wrapper = _get_wrapper(wrapper_key, device)
state = _plan_state.get(wrapper_key)
needs_plan = True
if state is not None:
needs_plan = (
state.get("batch_size") != batch_size
or state.get("len_indptr") != len_indptr
or state.get("num_kv_indices") != num_kv_indices
or state.get("sm_scale") != sm_scale
or state.get("kv_indptr_ptr") != kv_indptr.data_ptr()
or state.get("kv_indices_ptr") != kv_indices.data_ptr()
or state.get("last_page_ptr") != kv_last_page_len.data_ptr()
)
if needs_plan:
wrapper.plan(
indptr=kv_indptr,
indices=kv_indices,
last_page_len=kv_last_page_len,
num_qo_heads=num_qo_heads,
num_kv_heads=num_kv_heads,
head_dim=head_dim,
page_size=page_size,
pos_encoding_mode="NONE",
q_data_type=q.dtype,
kv_data_type=k_cache.dtype,
sm_scale=sm_scale,
)
_plan_state[wrapper_key] = {
"batch_size": batch_size,
"len_indptr": len_indptr,
"num_kv_indices": num_kv_indices,
"sm_scale": sm_scale,
"kv_indptr_ptr": kv_indptr.data_ptr(),
"kv_indices_ptr": kv_indices.data_ptr(),
"last_page_ptr": kv_last_page_len.data_ptr(),
}
output, lse = wrapper.run(
q,
(k_cache, v_cache),
return_lse=True,
)
return output, lse
scrolls · 94 lines total
Source code from FlashInfer-Bench (flashinfer-ai/flashinfer-trace) · Apache-2.0
Best evidence level for this revision: reproducible
JSON