Skip to content
KernelIndex
Search⌘K

flashinfer_wrapper_bcbabf

FlashInfer-Bench baselines · python · Apache-2.0

Use it

Vendorable · source mirrored · Apache-2.0View source →

No package. Vendor the mirrored source: 75 lines, Apache-2.0, pinned at da91508.

main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-flashinfer-wrapper-bcbabf?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

40 measurements across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
GQA paged decode h20 kv4 d128 ps64bf16 · [1, 20, 128] · num_pages=52 · num_kv_indices=51
NVIDIA B200
95.2µs
#1 of 2
2026-03-28
GQA paged decode h20 kv4 d128 ps64bf16 · [1, 20, 128] · num_pages=59 · num_kv_indices=58
NVIDIA B200
105.4µs
#1 of 2
2026-03-28
GQA paged decode h20 kv4 d128 ps64bf16 · [1, 20, 128] · num_pages=59 · num_kv_indices=58
NVIDIA B200
154.7µs
#2 of 2
2026-03-28
GQA paged decode h20 kv4 d128 ps64bf16 · [1, 20, 128] · num_pages=52 · num_kv_indices=51
NVIDIA B200
171.5µs
#2 of 2
2026-03-28
GQA paged decode h20 kv4 d128 ps64bf16 · [1, 20, 128] · num_pages=327 · num_kv_indices=296
NVIDIA B200
455.7µs
#1 of 2
2026-03-28
GQA paged decode h20 kv4 d128 ps64bf16 · [1, 20, 128] · num_pages=327 · num_kv_indices=296
NVIDIA B200
473.4µs
#2 of 2
2026-03-28
GQA paged decode h20 kv4 d128 ps64bf16 · [2, 20, 128] · num_pages=1166 · num_kv_indices=100
NVIDIA B200
1.44ms
#1 of 2
2026-03-28
GQA paged decode h20 kv4 d128 ps64bf16 · [2, 20, 128] · num_pages=1180 · num_kv_indices=114
NVIDIA B200
1.46ms
#1 of 2
2026-03-28
GQA paged decode h20 kv4 d128 ps64bf16 · [1, 20, 128] · num_pages=1122 · num_kv_indices=823
NVIDIA B200
1.47ms
#1 of 2
2026-03-28
GQA paged decode h20 kv4 d128 ps64bf16 · [2, 20, 128] · num_pages=1166 · num_kv_indices=100
NVIDIA B200
1.49ms
#2 of 2
2026-03-28
Show all 40 measurements ›
GQA paged decode h20 kv4 d128 ps64bf16 · [1, 20, 128] · num_pages=1122 · num_kv_indices=823
NVIDIA B200
1.51ms
#2 of 2
2026-03-28
GQA paged decode h20 kv4 d128 ps64bf16 · [2, 20, 128] · num_pages=1180 · num_kv_indices=114
NVIDIA B200
1.51ms
#2 of 2
2026-03-28
GQA paged decode h20 kv4 d128 ps64bf16 · [2, 20, 128] · num_pages=1624 · num_kv_indices=500
NVIDIA B200
2.03ms
#1 of 2
2026-03-28
GQA paged decode h20 kv4 d128 ps64bf16 · [2, 20, 128] · num_pages=1624 · num_kv_indices=500
NVIDIA B200
2.09ms
#2 of 2
2026-03-28
GQA paged decode h20 kv4 d128 ps64bf16 · [4, 20, 128] · num_pages=3961 · num_kv_indices=201
NVIDIA B200
4.81ms
#1 of 2
2026-03-28
GQA paged decode h20 kv4 d128 ps64bf16 · [2, 20, 128] · num_pages=3874 · num_kv_indices=2292
NVIDIA B200
4.90ms
#1 of 2
2026-03-28
GQA paged decode h20 kv4 d128 ps64bf16 · [4, 20, 128] · num_pages=3961 · num_kv_indices=201
NVIDIA B200
4.92ms
#2 of 2
2026-03-28
GQA paged decode h20 kv4 d128 ps64bf16 · [2, 20, 128] · num_pages=3874 · num_kv_indices=2292
NVIDIA B200
5.02ms
#2 of 2
2026-03-28
GQA paged decode h20 kv4 d128 ps64bf16 · [4, 20, 128] · num_pages=5157 · num_kv_indices=1280
NVIDIA B200
6.33ms
#1 of 2
2026-03-28
GQA paged decode h20 kv4 d128 ps64bf16 · [4, 20, 128] · num_pages=5157 · num_kv_indices=1280
NVIDIA B200
6.49ms
#2 of 2
2026-03-28
GQA paged decode h20 kv4 d128 ps64bf16 · [8, 20, 128] · num_pages=7357 · num_kv_indices=419
NVIDIA B200
9.32ms
#1 of 2
2026-03-28
GQA paged decode h20 kv4 d128 ps64bf16 · [4, 20, 128] · num_pages=7341 · num_kv_indices=3334
NVIDIA B200
9.55ms
#1 of 2
2026-03-28
GQA paged decode h20 kv4 d128 ps64bf16 · [8, 20, 128] · num_pages=7357 · num_kv_indices=419
NVIDIA B200
9.69ms
#2 of 2
2026-03-28
GQA paged decode h20 kv4 d128 ps64bf16 · [4, 20, 128] · num_pages=7341 · num_kv_indices=3334
NVIDIA B200
9.94ms
#2 of 2
2026-03-28
GQA paged decode h20 kv4 d128 ps64bf16 · [8, 20, 128] · num_pages=7837 · num_kv_indices=2426
NVIDIA B200
10.1ms
#1 of 2
2026-03-28
GQA paged decode h20 kv4 d128 ps64bf16 · [8, 20, 128] · num_pages=7837 · num_kv_indices=2426
NVIDIA B200
10.5ms
#2 of 2
2026-03-28
GQA paged decode h20 kv4 d128 ps64bf16 · [16, 20, 128] · num_pages=10211 · num_kv_indices=808
NVIDIA B200
12.9ms
#1 of 2
2026-03-28
GQA paged decode h20 kv4 d128 ps64bf16 · [16, 20, 128] · num_pages=10211 · num_kv_indices=808
NVIDIA B200
13.4ms
#2 of 2
2026-03-28
GQA paged decode h20 kv4 d128 ps64bf16 · [8, 20, 128] · num_pages=10178 · num_kv_indices=6728
NVIDIA B200
13.5ms
#1 of 2
2026-03-28
GQA paged decode h20 kv4 d128 ps64bf16 · [8, 20, 128] · num_pages=10178 · num_kv_indices=6728
NVIDIA B200
14.0ms
#2 of 2
2026-03-28
GQA paged decode h20 kv4 d128 ps64bf16 · [16, 20, 128] · num_pages=10929 · num_kv_indices=5106
NVIDIA B200
14.4ms
#1 of 2
2026-03-28
GQA paged decode h20 kv4 d128 ps64bf16 · [16, 20, 128] · num_pages=10929 · num_kv_indices=5106
NVIDIA B200
15.0ms
#2 of 2
2026-03-28
GQA paged decode h20 kv4 d128 ps64bf16 · [32, 20, 128] · num_pages=15356 · num_kv_indices=1645
NVIDIA B200
19.4ms
#1 of 2
2026-03-28
GQA paged decode h20 kv4 d128 ps64bf16 · [32, 20, 128] · num_pages=15356 · num_kv_indices=1645
NVIDIA B200
20.2ms
#2 of 2
2026-03-28
NVIDIA B200
20.3ms
#1 of 2
2026-03-28
GQA paged decode h20 kv4 d128 ps64bf16 · [32, 20, 128] · num_pages=15868 · num_kv_indices=10156
NVIDIA B200
20.8ms
#1 of 2
2026-03-28
GQA paged decode h20 kv4 d128 ps64bf16 · [16, 20, 128] · num_pages=15292 · num_kv_indices=13578
NVIDIA B200
21.1ms
#1 of 2
2026-03-28
NVIDIA B200
21.1ms
#2 of 2
2026-03-28
GQA paged decode h20 kv4 d128 ps64bf16 · [32, 20, 128] · num_pages=15868 · num_kv_indices=10156
NVIDIA B200
21.6ms
#2 of 2
2026-03-28
GQA paged decode h20 kv4 d128 ps64bf16 · [16, 20, 128] · num_pages=15292 · num_kv_indices=13578
NVIDIA B200
21.8ms
#2 of 2
2026-03-28

Reported · How evidence levels are derived →

Source and license

sourcehttps://huggingface.co/datasets/flashinfer-ai/flashinfer-trace
commitda915083d4c7c5e61aa3005e3d17ae488e0fc71c
revision digestsha256:1a081dab2e7f927149c8b3aa27e3bd7d401852544e3b7e2d2d8f6dddd0901d1c
license declaredApache-2.0
license concludedApache-2.0
authorsbaseline
imported2026-08-16

Kernel source

main.py75 lines
import torch
import flashinfer

# GQA group_size=5 (20 qo_heads / 4 kv_heads) is not a power-of-2 and is
# unsupported by FlashInfer's BatchDecodeWithPagedKVCacheWrapper.  Work-around:
# expand KV heads from 4 to 20 (repeat_interleave x5) so group_size=1 (MHA),
# which is mathematically equivalent.  We 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, kv_last_page_len, 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  # 5
    # Expand KV heads: [num_pages, page_size, 4, 128] -> [num_pages, page_size, 20, 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, 20, 128]
    expanded_heads = num_qo_heads  # 20
    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()
        or state["last_page_ptr"] != kv_last_page_len.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=kv_last_page_len,
            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(),
            "last_page_ptr": kv_last_page_len.data_ptr(),
        }
    output, lse = wrapper.run(q, paged_kv, return_lse=True)
    return output, lse
scrolls · 75 lines total

Source code from FlashInfer-Bench (flashinfer-ai/flashinfer-trace) · Apache-2.0

Best evidence level for this revision: reported

JSON