flashinfer_wrapper_023122
FlashInfer-Bench baselines · python · Apache-2.0
Kernel source · 60 lines ↓holds 10 records
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 60 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-flashinfer-wrapper-023122?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
10 measurements across 1 GPU, fastest first.
Operation / workload
Hardware
Latency
Rank
Observed
Reported · How evidence levels are derived →
Source and license
sourcehttps://huggingface.co/datasets/flashinfer-ai/flashinfer-trace
commitda915083d4c7c5e61aa3005e3d17ae488e0fc71c
revision digestsha256:9fdf07cc4db985f53d8d56be273c9c8277b2168d3e9337f5eea32a574764b251
license declaredApache-2.0
license concludedApache-2.0
authorsbaseline
imported2026-08-16
Kernel source
main.py60 lines
import torch
import flashinfer
# GQA group_size=5 (20 qo_heads / 4 kv_heads) is not a power-of-2 and is
# unsupported by FlashInfer kernels. Work-around: expand KV heads from 4
# to 20 (repeat_interleave x5) so 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)
buf = _workspace_cache.get(key)
if buf is None:
buf = torch.empty(_WORKSPACE_SIZE_BYTES, dtype=torch.uint8, device=device)
_workspace_cache[key] = buf
return buf
def _get_wrapper(key, device):
w = _wrapper_cache.get(key)
if w is None:
w = flashinfer.BatchPrefillWithRaggedKVCacheWrapper(_get_workspace(device), kv_layout="NHD")
_wrapper_cache[key] = w
return w
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
device = q.device
group_size = num_qo_heads // num_kv_heads # 5
# Expand KV heads: [total_kv, 4, 128] -> [total_kv, 20, 128]
k_exp = k.repeat_interleave(group_size, dim=1)
v_exp = v.repeat_interleave(group_size, dim=1)
expanded_heads = num_qo_heads # 20
wkey = (str(device), num_qo_heads, expanded_heads, head_dim, q.dtype, k.dtype)
wrapper = _get_wrapper(wkey, device)
state = _plan_state.get(wkey)
needs_plan = state is None or state["total_q"] != total_q or state["qo_ptr"] != qo_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=expanded_heads,
head_dim_qk=head_dim,
causal=True,
sm_scale=float(sm_scale),
q_data_type=q.dtype,
kv_data_type=k.dtype,
)
_plan_state[wkey] = {"total_q": total_q, "qo_ptr": qo_indptr.data_ptr()}
output, lse = wrapper.run(q, k_exp, v_exp, return_lse=True)
return output, lse
scrolls · 60 lines total
Source code from FlashInfer-Bench (flashinfer-ai/flashinfer-trace) · Apache-2.0
Best evidence level for this revision: reported
JSON