Skip to content
KernelIndex
Search⌘K

submission 721756

Sikuan Wang · 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-12.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-721756?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.7µs
#678 of 1143
2026-04-04

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:be37673435b1efc827230414901b4cc2e0e048a07f3c4519434ad4ce005ab638
license declaredunknown
license concludedunknown
authorsSikuan Wang
imported2026-08-26

Techniques

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

fp4Optimized MXFP4 GEMM: stride-aware Triton e8m0 shuffle.

Kernel source

submission-12.py84 lines
"""
Optimized MXFP4 GEMM: stride-aware Triton e8m0 shuffle.

The scale tensor from dynamic_mxfp4_quant is column-major (stride=(1, 8)).
The reference e8m0_shuffle does: alloc pad → copy → permute → .contiguous()
  = 2 allocations + 2 kernel launches.

This replaces it with a single stride-aware Triton kernel that reads from
the column-major input and writes the shuffled+padded output in one pass:
  = 1 allocation + 1 kernel launch.
"""
from __future__ import annotations

import torch
import triton
import triton.language as tl
import aiter
from aiter import dtypes
from aiter.ops.triton.quant import dynamic_mxfp4_quant
from task import input_t, output_t

_FP4 = dtypes.fp4x2
_E8M0 = dtypes.fp8_e8m0
_BF16 = dtypes.bf16


@triton.jit
def _e8m0_shuffle_k(
    in_ptr, out_ptr,
    M, K_SCALE, STRIDE_0, STRIDE_1, SN,
    TOTAL,
    BLOCK: tl.constexpr,
):
    pid = tl.program_id(0)
    offs = pid * BLOCK + tl.arange(0, BLOCK)
    mask = offs < TOTAL

    s0 = 32 * SN
    a = offs // s0
    r = offs % s0
    d = r // 256
    r = r % 256
    f = r // 64
    r = r % 64
    c = r // 4
    r = r % 4
    e = r // 2
    b = r % 2

    in_row = a * 32 + b * 16 + c
    in_col = d * 8 + e * 4 + f

    in_bounds = mask & (in_row < M) & (in_col < K_SCALE)
    in_idx = in_row * STRIDE_0 + in_col * STRIDE_1

    val = tl.load(in_ptr + in_idx, mask=in_bounds, other=0)
    tl.store(out_ptr + offs, val, mask=mask)


def _shuffle(scale: torch.Tensor) -> torch.Tensor:
    m, ks = scale.shape
    s0, s1 = scale.stride()
    sm = (m + 255) // 256 * 256
    sn = (ks + 7) // 8 * 8
    out = torch.empty(sm, sn, dtype=scale.dtype, device=scale.device)
    total = sm * sn
    _e8m0_shuffle_k[(triton.cdiv(total, 1024),)](
        scale, out, m, ks, s0, s1, sn, total, BLOCK=1024,
    )
    return out


@torch.inference_mode()
def custom_kernel(data: input_t) -> output_t:
    a = data[0]
    if not a.is_contiguous():
        a = a.contiguous()
    aq, asc = dynamic_mxfp4_quant(a)
    return aiter.gemm_a4w4(
        aq.view(_FP4), data[3],
        _shuffle(asc).view(_E8M0), data[4],
        dtype=_BF16, bpreshuffle=True,
    )
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