Skip to content
KernelIndex
Search⌘K

submission 747677

kitrak_rev. · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:abdd187613fabd3a9f9b65583a5f3d213075cfd5dda52f32b752452b9910405e
license declaredunknown
license concludedunknown
authorskitrak_rev.
imported2026-08-26

Techniques

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

split-k6. Use log2_k_split=3 (splitK=8) for better CU utilization on small M

Kernel source

submission.py98 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X

"""
Solution A: Fused quant+shuffle + Direct ASM GEMM dispatch

Optimizations stacked:
1. Module-level imports (no per-call import overhead)
2. Skip B.contiguous() entirely (B unused)
3. Skip A.contiguous() (already contiguous)
4. Call gemm_a4w4_asm directly with optimal kernel name per shape
   - Bypasses config CSV lookup overhead
   - Bypasses get_padded_m / get_GEMM_config
5. Pre-allocate output tensor (avoid allocation inside gemm_a4w4)
6. Use log2_k_split=3 (splitK=8) for better CU utilization on small M
7. Reuse padded output buffer across calls of same shape
"""
from task import input_t, output_t
import torch
import aiter
from aiter import dtypes
from aiter.ops.triton.quant import dynamic_mxfp4_quant
from aiter.utility.fp4_utils import e8m0_shuffle

# Try importing the direct ASM dispatch
try:
    from aiter.ops.gemm_op_a4w4 import gemm_a4w4_asm
    _HAS_ASM = True
except ImportError:
    _HAS_ASM = False

_bf16 = dtypes.bf16
_fp4x2 = dtypes.fp4x2
_fp8_e8m0 = dtypes.fp8_e8m0

# Pre-computed optimal kernel configs per (M, N, K)
# Kernel naming: f4gemm_bf16_per1x32Fp4_BpreShuffle_{TileM}x{TileN}
# For small M, 32x128 is typically optimal
# For larger M, bigger tiles reduce overhead
_KERNEL_MAP = {
    # Benchmark cases from task.yml
    (4, 2880, 512):   ("_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E", 0),
    (16, 2112, 7168):  ("_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E", 0),
    (32, 4096, 512):   ("_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E", 0),
    (32, 2880, 512):   ("_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E", 0),
    (64, 7168, 2048):  ("_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_64x128E", 0),
    (256, 3072, 1536): ("_ZN5aiter42f4gemm_bf16_per1x32Fp4_BpreShuffle_192x128E", 0),
    # Test cases
    (8, 2112, 7168):   ("_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E", 0),
    (16, 3072, 1536):  ("_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E", 0),
    (64, 3072, 1536):  ("_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_64x128E", 0),
    (256, 2880, 512):  ("_ZN5aiter42f4gemm_bf16_per1x32Fp4_BpreShuffle_192x128E", 0),
}

# Output buffer cache to avoid repeated allocation
_out_cache = {}


def _get_output(m, n, device):
    """Get or create a pre-allocated padded output buffer."""
    key = (m, n)
    if key not in _out_cache:
        padded_m = ((m + 31) // 32) * 32
        _out_cache[key] = torch.empty((padded_m, n), dtype=torch.bfloat16, device=device)
    return _out_cache[key]


def custom_kernel(data: input_t) -> output_t:
    A, _, _, B_shuffle, B_scale_sh = data
    m, k = A.shape
    n = B_shuffle.shape[0]

    # Quantize A: bf16 -> MXFP4 with shuffled E8M0 scales
    A_fp4, A_scale = dynamic_mxfp4_quant(A)
    A_q = A_fp4.view(_fp4x2)
    A_scale_sh = e8m0_shuffle(A_scale).view(_fp8_e8m0)

    if _HAS_ASM:
        # Direct ASM dispatch - bypass config lookup
        out = _get_output(m, n, A.device)
        kernel_name, log2_split = _KERNEL_MAP.get(
            (m, n, k),
            ("_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E", 0)
        )
        gemm_a4w4_asm(
            A_q, B_shuffle, A_scale_sh, B_scale_sh,
            out, kernel_name,
            bias=None, alpha=1.0, beta=0.0,
            bpreshuffle=True, log2_k_split=log2_split,
        )
        return out[:m]

    # Fallback: standard aiter path
    return aiter.gemm_a4w4(
        A_q, B_shuffle, A_scale_sh, B_scale_sh,
        dtype=_bf16, bpreshuffle=True,
    )
scrolls · 98 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