flashinfer / wrapper9b7f1e
flashinfer_wrapper_9b7f1e · FlashInfer-Bench baselines · python · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 32 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-flashinfer-wrapper-9b7f1e?include=source"interfacepython
revisionda915083d4c7
symbolrun
pathmain.py
Compatibility
declared hardwareNVIDIA B200, NVIDIA H100, NVIDIA H20, NVIDIA H200
architecturesunknown
dtypesbf16, fp32
Benchmark evidence
No published measurement for this revision.
No evidence · How evidence levels are derived →
Source and license
sourcehttps://huggingface.co/datasets/flashinfer-ai/flashinfer-trace
commitda915083d4c7c5e61aa3005e3d17ae488e0fc71c
revision digestsha256:551fe1a2bd517c89a275bf4d07e0a663714f7de1140bc3dfa7f04aa436813b33
license declaredApache-2.0
license concludedApache-2.0
authorsbaseline
imported2026-08-16
Kernel source
main.py32 lines
import math
import torch
from flashinfer.gdn_decode import gated_delta_rule_decode_pretranspose
def run(q, k, v, state, A_log, a, dt_bias, b, scale):
if isinstance(scale, torch.Tensor):
scale = float(scale.item())
else:
scale = float(scale)
if scale == 0.0:
scale = 1.0 / math.sqrt(q.shape[-1])
B, T, num_v_heads, head_size = v.shape
output = torch.empty(B, T, num_v_heads, head_size, dtype=q.dtype, device=q.device)
out, new_state = gated_delta_rule_decode_pretranspose(
q=q,
k=k,
v=v,
state=state,
A_log=A_log,
a=a,
dt_bias=dt_bias,
b=b,
scale=scale,
output=output,
use_qk_l2norm=False,
)
return out, new_state
scrolls · 32 lines total
Source code from the importing source · Apache-2.0
No published measurement for this revision
JSON