Skip to content
KernelIndex
Search⌘K

submission 617263

mega-dmitriy · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_v469_pro2.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-617263?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
8.61µs
#76 of 1143
2026-03-23

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:ebafd308fc625fca71268e0d0d48911e5263a8d8d4d217aaeb2ac5244f448864
license declaredunknown
license concludedunknown
authorsmega-dmitriy
imported2026-08-15

Techniques

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

num-warps = 4num_warps=4, num_stages=2, matrix_instr_nonkdim=16,
split-kdef _fused_splitk_kernel(
stages = 2SCALE_GP=16, GROUP_SIZE_M=8, NUM_STAGES=2,
tile-k = 16BM, BN, BK = 16, 64, 256
tile-n = 128RED_BN = 128

Kernel source

submission_v469_pro2.py391 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
"""
V469_pro2: Micro-optimizations on v466 for ~1% geomean improvement.
Changes:
1. Eliminate scale_bc = scale_f32 + tl.zeros() broadcast — use reshape trick
2. Pre-compute N-dependent scale shuffle terms outside K loop
3. BF16 amax: compute amax in bf16, only convert per-group max to f32
4. .wt cache modifier on quant intermediate stores (K=1536)
"""
import torch
import triton
import triton.language as tl
from task import input_t, output_t
import aiter

_ASM_KERNEL = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"
_cache = {}


@triton.jit
def _quant_shuffle_kernel(
    a_ptr, fp4_out_ptr, scale_out_ptr,
    stride_am, stride_ak,
    K_HALF: tl.constexpr, scaleN_pad: tl.constexpr,
    M_VAL: tl.constexpr, BLOCK_K: tl.constexpr,
):
    m = tl.program_id(0)
    blk_k = tl.program_id(1)
    if m >= M_VAL:
        return
    GPB: tl.constexpr = BLOCK_K // 32
    k_start = blk_k * BLOCK_K
    offs_k = tl.arange(0, BLOCK_K)
    # BF16 amax: skip bulk f32 conversion
    a_bf16 = tl.load(a_ptr + m * stride_am + (k_start + offs_k) * stride_ak)
    a_grouped_bf16 = tl.reshape(a_bf16, [GPB, 32])
    amax_bf16 = tl.max(tl.abs(a_grouped_bf16), axis=1, keep_dims=True)
    amax_f32 = amax_bf16.to(tl.float32)  # only GPB values converted
    amax_u32 = amax_f32.to(tl.uint32, bitcast=True)
    amax_u32 = (amax_u32 + 0x200000) & 0xFF800000
    amax_exp = (amax_u32 >> 23).to(tl.int32)
    scale_exp = tl.minimum(tl.maximum(amax_exp - 2, 0), 254)
    bs_e8m0 = scale_exp.to(tl.uint8)
    scale_f32 = (scale_exp.to(tl.uint32) << 23).to(tl.float32, bitcast=True)

    # Still need f32 for HW FP4 convert inputs
    a_f32 = a_bf16.to(tl.float32)
    a_grouped = tl.reshape(a_f32, [GPB, 32])
    a_pairs = tl.reshape(a_grouped, [GPB, 16, 2])
    a_even, a_odd = tl.split(a_pairs)

    # Scale broadcast via reshape: [GPB,1] -> repeat 16x -> [GPB,16] -> flatten
    scale_1d = tl.reshape(scale_f32, [GPB])
    scale_rep = scale_1d[:, None] * tl.full([1, 16], 1.0, dtype=tl.float32)

    a_even_flat = tl.reshape(a_even, [GPB * 16])
    a_odd_flat = tl.reshape(a_odd, [GPB * 16])
    scale_flat = tl.reshape(scale_rep, [GPB * 16])

    packed_u32 = tl.inline_asm_elementwise(
        asm="v_cvt_scalef32_pk_fp4_f32 $0, $1, $2, $3",
        constraints="=&v,v,v,v",
        args=[a_even_flat, a_odd_flat, scale_flat],
        dtype=tl.uint32, is_pure=True, pack=1,
    )
    packed_u8 = (packed_u32 & 0xFF).to(tl.uint8)
    fp4_flat = tl.reshape(packed_u8, [GPB * 16])
    fp4_offs = tl.arange(0, GPB * 16)
    tl.store(fp4_out_ptr + m * K_HALF + k_start // 2 + fp4_offs, fp4_flat, cache_modifier=".wt")
    g_base = k_start // 32
    g_vals = g_base + tl.arange(0, GPB)
    sh_off = ((m % 32 // 16) + (g_vals % 8 // 4) * 2 + (m % 16) * 4
              + (g_vals % 4) * 64 + (g_vals // 8) * 256
              + (m // 32) * (32 * scaleN_pad))
    scale_bytes = tl.reshape(bs_e8m0, [GPB])
    tl.store(scale_out_ptr + sh_off, scale_bytes, cache_modifier=".wt")


@triton.jit
def _fused_kernel(
    a_ptr, b_ptr, c_ptr, b_scales_ptr,
    N,
    stride_am, stride_ak, stride_bn, stride_cm, stride_cn,
    scaleN_pad_B,
    M_VAL: tl.constexpr, BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr,
    BLOCK_K: tl.constexpr, K_HALF: tl.constexpr, SCALE_GP: tl.constexpr,
    GROUP_SIZE_M: tl.constexpr, NUM_STAGES: tl.constexpr,
    MASK_M: tl.constexpr, MASK_N: tl.constexpr,
):
    pid = tl.program_id(0)
    NUM_PID_M: tl.constexpr = (M_VAL + BLOCK_M - 1) // BLOCK_M
    num_pid_n = tl.cdiv(N, BLOCK_N)
    num_pid_in_group = GROUP_SIZE_M * num_pid_n
    group_id = pid // num_pid_in_group
    first_pid_m = group_id * GROUP_SIZE_M
    group_size_m = min(NUM_PID_M - first_pid_m, GROUP_SIZE_M)
    pid_m = first_pid_m + ((pid % num_pid_in_group) % group_size_m)
    pid_n = (pid % num_pid_in_group) // group_size_m
    offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
    offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
    offs_k = tl.arange(0, BLOCK_K)
    NQB: tl.constexpr = BLOCK_K // SCALE_GP
    a_bf16_offs_k = tl.arange(0, BLOCK_K * 2)
    a_ptrs = a_ptr + offs_m[:, None] * stride_am
    b_ptrs = b_ptr + offs_n[None, :] * stride_bn + offs_k[:, None]
    offs_qs = tl.arange(0, NQB)
    accumulator = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)

    # Pre-compute N-dependent scale shuffle terms (loop-invariant)
    n_val = offs_n[:, None]
    n_sh_base = (n_val % 32 // 16) + (n_val % 16) * 4 + (n_val // 32) * (32 * scaleN_pad_B)

    for k_block in tl.range(0, K_HALF, BLOCK_K, num_stages=NUM_STAGES):
        if MASK_M:
            a_bf16 = tl.load(a_ptrs + (k_block * 2 + a_bf16_offs_k[None, :]) * stride_ak, mask=(offs_m[:, None] < M_VAL), other=0.0)
        else:
            a_bf16 = tl.load(a_ptrs + (k_block * 2 + a_bf16_offs_k[None, :]) * stride_ak)

        # BF16 amax: only convert per-group max to f32
        a_grouped_bf16 = tl.reshape(a_bf16, [BLOCK_M * NQB, 32])
        amax_bf16 = tl.max(tl.abs(a_grouped_bf16), axis=1, keep_dims=True)
        amax_f32 = amax_bf16.to(tl.float32)
        amax_u32 = amax_f32.to(tl.uint32, bitcast=True)
        amax_u32 = (amax_u32 + 0x200000) & 0xFF800000
        amax_exp = (amax_u32 >> 23).to(tl.int32)
        scale_exp = tl.minimum(tl.maximum(amax_exp - 2, 0), 254)
        bs_e8m0 = scale_exp.to(tl.uint8)
        scale_f32 = (scale_exp.to(tl.uint32) << 23).to(tl.float32, bitcast=True)

        # Convert to f32 for HW FP4 (still needed for F32 variant input)
        a_f32 = a_bf16.to(tl.float32)
        a_grouped = tl.reshape(a_f32, [BLOCK_M * NQB, 32])
        a_pairs = tl.reshape(a_grouped, [BLOCK_M * NQB, 16, 2])
        a_even, a_odd = tl.split(a_pairs)

        # Scale broadcast: reshape [G,1] -> [G] -> [:, None] * ones -> [G,16]
        scale_1d = tl.reshape(scale_f32, [BLOCK_M * NQB])
        scale_rep = scale_1d[:, None] * tl.full([1, 16], 1.0, dtype=tl.float32)

        a_even_flat = tl.reshape(a_even, [BLOCK_M * NQB * 16])
        a_odd_flat = tl.reshape(a_odd, [BLOCK_M * NQB * 16])
        scale_flat = tl.reshape(scale_rep, [BLOCK_M * NQB * 16])

        packed_u32 = tl.inline_asm_elementwise(
            asm="v_cvt_scalef32_pk_fp4_f32 $0, $1, $2, $3",
            constraints="=&v,v,v,v",
            args=[a_even_flat, a_odd_flat, scale_flat],
            dtype=tl.uint32, is_pure=True, pack=1,
        )
        packed_u8 = (packed_u32 & 0xFF).to(tl.uint8)
        a_quant = tl.reshape(packed_u8, [BLOCK_M, BLOCK_K])
        a_scales = tl.reshape(bs_e8m0, [BLOCK_M, NQB])

        if MASK_N:
            b = tl.load(b_ptrs + k_block, mask=(offs_n[None, :] < N), other=0)
        else:
            b = tl.load(b_ptrs + k_block)

        # B scale load with pre-computed N terms
        g_base = k_block // SCALE_GP
        g_val = g_base + offs_qs[None, :]
        sh_off = n_sh_base + (g_val % 8 // 4) * 2 + (g_val % 4) * 64 + (g_val // 8) * 256
        if MASK_N:
            b_sc = tl.load(b_scales_ptr + sh_off, mask=(offs_n[:, None] < N), other=0)
        else:
            b_sc = tl.load(b_scales_ptr + sh_off)
        accumulator += tl.dot_scaled(a_quant, a_scales, "e2m1", b, b_sc, "e2m1")

    c = accumulator.to(tl.bfloat16)
    c_mask = (offs_m[:, None] < M_VAL) & (offs_n[None, :] < N)
    c_ptrs = c_ptr + offs_m[:, None] * stride_cm + offs_n[None, :] * stride_cn
    tl.store(c_ptrs, c, mask=c_mask, cache_modifier=".wt")


@triton.jit
def _fused_splitk_kernel(
    a_ptr, b_ptr, workspace_ptr, b_scales_ptr,
    N, stride_am, stride_ak, stride_bn, stride_wm, stride_wn,
    scaleN_pad_B,
    M_VAL: tl.constexpr, BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr,
    BLOCK_K: tl.constexpr, SCALE_GP: tl.constexpr,
    K_PER_SPLIT: tl.constexpr, NUM_KSPLIT: tl.constexpr,
    MASK_M: tl.constexpr, MASK_N: tl.constexpr,
):
    pid_mn = tl.program_id(0)
    pid_k = tl.program_id(1)
    num_pid_m = tl.cdiv(M_VAL, BLOCK_M)
    num_pid_n = tl.cdiv(N, BLOCK_N)
    pid_m = pid_mn // num_pid_n
    pid_n = pid_mn % num_pid_n
    offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
    offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
    offs_k = tl.arange(0, BLOCK_K)
    NQB: tl.constexpr = BLOCK_K // SCALE_GP
    a_bf16_offs_k = tl.arange(0, BLOCK_K * 2)
    a_ptrs = a_ptr + offs_m[:, None] * stride_am
    b_ptrs = b_ptr + offs_n[None, :] * stride_bn + offs_k[:, None]
    offs_qs = tl.arange(0, NQB)
    k_start = pid_k * K_PER_SPLIT
    k_end = k_start + K_PER_SPLIT
    accumulator = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)

    # Pre-compute N-dependent scale shuffle terms
    n_val = offs_n[:, None]
    n_sh_base = (n_val % 32 // 16) + (n_val % 16) * 4 + (n_val // 32) * (32 * scaleN_pad_B)

    for k_block in tl.range(k_start, k_end, BLOCK_K):
        if MASK_M:
            a_bf16 = tl.load(a_ptrs + (k_block * 2 + a_bf16_offs_k[None, :]) * stride_ak, mask=(offs_m[:, None] < M_VAL), other=0.0)
        else:
            a_bf16 = tl.load(a_ptrs + (k_block * 2 + a_bf16_offs_k[None, :]) * stride_ak)

        # BF16 amax
        a_grouped_bf16 = tl.reshape(a_bf16, [BLOCK_M * NQB, 32])
        amax_bf16 = tl.max(tl.abs(a_grouped_bf16), axis=1, keep_dims=True)
        amax_f32 = amax_bf16.to(tl.float32)
        amax_u32 = amax_f32.to(tl.uint32, bitcast=True)
        amax_u32 = (amax_u32 + 0x200000) & 0xFF800000
        amax_exp = (amax_u32 >> 23).to(tl.int32)
        scale_exp = tl.minimum(tl.maximum(amax_exp - 2, 0), 254)
        bs_e8m0 = scale_exp.to(tl.uint8)
        scale_f32 = (scale_exp.to(tl.uint32) << 23).to(tl.float32, bitcast=True)

        a_f32 = a_bf16.to(tl.float32)
        a_grouped = tl.reshape(a_f32, [BLOCK_M * NQB, 32])
        a_pairs = tl.reshape(a_grouped, [BLOCK_M * NQB, 16, 2])
        a_even, a_odd = tl.split(a_pairs)

        scale_1d = tl.reshape(scale_f32, [BLOCK_M * NQB])
        scale_rep = scale_1d[:, None] * tl.full([1, 16], 1.0, dtype=tl.float32)

        a_even_flat = tl.reshape(a_even, [BLOCK_M * NQB * 16])
        a_odd_flat = tl.reshape(a_odd, [BLOCK_M * NQB * 16])
        scale_flat = tl.reshape(scale_rep, [BLOCK_M * NQB * 16])

        packed_u32 = tl.inline_asm_elementwise(
            asm="v_cvt_scalef32_pk_fp4_f32 $0, $1, $2, $3",
            constraints="=&v,v,v,v",
            args=[a_even_flat, a_odd_flat, scale_flat],
            dtype=tl.uint32, is_pure=True, pack=1,
        )
        packed_u8 = (packed_u32 & 0xFF).to(tl.uint8)
        a_quant = tl.reshape(packed_u8, [BLOCK_M, BLOCK_K])
        a_scales = tl.reshape(bs_e8m0, [BLOCK_M, NQB])

        if MASK_N:
            b = tl.load(b_ptrs + k_block, mask=(offs_n[None, :] < N), other=0)
        else:
            b = tl.load(b_ptrs + k_block)

        g_base = k_block // SCALE_GP
        g_val = g_base + offs_qs[None, :]
        sh_off = n_sh_base + (g_val % 8 // 4) * 2 + (g_val % 4) * 64 + (g_val // 8) * 256
        if MASK_N:
            b_sc = tl.load(b_scales_ptr + sh_off, mask=(offs_n[:, None] < N), other=0)
        else:
            b_sc = tl.load(b_scales_ptr + sh_off)
        accumulator += tl.dot_scaled(a_quant, a_scales, "e2m1", b, b_sc, "e2m1")

    w_mask = (offs_m[:, None] < M_VAL) & (offs_n[None, :] < N)
    w_ptrs = workspace_ptr + pid_k * stride_wm * M_VAL + offs_m[:, None] * stride_wm + offs_n[None, :] * stride_wn
    tl.store(w_ptrs, accumulator, mask=w_mask)


@triton.jit
def _reduce_kernel(
    workspace_ptr, c_ptr,
    M_VAL: tl.constexpr, N,
    stride_wm, stride_wn, stride_cm, stride_cn,
    NUM_KSPLIT: tl.constexpr, BLOCK_N: tl.constexpr,
):
    m = tl.program_id(0)
    pid_n = tl.program_id(1)
    offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
    n_mask = offs_n < N
    acc = tl.zeros((BLOCK_N,), dtype=tl.float32)
    for k in range(NUM_KSPLIT):
        vals = tl.load(workspace_ptr + k * M_VAL * stride_wm + m * stride_wm + offs_n * stride_wn, mask=n_mask, other=0.0)
        acc += vals
    tl.store(c_ptr + m * stride_cm + offs_n * stride_cn, acc.to(tl.bfloat16), mask=n_mask)


def custom_kernel(data: input_t) -> output_t:
    A, B, B_q, B_shuffle, B_scale_sh = data
    A = A.contiguous()
    M, K = A.shape
    N = B_q.shape[0]
    K_half = K // 2

    key = (M, N, K)
    if key not in _cache:
        out = torch.empty((M, N), dtype=torch.bfloat16, device=A.device)
        if K >= 4096:
            NUM_KSPLIT = K // 512
            workspace = torch.empty((NUM_KSPLIT, M, N), dtype=torch.float32, device=A.device)
            _cache[key] = (out, workspace, NUM_KSPLIT)
        elif K > 512 and K != 2048:
            scale_n = K // 32
            scale_n_pad = triton.cdiv(scale_n, 8) * 8
            m_pad = triton.cdiv(M, 256) * 256
            fp4_buf = torch.empty((M, K_half), dtype=torch.uint8, device=A.device)
            scale_buf = torch.zeros((m_pad, scale_n_pad), dtype=torch.uint8, device=A.device)
            _cache[key] = (out, fp4_buf, scale_buf, scale_n_pad, m_pad)
        else:
            _cache[key] = (out,)
    cached = _cache[key]
    out = cached[0]

    if K <= 512:
        B_q_u8 = B_q.view(torch.uint8)
        B_scale_u8 = B_scale_sh.view(torch.uint8)
        scaleN_pad_B = triton.cdiv(K // 32, 8) * 8
        BM, BN, BK = 16, 64, 256
        MASK_M = (M % BM) != 0
        MASK_N = (N % BN) != 0
        grid = (triton.cdiv(M, BM) * triton.cdiv(N, BN),)
        _fused_kernel[grid](
            A, B_q_u8, out, B_scale_u8, N,
            A.stride(0), A.stride(1), B_q_u8.stride(0),
            out.stride(0), out.stride(1), scaleN_pad_B,
            M_VAL=M, BLOCK_M=BM, BLOCK_N=BN, BLOCK_K=BK, K_HALF=K_half,
            SCALE_GP=16, GROUP_SIZE_M=8, NUM_STAGES=2,
            MASK_M=MASK_M, MASK_N=MASK_N,
            num_warps=4, num_stages=2, matrix_instr_nonkdim=16,
        )
    elif K == 2048:
        B_q_u8 = B_q.view(torch.uint8)
        B_scale_u8 = B_scale_sh.view(torch.uint8)
        scaleN_pad_B = triton.cdiv(K // 32, 8) * 8
        BM, BN, BK = 16, 128, 256
        MASK_M = (M % BM) != 0
        MASK_N = (N % BN) != 0
        grid = (triton.cdiv(M, BM) * triton.cdiv(N, BN),)
        _fused_kernel[grid](
            A, B_q_u8, out, B_scale_u8, N,
            A.stride(0), A.stride(1), B_q_u8.stride(0),
            out.stride(0), out.stride(1), scaleN_pad_B,
            M_VAL=M, BLOCK_M=BM, BLOCK_N=BN, BLOCK_K=BK, K_HALF=K_half,
            SCALE_GP=16, GROUP_SIZE_M=8, NUM_STAGES=2,
            MASK_M=MASK_M, MASK_N=MASK_N,
            num_warps=8, num_stages=2, matrix_instr_nonkdim=16,
        )
    elif K >= 4096:
        _, workspace, NUM_KSPLIT = cached
        B_q_u8 = B_q.view(torch.uint8)
        B_scale_u8 = B_scale_sh.view(torch.uint8)
        scaleN_pad_B = triton.cdiv(K // 32, 8) * 8
        BM, BN, BK = 16, 128, 256
        MASK_M = (M % BM) != 0
        MASK_N = (N % BN) != 0
        K_PER_SPLIT = K_half // NUM_KSPLIT
        num_mn_tiles = triton.cdiv(M, BM) * triton.cdiv(N, BN)
        grid_fused = (num_mn_tiles, NUM_KSPLIT)
        _fused_splitk_kernel[grid_fused](
            A, B_q_u8, workspace, B_scale_u8,
            N, A.stride(0), A.stride(1), B_q_u8.stride(0),
            workspace.stride(1), workspace.stride(2),
            scaleN_pad_B,
            M_VAL=M, BLOCK_M=BM, BLOCK_N=BN, BLOCK_K=BK,
            SCALE_GP=16, K_PER_SPLIT=K_PER_SPLIT, NUM_KSPLIT=NUM_KSPLIT,
            MASK_M=MASK_M, MASK_N=MASK_N,
            num_warps=8, num_stages=2, matrix_instr_nonkdim=16,
        )
        RED_BN = 128
        grid_reduce = (M, triton.cdiv(N, RED_BN))
        _reduce_kernel[grid_reduce](
            workspace, out, M_VAL=M, N=N,
            stride_wm=workspace.stride(1), stride_wn=workspace.stride(2),
            stride_cm=out.stride(0), stride_cn=out.stride(1),
            NUM_KSPLIT=NUM_KSPLIT, BLOCK_N=RED_BN,
            num_warps=4,
        )
    else:
        _, fp4_buf, scale_buf, scale_n_pad, m_pad = cached
        BLOCK_K = 256
        num_k_blocks = K // BLOCK_K
        grid = (M, num_k_blocks)
        _quant_shuffle_kernel[grid](
            A, fp4_buf, scale_buf,
            A.stride(0), A.stride(1),
            K_HALF=K_half, scaleN_pad=scale_n_pad,
            M_VAL=M, BLOCK_K=BLOCK_K,
            num_warps=4, num_stages=1,
        )
        A_q = fp4_buf.view(torch.float4_e2m1fn_x2)
        A_scale = scale_buf.view(torch.float8_e8m0fnu)
        aiter.gemm_a4w4_asm(A_q, B_shuffle, A_scale, B_scale_sh,
                            out, _ASM_KERNEL, bpreshuffle=True)
    return out
scrolls · 391 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