Skip to content
KernelIndex
Search⌘K

flashinfer / wrapperad4135

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

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-ad4135?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

18 measurements across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
GQA paged decode h32 kv8 d128 ps64bf16 · [4, 32, 128] · num_pages=349 · num_kv_indices=212
NVIDIA B200
28.8µs
#1 of 1
2026-04-30
GQA paged decode h32 kv8 d128 ps64bf16 · [4, 32, 128] · num_pages=389 · num_kv_indices=232
NVIDIA B200
31.4µs
#1 of 1
2026-04-30
GQA paged decode h32 kv8 d128 ps64bf16 · [4, 32, 128] · num_pages=398 · num_kv_indices=268
NVIDIA B200
34.4µs
#1 of 1
2026-04-30
GQA paged decode h32 kv8 d128 ps64bf16 · [8, 32, 128] · num_pages=839 · num_kv_indices=528
NVIDIA B200
57.8µs
#1 of 1
2026-04-30
GQA paged decode h32 kv8 d128 ps64bf16 · [8, 32, 128] · num_pages=902 · num_kv_indices=576
NVIDIA B200
58.5µs
#1 of 1
2026-04-30
GQA paged decode h32 kv8 d128 ps64bf16 · [8, 32, 128] · num_pages=910 · num_kv_indices=640
NVIDIA B200
64.0µs
#1 of 1
2026-04-30
GQA paged decode h32 kv8 d128 ps64bf16 · [16, 32, 128] · num_pages=1686 · num_kv_indices=1033
NVIDIA B200
102.9µs
#1 of 1
2026-04-30
GQA paged decode h32 kv8 d128 ps64bf16 · [16, 32, 128] · num_pages=1735 · num_kv_indices=1145
NVIDIA B200
115.6µs
#1 of 1
2026-04-30
GQA paged decode h32 kv8 d128 ps64bf16 · [16, 32, 128] · num_pages=1794 · num_kv_indices=1257
NVIDIA B200
116.0µs
#1 of 1
2026-04-30
GQA paged decode h32 kv8 d128 ps64bf16 · [32, 32, 128] · num_pages=15114 · num_kv_indices=13783
NVIDIA B200
1.25ms
#1 of 1
2026-04-30
Show all 18 measurements ›
NVIDIA B200
1.26ms
#1 of 1
2026-04-30
GQA paged decode h32 kv8 d128 ps64bf16 · [32, 32, 128] · num_pages=15178 · num_kv_indices=14103
NVIDIA B200
1.26ms
#1 of 1
2026-04-30
GQA paged decode h32 kv8 d128 ps64bf16 · [64, 32, 128] · num_pages=29189 · num_kv_indices=26582
NVIDIA B200
2.69ms
#1 of 1
2026-04-30
GQA paged decode h32 kv8 d128 ps64bf16 · [64, 32, 128] · num_pages=29441 · num_kv_indices=27222
NVIDIA B200
2.73ms
#1 of 1
2026-04-30
NVIDIA B200
2.85ms
#1 of 1
2026-04-30
NVIDIA B200
6.96ms
#1 of 1
2026-04-30
GQA paged decode h32 kv8 d128 ps64bf16 · [128, 32, 128] · num_pages=51792 · num_kv_indices=46414
NVIDIA B200
7.13ms
#1 of 1
2026-04-30
GQA paged decode h32 kv8 d128 ps64bf16 · [128, 32, 128] · num_pages=52033 · num_kv_indices=47438
NVIDIA B200
7.19ms
#1 of 1
2026-04-30

Reported · How evidence levels are derived →

Source and license

sourcehttps://huggingface.co/datasets/flashinfer-ai/flashinfer-trace
commitda915083d4c7c5e61aa3005e3d17ae488e0fc71c
revision digestsha256:4b3e387682a1d3046578b8c6593e5e44baeea6265c8f93fe446750b194be7463
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: reported

JSON