Skip to content
KernelIndex
Search⌘K

submission 672283

SomersBuchannan · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

No package. Vendor the mirrored source: 59 lines, June 9 Researcher Reciprocity License v1.0.

submission_x2.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-672283?include=source"
interfacepython
Compatibility
measured onAMD Instinct MI355X
declared hardwareAMD Instinct MI355X
architecturesgfx950
dtypesbf16, mxfp4

Benchmark evidence

1 measurement across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
AMD MXFP4 GEMMsuite of 6 cases
AMD Instinct MI355X
14.6µs
#519 of 1143
2026-03-30

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:6742ad13b8bfd3b510591fcae4f2ce3b67cb24773c69b4a48b4f5729ecde8fc0
license declaredunknown
license concludedunknown
authorsSomersBuchannan
imported2026-08-26

Techniques

Extracted from the mirrored source by pattern, never inferred. Each row cites its line.

fp4Optimized MXFP4 quant + GEMM submission using gemm_a16wfp4_preshuffle.

Kernel source

submission_x2.py59 lines
"""
Optimized MXFP4 quant + GEMM submission using gemm_a16wfp4_preshuffle.

Key insight: gemm_a16wfp4_preshuffle fuses bf16->fp4 quantization directly inside 
the GEMM kernel (tl.dot_scaled with inline _mxfp4_quant_op), eliminating the 
separate quant kernel + e8m0_shuffle kernel entirely.

The preshuffle kernel expects:
- x: bf16 [M, K] - quantized on-the-fly
- w: uint8 [N//16, K//2*16] - shuffled weight reshaped for preshuffle layout
- w_scales: uint8 [N//32, K//32*32] - shuffled scales reshaped for preshuffle layout
"""
import torch
from task import input_t, output_t


def custom_kernel(data: input_t) -> output_t:
    """
    Optimized MXFP4 quant + GEMM using gemm_a16wfp4_preshuffle.
    Single fused kernel: bf16 A quantized on-the-fly inside GEMM.
    """
    from aiter.ops.triton.gemm.basic.gemm_a16wfp4 import gemm_a16wfp4_preshuffle
    from aiter import dtypes

    A, B, B_q, B_shuffle, B_scale_sh = data

    m, k = A.shape
    
    # B_shuffle is [N, K//2] in fp4x2 dtype (shuffled to (16,16) tile coalesced)
    # gemm_a16wfp4_preshuffle expects w as [N//16, K//2*16] in uint8
    # B_scale_sh is [*, K//32] in fp8_e8m0 (already shuffled)
    # gemm_a16wfp4_preshuffle expects w_scales as [N//32, K//32*32] in uint8
    
    # Convert B_shuffle from fp4x2 to uint8 and reshape for preshuffle layout
    B_sh_uint8 = B_shuffle.view(torch.uint8)
    n = B_sh_uint8.shape[0]
    k_half = B_sh_uint8.shape[1]
    # Reshape: [N, K//2] -> [N//16, K//2 * 16]
    w_preshuffle = B_sh_uint8.reshape(n // 16, k_half * 16)
    
    # Convert B_scale_sh from fp8_e8m0 to uint8 and reshape for preshuffle layout
    B_sc_uint8 = B_scale_sh.view(torch.uint8)
    sm, sn = B_sc_uint8.shape
    # Reshape: [sm, sn] -> [sm//32, sn*32] where sm is padded M dimension
    # But for w_scales in preshuffle, it's [N//32, K//32*32]
    # B_scale_sh is already [padded_N, K//32] shuffled
    # We need [N//32, K//32 * 32]
    w_scales_preshuffle = B_sc_uint8.reshape(sm // 32, sn * 32)
    
    out = gemm_a16wfp4_preshuffle(
        A,                    # bf16 [M, K] - quantized on-the-fly
        w_preshuffle,         # uint8 [N//16, K//2*16] - preshuffle layout
        w_scales_preshuffle,  # uint8 [N//32, K//32*32] - preshuffle layout
        prequant=True,        # enable on-the-fly quantization
        dtype=dtypes.bf16,
    )

    return out
scrolls · 59 lines total

Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0

Best evidence level for this revision: reported

JSON