flashinfer / wrapperf4c6a8
flashinfer_wrapper_f4c6a8 · FlashInfer-Bench baselines · python · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 30 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-flashinfer-wrapper-f4c6a8?include=source"interfacepython
revisionda915083d4c7
symbolrun
pathmain.py
Compatibility
declared hardwareNVIDIA B200, NVIDIA H100, NVIDIA H20, NVIDIA H200
architecturesunknown
dtypesbf16, fp32, int32
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:c4d183de8e80497ac8f7af21d19b3072bbbf5fc4b5e6bcda6f0a8132df812016
license declaredApache-2.0
license concludedApache-2.0
authorsbaseline
imported2026-08-16
Kernel source
main.py30 lines
import math
import torch
from flashinfer.gdn_decode import gated_delta_rule_mtp
def run(q, k, v, initial_state, initial_state_indices, 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])
output, new_state = gated_delta_rule_mtp(
q=q,
k=k,
v=v,
initial_state=initial_state,
initial_state_indices=initial_state_indices,
A_log=A_log,
a=a,
dt_bias=dt_bias,
b=b,
scale=scale,
disable_state_update=False,
use_qk_l2norm=False,
)
return output, new_state
scrolls · 30 lines total
Source code from the importing source · Apache-2.0
No published measurement for this revision
JSON