group_gemm_mxfp4_flashinfer_g4_n4096_k2048
FlashInfer-Bench baselines · python · Apache-2.0
Kernel source · 35 lines ↓holds 2 records
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
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