Skip to content
KernelIndex
Search⌘K

submission 743751

Andrewxu313 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

v140h_shuffled_gemm.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-743751?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
14.0µs
#484 of 1143
2026-04-06

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:a3a6c0cfdf0877ac4999e502ae6ed0782cde343d7e5fc1e7489ac3f300654835
license declaredunknown
license concludedunknown
authorsAndrewxu313
imported2026-08-26

Techniques

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

tile-m = 16BLOCK_M=16, BLOCK_N=32,
tile-n = 32BLOCK_M=16, BLOCK_N=32,

Kernel source

v140h_shuffled_gemm.py358 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
"""v140h: Custom Triton GEMM reading shuffled B data directly.

Uses exact 6D permutation formulas from aiter source code:
- shuffle_weight: view(1, N//16, 16, Kh//32, 2, 16).permute(0,1,3,4,2,5)
- e8m0_shuffle: view(sm//32, 2, 16, sn//8, 2, 4).permute(0,3,5,2,4,1)
  with padding: sm = ceil(N/256)*256, sn = ceil(Ks/8)*8

NO dynamic_mxfp4_quant(B). NO e8m0_shuffle call. NO ASM GEMM.
"""

import torch
import triton
import triton.language as tl


# ─── FP4 E2M1 conversion ───

@triton.jit
def _to_e2m1(scaled_val):
    bits = scaled_val.to(tl.uint32, bitcast=True)
    sign = (bits >> 28) & 8
    pb = bits & 0x7FFFFFFF
    pf = pb.to(tl.float32, bitcast=True)
    dm = tl.full(pb.shape, 0x4A800000, dtype=tl.uint32)
    tmp_sub = pf + dm.to(tl.float32, bitcast=True)
    tb = tmp_sub.to(tl.uint32, bitcast=True)
    r_sub = (tb - dm) & 0xF
    mo = (pb >> 22) & 1
    r_norm = ((pb + 0xC11FFFFF + mo) >> 22) & 7
    r = tl.where(pf < 1.0, r_sub, r_norm)
    r = tl.where(pf >= 6.0, tl.full(r.shape, 7, dtype=tl.uint32), r)
    return (r | sign).to(tl.uint8)


# ─── A quant kernel (raw output) ───

@triton.jit
def _quant_kernel(A_ptr, Aq_ptr, As_ptr, M, K, Kh, Ks, BLOCK: tl.constexpr):
    pid = tl.program_id(0)
    row = pid // Ks
    bl = pid - row * Ks
    if row >= M:
        return
    k_base = bl * 32
    pair_off = tl.arange(0, BLOCK)
    base = A_ptr + row * K + k_base
    e0 = tl.load(base + pair_off * 2, mask=(k_base + pair_off * 2) < K, other=0.0).to(tl.float32)
    e1 = tl.load(base + pair_off * 2 + 1, mask=(k_base + pair_off * 2 + 1) < K, other=0.0).to(tl.float32)
    amax = tl.maximum(tl.max(tl.abs(e0), axis=0), tl.max(tl.abs(e1), axis=0))
    tmp = tl.where(pair_off == 0, amax, 0.0)
    amax_bits = tl.sum(tmp.to(tl.uint32, bitcast=True), axis=0)
    rounded = (amax_bits + 0x200000) & 0xFF800000
    biased_exp = (rounded >> 23) & 0xFF
    su = tl.where(biased_exp > 0, biased_exp.to(tl.int32) - 129, tl.full([], -127, dtype=tl.int32))
    su = tl.maximum(su, -127)
    su = tl.minimum(su, 127)
    e8 = (su + 127).to(tl.uint8)
    qs = tl.math.exp2((tl.maximum(tl.minimum(127 - su, 254), 1) - 127).to(tl.float32))
    packed = (_to_e2m1(e0 * qs).to(tl.uint8) & 0xF) | ((_to_e2m1(e1 * qs).to(tl.uint8) & 0xF) << 4)
    tl.store(Aq_ptr + row * Kh + bl * 16 + pair_off, packed, mask=pair_off < 16)
    tl.store(As_ptr + row * Ks + bl, e8)


# ─── Shuffled index helpers ───
# These compute flat byte offset into shuffled buffer for raw position (n, k).

@triton.jit
def _shuf_data_idx(n, kp, Kh_32):
    """shuffle_weight index: raw[n, kp] -> flat pos in shuffled buffer.
    Forward: view(1, N//16, 16, Kh//32, 2, 16).permute(0,1,3,4,2,5)
    Shuffled 6D = (0, n//16, kp//32, (kp//16)%2, n%16, kp%16)
    Flat = n_blk16 * (Kh_32 * 2 * 16 * 16) + kp_blk32 * (2*16*16) + kp_half * (16*16) + n_inner * 16 + kp_inner
    """
    n_blk16 = n // 16
    n_inner = n % 16
    kp_blk32 = kp // 32
    kp_half = (kp // 16) % 2
    kp_inner = kp % 16
    return n_blk16 * (Kh_32 * 512) + kp_blk32 * 512 + kp_half * 256 + n_inner * 16 + kp_inner


@triton.jit
def _shuf_scale_idx(n, ks, sn, sm_32):
    """e8m0_shuffle index: raw[n, ks] -> flat pos in shuffled buffer.
    Padding: sm = ceil(N/256)*256, sn_pad = ceil(Ks/8)*8.
    Forward: view(sm//32, 2, 16, sn//8, 2, 4).permute(0,3,5,2,4,1)
    6D raw decomposition of padded (n, ks):
      d0 = n // 32,  d1 = (n // 16) % 2,  d2 = n % 16
      d3 = ks // 8,  d4 = (ks // 4) % 2,  d5 = ks % 4
    Shuffled 6D = (d0, d3, d5, d2, d4, d1)
    Flat = d0*(sn_8 * 4*16*2*2) + d3*(4*16*2*2) + d5*(16*2*2) + d2*(2*2) + d4*2 + d1
    where sn_8 = sn // 8
    """
    d0 = n // 32
    d1 = (n // 16) % 2
    d2 = n % 16
    sn_8 = sn // 8
    d3 = ks // 8
    d4 = (ks // 4) % 2
    d5 = ks % 4
    return d0 * (sn_8 * 256) + d3 * 256 + d5 * 64 + d2 * 4 + d4 * 2 + d1


# ─── GEMM with shuffled B (no-mask, pipelined) ───

@triton.jit
def _gemm_shuf_nomask(
    A_ptr, B_shuf_ptr, C_ptr, A_scale_ptr, B_scale_shuf_ptr,
    stride_am, stride_ak, stride_as_m, stride_as_k,
    stride_cm, stride_cn,
    Kh_32, Ks_sn, Ks_sm_32,
    BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr,
    GROUP_SIZE_M: tl.constexpr, NUM_STAGES: tl.constexpr,
    M_CONST: tl.constexpr, N_CONST: tl.constexpr, K_CONST: tl.constexpr,
):
    pid = tl.program_id(0)
    NUM_PID_M: tl.constexpr = M_CONST // BLOCK_M
    NUM_PID_N: tl.constexpr = N_CONST // BLOCK_N
    num_pid_in_group: tl.constexpr = 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)
    acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)

    BKP: tl.constexpr = BLOCK_K // 2
    BKS: tl.constexpr = BLOCK_K // 32
    offs_kp = tl.arange(0, BKP)
    offs_ks = tl.arange(0, BKS)
    NUM_K: tl.constexpr = K_CONST // BLOCK_K

    for k_iter in tl.range(0, NUM_K, num_stages=NUM_STAGES):
        k_start = k_iter * BLOCK_K
        kp = k_start // 2
        ks = k_start // 32

        # A loads (raw layout)
        a = tl.load(A_ptr + offs_m[:, None] * stride_am + (kp + offs_kp[None, :]) * stride_ak)
        a_scale = tl.load(A_scale_ptr + offs_m[:, None] * stride_as_m + (ks + offs_ks[None, :]) * stride_as_k)

        # B data from shuffled layout: b[BKP, BN]
        kp_abs = kp + offs_kp  # [BKP]
        b_idx = _shuf_data_idx(offs_n[None, :], kp_abs[:, None], Kh_32)
        b = tl.load(B_shuf_ptr + b_idx)

        # B scale from shuffled layout: b_scale[BN, BKS]
        ks_abs = ks + offs_ks  # [BKS]
        bs_idx = _shuf_scale_idx(offs_n[:, None], ks_abs[None, :], Ks_sn, Ks_sm_32)
        b_scale = tl.load(B_scale_shuf_ptr + bs_idx)

        acc = tl.dot_scaled(a, a_scale, "e2m1", b, b_scale, "e2m1", acc=acc)

    c_ptrs = C_ptr + offs_m[:, None] * stride_cm + offs_n[None, :] * stride_cn
    tl.store(c_ptrs, acc.to(tl.bfloat16))


# ─── GEMM with shuffled B (masked) ───

@triton.jit
def _gemm_shuf_masked(
    A_ptr, B_shuf_ptr, C_ptr, A_scale_ptr, B_scale_shuf_ptr,
    M, N, K,
    stride_am, stride_ak, stride_as_m, stride_as_k,
    stride_cm, stride_cn,
    Kh_32, Ks_sn, Ks_sm_32,
    BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr,
    GROUP_SIZE_M: tl.constexpr,
):
    pid = tl.program_id(0)
    num_pid_m = tl.cdiv(M, 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)
    acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)

    BKP: tl.constexpr = BLOCK_K // 2
    BKS: tl.constexpr = BLOCK_K // 32
    offs_kp = tl.arange(0, BKP)
    offs_ks = tl.arange(0, BKS)

    for k_start in range(0, K, BLOCK_K):
        kp = k_start // 2
        ks = k_start // 32

        a = tl.load(A_ptr + offs_m[:, None] * stride_am + (kp + offs_kp[None, :]) * stride_ak,
                    mask=(offs_m[:, None] < M) & ((kp + offs_kp[None, :]) < K // 2), other=0)
        a_scale = tl.load(A_scale_ptr + offs_m[:, None] * stride_as_m + (ks + offs_ks[None, :]) * stride_as_k,
                          mask=(offs_m[:, None] < M) & ((ks + offs_ks[None, :]) < K // 32), other=127)

        kp_abs = kp + offs_kp
        b_idx = _shuf_data_idx(offs_n[None, :], kp_abs[:, None], Kh_32)
        b = tl.load(B_shuf_ptr + b_idx, mask=((kp_abs[:, None]) < K // 2) & (offs_n[None, :] < N), other=0)

        ks_abs = ks + offs_ks
        bs_idx = _shuf_scale_idx(offs_n[:, None], ks_abs[None, :], Ks_sn, Ks_sm_32)
        b_scale = tl.load(B_scale_shuf_ptr + bs_idx, mask=(offs_n[:, None] < N) & ((ks_abs[None, :]) < K // 32), other=127)

        acc = tl.dot_scaled(a, a_scale, "e2m1", b, b_scale, "e2m1", acc=acc)

    c_ptrs = C_ptr + offs_m[:, None] * stride_cm + offs_n[None, :] * stride_cn
    tl.store(c_ptrs, acc.to(tl.bfloat16), mask=(offs_m[:, None] < M) & (offs_n[None, :] < N))


# ─── Fused K=512 with shuffled B ───

@triton.jit
def _fused_k512_shuf(
    A_bf16_ptr, B_shuf_ptr, C_ptr, B_scale_shuf_ptr,
    M, N, Kh_32, Ks_sn, Ks_sm_32,
    stride_am, stride_ak, stride_cm, stride_cn,
    BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr,
):
    pid_m = tl.program_id(0)
    pid_n = 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)
    BK_PACKED: tl.constexpr = 256
    BK_SCALES: tl.constexpr = 16
    PAIRS_PER_GROUP: tl.constexpr = 16
    offs_kp = tl.arange(0, BK_PACKED)

    mask_m = offs_m[:, None] < M
    a_e_ptrs = A_bf16_ptr + offs_m[:, None] * stride_am + (offs_kp[None, :] * 2) * stride_ak
    a_o_ptrs = A_bf16_ptr + offs_m[:, None] * stride_am + (offs_kp[None, :] * 2 + 1) * stride_ak
    a_even = tl.load(a_e_ptrs, mask=mask_m, other=0.0).to(tl.float32)
    a_odd = tl.load(a_o_ptrs, mask=mask_m, other=0.0).to(tl.float32)
    abs_all = tl.maximum(tl.abs(a_even), tl.abs(a_odd))
    grouped = tl.reshape(abs_all, (BLOCK_M * BK_SCALES, PAIRS_PER_GROUP))
    amax = tl.max(grouped, axis=1)
    amax_bits = amax.to(tl.uint32, bitcast=True)
    rounded = (amax_bits + 0x200000) & 0xFF800000
    biased_exp = (rounded >> 23) & 0xFF
    su = tl.where(biased_exp > 0, biased_exp.to(tl.int32) - 129, -127)
    su = tl.maximum(su, -127)
    su = tl.minimum(su, 127)
    e8 = (su + 127).to(tl.uint8)
    a_scale = tl.reshape(e8, (BLOCK_M, BK_SCALES))
    qs_exp = tl.maximum(tl.minimum(127 - su, 254), 1)
    qs = tl.math.exp2((qs_exp - 127).to(tl.float32))
    qs_2d = tl.reshape(qs, (BLOCK_M * BK_SCALES, 1))
    qs_bc = tl.broadcast_to(qs_2d, (BLOCK_M * BK_SCALES, PAIRS_PER_GROUP))
    qs_flat = tl.reshape(qs_bc, (BLOCK_M, BK_PACKED))
    fp4_e = _to_e2m1(a_even * qs_flat)
    fp4_o = _to_e2m1(a_odd * qs_flat)
    a_q = (fp4_e.to(tl.uint8) & 0xF) | ((fp4_o.to(tl.uint8) & 0xF) << 4)

    # B from shuffled
    offs_ks = tl.arange(0, BK_SCALES)
    b_idx = _shuf_data_idx(offs_n[None, :], offs_kp[:, None], Kh_32)
    b = tl.load(B_shuf_ptr + b_idx, mask=offs_n[None, :] < N, other=0)

    bs_idx = _shuf_scale_idx(offs_n[:, None], offs_ks[None, :], Ks_sn, Ks_sm_32)
    b_scale = tl.load(B_scale_shuf_ptr + bs_idx, mask=offs_n[:, None] < N, other=127)

    acc = tl.dot_scaled(a_q, a_scale, "e2m1", b, b_scale, "e2m1")

    c_ptrs = C_ptr + offs_m[:, None] * stride_cm + offs_n[None, :] * stride_cn
    tl.store(c_ptrs, acc.to(tl.bfloat16), mask=(offs_m[:, None] < M) & (offs_n[None, :] < N))


# ─── Dispatch ───
_cache = {}


def custom_kernel(data):
    A = data[0]
    B_shuffle = data[3]
    B_scale_sh = data[4]
    m, k = A.shape
    n = B_shuffle.shape[0]
    Kh = k // 2
    Ks = k // 32

    # Shuffle params
    Kh_32 = Kh // 32  # for shuffle_weight: number of 32-byte blocks in Kh
    # e8m0_shuffle padding
    sm_pad = ((n + 255) // 256) * 256
    sn_pad = ((Ks + 7) // 8) * 8
    Ks_sm_32 = sm_pad // 32  # for scale shuffle formula

    B_sh_flat = B_shuffle.view(torch.uint8).reshape(-1)
    Bs_sh_flat = B_scale_sh.view(torch.uint8).reshape(-1)

    key = (m, n, k)
    if key not in _cache:
        dev = A.device
        _cache[key] = (
            torch.zeros(m, Kh, dtype=torch.uint8, device=dev),
            torch.zeros(m, Ks, dtype=torch.uint8, device=dev),
            torch.empty(m, n, dtype=torch.bfloat16, device=dev),
        )
    Aq, As, C = _cache[key]

    if k == 512 and m <= 32:
        grid = ((m + 15) // 16, (n + 31) // 32)
        _fused_k512_shuf[grid](
            A, B_sh_flat, C, Bs_sh_flat,
            m, n, Kh_32, sn_pad, Ks_sm_32,
            A.stride(0), A.stride(1), C.stride(0), C.stride(1),
            BLOCK_M=16, BLOCK_N=32,
        )
        return C

    _quant_kernel[(m * Ks,)](A, Aq, As, m, k, Kh, Ks, BLOCK=16)

    # Shape dispatch
    if k == 7168 and m == 16 and n == 2112:
        BM, BN, BK, GSM, NS = 16, 16, 1024, 1, 4
    elif k == 7168:
        BM, BN, BK, GSM, NS = 16, 16, 1024, 1, 3
    elif k == 2048 and m == 64 and n == 7168:
        BM, BN, BK, GSM, NS = 64, 32, 512, 4, 3
    elif k == 2048:
        BM, BN, BK, GSM, NS = 64, 32, 512, 4, 3
    elif k == 1536 and m == 256 and n == 3072:
        BM, BN, BK, GSM, NS = 128, 32, 512, 2, 3
    elif k == 512:
        BM, BN, BK, GSM, NS = 32, 64, 64, 4, 1
    else:
        BM, BN, BK, GSM, NS = 64, 64, 64, 8, 1

    if m % BM == 0 and n % BN == 0:
        grid = ((m // BM) * (n // BN),)
        _gemm_shuf_nomask[grid](
            Aq, B_sh_flat, C, As, Bs_sh_flat,
            Aq.stride(0), Aq.stride(1), As.stride(0), As.stride(1),
            C.stride(0), C.stride(1),
            Kh_32, sn_pad, Ks_sm_32,
            BLOCK_M=BM, BLOCK_N=BN, BLOCK_K=BK,
            GROUP_SIZE_M=GSM, NUM_STAGES=NS,
            M_CONST=m, N_CONST=n, K_CONST=k,
        )
    else:
        grid = (triton.cdiv(m, BM) * triton.cdiv(n, BN),)
        _gemm_shuf_masked[grid](
            Aq, B_sh_flat, C, As, Bs_sh_flat,
            m, n, k,
            Aq.stride(0), Aq.stride(1), As.stride(0), As.stride(1),
            C.stride(0), C.stride(1),
            Kh_32, sn_pad, Ks_sm_32,
            BLOCK_M=BM, BLOCK_N=BN, BLOCK_K=BK,
            GROUP_SIZE_M=GSM,
        )
    return C
scrolls · 358 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