flashinfer / wrapperea3787
flashinfer_wrapper_ea3787 · FlashInfer-Bench baselines · python · Apache-2.0
Kernel source · 101 lines ↓holds 38 records
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-ea3787?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
Show all 38 measurements ›Showing all 38 measurements ⌄
MLA paged prefill causal h16 ckv512 kpe64 ps1bf16 · [1954, 16, 64]
NVIDIA B200
68.9µs
#1 of 5
2025-10-21
MLA paged prefill causal h16 ckv512 kpe64 ps1bf16 · [1028, 16, 64]
NVIDIA B200
106.2µs
#1 of 5
2025-10-21
MLA paged prefill causal h16 ckv512 kpe64 ps1bf16 · [1187, 16, 64]
NVIDIA B200
123.3µs
#1 of 5
2025-10-21
MLA paged prefill causal h16 ckv512 kpe64 ps1bf16 · [3842, 16, 64]
NVIDIA B200
237.5µs
#1 of 5
2025-10-21
MLA paged prefill causal h16 ckv512 kpe64 ps1bf16 · [3024, 16, 64]
NVIDIA B200
507.6µs
#1 of 5
2025-10-21
MLA paged prefill causal h16 ckv512 kpe64 ps1bf16 · [6053, 16, 64]
NVIDIA B200
520.7µs
#1 of 5
2025-10-21
MLA paged prefill causal h16 ckv512 kpe64 ps1bf16 · [8987, 16, 64]
NVIDIA B200
849.5µs
#1 of 5
2025-10-21
MLA paged prefill causal h16 ckv512 kpe64 ps1bf16 · [15092, 16, 64]
NVIDIA B200
1.83ms
#1 of 5
2025-10-21
MLA paged prefill causal h16 ckv512 kpe64 ps1bf16 · [15883, 16, 64]
NVIDIA B200
2.54ms
#1 of 5
2025-10-21
MLA paged prefill causal h16 ckv512 kpe64 ps1bf16 · [10870, 16, 64]
NVIDIA B200
3.75ms
#1 of 4
2025-10-21
MLA paged prefill causal h16 ckv512 kpe64 ps1bf16 · [16384, 16, 64]
NVIDIA B200
16.5ms
#1 of 4
2025-10-21
Reported · How evidence levels are derived →
Source and license
sourcehttps://huggingface.co/datasets/flashinfer-ai/flashinfer-trace
commitda915083d4c7c5e61aa3005e3d17ae488e0fc71c
revision digestsha256:bcc5e6b5c47e98f95a4c85f871621f8a2976c0963b289e9e14d9d10f13604c8b
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.int8, 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.mla.BatchMLAPagedAttentionWrapper(workspace)
_wrapper_cache[key] = wrapper
return wrapper
def run(q_nope, q_pe, ckv_cache, kpe_cache, qo_indptr, kv_indptr, kv_indices, sm_scale):
total_q, num_qo_heads, head_dim_ckv = q_nope.shape
_, _, head_dim_kpe = q_pe.shape
page_size = ckv_cache.shape[1]
len_indptr = kv_indptr.shape[0]
num_kv_indices = kv_indices.shape[0]
batch_size = qo_indptr.shape[0] - 1
device = q_nope.device
wrapper_key = (
str(device),
num_qo_heads,
head_dim_ckv,
head_dim_kpe,
page_size,
q_nope.dtype,
q_pe.dtype,
ckv_cache.dtype,
kpe_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("total_q") != total_q
or 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("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:
kv_len_arr = (kv_indptr[1:] - kv_indptr[:-1]).to(torch.int32)
wrapper.plan(
qo_indptr=qo_indptr,
kv_indptr=kv_indptr,
kv_indices=kv_indices,
kv_len_arr=kv_len_arr,
num_heads=num_qo_heads,
head_dim_ckv=head_dim_ckv,
head_dim_kpe=head_dim_kpe,
page_size=page_size,
causal=True,
sm_scale=sm_scale,
q_data_type=q_nope.dtype,
kv_data_type=ckv_cache.dtype,
)
_plan_state[wrapper_key] = {
"total_q": total_q,
"batch_size": batch_size,
"len_indptr": len_indptr,
"num_kv_indices": num_kv_indices,
"sm_scale": sm_scale,
"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_nope,
q_pe,
ckv_cache,
kpe_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: reported
JSON