Skip to content
KernelIndex
Search⌘K

flashinfer / wrappera3d7c2

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

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-flashinfer-wrapper-a3d7c2?include=source"
interfacepython
revisionda915083d4c7
symbolrun
pathmain.py
Compatibility
declared hardwareNVIDIA 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:ed90bf90006cc3f47717c18b45de3146a2ca056f6445a1e963b6f5e1b2b5df51
license declaredApache-2.0
license concludedApache-2.0
authorsbaseline
imported2026-08-16

Kernel source

main.py31 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, intermediate_states_buffer=None):
    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, final_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,
        intermediate_states_buffer=intermediate_states_buffer,
        disable_state_update=True,
        use_qk_l2norm=False,
    )

    return output, final_state
scrolls · 31 lines total

Source code from the importing source · Apache-2.0

No published measurement for this revision

JSON