Skip to content
KernelIndex
Search⌘K

submission 577240

Arseni Ivanov · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_7.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-577240?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
9.33µs
#161 of 1143
2026-03-17

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:3463a073ca833e7abe0ebd255c52a457572fb66e99d4a11dbbf9e5e6a6316729
license declaredunknown
license concludedunknown
authorsArseni Ivanov
imported2026-08-15

Techniques

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

fp4"""Optimized MXFP4 quantization with reduced operations"""
split-kSPLIT_K: tl.constexpr,

Kernel source

submission_7.py315 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
import torch
import triton
import triton.language as tl
from task import input_t, output_t

@triton.jit
def inline_quantize_mxfp4_optimized(x, BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_K: tl.constexpr):
    """Optimized MXFP4 quantization with reduced operations"""
    x_reshaped = tl.reshape(x, [BLOCK_SIZE_M, BLOCK_SIZE_K // 32, 32])

    amax = tl.max(tl.abs(x_reshaped), axis=2)
    amax_bits = amax.to(tl.uint32, bitcast=True)

    amax_rounded = ((amax_bits + 0x200000) & 0xFF800000)
    exp_biased = (amax_rounded >> 23).to(tl.int32)
    scale_unbiased = tl.minimum(tl.maximum(exp_biased - 129, -127), 127)

    quant_scale_bits = ((127 - scale_unbiased).to(tl.uint32) << 23)
    quant_scale = quant_scale_bits.to(tl.float32, bitcast=True)

    qx = x_reshaped * tl.reshape(quant_scale, [BLOCK_SIZE_M, BLOCK_SIZE_K // 32, 1])
    qx_u32 = qx.to(tl.uint32, bitcast=True)

    sign = qx_u32 & 0x80000000
    qx_abs = (qx_u32 ^ sign).to(tl.float32, bitcast=True)

    denorm_val = ((qx_abs + 4194304.0).to(tl.int32, bitcast=True) - 0x4A800000)
    norm_bits = qx_abs.to(tl.int32, bitcast=True)
    mant_lsb = (norm_bits >> 22) & 1
    norm_val = ((norm_bits + (-1054867457 + mant_lsb)) >> 22)

    e2m1 = tl.where(qx_abs < 1.0, denorm_val, norm_val)
    e2m1 = tl.where(qx_abs >= 6.0, 7, e2m1)

    e2m1_packed = ((sign >> 28).to(tl.int32) | e2m1).to(tl.uint8)

    e2m1_pairs = tl.reshape(e2m1_packed, [BLOCK_SIZE_M, BLOCK_SIZE_K // 2, 2])
    evens, odds = tl.split(e2m1_pairs)
    x_fp4 = evens | (odds << 4)

    bs_e8m0 = (scale_unbiased + 127).to(tl.uint8)

    return x_fp4, bs_e8m0

HARDCODED_CONFIGS = {
    (4, 2880, 512):    {'BLOCK_M': 16, 'BLOCK_N': 64,  'BLOCK_K': 512, 'GROUP_M': 8, 'LOOP_STAGES': 2, 'num_warps': 4, 'num_stages': 1},
    (16, 2112, 7168):  {'BLOCK_M': 16, 'BLOCK_N': 128, 'BLOCK_K': 256, 'GROUP_M': 8, 'LOOP_STAGES': 2, 'num_warps': 4, 'num_stages': 2},
    (32, 4096, 512):   {'BLOCK_M': 16, 'BLOCK_N': 32,  'BLOCK_K': 256, 'GROUP_M': 8, 'LOOP_STAGES': 2, 'num_warps': 4, 'num_stages': 2},
    (32, 2880, 512):   {'BLOCK_M': 16, 'BLOCK_N': 64,  'BLOCK_K': 512, 'GROUP_M': 8, 'LOOP_STAGES': 2, 'num_warps': 4, 'num_stages': 1},
    (64, 7168, 2048):  {'BLOCK_M': 16, 'BLOCK_N': 256, 'BLOCK_K': 256, 'GROUP_M': 8, 'LOOP_STAGES': 2, 'num_warps': 8, 'num_stages': 2},
    (256, 3072, 1536): {'BLOCK_M': 16, 'BLOCK_N': 256, 'BLOCK_K': 256, 'GROUP_M': 8, 'LOOP_STAGES': 2, 'num_warps': 8, 'num_stages': 2},
}

def get_kernel_config(m, n, k):
    if (m, n, k) in HARDCODED_CONFIGS:
        return HARDCODED_CONFIGS[(m, n, k)]
    if k >= 4096:
        return {'BLOCK_M': 16, 'BLOCK_N': 128, 'BLOCK_K': 256, 'GROUP_M': 8, 'LOOP_STAGES': 2, 'num_warps': 4, 'num_stages': 2}
    elif k <= 512:
        return {'BLOCK_M': 16, 'BLOCK_N': 64,  'BLOCK_K': 512, 'GROUP_M': 8, 'LOOP_STAGES': 2, 'num_warps': 4, 'num_stages': 1}
    else:
        return {'BLOCK_M': 16, 'BLOCK_N': 128, 'BLOCK_K': 256, 'GROUP_M': 8, 'LOOP_STAGES': 2, 'num_warps': 8, 'num_stages': 2}

HARDCODED_REDUCE_CONFIGS = {
    (16, 2112): {'BLOCK_M': 16, 'BLOCK_N': 64, 'num_warps': 4},
    (64, 7168): {'BLOCK_M': 32, 'BLOCK_N': 64, 'num_warps': 8},
}

def get_reduce_config(m, n):
    if (m, n) in HARDCODED_REDUCE_CONFIGS:
        return HARDCODED_REDUCE_CONFIGS[(m, n)]
    if m >= 64:
        return {'BLOCK_M': 32, 'BLOCK_N': 64, 'num_warps': 8}
    else:
        return {'BLOCK_M': 16, 'BLOCK_N': 64, 'num_warps': 4}

# Removed @triton.heuristics completely!
@triton.jit
def fused_mxfp4_dot_scaled_kernel(
    A_ptr, B_ptr, B_scale_ptr, C_ptr, Workspace_ptr,
    M, N, K,
    stride_am, stride_ak,
    stride_bn, stride_bk,
    stride_bsn, stride_bsk,
    stride_cm, stride_cn,
    SPLIT_K: tl.constexpr,
    BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr, GROUP_M: tl.constexpr,
    EVEN_M: tl.constexpr, EVEN_N: tl.constexpr, EVEN_K: tl.constexpr, LOOP_STAGES: tl.constexpr,
    USE_SHUFFLED_B: tl.constexpr, EVICT_B_FIRST: tl.constexpr, USE_CG_C: tl.constexpr
):
    pid = tl.program_id(axis=0)
    pid_k = tl.program_id(axis=1)

    num_pid_m = tl.cdiv(M, BLOCK_M)
    num_pid_n = tl.cdiv(N, BLOCK_N)

    if K == 512:
        pid_m = pid % num_pid_m
        pid_n = pid // num_pid_m
    else:
        num_pid_in_group = GROUP_M * num_pid_n
        group_id = pid // num_pid_in_group
        first_pid_m = group_id * GROUP_M
        group_size_m = min(num_pid_m - first_pid_m, GROUP_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)

    a_ptrs = A_ptr + (offs_m[:, None] * stride_am + offs_k[None, :] * stride_ak)

    if USE_SHUFFLED_B:
        offs_bn_shuf = pid_n * (BLOCK_N // 16) + tl.arange(0, BLOCK_N // 16)
        offs_k_shuf = tl.arange(0, (BLOCK_K // 2) * 16)
        b_ptrs = B_ptr + (offs_bn_shuf[:, None] * stride_bn + offs_k_shuf[None, :] * stride_bk)
    else:
        offs_bk_q = tl.arange(0, BLOCK_K // 2)
        b_ptrs = B_ptr + (offs_bk_q[:, None] * stride_bk + offs_n[None, :] * stride_bn)

    offs_bsn = pid_n * (BLOCK_N // 32) + tl.arange(0, BLOCK_N // 32)
    offs_bks = tl.arange(0, (BLOCK_K // 32) * 32)
    b_scale_ptrs = B_scale_ptr + (offs_bsn[:, None] * stride_bsn + offs_bks[None, :] * stride_bsk)

    mask_m = offs_m < M
    mask_n = offs_n < N

    acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)

    a_ptrs += pid_k * BLOCK_K * stride_ak
    b_ptrs += pid_k * ((BLOCK_K // 2) * (16 if USE_SHUFFLED_B else 1)) * stride_bk
    b_scale_ptrs += pid_k * BLOCK_K * stride_bsk

    total_k_blocks = tl.cdiv(K, BLOCK_K)
    for k_idx in tl.range(pid_k, total_k_blocks, SPLIT_K, LOOP_STAGES):
        if EVEN_M and EVEN_K:
            a = tl.load(a_ptrs, eviction_policy="evict_last")
        elif EVEN_K:
            a = tl.load(a_ptrs, mask=mask_m[:, None], other=0.0, eviction_policy="evict_last")
        else:
            k_mask = (k_idx * BLOCK_K + offs_k) < K
            a = tl.load(a_ptrs, mask=(mask_m[:, None] & k_mask[None, :]), other=0.0, eviction_policy="evict_last")

        a_q, a_scale = inline_quantize_mxfp4_optimized(a.to(tl.float32), BLOCK_M, BLOCK_K)

        if USE_SHUFFLED_B:
            mask_bn = (pid_n * (BLOCK_N // 16) + tl.arange(0, BLOCK_N // 16)) < (N // 16)
            if EVICT_B_FIRST:
                if EVEN_N and EVEN_K:
                    b_raw = tl.load(b_ptrs, eviction_policy="evict_first")
                elif EVEN_K:
                    b_raw = tl.load(b_ptrs, mask=mask_bn[:, None], other=0, eviction_policy="evict_first")
                else:
                    b_raw = tl.load(b_ptrs, mask=mask_bn[:, None], other=0, eviction_policy="evict_first")
            else:
                if EVEN_N and EVEN_K:
                    b_raw = tl.load(b_ptrs)
                elif EVEN_K:
                    b_raw = tl.load(b_ptrs, mask=mask_bn[:, None], other=0)
                else:
                    b_raw = tl.load(b_ptrs, mask=mask_bn[:, None], other=0)
            b_q = (b_raw.reshape(1, BLOCK_N // 16, BLOCK_K // 64, 2, 16, 16)
                   .permute(0, 1, 4, 2, 3, 5).reshape(BLOCK_N, BLOCK_K // 2).trans(1, 0))
        else:
            if EVICT_B_FIRST:
                if EVEN_N and EVEN_K:
                    b_q = tl.load(b_ptrs, eviction_policy="evict_first")
                elif EVEN_K:
                    b_q = tl.load(b_ptrs, mask=mask_n[None, :], other=0, eviction_policy="evict_first")
                else:
                    bk_offs = (k_idx * BLOCK_K) // 2 + tl.arange(0, BLOCK_K // 2)
                    b_q = tl.load(b_ptrs, mask=(mask_n[None, :] & (bk_offs[:, None] < K // 2)), other=0, eviction_policy="evict_first")
            else:
                if EVEN_N and EVEN_K:
                    b_q = tl.load(b_ptrs)
                elif EVEN_K:
                    b_q = tl.load(b_ptrs, mask=mask_n[None, :], other=0)
                else:
                    bk_offs = (k_idx * BLOCK_K) // 2 + tl.arange(0, BLOCK_K // 2)
                    b_q = tl.load(b_ptrs, mask=(mask_n[None, :] & (bk_offs[:, None] < K // 2)), other=0)

        mask_bsn = (offs_bsn < (N // 32))
        if EVICT_B_FIRST:
            if EVEN_N and EVEN_K:
                b_scale_raw = tl.load(b_scale_ptrs, eviction_policy="evict_first")
            else:
                b_scale_raw = tl.load(b_scale_ptrs, mask=mask_bsn[:, None], other=0, eviction_policy="evict_first")
        else:
            if EVEN_N and EVEN_K:
                b_scale_raw = tl.load(b_scale_ptrs)
            else:
                b_scale_raw = tl.load(b_scale_ptrs, mask=mask_bsn[:, None], other=0)

        b_scale = (b_scale_raw.reshape(BLOCK_N // 32, BLOCK_K // 256, 4, 16, 2, 2, 1)
                   .permute(0, 5, 3, 1, 4, 2, 6).reshape(BLOCK_N, BLOCK_K // 32))

        acc = tl.dot_scaled(a_q, a_scale, 'e2m1', b_q, b_scale, 'e2m1', acc)

        a_ptrs += SPLIT_K * BLOCK_K * stride_ak
        b_ptrs += SPLIT_K * ((BLOCK_K // 2) * (16 if USE_SHUFFLED_B else 1)) * stride_bk
        b_scale_ptrs += SPLIT_K * BLOCK_K * stride_bsk

    if SPLIT_K == 1:
        c_ptrs = C_ptr + (offs_m[:, None] * stride_cm + offs_n[None, :] * stride_cn)
        c_mask = (offs_m[:, None] < M) & (offs_n[None, :] < N)
        tl.store(c_ptrs, acc.to(tl.bfloat16), mask=c_mask, cache_modifier=".cg" if USE_CG_C else "")
    else:
        ws_ptrs = Workspace_ptr + pid_k * (M * N) + offs_m[:, None] * N + offs_n[None, :]
        ws_mask = (offs_m[:, None] < M) & (offs_n[None, :] < N)
        tl.store(ws_ptrs, acc, mask=ws_mask, cache_modifier=".cg" if USE_CG_C else "")

# Removed @triton.heuristics here too
@triton.jit
def reduce_kernel(
    Workspace_ptr, C_ptr, M, N, stride_cm, stride_cn,
    SPLIT_K: tl.constexpr, BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, USE_CG_C: tl.constexpr
):
    pid_m, pid_n = tl.program_id(0), tl.program_id(1)
    offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
    offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
    mask = (offs_m[:, None] < M) & (offs_n[None, :] < N)

    acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
    for k in range(SPLIT_K):
        ptr = Workspace_ptr + k * M * N + offs_m[:, None] * N + offs_n[None, :]
        acc += tl.load(ptr, mask=mask, other=0.0)

    c_ptrs = C_ptr + offs_m[:, None] * stride_cm + offs_n[None, :] * stride_cn
    tl.store(c_ptrs, acc.to(tl.bfloat16), mask=mask, cache_modifier=".cg" if USE_CG_C else "")

#used to avoid re-initializing memory, not used to store or return results
_workspace_cache = {}
_c_cache = {}

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

    use_shuffled_b = (m == 64 and n == 7168 and k == 2048) or m == 256
    
    if use_shuffled_b:
        B_in_u8 = B_shuffle.view(torch.uint8)
        stride_bn = (k // 2) * 16
        stride_bk = 1
    else:
        B_in_u8 = B_q.view(torch.uint8)
        stride_bn = B_in_u8.stride(0)
        stride_bk = B_in_u8.stride(1)

    B_scale_sh_u8 = B_scale_sh.view(torch.uint8)
    stride_bsn = k
    stride_bsk = 1

    split_k = 16 if k >= 7168 else (2 if k == 2048 else 1)

    workspace = None
    if split_k > 1:
        ws_key = (A.device, m, n, split_k)
        if ws_key not in _workspace_cache:
            _workspace_cache[ws_key] = torch.empty((split_k, m, n), dtype=torch.float32, device=A.device)
        workspace = _workspace_cache[ws_key]

    c_key = (A.device, m, n)
    if c_key not in _c_cache:
        _c_cache[c_key] = torch.empty((m, n), device=A.device, dtype=torch.bfloat16)
    C = _c_cache[c_key]

    cfg = get_kernel_config(m, n, k)
    
    grid_m = (m + cfg['BLOCK_M'] - 1) // cfg['BLOCK_M']
    grid_n = (n + cfg['BLOCK_N'] - 1) // cfg['BLOCK_N']
    grid = (grid_m * grid_n, split_k)

    EVEN_M = (m % cfg['BLOCK_M'] == 0)
    EVEN_N = (n % cfg['BLOCK_N'] == 0)
    EVEN_K = (k % cfg['BLOCK_K'] == 0)
    EVICT_B_FIRST = (k >= 1536)
    USE_CG_C = ((m * n) > 128 * 128)

    fused_mxfp4_dot_scaled_kernel[grid](
        A, B_in_u8, B_scale_sh_u8, C, workspace, m, n, k,
        A.stride(0), A.stride(1), 
        stride_bn, stride_bk,
        stride_bsn, stride_bsk, 
        C.stride(0), C.stride(1),
        SPLIT_K=split_k,
        BLOCK_M=cfg['BLOCK_M'], BLOCK_N=cfg['BLOCK_N'], BLOCK_K=cfg['BLOCK_K'],
        GROUP_M=cfg['GROUP_M'], LOOP_STAGES=cfg['LOOP_STAGES'],
        EVEN_M=EVEN_M, EVEN_N=EVEN_N, EVEN_K=EVEN_K, 
        USE_SHUFFLED_B=use_shuffled_b, EVICT_B_FIRST=EVICT_B_FIRST, USE_CG_C=USE_CG_C,
        num_warps=cfg['num_warps'], num_stages=cfg['num_stages']
    )

    if split_k > 1:
        r_cfg = get_reduce_config(m, n)
        
        r_grid_m = (m + r_cfg['BLOCK_M'] - 1) // r_cfg['BLOCK_M']
        r_grid_n = (n + r_cfg['BLOCK_N'] - 1) // r_cfg['BLOCK_N']
        reduce_grid = (r_grid_m, r_grid_n)
        
        reduce_kernel[reduce_grid](
            workspace, C, m, n, 
            C.stride(0), C.stride(1), 
            SPLIT_K=split_k,
            BLOCK_M=r_cfg['BLOCK_M'], BLOCK_N=r_cfg['BLOCK_N'],
            USE_CG_C=USE_CG_C,
            num_warps=r_cfg['num_warps']
        )

    return C
scrolls · 315 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