Skip to content
KernelIndex
Search⌘K

submission 608622

Rakesh Jarupula · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

v3.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-608622?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.1µs
#934 of 1143
2026-03-22

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:ff1b0e882c1ea01dfe6fec893104b6541d023655ab7c025ff98682edffdfbcdb
license declaredunknown
license concludedunknown
authorsRakesh Jarupula
imported2026-08-26

Techniques

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

fp4Optimized FP4 GEMM kernel for AMD Instinct MI355X.

Kernel source

v3.py64 lines
"""
Optimized FP4 GEMM kernel for AMD Instinct MI355X.
bf16 A, MXFP4 B -> MXFP4 per-1x32 quant A -> gemm_a4w4 -> bf16 C.

Key optimizations:
1. Use dynamic_mxfp4_quant (patched #975) directly for A quantization
2. Use e8m0_shuffle for scale shuffling
3. Use aiter.gemm_a4w4 with bpreshuffle=True for the actual GEMM
4. Avoid redundant copies with .contiguous() only when necessary
"""

import torch
from task import input_t, output_t
from aiter import dtypes
import aiter
from aiter.ops.triton.quant import dynamic_mxfp4_quant
from aiter.utility.fp4_utils import e8m0_shuffle
from aiter.ops.shuffle import shuffle_weight

# Pre-import to avoid import overhead at runtime
_dynamic_mxfp4_quant = dynamic_mxfp4_quant
_e8m0_shuffle = e8m0_shuffle
_gemm_a4w4 = aiter.gemm_a4w4
_bf16 = dtypes.bf16
_fp4x2 = dtypes.fp4x2
_fp8_e8m0 = dtypes.fp8_e8m0


def _quant_a_mxfp4_shuffled(x: torch.Tensor):
    """Quantize x (bf16) to MXFP4 with shuffled E8M0 scales."""
    x_fp4, bs_e8m0 = _dynamic_mxfp4_quant(x)
    bs_e8m0_sh = _e8m0_shuffle(bs_e8m0)
    return x_fp4.view(_fp4x2), bs_e8m0_sh.view(_fp8_e8m0)


def custom_kernel(data: input_t) -> output_t:
    """
    Optimized MXFP4 GEMM:
      - A: [M, K] bf16  -> quantize to MXFP4 per-1x32 + shuffle scales
      - B_shuffle: [N, K/2] MXFP4 shuffled (16,16) tile coalesced  (precomputed)
      - B_scale_sh: [*, K/32] E8M0 shuffled                         (precomputed)
      - Output: [M, N] bf16
    """
    A, B, B_q, B_shuffle, B_scale_sh = data

    # Ensure A is contiguous for the quant kernel
    if not A.is_contiguous():
        A = A.contiguous()

    # Step 1: Quantize A to MXFP4 with shuffled E8M0 scales
    A_q, A_scale_sh = _quant_a_mxfp4_shuffled(A)

    # Step 2: GEMM a4w4 with pre-shuffled B and scales
    out = _gemm_a4w4(
        A_q,
        B_shuffle,
        A_scale_sh,
        B_scale_sh,
        dtype=_bf16,
        bpreshuffle=True,
    )

    return out
scrolls · 64 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