Skip to content
KernelIndex
Search⌘K

submission 612876

Danishlynx · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:fa934d86ac034b731227cdb3497c4c0349df401b7bf9a2ca87c0b040f53673a9
license declaredunknown
license concludedunknown
authorsDanishlynx
imported2026-08-15

Techniques

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

fp4fp4 = evens | (odds << 4)
num-warps = 4num_warps=4, num_stages=1,
split-kdef _fused_splitk_gemm(
stages = 1num_warps=c['NW'], num_stages=1,
tile-k = 256BLOCK_K = 256
tile-m = 16BLOCK_M = 16 if m <= 32 else max(16, min(32, triton.next_power_of_2(m)))
tile-n = 32BLOCK_N = 32

Kernel source

submission_gemm_v357.py429 lines
# /// script
# requires-python = ">=3.9"
# dependencies = []
# ///
# leaderboard = "amd-mxfp4-mm"

"""
v357: v354 without CDNA4 env vars (may hurt on ROCm 7.1).
Only HIP_FORCE_DEV_KERNARG=1 (HIP runtime level, not LLVM).
"""
import os, sys
os.environ["HIP_FORCE_DEV_KERNARG"] = "1"

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

P = lambda *a: print(*a, file=sys.stderr, flush=True)
_FP4X2 = _dt.fp4x2
_E8M0 = _dt.fp8_e8m0
_cache = {}


def _knl_name(tile_m, tile_n):
    base = f"f4gemm_bf16_per1x32Fp4_BpreShuffle_{tile_m}x{tile_n}"
    return f"_ZN5aiter{len(base)}{base}E"


@triton.jit
def _mxfp4_quant_op(x, BLOCK_K: tl.constexpr, BLOCK_M: tl.constexpr):
    EXP_BIAS_FP32: tl.constexpr = 127
    EXP_BIAS_FP4: tl.constexpr = 1
    MBITS_F32: tl.constexpr = 23
    MBITS_FP4: tl.constexpr = 1
    EBITS_F32: tl.constexpr = 8
    EBITS_FP4: tl.constexpr = 2
    max_normal: tl.constexpr = 6
    min_normal: tl.constexpr = 1
    QUANT: tl.constexpr = 32
    NUM_QB: tl.constexpr = BLOCK_K // QUANT
    x = x.reshape(BLOCK_M, NUM_QB, QUANT)
    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_unb = tl.log2(amax).floor() - 2
    scale_unb = tl.clamp(scale_unb, min=-127, max=127)
    bs = scale_unb.to(tl.uint8) + 127
    qscale = tl.exp2(-scale_unb)
    qx = x * qscale
    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)
    denorm_exp: tl.constexpr = (EXP_BIAS_FP32 - EXP_BIAS_FP4) + (MBITS_F32 - MBITS_FP4) + 1
    denorm_mask_int: tl.constexpr = denorm_exp << MBITS_F32
    denorm_mask_float: tl.constexpr = tl.cast(denorm_mask_int, tl.float32, bitcast=True)
    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.to(tl.int32)
    mant_odd = (normal_x >> (MBITS_F32 - MBITS_FP4)) & 1
    val_to_add: tl.constexpr = ((EXP_BIAS_FP4 - EXP_BIAS_FP32) << MBITS_F32) + (1 << 21) - 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 = tl.full(qx.type.get_block_shapes(), 0x7, dtype=tl.uint8)
    e2m1 = tl.where(normal_mask, normal_x, e2m1)
    e2m1 = tl.where(denormal_mask, denormal_x, e2m1)
    sign_lp = s >> (MBITS_F32 + EBITS_F32 - MBITS_FP4 - EBITS_FP4)
    sign_lp = sign_lp.to(tl.uint8)
    e2m1 = e2m1 | sign_lp
    e2m1 = tl.reshape(e2m1, [BLOCK_M, NUM_QB, QUANT // 2, 2])
    evens, odds = tl.split(e2m1)
    fp4 = evens | (odds << 4)
    fp4 = fp4.reshape(BLOCK_M, BLOCK_K // 2)
    return fp4, bs.reshape(BLOCK_M, NUM_QB)


@triton.jit
def xcd_swizzle(pid, domain_size, XCD_SWIZZLE: tl.constexpr):
    pids_per_group = domain_size // XCD_SWIZZLE
    extra_pid_groups = domain_size % XCD_SWIZZLE
    group = pid % XCD_SWIZZLE
    local_pid = pid // XCD_SWIZZLE
    new_pid = group * pids_per_group + tl.minimum(group, extra_pid_groups) + local_pid
    return new_pid


# ============ K<=1024: FUSED (same as v127) ============
@triton.jit
def _fused_quant_gemm_kernel(
    A_ptr, Bq_ptr, Bscale_sh_ptr, C_ptr,
    M, N, K: tl.constexpr,
    stride_a_m, stride_a_k, stride_bq_n, stride_bq_k,
    SN_DIV8_MUL256: tl.constexpr,
    stride_c_m, stride_c_n,
    BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: 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)
    acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
    QUANT: tl.constexpr = 32
    NSK: tl.constexpr = BLOCK_K // QUANT

    for ki in tl.range(0, K, BLOCK_K):
        a_offs_k = ki + tl.arange(0, BLOCK_K)
        a_ptrs = A_ptr + offs_m[:, None] * stride_a_m + a_offs_k[None, :] * stride_a_k
        a_mask = (offs_m < M)[:, None] & (a_offs_k < K)[None, :]
        a_tile = tl.load(a_ptrs, mask=a_mask, other=0.0).to(tl.float32)
        a_fp4, a_scale = _mxfp4_quant_op(a_tile, BLOCK_K, BLOCK_M)
        b_offs_k = ki // 2 + tl.arange(0, BLOCK_K // 2)
        b_ptrs = Bq_ptr + offs_n[None, :] * stride_bq_n + b_offs_k[:, None] * stride_bq_k
        b_mask = (offs_n < N)[None, :] & (b_offs_k < K // 2)[:, None]
        b_tile = tl.load(b_ptrs, mask=b_mask, other=0)
        bs_row = offs_n
        bs_col_base = ki // QUANT
        bs_col_offs = tl.arange(0, NSK)
        row = bs_row[:, None]
        col = (bs_col_base + bs_col_offs)[None, :]
        shuf_idx = (row // 32) * SN_DIV8_MUL256 + (col // 8) * 256 + (col % 4) * 64 + (row % 16) * 4 + ((col % 8) // 4) * 2 + ((row % 32) // 16)
        bs_mask = (offs_n[:, None] < N) & (bs_col_offs[None, :] < (K // QUANT - bs_col_base))
        b_scale = tl.load(Bscale_sh_ptr + shuf_idx, mask=bs_mask, other=0)
        acc = tl.dot_scaled(a_fp4, a_scale, "e2m1", b_tile, b_scale, "e2m1", acc)
    c_ptrs = C_ptr + offs_m[:, None] * stride_c_m + offs_n[None, :] * stride_c_n
    c_mask = (offs_m < M)[:, None] & (offs_n < N)[None, :]
    tl.store(c_ptrs, acc.to(tl.bfloat16), mask=c_mask)


# ============ FUSED SPLITK FOR K>1024, M<=32 ============
@triton.jit
def _fused_splitk_gemm(
    A_ptr, Bq_ptr, Bscale_sh_ptr, Y_ptr,
    M, N, K: tl.constexpr,
    stride_a_m, stride_a_k, stride_bq_n, stride_bq_k,
    SN_DIV8_MUL256: tl.constexpr,
    stride_y_k, stride_y_m, stride_y_n,
    grid_m, grid_n,
    BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr,
    SPLIT_K: tl.constexpr, XCD_SWIZZLE: tl.constexpr,
):
    pid = tl.program_id(0)
    total_tiles = grid_m * grid_n * SPLIT_K
    if XCD_SWIZZLE > 1:
        pid = xcd_swizzle(pid, total_tiles, XCD_SWIZZLE)
    pid_k = pid % SPLIT_K
    pid_mn = pid // SPLIT_K
    pid_m = pid_mn // grid_n
    pid_n = pid_mn % grid_n

    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)
    QUANT: tl.constexpr = 32
    NSK: tl.constexpr = BLOCK_K // QUANT

    k_per_split = (K + SPLIT_K - 1) // SPLIT_K
    k_per_split = ((k_per_split + BLOCK_K - 1) // BLOCK_K) * BLOCK_K
    k_start = pid_k * k_per_split
    k_end = min(k_start + k_per_split, K)

    for ki in tl.range(k_start, k_end, BLOCK_K):
        a_offs_k = ki + tl.arange(0, BLOCK_K)
        a_ptrs = A_ptr + offs_m[:, None] * stride_a_m + a_offs_k[None, :] * stride_a_k
        a_mask = (offs_m < M)[:, None] & (a_offs_k < K)[None, :]
        a_tile = tl.load(a_ptrs, mask=a_mask, other=0.0).to(tl.float32)
        a_fp4, a_scale = _mxfp4_quant_op(a_tile, BLOCK_K, BLOCK_M)
        b_offs_k = ki // 2 + tl.arange(0, BLOCK_K // 2)
        b_ptrs = Bq_ptr + offs_n[None, :] * stride_bq_n + b_offs_k[:, None] * stride_bq_k
        b_mask = (offs_n < N)[None, :] & (b_offs_k < K // 2)[:, None]
        b_tile = tl.load(b_ptrs, mask=b_mask, other=0)
        bs_row = offs_n
        bs_col_base = ki // QUANT
        bs_col_offs = tl.arange(0, NSK)
        row = bs_row[:, None]
        col = (bs_col_base + bs_col_offs)[None, :]
        shuf_idx = (row // 32) * SN_DIV8_MUL256 + (col // 8) * 256 + (col % 4) * 64 + (row % 16) * 4 + ((col % 8) // 4) * 2 + ((row % 32) // 16)
        bs_mask = (offs_n[:, None] < N) & (bs_col_offs[None, :] < (K // QUANT - bs_col_base))
        b_scale = tl.load(Bscale_sh_ptr + shuf_idx, mask=bs_mask, other=0)
        acc = tl.dot_scaled(a_fp4, a_scale, "e2m1", b_tile, b_scale, "e2m1", acc, fast_math=True)

    y_ptrs = Y_ptr + pid_k * stride_y_k + offs_m[:, None] * stride_y_m + offs_n[None, :] * stride_y_n
    y_mask = (offs_m < M)[:, None] & (offs_n < N)[None, :]
    if SPLIT_K > 1:
        tl.store(y_ptrs, acc, mask=y_mask)  # FP32 partials
    else:
        tl.store(y_ptrs, acc.to(tl.bfloat16), mask=y_mask)


@triton.jit
def _reduce_splitk(
    Y_ptr, Out_ptr, M, N,
    stride_y_k, stride_y_m, stride_y_n,
    stride_o_m, stride_o_n,
    SPLIT_K: tl.constexpr, BLOCK_N: tl.constexpr,
):
    pid_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 tl.range(0, SPLIT_K):
        vals = tl.load(Y_ptr + k * stride_y_k + pid_m * stride_y_m + offs_n * stride_y_n,
                       mask=n_mask, other=0.0)
        acc += vals.to(tl.float32)
    out_ptrs = Out_ptr + pid_m * stride_o_m + offs_n * stride_o_n
    tl.store(out_ptrs, acc.to(tl.bfloat16), mask=n_mask)


# ============ QUANT KERNEL FOR CK ASM PATH ============
@triton.jit
def _fused_quant_shuffle_kernel(
    x_ptr, x_fp4_ptr, bs_shuf_ptr,
    stride_x_m, stride_x_n, stride_fp4_m, stride_fp4_n,
    M, N, SN_DIV8_MUL256, SCALE_COLS,
    BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr,
    NUM_ITER: tl.constexpr, NUM_STAGES: tl.constexpr,
):
    pid_m = tl.program_id(0)
    start_n = tl.program_id(1) * NUM_ITER
    QUANT: tl.constexpr = 32
    NUM_QB: tl.constexpr = BLOCK_SIZE_N // QUANT
    for pid_n in tl.range(start_n, min(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
        x_mask = (x_offs_m < M)[:, None] & (x_offs_n < N)[None, :]
        x = tl.load(x_ptr + x_offs, mask=x_mask, other=0.0, cache_modifier=".cg").to(tl.float32)
        out_tensor, bs_e8m0 = _mxfp4_quant_op(x, BLOCK_SIZE_N, BLOCK_SIZE_M)
        out_offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
        out_offs_n = pid_n * BLOCK_SIZE_N // 2 + tl.arange(0, BLOCK_SIZE_N // 2)
        out_offs = out_offs_m[:, None] * stride_fp4_m + out_offs_n[None, :] * stride_fp4_n
        out_mask = (out_offs_m < M)[:, None] & (out_offs_n < (N // 2))[None, :]
        tl.store(x_fp4_ptr + out_offs, out_tensor, mask=out_mask)
        bs_row = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
        bs_col = pid_n * NUM_QB + tl.arange(0, NUM_QB)
        row = bs_row[:, None]
        col = bs_col[None, :]
        shuf_idx = (row // 32) * SN_DIV8_MUL256 + (col // 8) * 256 + (col % 4) * 64 + (row % 16) * 4 + ((col % 8) // 4) * 2 + ((row % 32) // 16)
        bs_mask = (bs_row[:, None] < M) & (bs_col[None, :] < SCALE_COLS)
        tl.store(bs_shuf_ptr + shuf_idx, bs_e8m0, mask=bs_mask)


def _init(m, k, n, device):
    QUANT = 32
    scale_cols = (k + QUANT - 1) // QUANT
    sn = ((scale_cols + 7) // 8) * 8
    sn_div8_mul256 = (sn // 8) * 256

    if k <= 1024:
        # Path A: Fused kernel (v127)
        BLOCK_K = max(32, triton.next_power_of_2(k))
        BLOCK_M = 16 if m <= 32 else max(16, min(32, triton.next_power_of_2(m)))
        BLOCK_N = 32
        NW = 4
        grid = (triton.cdiv(m, BLOCK_M), triton.cdiv(n, BLOCK_N))
        out = torch.empty(m, n, dtype=torch.bfloat16, device=device)
        return {
            'mode': 'fused', 'out': out,
            'grid': grid, 'BM': BLOCK_M, 'BN': BLOCK_N, 'BK': BLOCK_K,
            'sn_div8_mul256': sn_div8_mul256, 'NW': NW,
        }
    elif m <= 32:
        # Path B: Fused SplitK for small-M K>1024 (shape 2)
        BLOCK_K = 256
        BLOCK_N = 64
        BLOCK_M = 16 if m <= 16 else 32

        m_tiles = triton.cdiv(m, BLOCK_M)
        n_tiles = triton.cdiv(n, BLOCK_N)
        total_mn = m_tiles * n_tiles

        # Target ~256 WGs for 256 CUs
        SPLIT_K = max(1, min(16, 256 // max(1, total_mn)))
        # Cap at available K-iterations / 2
        k_iters = triton.cdiv(k, BLOCK_K)
        while SPLIT_K > 1 and k_iters < SPLIT_K * 2:
            SPLIT_K //= 2
        # Round to power of 2
        SPLIT_K = 1 << (SPLIT_K - 1).bit_length() if SPLIT_K > 1 else 1

        total_wgs = m_tiles * n_tiles * SPLIT_K
        XCD_SWIZZLE = 8 if total_wgs >= 16 else 1
        out = torch.empty(m, n, dtype=torch.bfloat16, device=device)

        if SPLIT_K > 1:
            scratch = torch.empty(SPLIT_K, m, n, dtype=torch.float32, device=device)
            reduce_grid = (m, triton.cdiv(n, 128))
        else:
            scratch = None
            reduce_grid = None

        return {
            'mode': 'splitk', 'out': out, 'scratch': scratch,
            'BM': BLOCK_M, 'BN': BLOCK_N, 'BK': BLOCK_K,
            'SPLIT_K': SPLIT_K, 'XCD_SWIZZLE': XCD_SWIZZLE,
            'grid_m': m_tiles, 'grid_n': n_tiles, 'total_wgs': total_wgs,
            'sn_div8_mul256': sn_div8_mul256, 'reduce_grid': reduce_grid,
        }
    else:
        # Path C: CK ASM for large-M K>1024 (shapes 5, 6)
        sm = ((m + 255) // 256) * 256
        x_fp4 = torch.empty((m, k // 2), dtype=torch.uint8, device=device)
        bs_shuffled = torch.empty(sm, sn, dtype=torch.uint8, device=device)
        out = torch.empty(m, n, dtype=torch.bfloat16, device=device)

        if m <= 64:
            BSM = triton.next_power_of_2(m)
            NUM_ITER, BSN, NW, NS = 1, 128, 4, 1
        else:
            NUM_ITER, BSM, BSN, NW, NS = 2, 32, 64, 4, 2

        grid = (triton.cdiv(m, BSM), triton.cdiv(k, BSN * NUM_ITER))
        knl = _knl_name(32, 128)
        l2ks = None
        gemm_wgs = triton.cdiv(m, 32) * triton.cdiv(n, 128)
        if gemm_wgs < 32:
            l2ks = 3
        elif gemm_wgs < 64:
            l2ks = 2

        return {
            'mode': 'asm',
            'x_fp4': x_fp4, 'bs_shuffled': bs_shuffled, 'out': out,
            'sc': scale_cols, 'sn_div8_mul256': sn_div8_mul256,
            'grid': grid, 'BSM': BSM, 'BSN': BSN,
            'NW': NW, 'NS': NS, 'NI': NUM_ITER,
            'knl': knl, 'l2ks': l2ks,
        }


def custom_kernel(data: input_t) -> output_t:
    A, B, B_q, B_shuffle, B_scale_sh = data
    m, k = A.shape
    n = B_q.shape[0]
    key = (m, k, n)
    if key not in _cache:
        _cache[key] = _init(m, k, n, A.device)
    c = _cache[key]

    if c['mode'] == 'fused':
        Bq_uint8 = B_q.view(torch.uint8)
        Bscale_uint8 = B_scale_sh.view(torch.uint8)
        _fused_quant_gemm_kernel[c['grid']](
            A, Bq_uint8, Bscale_uint8, c['out'],
            m, n, k,
            A.stride(0), A.stride(1),
            Bq_uint8.stride(0), Bq_uint8.stride(1),
            c['sn_div8_mul256'],
            c['out'].stride(0), c['out'].stride(1),
            BLOCK_M=c['BM'], BLOCK_N=c['BN'], BLOCK_K=c['BK'],
            num_warps=c['NW'], num_stages=1,
        )
        return c['out']

    elif c['mode'] == 'splitk':
        Bq_uint8 = B_q.view(torch.uint8)
        Bscale_uint8 = B_scale_sh.view(torch.uint8)
        SPLIT_K = c['SPLIT_K']
        if SPLIT_K == 1:
            _fused_splitk_gemm[(c['total_wgs'],)](
                A, Bq_uint8, Bscale_uint8, c['out'],
                m, n, k,
                A.stride(0), A.stride(1),
                Bq_uint8.stride(0), Bq_uint8.stride(1),
                c['sn_div8_mul256'],
                0, c['out'].stride(0), c['out'].stride(1),
                c['grid_m'], c['grid_n'],
                BLOCK_M=c['BM'], BLOCK_N=c['BN'], BLOCK_K=c['BK'],
                SPLIT_K=1, XCD_SWIZZLE=c['XCD_SWIZZLE'],
                num_warps=4, num_stages=1,
            )
            return c['out']
        else:
            scratch = c['scratch']
            _fused_splitk_gemm[(c['total_wgs'],)](
                A, Bq_uint8, Bscale_uint8, scratch,
                m, n, k,
                A.stride(0), A.stride(1),
                Bq_uint8.stride(0), Bq_uint8.stride(1),
                c['sn_div8_mul256'],
                scratch.stride(0), scratch.stride(1), scratch.stride(2),
                c['grid_m'], c['grid_n'],
                BLOCK_M=c['BM'], BLOCK_N=c['BN'], BLOCK_K=c['BK'],
                SPLIT_K=SPLIT_K, XCD_SWIZZLE=c['XCD_SWIZZLE'],
                num_warps=4, num_stages=1,
            )
            _reduce_splitk[c['reduce_grid']](
                scratch, c['out'], m, n,
                scratch.stride(0), scratch.stride(1), scratch.stride(2),
                c['out'].stride(0), c['out'].stride(1),
                SPLIT_K=SPLIT_K, BLOCK_N=128,
                num_warps=4,
            )
            return c['out']

    else:  # asm
        x_fp4 = c['x_fp4']
        bs_shuf = c['bs_shuffled']
        _fused_quant_shuffle_kernel[c['grid']](
            A, x_fp4, bs_shuf,
            A.stride(0), A.stride(1),
            x_fp4.stride(0), x_fp4.stride(1),
            m, k,
            c['sn_div8_mul256'], c['sc'],
            BLOCK_SIZE_M=c['BSM'], BLOCK_SIZE_N=c['BSN'],
            NUM_ITER=c['NI'], NUM_STAGES=c['NS'],
            num_warps=c['NW'], waves_per_eu=0, num_stages=1,
        )
        out = c['out']
        aiter.gemm_a4w4_asm(
            x_fp4.view(_FP4X2), B_shuffle,
            bs_shuf.view(_E8M0), B_scale_sh,
            out, c['knl'],
            bpreshuffle=True,
            log2_k_split=c['l2ks'],
        )
        return out
scrolls · 429 lines total

Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0

Changes from previous submission

Against this author's previous submission submission 566493.

⋯ 4 unchanged lines
# leaderboard = "amd-mxfp4-mm"
"""
- v116: Hybrid with BN=32 for K<=1024 (more WGs) + ASM for K>1024.
- v110 showed BN=64 improved shape1 to 6.56us. BN=32 may push further:
- Shape 1: grid(1,90)=90 WGs (was 45 with BN=64)
- Shape 3: grid(1,128)=128 WGs (was 64)
- Shape 4: grid(1,90)=90 WGs (was 45)
- v111 confirmed fused Triton is 3x slower for K>1024 -> keep ASM path.
+ v357: v354 without CDNA4 env vars (may hurt on ROCm 7.1).
+ Only HIP_FORCE_DEV_KERNARG=1 (HIP runtime level, not LLVM).
"""
+ import os, sys
+ os.environ["HIP_FORCE_DEV_KERNARG"] = "1"
+
from task import input_t, output_t
import torch
import triton
⋯ 1 unchanged lines
import aiter
from aiter import dtypes as _dt
+ P = lambda *a: print(*a, file=sys.stderr, flush=True)
_FP4X2 = _dt.fp4x2
_E8M0 = _dt.fp8_e8m0
- _BF16 = _dt.bf16
_cache = {}
⋯ 59 unchanged lines
@triton.jit
+ def xcd_swizzle(pid, domain_size, XCD_SWIZZLE: tl.constexpr):
+ pids_per_group = domain_size // XCD_SWIZZLE
+ extra_pid_groups = domain_size % XCD_SWIZZLE
+ group = pid % XCD_SWIZZLE
+ local_pid = pid // XCD_SWIZZLE
+ new_pid = group * pids_per_group + tl.minimum(group, extra_pid_groups) + local_pid
+ return new_pid
+
+
+ # ============ K<=1024: FUSED (same as v127) ============
+ @triton.jit
def _fused_quant_gemm_kernel(
A_ptr, Bq_ptr, Bscale_sh_ptr, C_ptr,
M, N, K: tl.constexpr,
- stride_a_m, stride_a_k,
- stride_bq_n, stride_bq_k,
+ stride_a_m, stride_a_k, stride_bq_n, stride_bq_k,
SN_DIV8_MUL256: tl.constexpr,
stride_c_m, stride_c_n,
BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr,
⋯ 12 unchanged lines
a_mask = (offs_m < M)[:, None] & (a_offs_k < K)[None, :]
a_tile = tl.load(a_ptrs, mask=a_mask, other=0.0).to(tl.float32)
a_fp4, a_scale = _mxfp4_quant_op(a_tile, BLOCK_K, BLOCK_M)
-
b_offs_k = ki // 2 + tl.arange(0, BLOCK_K // 2)
b_ptrs = Bq_ptr + offs_n[None, :] * stride_bq_n + b_offs_k[:, None] * stride_bq_k
b_mask = (offs_n < N)[None, :] & (b_offs_k < K // 2)[:, None]
b_tile = tl.load(b_ptrs, mask=b_mask, other=0)
-
bs_row = offs_n
bs_col_base = ki // QUANT
bs_col_offs = tl.arange(0, NSK)
⋯ 2 unchanged lines
shuf_idx = (row // 32) * SN_DIV8_MUL256 + (col // 8) * 256 + (col % 4) * 64 + (row % 16) * 4 + ((col % 8) // 4) * 2 + ((row % 32) // 16)
bs_mask = (offs_n[:, None] < N) & (bs_col_offs[None, :] < (K // QUANT - bs_col_base))
b_scale = tl.load(Bscale_sh_ptr + shuf_idx, mask=bs_mask, other=0)
-
acc = tl.dot_scaled(a_fp4, a_scale, "e2m1", b_tile, b_scale, "e2m1", acc)
-
c_ptrs = C_ptr + offs_m[:, None] * stride_c_m + offs_n[None, :] * stride_c_n
c_mask = (offs_m < M)[:, None] & (offs_n < N)[None, :]
tl.store(c_ptrs, acc.to(tl.bfloat16), mask=c_mask)
+ # ============ FUSED SPLITK FOR K>1024, M<=32 ============
@triton.jit
+ def _fused_splitk_gemm(
+ A_ptr, Bq_ptr, Bscale_sh_ptr, Y_ptr,
+ M, N, K: tl.constexpr,
+ stride_a_m, stride_a_k, stride_bq_n, stride_bq_k,
+ SN_DIV8_MUL256: tl.constexpr,
+ stride_y_k, stride_y_m, stride_y_n,
+ grid_m, grid_n,
+ BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr,
+ SPLIT_K: tl.constexpr, XCD_SWIZZLE: tl.constexpr,
+ ):
+ pid = tl.program_id(0)
+ total_tiles = grid_m * grid_n * SPLIT_K
+ if XCD_SWIZZLE > 1:
+ pid = xcd_swizzle(pid, total_tiles, XCD_SWIZZLE)
+ pid_k = pid % SPLIT_K
+ pid_mn = pid // SPLIT_K
+ pid_m = pid_mn // grid_n
+ pid_n = pid_mn % grid_n
+
+ 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)
+ QUANT: tl.constexpr = 32
+ NSK: tl.constexpr = BLOCK_K // QUANT
+
+ k_per_split = (K + SPLIT_K - 1) // SPLIT_K
+ k_per_split = ((k_per_split + BLOCK_K - 1) // BLOCK_K) * BLOCK_K
+ k_start = pid_k * k_per_split
+ k_end = min(k_start + k_per_split, K)
+
+ for ki in tl.range(k_start, k_end, BLOCK_K):
+ a_offs_k = ki + tl.arange(0, BLOCK_K)
+ a_ptrs = A_ptr + offs_m[:, None] * stride_a_m + a_offs_k[None, :] * stride_a_k
+ a_mask = (offs_m < M)[:, None] & (a_offs_k < K)[None, :]
+ a_tile = tl.load(a_ptrs, mask=a_mask, other=0.0).to(tl.float32)
+ a_fp4, a_scale = _mxfp4_quant_op(a_tile, BLOCK_K, BLOCK_M)
+ b_offs_k = ki // 2 + tl.arange(0, BLOCK_K // 2)
+ b_ptrs = Bq_ptr + offs_n[None, :] * stride_bq_n + b_offs_k[:, None] * stride_bq_k
+ b_mask = (offs_n < N)[None, :] & (b_offs_k < K // 2)[:, None]
+ b_tile = tl.load(b_ptrs, mask=b_mask, other=0)
+ bs_row = offs_n
+ bs_col_base = ki // QUANT
+ bs_col_offs = tl.arange(0, NSK)
+ row = bs_row[:, None]
+ col = (bs_col_base + bs_col_offs)[None, :]
+ shuf_idx = (row // 32) * SN_DIV8_MUL256 + (col // 8) * 256 + (col % 4) * 64 + (row % 16) * 4 + ((col % 8) // 4) * 2 + ((row % 32) // 16)
+ bs_mask = (offs_n[:, None] < N) & (bs_col_offs[None, :] < (K // QUANT - bs_col_base))
+ b_scale = tl.load(Bscale_sh_ptr + shuf_idx, mask=bs_mask, other=0)
+ acc = tl.dot_scaled(a_fp4, a_scale, "e2m1", b_tile, b_scale, "e2m1", acc, fast_math=True)
+
+ y_ptrs = Y_ptr + pid_k * stride_y_k + offs_m[:, None] * stride_y_m + offs_n[None, :] * stride_y_n
+ y_mask = (offs_m < M)[:, None] & (offs_n < N)[None, :]
+ if SPLIT_K > 1:
+ tl.store(y_ptrs, acc, mask=y_mask) # FP32 partials
+ else:
+ tl.store(y_ptrs, acc.to(tl.bfloat16), mask=y_mask)
+
+
+ @triton.jit
+ def _reduce_splitk(
+ Y_ptr, Out_ptr, M, N,
+ stride_y_k, stride_y_m, stride_y_n,
+ stride_o_m, stride_o_n,
+ SPLIT_K: tl.constexpr, BLOCK_N: tl.constexpr,
+ ):
+ pid_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 tl.range(0, SPLIT_K):
+ vals = tl.load(Y_ptr + k * stride_y_k + pid_m * stride_y_m + offs_n * stride_y_n,
+ mask=n_mask, other=0.0)
+ acc += vals.to(tl.float32)
+ out_ptrs = Out_ptr + pid_m * stride_o_m + offs_n * stride_o_n
+ tl.store(out_ptrs, acc.to(tl.bfloat16), mask=n_mask)
+
+
+ # ============ QUANT KERNEL FOR CK ASM PATH ============
+ @triton.jit
def _fused_quant_shuffle_kernel(
x_ptr, x_fp4_ptr, bs_shuf_ptr,
- stride_x_m, stride_x_n,
- stride_fp4_m, stride_fp4_n,
- M, N,
- SN_DIV8_MUL256, SCALE_COLS,
+ stride_x_m, stride_x_n, stride_fp4_m, stride_fp4_n,
+ M, N, SN_DIV8_MUL256, SCALE_COLS,
BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr,
NUM_ITER: tl.constexpr, NUM_STAGES: tl.constexpr,
):
⋯ 1 unchanged lines
start_n = tl.program_id(1) * NUM_ITER
QUANT: tl.constexpr = 32
NUM_QB: tl.constexpr = BLOCK_SIZE_N // QUANT
-
for pid_n in tl.range(start_n, min(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
x_mask = (x_offs_m < M)[:, None] & (x_offs_n < N)[None, :]
x = tl.load(x_ptr + x_offs, mask=x_mask, other=0.0, cache_modifier=".cg").to(tl.float32)
-
out_tensor, bs_e8m0 = _mxfp4_quant_op(x, BLOCK_SIZE_N, BLOCK_SIZE_M)
-
out_offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
out_offs_n = pid_n * BLOCK_SIZE_N // 2 + tl.arange(0, BLOCK_SIZE_N // 2)
out_offs = out_offs_m[:, None] * stride_fp4_m + out_offs_n[None, :] * stride_fp4_n
out_mask = (out_offs_m < M)[:, None] & (out_offs_n < (N // 2))[None, :]
tl.store(x_fp4_ptr + out_offs, out_tensor, mask=out_mask)
-
bs_row = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
bs_col = pid_n * NUM_QB + tl.arange(0, NUM_QB)
row = bs_row[:, None]
⋯ 6 unchanged lines
def _init(m, k, n, device):
QUANT = 32
scale_cols = (k + QUANT - 1) // QUANT
+ sn = ((scale_cols + 7) // 8) * 8
+ sn_div8_mul256 = (sn // 8) * 256
if k <= 1024:
+ # Path A: Fused kernel (v127)
BLOCK_K = max(32, triton.next_power_of_2(k))
- BLOCK_M = max(16, min(32, triton.next_power_of_2(m)))
+ BLOCK_M = 16 if m <= 32 else max(16, min(32, triton.next_power_of_2(m)))
BLOCK_N = 32
NW = 4
grid = (triton.cdiv(m, BLOCK_M), triton.cdiv(n, BLOCK_N))
out = torch.empty(m, n, dtype=torch.bfloat16, device=device)
- sn = ((scale_cols + 7) // 8) * 8
- sn_div8_mul256 = (sn // 8) * 256
return {
'mode': 'fused', 'out': out,
'grid': grid, 'BM': BLOCK_M, 'BN': BLOCK_N, 'BK': BLOCK_K,
'sn_div8_mul256': sn_div8_mul256, 'NW': NW,
}
+ elif m <= 32:
+ # Path B: Fused SplitK for small-M K>1024 (shape 2)
+ BLOCK_K = 256
+ BLOCK_N = 64
+ BLOCK_M = 16 if m <= 16 else 32
+
+ m_tiles = triton.cdiv(m, BLOCK_M)
+ n_tiles = triton.cdiv(n, BLOCK_N)
+ total_mn = m_tiles * n_tiles
+
+ # Target ~256 WGs for 256 CUs
+ SPLIT_K = max(1, min(16, 256 // max(1, total_mn)))
+ # Cap at available K-iterations / 2
+ k_iters = triton.cdiv(k, BLOCK_K)
+ while SPLIT_K > 1 and k_iters < SPLIT_K * 2:
+ SPLIT_K //= 2
+ # Round to power of 2
+ SPLIT_K = 1 << (SPLIT_K - 1).bit_length() if SPLIT_K > 1 else 1
+
+ total_wgs = m_tiles * n_tiles * SPLIT_K
+ XCD_SWIZZLE = 8 if total_wgs >= 16 else 1
+ out = torch.empty(m, n, dtype=torch.bfloat16, device=device)
+
+ if SPLIT_K > 1:
+ scratch = torch.empty(SPLIT_K, m, n, dtype=torch.float32, device=device)
+ reduce_grid = (m, triton.cdiv(n, 128))
+ else:
+ scratch = None
+ reduce_grid = None
+
+ return {
+ 'mode': 'splitk', 'out': out, 'scratch': scratch,
+ 'BM': BLOCK_M, 'BN': BLOCK_N, 'BK': BLOCK_K,
+ 'SPLIT_K': SPLIT_K, 'XCD_SWIZZLE': XCD_SWIZZLE,
+ 'grid_m': m_tiles, 'grid_n': n_tiles, 'total_wgs': total_wgs,
+ 'sn_div8_mul256': sn_div8_mul256, 'reduce_grid': reduce_grid,
+ }
else:
+ # Path C: CK ASM for large-M K>1024 (shapes 5, 6)
sm = ((m + 255) // 256) * 256
- sn = ((scale_cols + 7) // 8) * 8
x_fp4 = torch.empty((m, k // 2), dtype=torch.uint8, device=device)
bs_shuffled = torch.empty(sm, sn, dtype=torch.uint8, device=device)
out = torch.empty(m, n, dtype=torch.bfloat16, device=device)
- if m <= 32:
+ if m <= 64:
BSM = triton.next_power_of_2(m)
NUM_ITER, BSN, NW, NS = 1, 128, 4, 1
- elif m <= 64:
- NUM_ITER, BSM, BSN, NW, NS = 2, 32, 64, 4, 2
else:
NUM_ITER, BSM, BSN, NW, NS = 2, 32, 64, 4, 2
grid = (triton.cdiv(m, BSM), triton.cdiv(k, BSN * NUM_ITER))
knl = _knl_name(32, 128)
+ l2ks = None
+ gemm_wgs = triton.cdiv(m, 32) * triton.cdiv(n, 128)
+ if gemm_wgs < 32:
+ l2ks = 3
+ elif gemm_wgs < 64:
+ l2ks = 2
+
return {
'mode': 'asm',
'x_fp4': x_fp4, 'bs_shuffled': bs_shuffled, 'out': out,
- 'sc': scale_cols, 'sn_div8_mul256': (sn // 8) * 256,
+ 'sc': scale_cols, 'sn_div8_mul256': sn_div8_mul256,
'grid': grid, 'BSM': BSM, 'BSN': BSN,
'NW': NW, 'NS': NS, 'NI': NUM_ITER,
- 'knl': knl,
+ 'knl': knl, 'l2ks': l2ks,
}
def custom_kernel(data: input_t) -> output_t:
A, B, B_q, B_shuffle, B_scale_sh = data
- m = A.shape[0]
- k = A.shape[1]
+ m, k = A.shape
n = B_q.shape[0]
-
key = (m, k, n)
if key not in _cache:
_cache[key] = _init(m, k, n, A.device)
⋯ 13 unchanged lines
num_warps=c['NW'], num_stages=1,
)
return c['out']
- else:
+
+ elif c['mode'] == 'splitk':
+ Bq_uint8 = B_q.view(torch.uint8)
+ Bscale_uint8 = B_scale_sh.view(torch.uint8)
+ SPLIT_K = c['SPLIT_K']
+ if SPLIT_K == 1:
+ _fused_splitk_gemm[(c['total_wgs'],)](
+ A, Bq_uint8, Bscale_uint8, c['out'],
+ m, n, k,
+ A.stride(0), A.stride(1),
+ Bq_uint8.stride(0), Bq_uint8.stride(1),
+ c['sn_div8_mul256'],
+ 0, c['out'].stride(0), c['out'].stride(1),
+ c['grid_m'], c['grid_n'],
+ BLOCK_M=c['BM'], BLOCK_N=c['BN'], BLOCK_K=c['BK'],
+ SPLIT_K=1, XCD_SWIZZLE=c['XCD_SWIZZLE'],
+ num_warps=4, num_stages=1,
+ )
+ return c['out']
+ else:
+ scratch = c['scratch']
+ _fused_splitk_gemm[(c['total_wgs'],)](
+ A, Bq_uint8, Bscale_uint8, scratch,
+ m, n, k,
+ A.stride(0), A.stride(1),
+ Bq_uint8.stride(0), Bq_uint8.stride(1),
+ c['sn_div8_mul256'],
+ scratch.stride(0), scratch.stride(1), scratch.stride(2),
+ c['grid_m'], c['grid_n'],
+ BLOCK_M=c['BM'], BLOCK_N=c['BN'], BLOCK_K=c['BK'],
+ SPLIT_K=SPLIT_K, XCD_SWIZZLE=c['XCD_SWIZZLE'],
+ num_warps=4, num_stages=1,
+ )
+ _reduce_splitk[c['reduce_grid']](
+ scratch, c['out'], m, n,
+ scratch.stride(0), scratch.stride(1), scratch.stride(2),
+ c['out'].stride(0), c['out'].stride(1),
+ SPLIT_K=SPLIT_K, BLOCK_N=128,
+ num_warps=4,
+ )
+ return c['out']
+
+ else: # asm
x_fp4 = c['x_fp4']
bs_shuf = c['bs_shuffled']
_fused_quant_shuffle_kernel[c['grid']](
⋯ 12 unchanged lines
bs_shuf.view(_E8M0), B_scale_sh,
out, c['knl'],
bpreshuffle=True,
+ log2_k_split=c['l2ks'],
)
return out
scrolls · 361 diff lines total

Best evidence level for this revision: reported

JSON