Skip to content
KernelIndex
Search⌘K

submission 639085

Amanpreet Singh · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:2d721db00a6e5db1c29f0b0f86fbda1851f041a267877b4e718b447edc09d123
license declaredunknown
license concludedunknown
authorsAmanpreet Singh
imported2026-08-26

Techniques

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

mmaacc += tl.dot(a_scaled, b_scaled)

Kernel source

submission.py147 lines
import torch
import triton
import triton.language as tl
from task import input_t, output_t
from utils import make_match_reference
from aiter import dtypes
from aiter.ops.shuffle import shuffle_weight
from aiter.ops.triton.quant import dynamic_mxfp4_quant
from aiter.utility.fp4_utils import e8m0_shuffle

SCALE_GROUP_SIZE = 32


def _quant_mxfp4_shuffled(x):
    x_fp4, bs_e8m0 = dynamic_mxfp4_quant(x)
    bs_e8m0 = e8m0_shuffle(bs_e8m0)
    return x_fp4.view(dtypes.fp4x2), bs_e8m0.view(dtypes.fp8_e8m0)


@triton.jit
def _fp4_unpack_and_scale_dot(
    a_ptr, a_scale_ptr,
    b_ptr, b_scale_ptr,
    c_ptr,
    M, N, K,
    stride_am, stride_ak,
    stride_bk, stride_bn,
    stride_cm, stride_cn,
    BLOCK_M: tl.constexpr,
    BLOCK_N: tl.constexpr,
    BLOCK_K: tl.constexpr,
):
    pid_m = tl.program_id(0)
    pid_n = tl.program_id(1)

    offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
    offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
    offs_k = tl.arange(0, BLOCK_K)

    acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)

    num_k_blocks = tl.cdiv(K, BLOCK_K)

    for k_block in range(num_k_blocks):
        k_start = k_block * BLOCK_K

        a_offs = offs_m[:, None] * stride_am + ((k_start + offs_k)[None, :] // 2) * stride_ak
        mask_a = (offs_m[:, None] < M) & ((k_start + offs_k)[None, :] < K)
        a_packed = tl.load(a_ptr + a_offs, mask=mask_a, other=0).to(tl.uint8)

        b_offs = ((k_start + offs_k)[:, None] // 2) * stride_bk + offs_n[None, :] * stride_bn
        mask_b = ((k_start + offs_k)[:, None] < K) & (offs_n[None, :] < N)
        b_packed = tl.load(b_ptr + b_offs, mask=mask_b, other=0).to(tl.uint8)

        a_lo = (a_packed & 0x0F).to(tl.float32)
        a_hi = ((a_packed >> 4) & 0x0F).to(tl.float32)

        b_lo = (b_packed & 0x0F).to(tl.float32)
        b_hi = ((b_packed >> 4) & 0x0F).to(tl.float32)

        a_f32 = tl.interleave(a_lo, a_hi)
        b_f32 = tl.interleave(b_lo, b_hi)

        k_scale = (k_start + tl.arange(0, BLOCK_K)) // SCALE_GROUP_SIZE
        a_scale_offs = offs_m[:, None] * tl.cdiv(K, SCALE_GROUP_SIZE) + k_scale[None, :]
        a_scale_mask = (offs_m[:, None] < M) & (k_scale[None, :] < tl.cdiv(K, SCALE_GROUP_SIZE))
        a_scale = tl.load(a_scale_ptr + a_scale_offs, mask=a_scale_mask, other=0).to(tl.uint8)
        a_exp = a_scale.to(tl.float32) - 127.0
        a_scale_f32 = tl.exp2(a_exp)

        b_scale_offs = offs_n[None, :] * tl.cdiv(K, SCALE_GROUP_SIZE) + k_scale[:, None]
        b_scale_mask = (offs_n[None, :] < N) & (k_scale[:, None] < tl.cdiv(K, SCALE_GROUP_SIZE))
        b_scale = tl.load(b_scale_ptr + b_scale_offs, mask=b_scale_mask, other=0).to(tl.uint8)
        b_exp = b_scale.to(tl.float32) - 127.0
        b_scale_f32 = tl.exp2(b_exp)

        a_scaled = a_f32 * a_scale_f32
        b_scaled = b_f32 * b_scale_f32

        acc += tl.dot(a_scaled, b_scaled)

    c_offs = offs_m[:, None] * stride_cm + offs_n[None, :] * stride_cn
    mask_c = (offs_m[:, None] < M) & (offs_n[None, :] < N)
    tl.store(c_ptr + c_offs, acc.to(tl.bfloat16), mask=mask_c)


_SHAPE_CONFIGS = {
    (4,   2880, 512):  {"BLOCK_M": 16,  "BLOCK_N": 64,  "BLOCK_K": 64,  "num_warps": 2, "num_stages": 2},
    (16,  2112, 7168): {"BLOCK_M": 16,  "BLOCK_N": 128, "BLOCK_K": 64,  "num_warps": 4, "num_stages": 2},
    (32,  4096, 512):  {"BLOCK_M": 32,  "BLOCK_N": 128, "BLOCK_K": 64,  "num_warps": 4, "num_stages": 2},
    (32,  2880, 512):  {"BLOCK_M": 32,  "BLOCK_N": 64,  "BLOCK_K": 64,  "num_warps": 4, "num_stages": 2},
    (64,  7168, 2048): {"BLOCK_M": 64,  "BLOCK_N": 128, "BLOCK_K": 64,  "num_warps": 4, "num_stages": 2},
    (256, 3072, 1536): {"BLOCK_M": 64,  "BLOCK_N": 128, "BLOCK_K": 64,  "num_warps": 4, "num_stages": 3},
}

_DEFAULT_CONFIG = {"BLOCK_M": 32, "BLOCK_N": 64, "BLOCK_K": 64, "num_warps": 4, "num_stages": 2}


def _get_config(m, n, k):
    return _SHAPE_CONFIGS.get((m, n, k), _DEFAULT_CONFIG)


def _custom_quant_gemm(A, B_shuffle, B_scale_sh):
    import aiter
    A_q, A_scale_sh = _quant_mxfp4_shuffled(A)
    out = aiter.gemm_a4w4(
        A_q,
        B_shuffle,
        A_scale_sh,
        B_scale_sh,
        dtype=dtypes.bf16,
        bpreshuffle=True,
    )
    return out


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_shuffle.shape[0] if hasattr(B_shuffle, 'shape') else B_q.shape[0]

    return _custom_quant_gemm(A, B_shuffle, B_scale_sh)


def _quant_mxfp4_noshuf(x):
    x_fp4, bs_e8m0 = dynamic_mxfp4_quant(x)
    return x_fp4.view(dtypes.fp4x2), bs_e8m0.view(dtypes.fp8_e8m0)


def ref_kernel(data: input_t) -> output_t:
    import aiter
    A, B, B_q, B_shuffle, B_scale_sh = data
    A = A.contiguous()
    A_q, A_scale_sh = _quant_mxfp4_shuffled(A)
    out = aiter.gemm_a4w4(
        A_q,
        B_shuffle,
        A_scale_sh,
        B_scale_sh,
        dtype=dtypes.bf16,
        bpreshuffle=True,
    )
    return out


check_implementation = make_match_reference(ref_kernel, rtol=1e-02, atol=1e-02)
scrolls · 147 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