Skip to content
KernelIndex
Search⌘K

flashinfer / wrapper8cad92

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

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

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

38 measurements across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
NVIDIA B200
174.8µs
#4 of 7
2025-10-21
NVIDIA B200
177.4µs
#1 of 1
2025-10-21
NVIDIA B200
178.4µs
#1 of 1
2025-10-21
NVIDIA B200
178.4µs
#1 of 1
2025-10-21
NVIDIA B200
179.9µs
#1 of 1
2025-10-21
NVIDIA B200
181.4µs
#1 of 1
2025-10-21
NVIDIA B200
181.7µs
#1 of 1
2025-10-21
NVIDIA B200
182.7µs
#1 of 1
2025-10-21
NVIDIA B200
182.9µs
#1 of 1
2025-10-21
NVIDIA B200
184.8µs
#1 of 1
2025-10-21
Show all 38 measurements ›
NVIDIA B200
186.4µs
#1 of 1
2025-10-21
NVIDIA B200
189.5µs
#1 of 1
2025-10-21
NVIDIA B200
190.0µs
#1 of 1
2025-10-21
NVIDIA B200
190.5µs
#4 of 7
2025-10-21
NVIDIA B200
191.1µs
#1 of 1
2025-10-21
NVIDIA B200
197.0µs
#1 of 1
2025-10-21
NVIDIA B200
200.1µs
#1 of 1
2025-10-21
NVIDIA B200
225.6µs
#1 of 1
2025-10-21
NVIDIA B200
247.3µs
#1 of 1
2025-10-21
NVIDIA B200
254.7µs
#1 of 1
2025-10-21
NVIDIA B200
254.9µs
#1 of 1
2025-10-21
NVIDIA B200
271.1µs
#1 of 1
2025-10-21
NVIDIA B200
278.4µs
#1 of 1
2025-10-21
NVIDIA B200
286.8µs
#1 of 1
2025-10-21
NVIDIA B200
293.5µs
#1 of 1
2025-10-21
NVIDIA B200
301.1µs
#1 of 1
2025-10-21
NVIDIA B200
319.2µs
#1 of 1
2025-10-21
NVIDIA B200
325.4µs
#1 of 1
2025-10-21
NVIDIA B200
337.3µs
#1 of 1
2025-10-21
NVIDIA B200
357.8µs
#1 of 1
2025-10-21
NVIDIA B200
366.6µs
#1 of 1
2025-10-21
NVIDIA B200
373.9µs
#1 of 1
2025-10-21
NVIDIA B200
377.3µs
#1 of 1
2025-10-21
NVIDIA B200
379.0µs
#1 of 1
2025-10-21
NVIDIA B200
417.1µs
#1 of 1
2025-10-21
NVIDIA B200
424.5µs
#1 of 1
2025-10-21
NVIDIA B200
506.8µs
#1 of 1
2025-10-21
NVIDIA B200
928.5µs
#1 of 1
2025-10-21

Reproduction-ready · How evidence levels are derived →

Source and license

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

Kernel source

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


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 = qo_indptr.shape[0] - 1
    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)

    if isinstance(sm_scale, torch.Tensor):
        sm_scale_value = float(sm_scale.item())
    else:
        sm_scale_value = float(sm_scale)

    needs_plan = True
    if state is not None:
        needs_plan = (
            state.get("total_q") != total_q
            or state.get("batch_size") != batch_size
            or state.get("num_kv_indices") != num_kv_indices
            or state.get("sm_scale") != sm_scale_value
            or state.get("qo_indptr_ptr") != qo_indptr.data_ptr()
            or state.get("kv_indptr_ptr") != kv_indptr.data_ptr()
            or state.get("kv_indices_ptr") != kv_indices.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=sm_scale,
            q_data_type=q.dtype,
            kv_data_type=k_cache.dtype,
        )
        _plan_state[wrapper_key] = {
            "total_q": total_q,
            "batch_size": batch_size,
            "num_kv_indices": num_kv_indices,
            "sm_scale": sm_scale_value,
            "qo_indptr_ptr": qo_indptr.data_ptr(),
            "kv_indptr_ptr": kv_indptr.data_ptr(),
            "kv_indices_ptr": kv_indices.data_ptr(),
        }

    output, lse = wrapper.run(
        q,
        (k_cache, v_cache),
        return_lse=True,
    )

    return output, lse
scrolls · 101 lines total

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

Best evidence level for this revision: reproducible

JSON