Skip to content
KernelIndex
Search⌘K

submission 517426

Ryan Mathieu · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_reuse_triton_k512m32_bs32.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-517426?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
13.6µs
#454 of 1143
2026-03-08

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:2516976a2578b238abd359ad7e2f95a4ba82fae339164cc2862d6a5553fc5c0a
license declaredunknown
license concludedunknown
authorsRyan Mathieu
imported2026-08-26

Kernel source

submission_reuse_triton_k512m32_bs32.py84 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X

from __future__ import annotations

import torch
import triton
from task import input_t, output_t
import aiter
from aiter import dtypes
from aiter.utility import fp4_utils

ASM = aiter.gemm_a4w4_asm
K32 = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"
QK = fp4_utils._dynamic_mxfp4_quant_kernel_asm_layout
_AQ_CACHE = {}
_SCALE_CACHE = {}
_OUT_CACHE = {}


def _aq_buf(m: int, k: int, device: torch.device):
    key = (m, k, device)
    buf = _AQ_CACHE.get(key)
    if buf is None:
        buf = torch.empty((m, k // 2), dtype=torch.uint8, device=device)
        _AQ_CACHE[key] = buf
    return buf


def _scale_buf(m: int, k: int, device: torch.device):
    key = (m, k, device)
    buf = _SCALE_CACHE.get(key)
    if buf is None:
        scale_n = triton.cdiv(k // 32, 8) * 8
        buf = torch.empty((triton.cdiv(m, 256) * 256, scale_n), dtype=torch.uint8, device=device)
        _SCALE_CACHE[key] = buf
    return buf


def _out_buf(m: int, n: int, device: torch.device):
    key = (m, n, device)
    buf = _OUT_CACHE.get(key)
    if buf is None:
        buf = torch.empty((((m + 31) // 32) * 32, n), dtype=dtypes.bf16, device=device)
        _OUT_CACHE[key] = buf
    return buf


def _quant_into(x: torch.Tensor, x_fp4: torch.Tensor, scale: torch.Tensor):
    m, n = x.shape
    block_size = 32 if (n == 512 and m <= 32) else 128
    scale_m = triton.cdiv(m, 32) * 32
    scale_n_valid = triton.cdiv(n, 32)
    scale_n_pad = triton.cdiv(scale_n_valid, 8) * 8
    QK[(triton.cdiv(m, block_size), scale_n_pad)](
        x,
        x_fp4,
        scale,
        *x.stride(),
        *x_fp4.stride(),
        *scale.stride(),
        M=m,
        N=n,
        scaleN=scale_n_valid,
        scaleM_pad=scale_m,
        scaleN_pad=scale_n_pad,
        BLOCK_SIZE=block_size,
        MXFP4_QUANT_BLOCK_SIZE=32,
        SCALING_MODE=0,
        SHUFFLE=True,
    )


@torch.inference_mode()
def custom_kernel(data: input_t) -> output_t:
    A, _B, _B_q, B_shuffle, B_scale_sh = data
    m, k = A.shape
    aq_u8 = _aq_buf(m, k, A.device)
    a_scale_u8 = _scale_buf(m, k, A.device)
    out = _out_buf(m, B_shuffle.shape[0], A.device)
    _quant_into(A, aq_u8, a_scale_u8)
    ASM(aq_u8.view(dtypes.fp4x2), B_shuffle, a_scale_u8.view(dtypes.fp8_e8m0), B_scale_sh, out, K32, None, 1.0, 0.0, True, 0)
    return out[:m]
scrolls · 84 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