Skip to content
KernelIndex
Search⌘K

submission 693087

LiangSu8899 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:53859815c26e0300b9cea659ffb40159329b5a7ae3cab7e9e27ec66e6ae492bd
license declaredunknown
license concludedunknown
authorsLiangSu8899
imported2026-08-26

Techniques

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

split-kcustom_kernel._get_splitk = gmod.get_splitk
stages = 1num_warps=NW_q, waves_per_eu=0, num_stages=1)

Kernel source

submission.py206 lines
"""GEMM v162: Manually tuned fused kernel configs for k<=512 shapes.
v158 diagnostic showed default _get_config returns:
  m=4: BSM=4, BSN=128, warps=4, grid=23
  m=32: BSM=8, BSN=128, warps=8, grid=128/92
  m=64: BSM=16, BSN=128, warps=8, stages=2, grid=224
  m=256: BSM=8, BSN=128, warps=8, grid=768

Try BSN=64 for more parallelism, fewer warps for less overhead.
Keep quant+ASM for k>512 (proven best).
"""
from task import input_t, output_t
import triton


def custom_kernel(data: input_t) -> output_t:
    import torch
    import aiter
    from aiter import dtypes
    from aiter.ops.triton.gemm.basic import gemm_a16wfp4 as gmod

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

    if not hasattr(custom_kernel, '_init'):
        custom_kernel._init = True
        custom_kernel._b = {}
        custom_kernel._fused_kernel = gmod._gemm_a16wfp4_preshuffle_kernel
        custom_kernel._get_splitk = gmod.get_splitk
        custom_kernel._get_config = gmod._get_config
        custom_kernel._fp4x2 = dtypes.fp4x2
        custom_kernel._e8m0 = dtypes.fp8_e8m0

        from aiter.ops.triton.quant.quant import _mxfp4_quant_op
        import triton.language as tl

        @triton.heuristics({
            "EVEN_M_N": lambda args: (
                args["M"] % args["BLOCK_SIZE_M"] == 0 and
                args["N"] % (args["BLOCK_SIZE_N"] * args["NUM_ITER"]) == 0
            ),
        })
        @triton.jit
        def _fq(
            x_ptr, x_fp4_ptr, bs_ptr,
            stride_x_m_in, stride_x_n_in,
            stride_x_fp4_m_in, stride_x_fp4_n_in,
            M, N, SN,
            BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr,
            NUM_ITER: tl.constexpr, NUM_STAGES: tl.constexpr,
            MXFP4_QUANT_BLOCK_SIZE: tl.constexpr,
            EVEN_M_N: tl.constexpr, SCALING_MODE: tl.constexpr,
        ):
            pid_m = tl.program_id(0)
            start_n = tl.program_id(1) * NUM_ITER
            stride_x_m = tl.cast(stride_x_m_in, tl.int64)
            stride_x_n = tl.cast(stride_x_n_in, tl.int64)
            stride_fp4_m = tl.cast(stride_x_fp4_m_in, tl.int64)
            stride_fp4_n = tl.cast(stride_x_fp4_n_in, tl.int64)
            NUM_QB: tl.constexpr = BLOCK_SIZE_N // MXFP4_QUANT_BLOCK_SIZE
            for pid_n in tl.range(start_n, tl.minimum(start_n + NUM_ITER, N),
                                  num_stages=NUM_STAGES):
                x_offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
                x_offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)
                x_offs = x_offs_m[:, None] * stride_x_m + x_offs_n[None, :] * stride_x_n
                if EVEN_M_N:
                    x = tl.load(x_ptr + x_offs, cache_modifier=".cg").to(tl.float32)
                else:
                    x_mask = (x_offs_m < M)[:, None] & (x_offs_n < N)[None, :]
                    x = tl.load(x_ptr + x_offs, mask=x_mask, cache_modifier=".cg").to(tl.float32)
                out_t, bs_e8 = _mxfp4_quant_op(x, BLOCK_SIZE_N, BLOCK_SIZE_M, MXFP4_QUANT_BLOCK_SIZE)
                o_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
                o_n = pid_n * BLOCK_SIZE_N // 2 + tl.arange(0, BLOCK_SIZE_N // 2)
                o_offs = o_m[:, None] * stride_fp4_m + o_n[None, :] * stride_fp4_n
                if EVEN_M_N:
                    tl.store(x_fp4_ptr + o_offs, out_t)
                else:
                    tl.store(x_fp4_ptr + o_offs, out_t, mask=(o_m < M)[:, None] & (o_n < N // 2)[None, :])
                bs_r = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
                bs_c = pid_n * NUM_QB + tl.arange(0, NUM_QB)
                n_sc = (N + MXFP4_QUANT_BLOCK_SIZE - 1) // MXFP4_QUANT_BLOCK_SIZE
                d0 = bs_r // 32; d5 = (bs_r % 32) // 16; d3 = bs_r % 16
                d1 = bs_c // 8; d4 = (bs_c % 8) // 4; d2 = bs_c % 4
                SN64 = tl.cast(SN, tl.int64)
                sh_offs = (d0[:, None].to(tl.int64) * (SN64 * 32)
                           + d1[None, :].to(tl.int64) * 256
                           + d2[None, :].to(tl.int64) * 64
                           + d3[:, None].to(tl.int64) * 4
                           + d4[None, :].to(tl.int64) * 2
                           + d5[:, None].to(tl.int64))
                if EVEN_M_N:
                    tl.store(bs_ptr + sh_offs, bs_e8)
                else:
                    tl.store(bs_ptr + sh_offs, bs_e8, mask=(bs_r < M)[:, None] & (bs_c < n_sc)[None, :])

        custom_kernel._fq = _fq
        custom_kernel._asm = aiter.gemm_a4w4_asm

    b = custom_kernel._b
    key = (m, n, k)
    if key not in b:
        if k <= 512:
            K_val = k // 2

            # Manually tuned configs per shape
            # Default from _get_config: BSM=4/8, BSN=128, BSK=512
            # Try BSN=64 for more tiles, lower warps for less overhead
            if m <= 4:
                BSM, BSN, BSK = 4, 64, 512
                GSM, NKS = 1, 1
                nw, ns, wpe, cm = 4, 1, 2, '.cg'
            elif m <= 16:
                BSM, BSN, BSK = 4, 128, 512
                GSM, NKS = 1, 1
                nw, ns, wpe, cm = 4, 1, 2, '.cg'
            elif m <= 32:
                BSM, BSN, BSK = 8, 64, 512
                GSM, NKS = 1, 1
                nw, ns, wpe, cm = 4, 1, 2, '.cg'
            else:
                BSM, BSN, BSK = 8, 128, 512
                GSM, NKS = 1, 1
                nw, ns, wpe, cm = 8, 1, 2, '.cg'

            SPLITK_BS, BSK, NKS = custom_kernel._get_splitk(K_val, BSK, NKS)
            grid_mn = triton.cdiv(m, BSM) * triton.cdiv(n, BSN)
            out = torch.empty(m, n, dtype=torch.bfloat16, device=A.device)
            EVEN_K = (K_val % (BSK // 2) == 0 and SPLITK_BS % BSK == 0 and K_val % (SPLITK_BS // 2) == 0)
            meta = {
                'BLOCK_SIZE_M': BSM, 'BLOCK_SIZE_N': BSN, 'BLOCK_SIZE_K': BSK,
                'GROUP_SIZE_M': GSM, 'NUM_KSPLIT': NKS, 'SPLITK_BLOCK_SIZE': SPLITK_BS,
                'num_warps': nw, 'num_stages': ns, 'waves_per_eu': wpe,
                'matrix_instr_nonkdim': 16, 'cache_modifier': cm,
                'GRID_MN': grid_mn, 'PREQUANT': True,
            }
            b[key] = ('fused', out, K_val, EVEN_K, meta, grid_mn, NKS)
        else:
            SG = 32
            n_sc = (k + SG - 1) // SG
            sm = (m + 255) // 256 * 256
            sn = (n_sc + 7) // 8 * 8
            BSM_q = 16
            BSN_q = 32
            NI_q, NW_q, NS_q = 1, 2, 2
            if m >= 64:
                BSM_q = 64
            grid_q = (triton.cdiv(m, BSM_q), triton.cdiv(k, BSN_q * NI_q))
            x_fp4 = torch.empty((m, k // 2), dtype=torch.uint8, device=A.device)
            scale_sh = torch.zeros(sm * sn, dtype=torch.uint8, device=A.device)
            A_q_view = x_fp4.view(custom_kernel._fp4x2)
            A_s_view = scale_sh.view(sm, sn).view(custom_kernel._e8m0)
            out = torch.empty(m, n, dtype=torch.bfloat16, device=A.device)
            tile_m, tile_n = 32, 128
            base_name = f"f4gemm_bf16_per1x32Fp4_BpreShuffle_{tile_m}x{tile_n}"
            kname = f"_ZN5aiter{len(base_name)}{base_name}E"
            b[key] = ('asm', x_fp4, scale_sh, grid_q,
                      BSM_q, BSN_q, NI_q, NS_q, NW_q,
                      sm, sn, A_q_view, A_s_view, out, kname)

    entry = b[key]

    if entry[0] == 'fused':
        _, out, K_val, EVEN_K, meta, grid_mn, NKS = entry
        kernel = custom_kernel._fused_kernel
        b_data = B_shuffle.view(torch.uint8)
        b_scale = B_scale_sh.view(torch.uint8)
        stride_bn = 16 * b_data.stride(0)
        stride_bk = b_data.stride(1)
        stride_bsn = 32 * b_scale.stride(0)
        stride_bsk = b_scale.stride(1)
        stride_ck, stride_cm, stride_cn = 0, out.stride(0), out.stride(1)
        grid = (grid_mn * NKS,)
        kernel[grid](
            A, b_data, out, b_scale,
            m, n, K_val,
            A.stride(0), A.stride(1),
            stride_bn, stride_bk,
            stride_ck, stride_cm, stride_cn,
            stride_bsn, stride_bsk,
            EVEN_K=EVEN_K, **meta,
        )
        if NKS > 1:
            return out.sum(dim=0).to(torch.bfloat16)
        return out

    else:
        (_, x_fp4, scale_sh, grid_q,
         BSM_q, BSN_q, NI_q, NS_q, NW_q,
         sm, sn, A_q_view, A_s_view, out, kname) = entry

        custom_kernel._fq[grid_q](
            A, x_fp4, scale_sh,
            A.stride(0), A.stride(1),
            x_fp4.stride(0), x_fp4.stride(1),
            M=m, N=k, SN=sn,
            BLOCK_SIZE_M=BSM_q, BLOCK_SIZE_N=BSN_q,
            NUM_ITER=NI_q, NUM_STAGES=NS_q,
            MXFP4_QUANT_BLOCK_SIZE=32, SCALING_MODE=0,
            num_warps=NW_q, waves_per_eu=0, num_stages=1)

        custom_kernel._asm(
            A_q_view, B_shuffle, A_s_view, B_scale_sh,
            out, kname, bpreshuffle=True)

        return out
scrolls · 206 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