Skip to content
KernelIndex
Search⌘K

submission 533545

josusanmartin · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:48d7a8501210e3d4b0bdb38d6e8b55e5eeafc82ccd92a823db9432e7e538cde3
license declaredunknown
license concludedunknown
authorsjosusanmartin
imported2026-08-15

Kernel source

submission.py287 lines
# Write your code here# Write your code here# Write your code here#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X

from __future__ import annotations

import torch
import triton
import triton.language as tl
import aiter
from aiter import dtypes
from aiter.ops.gemm_op_a4w4 import gemm_a4w4_asm
from aiter.ops.triton.quant import dynamic_mxfp4_quant
from aiter.ops.triton._triton_kernels.quant.quant import _mxfp4_quant_op as _mxfp4_quant_op_even

from task import input_t, output_t

import aiter.ops.triton._triton_kernels.gemm.basic.gemm_a16wfp4 as _kernel_module

_kernel_module._mxfp4_quant_op = _mxfp4_quant_op_even

from aiter.ops.triton.gemm.basic.gemm_a16wfp4 import gemm_a16wfp4_preshuffle

_BF16 = dtypes.bf16
_FP4X2 = dtypes.fp4x2
_FP8_E8M0 = dtypes.fp8_e8m0
_KERNEL_32X128 = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"

_PUBLIC_SMALL = {
    (4, 2880, 512): {
        "BLOCK_SIZE_M": 8,
        "BLOCK_SIZE_N": 64,
        "BLOCK_SIZE_K": 512,
        "GROUP_SIZE_M": 1,
        "num_warps": 4,
        "num_stages": 2,
        "waves_per_eu": 2,
        "matrix_instr_nonkdim": 16,
        "cache_modifier": ".cg",
        "NUM_KSPLIT": 1,
    },
    (16, 2112, 7168): {
        "BLOCK_SIZE_M": 16,
        "BLOCK_SIZE_N": 64,
        "BLOCK_SIZE_K": 512,
        "GROUP_SIZE_M": 1,
        "num_warps": 4,
        "num_stages": 2,
        "waves_per_eu": 1,
        "matrix_instr_nonkdim": 16,
        "cache_modifier": ".cg",
        "NUM_KSPLIT": 8,
    },
    (32, 4096, 512): {
        "BLOCK_SIZE_M": 8,
        "BLOCK_SIZE_N": 64,
        "BLOCK_SIZE_K": 512,
        "GROUP_SIZE_M": 1,
        "num_warps": 4,
        "num_stages": 3,
        "waves_per_eu": 1,
        "matrix_instr_nonkdim": 16,
        "cache_modifier": ".cg",
        "NUM_KSPLIT": 1,
    },
    (32, 2880, 512): {
        "BLOCK_SIZE_M": 8,
        "BLOCK_SIZE_N": 64,
        "BLOCK_SIZE_K": 512,
        "GROUP_SIZE_M": 1,
        "num_warps": 4,
        "num_stages": 2,
        "waves_per_eu": 1,
        "matrix_instr_nonkdim": 16,
        "cache_modifier": None,
        "NUM_KSPLIT": 1,
    },
}

_PUBLIC_LARGE = {
    (64, 7168, 2048): 2,
    (256, 3072, 1536): 1,
}

_HIDDEN_SHAPES = {
    (8, 2112, 7168),
    (16, 3072, 1536),
    (64, 3072, 1536),
    (256, 2880, 512),
}

_QUANT_BLOCK = 32
_QUANT_TILE = 128
_OUT_PAD_BF16 = 32
_BUFS = {}


@triton.jit
def _dynamic_mxfp4_quant_kernel_even_asm_layout(
    x_ptr,
    x_fp4_ptr,
    bs_ptr,
    stride_x_m,
    stride_x_n,
    stride_x_fp4_m,
    stride_x_fp4_n,
    stride_bs_m,
    stride_bs_n,
    M: tl.constexpr,
    N: tl.constexpr,
    scaleN: tl.constexpr,
    scaleM_pad: tl.constexpr,
    scaleN_pad: tl.constexpr,
    BLOCK_SIZE: tl.constexpr,
    MXFP4_QUANT_BLOCK_SIZE: tl.constexpr,
    SHUFFLE: tl.constexpr,
):
    pid_m = tl.program_id(0)
    pid_n = tl.program_id(1)

    stride_x_m = tl.cast(stride_x_m, tl.int64)
    stride_x_n = tl.cast(stride_x_n, tl.int64)
    stride_x_fp4_m = tl.cast(stride_x_fp4_m, tl.int64)
    stride_x_fp4_n = tl.cast(stride_x_fp4_n, tl.int64)

    x_offs_m = pid_m * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
    x_offs_n = pid_n * MXFP4_QUANT_BLOCK_SIZE + tl.arange(0, MXFP4_QUANT_BLOCK_SIZE)
    x_offs = x_offs_m[:, None] * stride_x_m + x_offs_n[None, :] * stride_x_n
    x_mask = (x_offs_m < M)[:, None] & (x_offs_n < N)[None, :]
    x = tl.load(x_ptr + x_offs, mask=x_mask).to(tl.float32)

    out_tensor, bs_e8m0 = _mxfp4_quant_op_even(
        x,
        MXFP4_QUANT_BLOCK_SIZE,
        BLOCK_SIZE,
        MXFP4_QUANT_BLOCK_SIZE,
    )

    out_offs_m = pid_m * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
    out_offs_n = pid_n * MXFP4_QUANT_BLOCK_SIZE // 2 + tl.arange(
        0, MXFP4_QUANT_BLOCK_SIZE // 2
    )
    out_offs = (
        out_offs_m[:, None] * stride_x_fp4_m + out_offs_n[None, :] * stride_x_fp4_n
    )
    out_mask = (out_offs_m < M)[:, None] & (out_offs_n < (N // 2))[None, :]
    tl.store(x_fp4_ptr + out_offs, out_tensor, mask=out_mask)

    bs_offs_m = pid_m * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
    bs_offs_n = pid_n

    if SHUFFLE:
        bs_offs_0 = bs_offs_m[:, None] // 32
        bs_offs_1 = bs_offs_m[:, None] % 32
        bs_offs_2 = bs_offs_1 % 16
        bs_offs_1 = bs_offs_1 // 16
        bs_offs_3 = bs_offs_n[None, :] // 8
        bs_offs_4 = bs_offs_n[None, :] % 8
        bs_offs_5 = bs_offs_4 % 4
        bs_offs_4 = bs_offs_4 // 4
        bs_offs = (
            bs_offs_1
            + bs_offs_4 * 2
            + bs_offs_2 * 4
            + bs_offs_5 * 64
            + bs_offs_3 * 256
            + bs_offs_0 * 32 * scaleN
        )
        bs_mask1 = (bs_offs_m < M)[:, None] & (bs_offs_n < scaleN)[None, :]
        bs_mask2 = (bs_offs_m < scaleM_pad)[:, None] & (bs_offs_n < scaleN_pad)[None, :]
        bs_e8m0 = tl.where(bs_mask1, bs_e8m0, 127)
        tl.store(bs_ptr + bs_offs, bs_e8m0, mask=bs_mask2)
    else:
        bs_offs = bs_offs_m[:, None] * stride_bs_m + bs_offs_n[None, :] * stride_bs_n
        bs_mask = (bs_offs_m < M)[:, None] & (bs_offs_n < N)[None, :]
        tl.store(bs_ptr + bs_offs, bs_e8m0, mask=bs_mask)


def _e8m0_shuffle_safe(scale: torch.Tensor) -> torch.Tensor:
    m, n = scale.shape
    scale_padded = torch.empty(
        ((m + 255) // 256) * 256,
        ((n + 7) // 8) * 8,
        dtype=scale.dtype,
        device=scale.device,
    )
    scale_padded.fill_(0x7F)
    scale_padded[:m, :n] = scale
    sm, sn = scale_padded.shape
    return (
        scale_padded.view(sm // 32, 2, 16, sn // 8, 2, 4)
        .permute(0, 3, 5, 2, 4, 1)
        .contiguous()
        .view(sm, sn)
    )


def _safe_wrapper(a: torch.Tensor, b_shuffle: torch.Tensor, b_scale_sh: torch.Tensor):
    a_q_raw, a_scale = dynamic_mxfp4_quant(a.contiguous())
    a_scale_sh = _e8m0_shuffle_safe(a_scale)
    return aiter.gemm_a4w4(
        a_q_raw.view(_FP4X2),
        b_shuffle,
        a_scale_sh.view(_FP8_E8M0),
        b_scale_sh,
        dtype=_BF16,
        bpreshuffle=True,
    )


def _get_large_bufs(m: int, k: int, n: int, device):
    x_fp4 = torch.empty((m, k >> 1), dtype=torch.uint8, device=device)
    scale_n = (k + _QUANT_BLOCK - 1) // _QUANT_BLOCK
    scale_n_pad = ((scale_n + 7) >> 3) << 3
    scale_m_pad = ((m + 255) >> 8) << 8
    scale = torch.empty((scale_m_pad, scale_n_pad), dtype=torch.uint8, device=device)
    padded_m = ((m + 31) >> 5) << 5
    out = torch.empty_strided(
        (padded_m, n),
        (n + _OUT_PAD_BF16, 1),
        dtype=_BF16,
        device=device,
    )
    return x_fp4, scale, scale_n, scale_n_pad, scale_m_pad, out


@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
    n = b.shape[0]
    key = (int(m), int(n), int(k))

    if key in _HIDDEN_SHAPES or (key not in _PUBLIC_SMALL and key not in _PUBLIC_LARGE):
        return _safe_wrapper(a, b_shuffle, b_scale_sh)

    if key in _PUBLIC_LARGE:
        if key not in _BUFS:
            _BUFS[key] = ("large", _get_large_bufs(m, k, n, a.device))
        _, (x_fp4, scale, scale_n, scale_n_pad, scale_m_pad, out) = _BUFS[key]
        grid = ((m + _QUANT_TILE - 1) // _QUANT_TILE, scale_n_pad)
        _dynamic_mxfp4_quant_kernel_even_asm_layout[grid](
            a,
            x_fp4,
            scale,
            a.stride(0),
            a.stride(1),
            x_fp4.stride(0),
            x_fp4.stride(1),
            scale.stride(0),
            scale.stride(1),
            M=m,
            N=k,
            scaleN=scale_n,
            scaleM_pad=scale_m_pad,
            scaleN_pad=scale_n_pad,
            BLOCK_SIZE=_QUANT_TILE,
            MXFP4_QUANT_BLOCK_SIZE=_QUANT_BLOCK,
            SHUFFLE=True,
        )
        gemm_a4w4_asm(
            x_fp4.view(_FP4X2),
            b_shuffle,
            scale.view(_FP8_E8M0),
            b_scale_sh,
            out,
            _KERNEL_32X128,
            bpreshuffle=True,
            log2_k_split=_PUBLIC_LARGE[key],
        )
        return out[:m]

    if key not in _BUFS:
        _BUFS[key] = ("small", torch.empty((m, n), dtype=_BF16, device=a.device))
    _, out = _BUFS[key]

    w = b_shuffle.view(torch.uint8).reshape(n // 16, k // 2 * 16)
    sm, sn = b_scale_sh.shape
    w_scales = b_scale_sh.view(torch.uint8).reshape(sm // 32, sn * 32)
    return gemm_a16wfp4_preshuffle(
        a,
        w,
        w_scales,
        prequant=True,
        y=out,
        config=_PUBLIC_SMALL.get(key),
    )
scrolls · 287 lines total

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

Changes from previous submission

Against this author's previous submission submission 532764.

- #!POPCORN leaderboard amd-mxfp4-mm
+ # Write your code here# Write your code here# Write your code here#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
from __future__ import annotations
⋯ 43 unchanged lines
"waves_per_eu": 1,
"matrix_instr_nonkdim": 16,
"cache_modifier": ".cg",
- "NUM_KSPLIT": 7,
+ "NUM_KSPLIT": 8,
},
(32, 4096, 512): {
"BLOCK_SIZE_M": 8,
⋯ 13 unchanged lines
"BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1,
"num_warps": 4,
- "num_stages": 3,
+ "num_stages": 2,
"waves_per_eu": 1,
"matrix_instr_nonkdim": 16,
- "cache_modifier": ".cg",
+ "cache_modifier": None,
"NUM_KSPLIT": 1,
},
}
⋯ 12 unchanged lines
_QUANT_BLOCK = 32
_QUANT_TILE = 128
+ _OUT_PAD_BF16 = 32
_BUFS = {}
⋯ 117 unchanged lines
scale_m_pad = ((m + 255) >> 8) << 8
scale = torch.empty((scale_m_pad, scale_n_pad), dtype=torch.uint8, device=device)
padded_m = ((m + 31) >> 5) << 5
- out = torch.empty((padded_m, n), dtype=_BF16, device=device)
+ out = torch.empty_strided(
+ (padded_m, n),
+ (n + _OUT_PAD_BF16, 1),
+ dtype=_BF16,
+ device=device,
+ )
return x_fp4, scale, scale_n, scale_n_pad, scale_m_pad, out
scrolls · 49 diff lines total

Best evidence level for this revision: reported

JSON