Skip to content
KernelIndex
Search⌘K

submission 588876

StephenCao422 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:acafa11c07e3ed46bb5685b47bacfe57718af38b5ea71b6aed428c86050382dc
license declaredunknown
license concludedunknown
authorsStephenCao422
imported2026-08-26

Techniques

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

num-warps = 4num_warps = 4
persistent-kernelnum_sk = tl.num_programs(axis=1) if SPLIT_K > 1 else 1
split-kSPLIT_K: tl.constexpr,
stages = 3num_stages=3
tile-k = 256BLOCK_SIZE_K = 256
tile-m = 16BLOCK_SIZE_M = 16
tile-n = 128BLOCK_SIZE_N = 128

Kernel source

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

import aiter
import torch
import triton
import triton.language as tl
from task import input_t, output_t

@triton.jit
def _fused_quant_gemm_kernel(
    a_ptr, b_ptr, c_ptr, b_scales_ptr,
    M, N, K,
    stride_am, stride_ak,
    stride_bn, stride_bk,
    stride_cm, stride_cn,
    SPLIT_K: tl.constexpr,
    BLOCK_SIZE_M: tl.constexpr,
    BLOCK_SIZE_N: tl.constexpr,
    BLOCK_SIZE_K: tl.constexpr,
    GROUP_SIZE_M: tl.constexpr,
):
    pid = tl.program_id(axis=0)
    pid_sk = tl.program_id(axis=1) if SPLIT_K > 1 else 0
    num_sk = tl.num_programs(axis=1) if SPLIT_K > 1 else 1

    num_pid_m = tl.cdiv(M, BLOCK_SIZE_M)
    num_pid_n = tl.cdiv(N, BLOCK_SIZE_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 % group_size_m)
    pid_n = (pid % num_pid_in_group) // group_size_m
    
    offs_am = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
    offs_k = tl.arange(0, BLOCK_SIZE_K)
    a_ptrs = a_ptr + (offs_am[:, None] * stride_am + offs_k[None, :] * stride_ak)
    
    offs_bn = (pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)) % N
    offs_k_b = tl.arange(0, BLOCK_SIZE_K // 2)
    b_ptrs = b_ptr + (offs_bn[:, None] * stride_bn + offs_k_b[None, :] * stride_bk)
    
    a_ptrs += pid_sk * BLOCK_SIZE_K * stride_ak

    SCALE_GROUP_SIZE: tl.constexpr = 32
    NUM_QUANT_BLOCKS: tl.constexpr = BLOCK_SIZE_K // SCALE_GROUP_SIZE
    
    # Mathematical un-shuffling mappings for B_scale_sh (6D coordinate inversion)
    n_idx_s = (pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)) % N
    i0 = n_idx_s // 32
    i1 = (n_idx_s % 32) // 16
    i2 = n_idx_s % 16

    accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
    
    # Precompute mask for A if M is not a multiple of BLOCK_SIZE_M
    mask_am = offs_am < M

    EXP_BIAS_FP4: tl.constexpr = 1
    EXP_BIAS_FP32: tl.constexpr = 127
    MBITS_F32: tl.constexpr = 23
    MBITS_FP4: tl.constexpr = 1
    EBITS_F32: tl.constexpr = 8
    EBITS_FP4: tl.constexpr = 2
    
    max_int: tl.constexpr = (1 << (EBITS_FP4 + MBITS_FP4)) - 1
    max_normal = 2 ** (3 - EXP_BIAS_FP4) * (3 / 2)
    min_normal = 2 ** (1 - EXP_BIAS_FP4)
    denorm_exp = (EXP_BIAS_FP32 - EXP_BIAS_FP4) + (MBITS_F32 - MBITS_FP4) + 1
    denorm_mask_int = denorm_exp << MBITS_F32
    denorm_mask_float: tl.constexpr = tl.cast(denorm_mask_int, tl.float32, bitcast=True)
    val_to_add = ((EXP_BIAS_FP4 - EXP_BIAS_FP32) << MBITS_F32) + (1 << 21) - 1

    for k in range(pid_sk, tl.cdiv(K, BLOCK_SIZE_K), num_sk):
        a = tl.load(a_ptrs, mask=mask_am[:, None] & (offs_k[None, :] < K - k * BLOCK_SIZE_K), other=0.0).to(tl.float32)
        
        # ---------------- A Dynamic Quantization to MXFP4 ---------------- #
        x = a.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS, SCALE_GROUP_SIZE)
        amax = tl.max(tl.abs(x), axis=-1, keep_dims=True)
        amax = amax.to(tl.int32, bitcast=True)
        amax = (amax + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000
        amax = amax.to(tl.float32, bitcast=True)
        scale_e8m0_unbiased = tl.log2(amax).floor() - 2
        scale_e8m0_unbiased = tl.clamp(scale_e8m0_unbiased, min=-127, max=127)
        a_scales_e8m0 = scale_e8m0_unbiased.to(tl.uint8) + 127
        quant_scale = tl.exp2(-scale_e8m0_unbiased)
        
        qx = x * quant_scale
        qx = qx.to(tl.uint32, bitcast=True)
        s = qx & 0x80000000
        qx = qx ^ s
        qx_fp32 = qx.to(tl.float32, bitcast=True)
        saturate_mask = qx_fp32 >= max_normal
        denormal_mask = (not saturate_mask) & (qx_fp32 < min_normal)
        normal_mask = not (saturate_mask | denormal_mask)
        
        denormal_x = qx_fp32 + denorm_mask_float
        denormal_x = denormal_x.to(tl.uint32, bitcast=True)
        denormal_x -= denorm_mask_int
        denormal_x = denormal_x.to(tl.uint8)
        
        normal_x = qx
        mant_odd = (normal_x >> (MBITS_F32 - MBITS_FP4)) & 1
        normal_x += val_to_add
        normal_x += mant_odd
        normal_x = normal_x >> (MBITS_F32 - MBITS_FP4)
        normal_x = normal_x.to(tl.uint8)
        
        e2m1_value = tl.full(qx.type.get_block_shapes(), 0x7, dtype=tl.uint8)
        e2m1_value = tl.where(normal_mask, normal_x, e2m1_value)
        e2m1_value = tl.where(denormal_mask, denormal_x, e2m1_value)
        
        sign_lp = s >> (MBITS_F32 + EBITS_F32 - MBITS_FP4 - EBITS_FP4)
        sign_lp = sign_lp.to(tl.uint8)
        e2m1_value = e2m1_value | sign_lp
        e2m1_value = tl.reshape(e2m1_value, [BLOCK_SIZE_M, NUM_QUANT_BLOCKS, SCALE_GROUP_SIZE // 2, 2])
        evens, odds = tl.split(e2m1_value)
        a_fp4 = evens | (odds << 4)
        a_fp4 = a_fp4.reshape(BLOCK_SIZE_M, BLOCK_SIZE_K // 2)
        a_scales_e8m0 = a_scales_e8m0.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS)
        # ----------------------------------------------------------------- #
        
        # ---------------- B Load Native ---------------- #
        k_b_offset = k * (BLOCK_SIZE_K // 2) * stride_bk
        b_T = tl.load(b_ptrs + k_b_offset, mask=offs_k_b[None, :] < (K - k * BLOCK_SIZE_K) // 2, other=0.0)
        b = b_T.trans(1, 0)
        
        # ---------------- B_scale Decode ---------------- #
        k_idx_base = k * (BLOCK_SIZE_K // 32)
        k_idx = k_idx_base + tl.arange(0, BLOCK_SIZE_K // 32)
        
        i3 = k_idx // 8
        i4 = (k_idx % 8) // 4
        i5 = k_idx % 4
        
        flat_idx = i0[:, None] * K + i3[None, :] * 256 + i5[None, :] * 64 + i2[:, None] * 4 + i4[None, :] * 2 + i1[:, None]
        
        b_scales_e8m0 = tl.load(b_scales_ptr + flat_idx, mask=k_idx[None, :] < K // 32, other=0)
        # ----------------------------------------------------------------- #
        
        accumulator = tl.dot_scaled(a_fp4, a_scales_e8m0, "e2m1", b, b_scales_e8m0, "e2m1", accumulator)
        
        a_ptrs += num_sk * BLOCK_SIZE_K * stride_ak
        
    c = accumulator.to(c_ptr.type.element_ty)
    offs_cm = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M).to(tl.int64)
    offs_cn = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N).to(tl.int64)
    c_ptrs = c_ptr + stride_cm * offs_cm[:, None] + stride_cn * offs_cn[None, :]
    c_mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < N)
    
    if SPLIT_K > 1:
        tl.atomic_add(c_ptrs, c, mask=c_mask)
    else:
        tl.store(c_ptrs, c, mask=c_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.shape
    
    # Dynamic Tuning for extremely skewed M shapes in MoE
    BLOCK_SIZE_N = 128
    BLOCK_SIZE_K = 256
    GROUP_SIZE_M = 8

    # Dynamic Tuning for extremely skewed M shapes in MoE
    if M <= 16:
        BLOCK_SIZE_M = 16
        num_warps = 4
    elif M <= 32:
        BLOCK_SIZE_M = 32
        num_warps = 4
    elif M <= 64:
        BLOCK_SIZE_M = 64
        num_warps = 8
    else:
        BLOCK_SIZE_M = 128
        num_warps = 8
        
    TOTAL_SPATIAL_BLOCKS = triton.cdiv(M, BLOCK_SIZE_M) * triton.cdiv(N, BLOCK_SIZE_N)
    
    if K <= 512:
        SPLIT_K = 1
    elif TOTAL_SPATIAL_BLOCKS < 120:
        desired_sk = 120 // TOTAL_SPATIAL_BLOCKS
        k_chunks = triton.cdiv(K, BLOCK_SIZE_K)
        SPLIT_K = max(1, min(desired_sk, k_chunks))
        SPLIT_K = min(SPLIT_K, 16)
    else:
        SPLIT_K = 1

    if SPLIT_K > 1:
        C_out = torch.zeros((M, N), device=A.device, dtype=torch.float32)
    else:
        C_out = torch.empty((M, N), device=A.device, dtype=torch.bfloat16)
        
    grid = lambda META: (triton.cdiv(M, META['BLOCK_SIZE_M']) * triton.cdiv(N, META['BLOCK_SIZE_N']), META['SPLIT_K'])

    _fused_quant_gemm_kernel[grid](
        A, B_q.view(torch.uint8), C_out, B_scale_sh.view(torch.uint8),
        M, N, K,
        A.stride(0), A.stride(1),
        B_q.stride(0), B_q.stride(1),
        C_out.stride(0), C_out.stride(1),
        SPLIT_K=SPLIT_K,
        BLOCK_SIZE_M=BLOCK_SIZE_M,
        BLOCK_SIZE_N=BLOCK_SIZE_N,
        BLOCK_SIZE_K=BLOCK_SIZE_K,
        GROUP_SIZE_M=GROUP_SIZE_M,
        num_warps=num_warps,
        num_stages=3
    )
    
    if SPLIT_K > 1:
        return C_out.to(torch.bfloat16)
    return C_out
scrolls · 220 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