Skip to content
KernelIndex
Search⌘K

flashinfer / wrapper5222a7

flashinfer_wrapper_5222a7 · FlashInfer-Bench baselines · python · Apache-2.0

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-flashinfer-wrapper-5222a7?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
NVIDIA B200
16.7µs
#1 of 1
2026-04-07
GQA paged prefill causal h40 kv10 d128 ps1bf16 · [77, 40, 128] · num_pages=874
NVIDIA B200
23.5µs
#1 of 1
2026-04-07
NVIDIA B200
35.9µs
#1 of 1
2026-04-07
GQA paged prefill causal h40 kv10 d128 ps1bf16 · [563, 40, 128] · num_pages=719
NVIDIA B200
50.1µs
#1 of 1
2026-04-07
NVIDIA B200
64.5µs
#1 of 1
2026-04-07
GQA paged prefill causal h40 kv10 d128 ps1bf16 · [287, 40, 128] · num_pages=4518
NVIDIA B200
76.8µs
#1 of 1
2026-04-07
NVIDIA B200
97.3µs
#1 of 1
2026-04-07
GQA paged prefill causal h40 kv10 d128 ps1bf16 · [1114, 40, 128] · num_pages=3116
NVIDIA B200
101.3µs
#1 of 1
2026-04-07
GQA paged prefill causal h40 kv10 d128 ps1bf16 · [77, 40, 128] · num_pages=16681
NVIDIA B200
105.5µs
#1 of 1
2026-04-07
NVIDIA B200
128.0µs
#1 of 1
2026-04-07
Show all 20 measurements ›
GQA paged prefill causal h40 kv10 d128 ps1bf16 · [1114, 40, 128] · num_pages=8799
NVIDIA B200
130.4µs
#1 of 1
2026-04-07
GQA paged prefill causal h40 kv10 d128 ps1bf16 · [287, 40, 128] · num_pages=18138
NVIDIA B200
146.4µs
#1 of 1
2026-04-07
NVIDIA B200
181.2µs
#1 of 1
2026-04-07
NVIDIA B200
199.9µs
#1 of 1
2026-04-07
NVIDIA B200
209.9µs
#1 of 1
2026-04-07
NVIDIA B200
227.0µs
#1 of 1
2026-04-07
GQA paged prefill causal h40 kv10 d128 ps1bf16 · [563, 40, 128] · num_pages=47098
NVIDIA B200
286.1µs
#1 of 1
2026-04-07
NVIDIA B200
405.4µs
#1 of 1
2026-04-07
NVIDIA B200
426.5µs
#1 of 1
2026-04-07
NVIDIA B200
620.8µs
#1 of 1
2026-04-07

Reproduction-ready · How evidence levels are derived →

Source and license

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

Kernel source

main.py59 lines
import torch
import flashinfer

# GQA group_size = 40/10 = 4, which is a power of 2.
# FlashInfer natively supports GQA with group_size=4 without KV head expansion.

_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, qo_indptr, kv_indptr, kv_indices, sm_scale):
    total_q, num_qo_heads, head_dim = q.shape
    _, page_size, num_kv_heads, _ = k_cache.shape
    batch_size = kv_indptr.shape[0] - 1
    device = q.device
    paged_kv = torch.stack([k_cache, v_cache], dim=1)  # [num_pages, 2, page_size, num_kv_heads, head_dim]
    wkey = (str(device), num_qo_heads, num_kv_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["total_q"] != total_q or state["qo_ptr"] != qo_indptr.data_ptr()
    if needs_plan:
        last_page_len = torch.ones(batch_size, 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=last_page_len,
            num_qo_heads=num_qo_heads,
            num_kv_heads=num_kv_heads,
            head_dim_qk=head_dim,
            page_size=page_size,
            causal=True,
            sm_scale=float(sm_scale),
            q_data_type=q.dtype,
            kv_data_type=k_cache.dtype,
        )
        _plan_state[wkey] = {"total_q": total_q, "qo_ptr": qo_indptr.data_ptr()}
    output, lse = wrapper.run(q, paged_kv, return_lse=True)
    return output, lse
scrolls · 59 lines total

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

Best evidence level for this revision: reproducible

JSON