Skip to content
KernelIndex
Search⌘K

flashinfer / wrapperf9a07b

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

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-flashinfer-wrapper-f9a07b?include=source"
interfacepython
revisionda915083d4c7
symbolrun
pathmain.py
Compatibility
measured onNVIDIA B200
declared hardwareNVIDIA A100, NVIDIA B200, NVIDIA GeForce RTX 4090, NVIDIA H100, NVIDIA H20, NVIDIA H200
architecturesunknown
dtypesbf16, fp32, int32

Benchmark evidence

21 measurements across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
GQA ragged prefill causal h32 kv8 d128bf16 · [1, 8, 128] · #162cf6
NVIDIA B200
10.2µs
#1 of 10
2025-10-21
GQA ragged prefill causal h32 kv8 d128bf16 · [1, 8, 128] · #162cf6
NVIDIA B200
10.2µs
#2 of 10
2025-10-21
GQA ragged prefill causal h32 kv8 d128bf16 · [1, 8, 128] · #5fa6fa
NVIDIA B200
10.3µs
#1 of 5
2025-10-21
GQA ragged prefill causal h32 kv8 d128bf16 · [7, 8, 128] · #227ae2
NVIDIA B200
10.3µs
#1= of 10
2025-10-21
GQA ragged prefill causal h32 kv8 d128bf16 · [7, 8, 128] · #227ae2
NVIDIA B200
10.3µs
#1 of 10
2025-10-21
GQA ragged prefill causal h32 kv8 d128bf16 · [7, 8, 128] · #cb4b00
NVIDIA B200
10.3µs
#1 of 5
2025-10-21
NVIDIA B200
10.4µs
#1 of 10
2025-10-21
NVIDIA B200
10.5µs
#2 of 10
2025-10-21
NVIDIA B200
10.6µs
#1 of 5
2025-10-21
GQA ragged prefill causal h32 kv8 d128bf16 · [35, 8, 128] · #10e83c
NVIDIA B200
11.0µs
#1 of 20
2025-10-21
Show all 21 measurements ›
GQA ragged prefill causal h32 kv8 d128bf16 · [35, 8, 128] · #10e83c
NVIDIA B200
11.1µs
#2 of 20
2025-10-21
NVIDIA B200
11.6µs
#1 of 5
2025-10-21
GQA ragged prefill causal h32 kv8 d128bf16 · [35, 8, 128] · #10e83c
NVIDIA B200
12.1µs
#3 of 20
2025-10-21
GQA ragged prefill causal h32 kv8 d128bf16 · [35, 8, 128] · #f3d59b
NVIDIA B200
12.1µs
#1 of 5
2025-10-21
GQA ragged prefill causal h32 kv8 d128bf16 · [35, 8, 128] · #10e83c
NVIDIA B200
12.3µs
#4 of 20
2025-10-21
NVIDIA B200
12.3µs
#1 of 10
2025-10-21
NVIDIA B200
12.5µs
#2 of 10
2025-10-21
NVIDIA B200
28.5µs
#1 of 5
2025-10-21
NVIDIA B200
525.3µs
#1 of 5
2025-10-21
NVIDIA B200
529.2µs
#1 of 5
2025-10-21
NVIDIA B200
536.5µs
#1 of 5
2025-10-21

Reported · How evidence levels are derived →

Source and license

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

Kernel source

main.py90 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.BatchPrefillWithRaggedKVCacheWrapper(
            workspace,
            kv_layout="NHD",
        )
        _wrapper_cache[key] = wrapper
    return wrapper


def run(q, k, v, qo_indptr, kv_indptr, sm_scale):
    total_q, num_qo_heads, head_dim = q.shape
    total_kv, num_kv_heads, _ = k.shape
    batch_size = qo_indptr.shape[0] - 1

    device = q.device
    wrapper_key = (
        str(device),
        num_qo_heads,
        num_kv_heads,
        head_dim,
        q.dtype,
        k.dtype,
        v.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("total_q") != total_q
            or state.get("total_kv") != total_kv
            or state.get("batch_size") != batch_size
            or state.get("sm_scale") != sm_scale
            or state.get("qo_indptr_ptr") != qo_indptr.data_ptr()
            or state.get("kv_indptr_ptr") != kv_indptr.data_ptr()
        )

    if needs_plan:
        wrapper.plan(
            qo_indptr=qo_indptr,
            kv_indptr=kv_indptr,
            num_qo_heads=num_qo_heads,
            num_kv_heads=num_kv_heads,
            head_dim_qk=head_dim,
            causal=True,
            sm_scale=sm_scale,
            q_data_type=q.dtype,
            kv_data_type=k.dtype,
        )
        _plan_state[wrapper_key] = {
            "total_q": total_q,
            "total_kv": total_kv,
            "batch_size": batch_size,
            "sm_scale": sm_scale,
            "qo_indptr_ptr": qo_indptr.data_ptr(),
            "kv_indptr_ptr": kv_indptr.data_ptr(),
        }

    output, lse = wrapper.run(
        q,
        k,
        v,
        return_lse=True,
    )

    return output, lse
scrolls · 90 lines total

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

Best evidence level for this revision: reported

JSON