Skip to content
KernelIndex
Search⌘K

submission 671354

sizezheng_94252 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-671354?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
17.5µs
#671 of 1143
2026-03-30

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:9488cdf75e3b97bca304536bb7e0ec9d084798d9d6d1d321b8953a6ce327d9b8
license declaredunknown
license concludedunknown
authorssizezheng_94252
imported2026-08-26

Techniques

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

fp4Optimized MXFP4 GEMM: gemm_a16wfp4 with fused quant + fast B scale computation.
tile-n = 64BLOCK_N = 64

Kernel source

submission.py127 lines
"""
Optimized MXFP4 GEMM: gemm_a16wfp4 with fused quant + fast B scale computation.
Key: eliminates full B re-quantization. Only computes B scales (tiny kernel).
The Triton GEMM kernel handles A quant internally (fused).
"""
from task import input_t, output_t
import torch
import triton
import triton.language as tl
from aiter import dtypes
from aiter.ops.triton.gemm.basic.gemm_a16wfp4 import gemm_a16wfp4
from aiter.ops.triton.quant import dynamic_mxfp4_quant


@triton.jit
def _compute_b_scale_kernel(
    B_ptr, Scale_ptr,
    N, K,
    stride_bn, stride_bk,
    BLOCK_N: tl.constexpr,
    GROUP_SIZE: tl.constexpr,
):
    """Compute e8m0 per-group-of-32 scale from bf16 B. No quantization of data."""
    pid = tl.program_id(0)
    n_start = pid * BLOCK_N
    n_offsets = n_start + tl.arange(0, BLOCK_N)
    n_mask = n_offsets < N

    num_groups = K // GROUP_SIZE
    for g in range(num_groups):
        k_start = g * GROUP_SIZE
        # Load bf16 group and find abs max
        amax = tl.zeros([BLOCK_N], dtype=tl.float32)
        for k_off in range(0, GROUP_SIZE, 16):
            k_offsets = k_start + k_off + tl.arange(0, 16)
            ptrs = B_ptr + n_offsets[:, None] * stride_bn + k_offsets[None, :]
            mask = n_mask[:, None] & (k_offsets[None, :] < K)
            vals = tl.load(ptrs, mask=mask, other=0.0).to(tl.float32)
            amax = tl.maximum(amax, tl.max(tl.abs(vals), axis=1))

        # Match aiter's exact formula:
        # 1. Round amax UP to nearest power of 2 (via float32 bit manipulation)
        # 2. scale_unbiased = floor(log2(amax_rounded)) - 2
        # 3. e8m0 = scale_unbiased + 127
        amax = tl.maximum(amax, 1e-12)
        # Round up to power of 2: add 0x200000 (half ULP of mantissa) then mask
        amax_i32 = amax.to(tl.int32, bitcast=True)
        amax_i32 = (amax_i32 + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000
        amax_rounded = amax_i32.to(tl.float32, bitcast=True)
        scale_unbiased = tl.math.floor(tl.math.log2(amax_rounded)) - 2
        scale_unbiased = tl.maximum(tl.minimum(scale_unbiased, 127), -127)
        e8m0 = (scale_unbiased.to(tl.int32) + 127).to(tl.uint8)

        # Store: row-major (N, K//32)
        s_ptrs = Scale_ptr + n_offsets * num_groups + g
        tl.store(s_ptrs, e8m0, mask=n_mask)


def compute_b_scale(B, K):
    """Compute e8m0 scales from bf16 B. Much faster than full quant."""
    N = B.shape[0]
    num_groups = K // 32
    scale = torch.empty((N, num_groups), dtype=torch.uint8, device=B.device)
    BLOCK_N = 64
    grid = ((N + BLOCK_N - 1) // BLOCK_N,)
    _compute_b_scale_kernel[grid](
        B, scale, N, K,
        B.stride(0), B.stride(1),
        BLOCK_N=BLOCK_N, GROUP_SIZE=32,
    )
    return scale


# Per-shape optimal configs
_SHAPE_CONFIGS = {
    (4, 2880, 512):    {'BLOCK_SIZE_M': 4, 'BLOCK_SIZE_N': 128, 'BLOCK_SIZE_K': 256,
                        'GROUP_SIZE_M': 1, 'num_warps': 4, 'num_stages': 2,
                        'waves_per_eu': 3, 'matrix_instr_nonkdim': 16,
                        'cache_modifier': None, 'NUM_KSPLIT': 1},
    (16, 2112, 7168):  {'BLOCK_SIZE_M': 4, 'BLOCK_SIZE_N': 128, 'BLOCK_SIZE_K': 256,
                        'GROUP_SIZE_M': 1, 'num_warps': 4, 'num_stages': 2,
                        'waves_per_eu': 3, 'matrix_instr_nonkdim': 16,
                        'cache_modifier': None, 'NUM_KSPLIT': 16},
    (32, 4096, 512):   {'BLOCK_SIZE_M': 4, 'BLOCK_SIZE_N': 128, 'BLOCK_SIZE_K': 512,
                        'GROUP_SIZE_M': 1, 'num_warps': 4, 'num_stages': 2,
                        'waves_per_eu': 2, 'matrix_instr_nonkdim': 16,
                        'cache_modifier': None, 'NUM_KSPLIT': 1},
    (32, 2880, 512):   {'BLOCK_SIZE_M': 4, 'BLOCK_SIZE_N': 128, 'BLOCK_SIZE_K': 512,
                        'GROUP_SIZE_M': 1, 'num_warps': 4, 'num_stages': 2,
                        'waves_per_eu': 2, 'matrix_instr_nonkdim': 16,
                        'cache_modifier': None, 'NUM_KSPLIT': 1},
    (64, 7168, 2048):  {'BLOCK_SIZE_M': 16, 'BLOCK_SIZE_N': 128, 'BLOCK_SIZE_K': 256,
                        'GROUP_SIZE_M': 1, 'num_warps': 4, 'num_stages': 2,
                        'waves_per_eu': 2, 'matrix_instr_nonkdim': 16,
                        'cache_modifier': None, 'NUM_KSPLIT': 1},
    (256, 3072, 1536): {'BLOCK_SIZE_M': 16, 'BLOCK_SIZE_N': 256, 'BLOCK_SIZE_K': 256,
                        'GROUP_SIZE_M': 1, 'num_warps': 4, 'num_stages': 2,
                        'waves_per_eu': 2, 'matrix_instr_nonkdim': 16,
                        'cache_modifier': None, 'NUM_KSPLIT': 1},
}
_DEFAULT_CFG = {'BLOCK_SIZE_M': 4, 'BLOCK_SIZE_N': 128, 'BLOCK_SIZE_K': 256,
                'GROUP_SIZE_M': 1, 'num_warps': 4, 'num_stages': 2,
                'waves_per_eu': 3, 'matrix_instr_nonkdim': 16,
                'cache_modifier': None, 'NUM_KSPLIT': 1}

_b_cache = {}


def custom_kernel(data: input_t) -> output_t:
    A, B, B_q, B_shuffle, B_scale_sh = data
    A = A.contiguous()
    m, k = A.shape
    n = B.shape[0]

    # Cache B_scale
    b_ptr = B.data_ptr()
    if b_ptr not in _b_cache:
        _b_cache.clear()
        _, B_scale = dynamic_mxfp4_quant(B.contiguous())
        _b_cache[b_ptr] = (B_scale, B)

    B_scale, _ = _b_cache[b_ptr]
    B_q_u8 = B_q.view(torch.uint8)

    cfg = _SHAPE_CONFIGS.get((m, n, k), _DEFAULT_CFG)
    return gemm_a16wfp4(A, B_q_u8, B_scale, False, dtype=dtypes.bf16, config=cfg)
scrolls · 127 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