Skip to content
KernelIndex
Search⌘K

group_gemm_mxfp4_flashinfer_g4_n4096_k2048

FlashInfer-Bench baselines · python · Apache-2.0

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-group-gemm-mxfp4-flashinfer-g4-n4096-k2048?include=source"
interfacepython
revisionda915083d4c7
symbolrun
pathmain.py
Compatibility
measured onNVIDIA B200
declared hardwareNVIDIA B200
architecturessm_100
dtypesfp8_e4m3, int32, int8

Benchmark evidence

2 measurements across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
NVIDIA B200
321.8µs
#1 of 1
2026-06-06
NVIDIA B200
468.2µs
#1 of 1
2026-06-06

Reported · How evidence levels are derived →

Source and license

sourcehttps://huggingface.co/datasets/flashinfer-ai/flashinfer-trace
commitda915083d4c7c5e61aa3005e3d17ae488e0fc71c
revision digestsha256:f38fc9c3c93578e4f9111d1814c91e215dbcf26c13bbc84c4cdba5e95afabef3
license declaredApache-2.0
license concludedApache-2.0
authorsbaseline
imported2026-08-16

Kernel source

main.py35 lines
import os
os.environ["FLASHINFER_DISABLE_VERSION_CHECK"] = "1"
import torch
from flashinfer.gemm import group_gemm_mxfp4_nt_groupwise
from flashinfer.fp4_quantization import _pad_scale_factors, get_fp4_quantization_module
from flashinfer.utils import get_compute_capability


def _swizzle_blockscale(unswizzled_sf, b, m, n, sf_vec_size=32):
    padded = torch.stack([_pad_scale_factors(unswizzled_sf[i], m, n, sf_vec_size) for i in range(b)])
    major, minor = get_compute_capability(unswizzled_sf.device)
    out = get_fp4_quantization_module(f"{major}{minor}").block_scale_interleave_sm100(padded)
    return out.view(padded.shape)


def run(a_fp8, a_scale, b_fp4, b_scale, m_indptr):
    G = b_fp4.shape[0]
    n = b_fp4.shape[1]
    k_padded = b_fp4.shape[2] * 2
    tile = 32
    m = a_fp8.shape[0] // G
    a_sc = a_scale.view(torch.uint8)
    b_sc = b_scale.view(torch.uint8)
    b_u8 = b_fp4.view(torch.uint8)
    a_sw = _swizzle_blockscale(a_sc.unflatten(0, (G, m)), G, m, k_padded, tile).flatten(0, 1)
    ga = torch.arange(0, G + 1, dtype=torch.int32, device=a_fp8.device)
    a_sf = 128
    mip = (m_indptr + ga * (a_sf - 1)) // a_sf * a_sf
    m_sf = mip[1:] - mip[:-1]
    ch = a_sw.chunk(G, dim=0)
    ch = [torch.cat([x, torch.zeros(int(m_sf[i]) - x.shape[0], *x.shape[1:], dtype=x.dtype, device=x.device)]) for i, x in enumerate(ch)]
    a_sw = torch.cat(ch)
    b_sw = _swizzle_blockscale(b_sc, G, n, k_padded, tile)
    return group_gemm_mxfp4_nt_groupwise(a_fp8, b_u8, a_sw, b_sw, m_indptr, out_dtype=torch.bfloat16)[:, :n]
scrolls · 35 lines total

Source code from FlashInfer-Bench (flashinfer-ai/flashinfer-trace) · Apache-2.0

Best evidence level for this revision: reported

JSON