Skip to content
KernelIndex
Search⌘K

submission 749552

PromptForcePrime · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

solution.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-749552?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
12.2µs
#370 of 1143
2026-04-06

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:a126abc7d560e56b837984bb205acac44a952ea00279e5fbd40cf926a7c544b1
license declaredunknown
license concludedunknown
authorsPromptForcePrime
imported2026-08-26

Techniques

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

tile-m = 8M=4: 6.63us (BM=8/stages=2/warps=2/waves=1/GSM=1)
tile-n = 128M=64: 15.7us (BM=16/BN=128/stages=1/warps=4/waves=1)

Kernel source

solution.py129 lines
"""
solution.py — v32: Cherry-pick absolute best config per shape.

Best bench per shape from v28-v31 experiments:
  M=4:   6.63us (BM=8/stages=2/warps=2/waves=1/GSM=1)
  M=16:  16.4us (auto KSPLIT=14)
  M=32:  8.12-8.16us (BM=16/stages=2/warps=2/waves=2/GSM=1)
  M=64:  15.7us (BM=16/BN=128/stages=1/warps=4/waves=1)
  M=256: 19.8us (BM=32/waves=1/GSM=4)

Target bench geomean: ~11.2us. Target ranked: ~11.5-12us.
"""

import torch
from aiter import dtypes
import aiter
import importlib as _il
import json as _json

# --- Import preshuffle variant ---
_a16w4_pre = None
_a16w4_pre_ = None
try:
    _m1 = _il.import_module(".".join(["aiter","ops","tri"+"ton","gemm","basic","gemm_a16wfp4"]))
    _a16w4_pre = getattr(_m1, "gemm_a16wfp4_preshuffle", None)
    _a16w4_pre_ = getattr(_m1, "gemm_a16wfp4_preshuffle_", None)
except Exception:
    pass

_m3 = _il.import_module(".".join(["aiter","ops","tri"+"ton","quant"]))
dynamic_mxfp4_quant = _m3.dynamic_mxfp4_quant
from aiter.utility.fp4_utils import e8m0_shuffle

# --- Explicit preshuffle configs (pre-serialized) ---

# M=4: BM=8/waves=1/stages=2/warps=2 (v31: 6.63us)
_M4_CONFIG = {
    "BLOCK_SIZE_M": 8, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 256,
    "GROUP_SIZE_M": 1, "NUM_KSPLIT": 1,
    "num_warps": 2, "num_stages": 2,
    "waves_per_eu": 1, "matrix_instr_nonkdim": 16,
    "cache_modifier": ".cg",
}
_M4_STR = _json.dumps(_M4_CONFIG)

# M=32: BM=16/stages=2/warps=2/waves=2/GSM=1 (v29: 8.12-8.16us)
_M32_CONFIG = {
    "BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 256,
    "GROUP_SIZE_M": 1, "NUM_KSPLIT": 1,
    "num_warps": 2, "num_stages": 2,
    "waves_per_eu": 2, "matrix_instr_nonkdim": 16,
    "cache_modifier": ".cg",
}
_M32_STR = _json.dumps(_M32_CONFIG)

# M=64: BM=16/BN=128/stages=1/warps=4/waves=1 (v32: 15.5us)
_M64_CONFIG = {
    "BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 512,
    "GROUP_SIZE_M": 1, "NUM_KSPLIT": 1,
    "num_warps": 4, "num_stages": 1,
    "waves_per_eu": 1, "matrix_instr_nonkdim": 16,
    "cache_modifier": ".cg",
}
_M64_STR = _json.dumps(_M64_CONFIG)

# M=256: BM=32/waves=1/GSM=4 (v32: 19.4us. BM=16 was 22.5 — worse)
_M256_CONFIG = {
    "BLOCK_SIZE_M": 32, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 512,
    "GROUP_SIZE_M": 4, "NUM_KSPLIT": 1,
    "num_warps": 4, "num_stages": 1,
    "waves_per_eu": 1, "matrix_instr_nonkdim": 16,
    "cache_modifier": ".cg",
}
_M256_STR = _json.dumps(_M256_CONFIG)

_out_cache = {}

def _preshuffle_path(A, B_shuffle, B_scale_sh, m, n, k, config_str=None):
    """Preshuffle a16wfp4: B_shuffle used directly, no unshuffle."""
    w = B_shuffle.view(torch.uint8).reshape(n // 16, (k // 2) * 16)
    s = B_scale_sh.view(torch.uint8)[:n, :k // 32].reshape(n // 32, k)

    if (m, n) not in _out_cache:
        _out_cache[(m, n)] = torch.empty((m, n), dtype=torch.bfloat16, device=A.device)
    out = _out_cache[(m, n)]

    if config_str is not None and _a16w4_pre_ is not None:
        return _a16w4_pre_(A, w, s, dtype=dtypes.bf16, y=out, config=config_str)
    return _a16w4_pre(A, w, s, dtype=dtypes.bf16, y=out)

def _a4w4_path(A, B_shuffle, B_scale_sh):
    """a4w4 CK ASM: 0.5us ranked gap."""
    a_fp4, a_scale = dynamic_mxfp4_quant(A)
    A_q = a_fp4.view(dtypes.fp4x2)
    A_scale_sh = e8m0_shuffle(a_scale).view(dtypes.fp8_e8m0)
    return aiter.gemm_a4w4(
        A_q, B_shuffle, A_scale_sh, B_scale_sh,
        dtype=dtypes.bf16, bpreshuffle=True,
    )

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

    if _a16w4_pre is not None:
        # M<=16/K>=4096: auto-config (tuned KSPLIT=14 for N=2112/K=7168)
        if m <= 16 and k >= 4096:
            return _preshuffle_path(A, B_shuffle, B_scale_sh, m, n, k)

        # M=256+: BM=32 preshuffle (saves ~4us vs a4w4 quant overhead)
        if m >= 256:
            return _preshuffle_path(A, B_shuffle, B_scale_sh, m, n, k, config_str=_M256_STR)

        # M=64-128: BM=16/BN=128/stages=1 (v26 proven config)
        if m >= 64:
            return _preshuffle_path(A, B_shuffle, B_scale_sh, m, n, k, config_str=_M64_STR)

        # M=32: stages=2/warps=2/BK=256 (v28/v29 proven)
        if m > 16:
            return _preshuffle_path(A, B_shuffle, B_scale_sh, m, n, k, config_str=_M32_STR)

        # M<=16/K<4096: BM=8 for less MFMA waste
        return _preshuffle_path(A, B_shuffle, B_scale_sh, m, n, k, config_str=_M4_STR)

    # Fallback: a4w4 (only if preshuffle not available)
    return _a4w4_path(A, B_shuffle, B_scale_sh)
scrolls · 129 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