Skip to content
KernelIndex
Search⌘K

submission 715655

xiehuanyi · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_v18.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-715655?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
20.3µs
#723 of 1143
2026-04-03

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:6342cf42287e73bf012cfd933f20e900fa150aeaaac721e6a829cbc52b9c4837
license declaredunknown
license concludedunknown
authorsxiehuanyi
imported2026-08-26

Techniques

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

fp4MXFP4 GEMM v18: Optimal hybrid - Triton for small K+M, ASM for rest.

Kernel source

submission_v18.py89 lines
"""
MXFP4 GEMM v18: Optimal hybrid - Triton for small K+M, ASM for rest.

Based on empirical benchmarks:
- K <= 2048 and M <= 32: Triton gemm_afp4wfp4 is 15-21% faster
- K > 2048 or M > 32: ASM gemm_a4w4 with log2_k_split is better
"""
from task import input_t, output_t
import torch
import sys

_first_call = True


def custom_kernel(data: input_t) -> output_t:
    global _first_call
    import aiter
    from aiter import dtypes
    from aiter.ops.triton.quant import dynamic_mxfp4_quant
    from aiter.utility.fp4_utils import e8m0_shuffle

    A, B, B_q, B_shuffle, B_scale_sh = data
    A = A.contiguous()
    B = B.contiguous()
    m, k = A.shape
    n = B.shape[0]

    A_fp4_raw, A_scale_raw = dynamic_mxfp4_quant(A)

    # Trigger JIT on first call
    if _first_call:
        _first_call = False
        A_scale_sh = e8m0_shuffle(A_scale_raw).view(dtypes.fp8_e8m0)
        _ = aiter.gemm_a4w4(
            A_fp4_raw.view(dtypes.fp4x2), B_shuffle,
            A_scale_sh, B_scale_sh,
            dtype=dtypes.bf16, bpreshuffle=True,
        )

    # Decision: Triton for small K+M, ASM for rest
    use_triton = (m <= 32 and k <= 2048)

    if use_triton:
        _, B_scale_raw = dynamic_mxfp4_quant(B)
        try:
            from aiter.ops.triton.gemm.basic.gemm_afp4wfp4 import gemm_afp4wfp4
            n_blocks = (n + 63) // 64
            num_ksplit = max(1, min(304 // max(n_blocks, 1), k // 128))
            config = {
                'BLOCK_SIZE_M': 32, 'BLOCK_SIZE_N': 64, 'BLOCK_SIZE_K': 128,
                'GROUP_SIZE_M': 8, 'NUM_KSPLIT': num_ksplit,
                'num_warps': 4, 'num_stages': 2, 'waves_per_eu': 0,
                'matrix_instr_nonkdim': 16, 'cache_modifier': '.cg',
            }
            out = gemm_afp4wfp4(
                A_fp4_raw, B_q.view(A_fp4_raw.dtype),
                A_scale_raw, B_scale_raw,
                dtype=torch.bfloat16, config=config,
            )
            return out[:m, :n]
        except Exception as e:
            print(f"[OPT] Triton failed: {e}", file=sys.stderr)

    # ASM path with k_split for small M, default for large M
    A_fp4 = A_fp4_raw.view(dtypes.fp4x2)
    A_scale_sh = e8m0_shuffle(A_scale_raw).view(dtypes.fp8_e8m0)

    if m < 64:
        kname = '_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E'
        n_blocks = (n + 127) // 128
        best_ks = 0
        for ks in range(1, 8):
            if (k >> ks) < 64:
                break
            best_ks = ks
            if n_blocks * (1 << ks) >= 304:
                break
        out = torch.empty((m, n), dtype=torch.bfloat16, device='cuda')
        torch.ops.aiter.gemm_a4w4_asm(
            A_fp4, B_shuffle, A_scale_sh, B_scale_sh,
            out, kname, bpreshuffle=True, log2_k_split=best_ks,
        )
        return out

    return aiter.gemm_a4w4(
        A_fp4, B_shuffle, A_scale_sh, B_scale_sh,
        dtype=dtypes.bf16, bpreshuffle=True,
    )
scrolls · 89 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