Skip to content
KernelIndex
Search⌘K

flashinfer / wrapper1b7890

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

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-flashinfer-wrapper-1b7890?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 kv4 d128 ps64bf16 · [1, 24, 128] · num_pages=425 · num_kv_indices=361
NVIDIA B200
359.8µs
#1 of 1
2026-03-31
NVIDIA B200
962.8µs
#1 of 1
2026-03-31
GQA paged decode h24 kv4 d128 ps64bf16 · [1, 24, 128] · num_pages=1266 · num_kv_indices=882
NVIDIA B200
1.01ms
#1 of 1
2026-03-31
GQA paged decode h24 kv4 d128 ps64bf16 · [4, 24, 128] · num_pages=1610 · num_kv_indices=247
NVIDIA B200
1.16ms
#1 of 1
2026-03-31
GQA paged decode h24 kv4 d128 ps64bf16 · [4, 24, 128] · num_pages=1729 · num_kv_indices=275
NVIDIA B200
1.26ms
#1 of 1
2026-03-31
GQA paged decode h24 kv4 d128 ps64bf16 · [8, 24, 128] · num_pages=6458 · num_kv_indices=471
NVIDIA B200
4.57ms
#1 of 1
2026-03-31
GQA paged decode h24 kv4 d128 ps64bf16 · [8, 24, 128] · num_pages=6785 · num_kv_indices=527
NVIDIA B200
4.81ms
#1 of 1
2026-03-31
GQA paged decode h24 kv4 d128 ps64bf16 · [4, 24, 128] · num_pages=6180 · num_kv_indices=4583
NVIDIA B200
4.85ms
#1 of 1
2026-03-31
GQA paged decode h24 kv4 d128 ps64bf16 · [8, 24, 128] · num_pages=8888 · num_kv_indices=3176
NVIDIA B200
6.58ms
#1 of 1
2026-03-31
GQA paged decode h24 kv4 d128 ps64bf16 · [16, 24, 128] · num_pages=14913 · num_kv_indices=993
NVIDIA B200
10.5ms
#1 of 1
2026-03-31
Show all 20 measurements ›
GQA paged decode h24 kv4 d128 ps64bf16 · [8, 24, 128] · num_pages=14271 · num_kv_indices=7941
NVIDIA B200
10.9ms
#1 of 1
2026-03-31
GQA paged decode h24 kv4 d128 ps64bf16 · [16, 24, 128] · num_pages=14977 · num_kv_indices=6068
NVIDIA B200
11.2ms
#1 of 1
2026-03-31
GQA paged decode h24 kv4 d128 ps64bf16 · [32, 24, 128] · num_pages=20673 · num_kv_indices=1954
NVIDIA B200
14.6ms
#1 of 1
2026-03-31
GQA paged decode h24 kv4 d128 ps64bf16 · [32, 24, 128] · num_pages=21889 · num_kv_indices=2178
NVIDIA B200
15.5ms
#1 of 1
2026-03-31
GQA paged decode h24 kv4 d128 ps64bf16 · [32, 24, 128] · num_pages=20737 · num_kv_indices=12162
NVIDIA B200
15.7ms
#1 of 1
2026-03-31
GQA paged decode h24 kv4 d128 ps64bf16 · [16, 24, 128] · num_pages=19713 · num_kv_indices=16593
NVIDIA B200
16.4ms
#1 of 1
2026-03-31
GQA paged decode h24 kv4 d128 ps64bf16 · [64, 24, 128] · num_pages=22913 · num_kv_indices=3939
NVIDIA B200
16.4ms
#1 of 1
2026-03-31
GQA paged decode h24 kv4 d128 ps64bf16 · [64, 24, 128] · num_pages=25089 · num_kv_indices=4387
NVIDIA B200
17.9ms
#1 of 1
2026-03-31
GQA paged decode h24 kv4 d128 ps64bf16 · [128, 24, 128] · num_pages=27201 · num_kv_indices=7865
NVIDIA B200
19.7ms
#1 of 1
2026-03-31
GQA paged decode h24 kv4 d128 ps64bf16 · [128, 24, 128] · num_pages=31809 · num_kv_indices=8761
NVIDIA B200
23.0ms
#1 of 1
2026-03-31

Reported · How evidence levels are derived →

Source and license

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

Kernel source

main.py81 lines
import torch
import flashinfer

# group_size=6 (24 qo_heads / 4 kv_heads) is not natively supported.
# Work-around: expand KV heads from 4 to 24 (repeat_interleave x6)
# so the wrapper sees group_size=1 (MHA), which is mathematically equivalent.

_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
    # kv_last_page_len may have an extra trailing element from workload capture;
    # always clamp to batch_size.
    kv_last_page_len = kv_last_page_len[:batch_size]
    group_size = num_qo_heads // num_kv_heads
    # Expand KV heads: [pages, page_size, 4, head_dim] -> [pages, page_size, 24, head_dim]
    k_exp = k_cache.repeat_interleave(group_size, dim=2)
    v_exp = v_cache.repeat_interleave(group_size, dim=2)
    expanded_kv_heads = num_qo_heads  # 24

    device = q.device
    wkey = (str(device), num_qo_heads, expanded_kv_heads, head_dim, page_size, q.dtype, k_exp.dtype)
    wrapper = _get_wrapper(wkey, device)
    state = _plan_state.get(wkey)

    needs_plan = True
    if state is not None:
        needs_plan = (
            state.get("batch_size") != batch_size
            or state.get("kv_indptr_ptr") != kv_indptr.data_ptr()
            or state.get("kv_indices_ptr") != kv_indices.data_ptr()
            or state.get("sm_scale") != sm_scale
        )

    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=expanded_kv_heads,
            head_dim=head_dim,
            page_size=page_size,
            pos_encoding_mode="NONE",
            q_data_type=q.dtype,
            kv_data_type=k_exp.dtype,
            sm_scale=sm_scale,
        )
        _plan_state[wkey] = {
            "batch_size": batch_size,
            "kv_indptr_ptr": kv_indptr.data_ptr(),
            "kv_indices_ptr": kv_indices.data_ptr(),
            "sm_scale": sm_scale,
        }

    output, lse = wrapper.run(q, (k_exp, v_exp), return_lse=True)
    return output, lse
scrolls · 81 lines total

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

Best evidence level for this revision: reported

JSON