Skip to content
KernelIndex
Search⌘K

submission 755042

musicofhel · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

gemm_v4.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-755042?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
23.7µs
#802 of 1143
2026-04-07

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:2b622658d7d671e365faca9a1620bf45df357efdc175e0bf21efed25298db9bd
license declaredunknown
license concludedunknown
authorsmusicofhel
imported2026-08-26

Techniques

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

split-kGEMM v4: Shape-specific ASM kernel tile selection + splitK.

Kernel source

gemm_v4.py109 lines
"""
GEMM v4: Shape-specific ASM kernel tile selection + splitK.

Forces optimal tile size AND splitK per benchmark shape.
Available .co tiles: 32x{128-1024}, 64x{128-1024}, 96x128,
128x{128-512}, 160x{128-384}, 192x{128-256}, 224-256x{128-256}

Strategy:
- tile_M >= M (avoid wasted rows from padding to tile boundary)
- tile_N chosen to create ~256/splitK total tiles for good CU util
- splitK fills remaining idle CUs along K dimension
"""
import torch
from task import input_t, output_t
import aiter
from aiter import dtypes
from aiter.ops.triton.quant import dynamic_mxfp4_quant
from aiter.utility.fp4_utils import e8m0_shuffle

try:
    from aiter.ops.gemm_op_a4w4 import gemm_a4w4_asm
except ImportError:
    try:
        from aiter import gemm_a4w4_asm
    except ImportError:
        gemm_a4w4_asm = None

_dyn_quant = dynamic_mxfp4_quant
_shuffle = e8m0_shuffle
_asm_gemm = gemm_a4w4_asm
_fallback_gemm = aiter.gemm_a4w4
_fp4x2 = dtypes.fp4x2
_e8m0 = dtypes.fp8_e8m0

def _mangle(tile_m, tile_n):
    """Construct C++ mangled kernel name for given tile size."""
    name = f"f4gemm_bf16_per1x32Fp4_BpreShuffle_{tile_m}x{tile_n}"
    return f"_ZN5aiter{len(name)}{name}E"

# Pre-computed optimal configs per (M, N, K)
# Format: (kernelName, log2_k_split)
# tile_M: 32 for M<=32, 64 for M<=64, 128/192/256 for larger M
# splitK: fill idle CUs
_SHAPE_CONFIGS = {
    # splitK=0 for all (splitK causes accuracy failures with FP4!)
    # M=4: 32x128 is the smallest M-tile available
    (4, 2880, 512): (_mangle(32, 128), 0),
    # M=16: 32x128 (exact tile match not available)
    (16, 2112, 7168): (_mangle(32, 128), 0),
    # M=32: 32x128 (exact M match)
    (32, 4096, 512): (_mangle(32, 128), 0),
    (32, 2880, 512): (_mangle(32, 128), 0),
    # M=64: force 64x128 (exact M match, less waste than 192x128)
    (64, 7168, 2048): (_mangle(64, 128), 0),
    # M=256: force 256x128 (exact M match, less waste than 192x128)
    (256, 3072, 1536): (_mangle(256, 128), 0),
    # Test shapes:
    (8, 2112, 7168): (_mangle(32, 128), 0),
    (16, 3072, 1536): (_mangle(32, 128), 0),
    (64, 3072, 1536): (_mangle(64, 128), 0),
    (256, 2880, 512): (_mangle(256, 128), 0),
}

# Default fallback config
_DEFAULT_SPLITK = 0


def custom_kernel(data: input_t) -> output_t:
    A, B, B_q, B_shuffle, B_scale_sh = data

    m, k = A.shape
    n = B.shape[0]

    # Quantize A to MXFP4
    A_q, A_scale = _dyn_quant(A.contiguous())
    A_scale = _shuffle(A_scale)

    A_q_fp4 = A_q.view(_fp4x2)
    A_scale_e8m0 = A_scale.view(_e8m0)

    if _asm_gemm is None:
        return _fallback_gemm(A_q_fp4, B_shuffle, A_scale_e8m0, B_scale_sh,
                              dtype=torch.bfloat16, bpreshuffle=True)

    # Get shape-specific config
    config = _SHAPE_CONFIGS.get((m, n, k))
    if config is not None:
        kernel_name, splitK = config
    else:
        kernel_name, splitK = "", _DEFAULT_SPLITK

    padded_m = (m + 31) // 32 * 32
    out = torch.empty((padded_m, n), dtype=torch.bfloat16, device='cuda')

    _asm_gemm(
        A_q_fp4,
        B_shuffle,
        A_scale_e8m0,
        B_scale_sh,
        out,
        kernel_name,
        None,   # bias
        1.0,    # alpha
        0.0,    # beta
        True,   # bpreshuffle
        splitK,
    )
    return out[:m]
scrolls · 109 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