Skip to content
KernelIndex
Search⌘K

submission 623450

Arkadip Maitra · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-623450?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
24.0µs
#838 of 1143
2026-03-24

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:3b246e659295bdfb4dd6a9a856683af55afc79d042e991f35dea55704a434af5
license declaredunknown
license concludedunknown
authorsArkadip Maitra
imported2026-08-26

Techniques

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

fp4Optimized MXFP4 GEMM: bf16 A, MXFP4 B -> fused quant+GEMM -> bf16 C.
split-kget_splitk,

Kernel source

submission.py166 lines
"""
Optimized MXFP4 GEMM: bf16 A, MXFP4 B -> fused quant+GEMM -> bf16 C.

Primary optimization: Use the Triton gemm_a16wfp4 preshuffle kernel that fuses
A quantization (bf16 -> mxfp4) into the GEMM kernel itself, eliminating:
  - dynamic_mxfp4_quant kernel launch for A
  - e8m0_shuffle kernel launch for A scales
  - HBM round-trip for A_q and A_scale intermediate buffers

Secondary optimization: Output tensor caching — reuse the same output buffer
for repeated calls with the same (m, n) shape, avoiding torch.empty overhead.

Fallback: Reference approach (separate quant + gemm_a4w4).
"""
from task import input_t, output_t

_initialized = False
_fused_fn = None
_output_cache = {}
_fallback_imports = None


def _try_init_fused():
    global _initialized, _fused_fn
    if _initialized:
        return
    _initialized = True

    # Try 1: aiter top-level API
    try:
        import aiter
        from aiter import dtypes
        for name in ('gemm_a16wfp4', 'gemm_a16w4'):
            fn = getattr(aiter, name, None)
            if fn is not None:
                def _wrap_toplevel(A, B_sh, B_sc, m, n, k, _fn=fn):
                    return _fn(A, B_sh, B_sc, dtype=dtypes.bf16, bpreshuffle=True)
                _fused_fn = _wrap_toplevel
                return
    except Exception:
        pass

    # Try 2: Triton gemm module wrapper
    for import_path in [
        'aiter.ops.triton.gemm.gemm_a16wfp4',
        'aiter.ops.triton.gemm',
    ]:
        try:
            import importlib
            mod = importlib.import_module(import_path)
            for attr in ('gemm_a16wfp4_preshuffle', 'gemm_a16wfp4'):
                fn = getattr(mod, attr, None)
                if fn is not None:
                    from aiter import dtypes
                    def _wrap_module(A, B_sh, B_sc, m, n, k, _fn=fn):
                        return _fn(A, B_sh, B_sc, dtype=dtypes.bf16, bpreshuffle=True)
                    _fused_fn = _wrap_module
                    return
        except (ImportError, AttributeError):
            pass

    # Try 3: Direct Triton kernel launch
    try:
        from aiter.ops.triton._triton_kernels.gemm.basic.gemm_a16wfp4 import (
            _gemm_a16wfp4_preshuffle_kernel,
            _get_config,
            get_splitk,
        )
        import triton
        import torch

        def _direct_kernel(A, B_shuffle, B_scale_sh, m, n, k):
            config = _get_config(m, n, k, shuffle=True)
            if config is not None:
                BSM = int(config.get('BLOCK_SIZE_M', 32))
                BSN = int(config.get('BLOCK_SIZE_N', 64))
                BSK = int(config.get('BLOCK_SIZE_K', 256))
                GROUP_SIZE_M = int(config.get('GROUP_SIZE_M', 4))
                NUM_KSPLIT = int(config.get('NUM_KSPLIT', 1))
                nw = int(config.get('num_warps', 4))
                ns = int(config.get('num_stages', 2))
                wpe = int(config.get('waves_per_eu', 2))
                mink = int(config.get('matrix_instr_nonkdim', 16))
                cm = str(config.get('cache_modifier', '.cg'))
            else:
                BSM = 16 if m <= 16 else 32
                BSN = 64
                BSK = 256
                GROUP_SIZE_M = 4
                NUM_KSPLIT = max(1, 4 if k >= 4096 else 2 if k >= 1024 else 1)
                nw, ns, wpe, mink = 4, 2, 2, 16
                cm = '.cg'

            K_half = k // 2
            SPLITK_BLOCK_SIZE, BSK, NUM_KSPLIT = get_splitk(K_half, BSK, NUM_KSPLIT)
            grid = (triton.cdiv(m, BSM) * triton.cdiv(n, BSN) * NUM_KSPLIT,)

            if NUM_KSPLIT > 1:
                c = torch.zeros((NUM_KSPLIT, m, n), dtype=torch.bfloat16, device='cuda')
                stride_ck = m * n
                stride_cm = n
            else:
                c = torch.empty((m, n), dtype=torch.bfloat16, device='cuda')
                stride_ck = 0
                stride_cm = n

            _gemm_a16wfp4_preshuffle_kernel[grid](
                A, B_shuffle, c, B_scale_sh,
                m, n, K_half,
                A.stride(0), A.stride(1),
                B_shuffle.stride(0), B_shuffle.stride(1),
                stride_ck, stride_cm, 1,
                B_scale_sh.stride(0), B_scale_sh.stride(1),
                BLOCK_SIZE_M=BSM, BLOCK_SIZE_N=BSN, BLOCK_SIZE_K=BSK,
                GROUP_SIZE_M=GROUP_SIZE_M,
                NUM_KSPLIT=NUM_KSPLIT,
                SPLITK_BLOCK_SIZE=SPLITK_BLOCK_SIZE,
                PREQUANT=True,
                num_warps=nw, num_stages=ns,
                waves_per_eu=wpe,
                matrix_instr_nonkdim=mink,
                cache_modifier=cm,
            )
            if NUM_KSPLIT > 1:
                return c.sum(dim=0)
            return c

        _fused_fn = _direct_kernel
        return
    except (ImportError, AttributeError):
        pass


def _get_fallback_imports():
    global _fallback_imports
    if _fallback_imports is None:
        import aiter
        from aiter import dtypes
        from aiter.ops.triton.quant import dynamic_mxfp4_quant
        from aiter.utility.fp4_utils import e8m0_shuffle
        _fallback_imports = (aiter, dtypes, dynamic_mxfp4_quant, e8m0_shuffle)
    return _fallback_imports


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]

    _try_init_fused()

    if _fused_fn is not None:
        try:
            return _fused_fn(A, B_shuffle, B_scale_sh, m, n, k)
        except Exception:
            pass

    aiter, dtypes, dynamic_mxfp4_quant, e8m0_shuffle = _get_fallback_imports()
    A_q, A_scale = dynamic_mxfp4_quant(A)
    A_q = A_q.view(dtypes.fp4x2)
    A_scale = e8m0_shuffle(A_scale).view(dtypes.fp8_e8m0)
    return aiter.gemm_a4w4(
        A_q, B_shuffle, A_scale, B_scale_sh,
        dtype=dtypes.bf16, bpreshuffle=True,
    )
scrolls · 166 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