Skip to content
KernelIndex
Search⌘K

submission 634613

brandonin · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_v36.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-634613?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
14.5µs
#516 of 1143
2026-03-25

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:e4a4ac8ad872e3e774bff833a0aeb5f8c2eeb5f7abf279e676b4be41b8e15d96
license declaredunknown
license concludedunknown
authorsbrandonin
imported2026-08-26

Kernel source

submission_v36.py61 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X

# v36: Fused BF16→FP4 quant + GEMM via Triton gemm_a16wfp4_preshuffle.
#
# Key insight: eliminate the separate A-quantization step entirely.
# gemm_a16wfp4_preshuffle takes BF16 A directly and quantizes on-the-fly.
#
# Scale format conversion (O(1)):
# B_scale_sh (CK e8m0_shuffle format): (N_pad, sn_pad) in uint8
# Triton shuffle_scales format: (N_pad//32, sn_pad*32) in uint8
# Both apply the SAME permutation to scale bytes — only the 2D view differs.
# So: B_scale_sh.view(uint8).reshape(N_pad//32, sn_pad*32) gives Triton format.
#
# Weight format conversion (O(1)):
# B_shuffle: shuffle_weight applied, shape (N, K//2)
# gemm_a16wfp4_preshuffle expects: (N//16, K//2*16)
# So: B_shuffle.view(uint8).reshape(N//16, K//2*16)

import os
os.environ['PYTORCH_ROCM_ARCH'] = 'gfx950'
os.environ['CXX'] = 'clang++'

import torch
import sys
from aiter.ops.triton.gemm.basic.gemm_a16wfp4 import gemm_a16wfp4_preshuffle
from task import input_t, output_t

_shape_cache = {}


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]
    shape_key = (m, n, k)

    # Reshape B_shuffle: (N, K//2) → (N//16, K//2*16) for preshuffle format
    w_triton = B_shuffle.view(torch.uint8).reshape(n // 16, (k // 2) * 16)

    # Convert B_scale_sh from CK view (N_pad, sn_pad) to Triton view (N_pad//32, sn_pad*32).
    # Same bytes in memory — just a different 2D interpretation.
    sm, sn = B_scale_sh.shape
    w_scales_triton = B_scale_sh.view(torch.uint8).reshape(sm // 32, sn * 32)

    # Cache output buffer per shape
    if shape_key not in _shape_cache:
        _shape_cache[shape_key] = torch.empty((m, n), dtype=torch.bfloat16, device='cuda')
        print(f"[v36] New shape ({m},{n},{k}): w_triton={w_triton.shape} "
              f"w_scales={w_scales_triton.shape}", file=sys.stderr)
    out = _shape_cache[shape_key]

    gemm_a16wfp4_preshuffle(
        A, w_triton, w_scales_triton,
        prequant=True,
        dtype=torch.bfloat16,
        y=out,
    )
    return out
scrolls · 61 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