Skip to content
KernelIndex
Search⌘K

submission 728249

TraceByWind · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-728249?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.5µs
#512 of 1143
2026-04-04

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:b7189954d4d0d5b77285b029456c4eac5f7d15560c49f26e6c28819fd8117e66
license declaredunknown
license concludedunknown
authorsTraceByWind
imported2026-08-26

Techniques

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

fp4A16WFP4 GEMM with BPRESHUFFLE: BF16 A + MXFP4 B (shuffled) -> BF16 C.
split-kand (args["SPLITK_BLOCK_SIZE"] % args["BLOCK_SIZE_K"] == 0)
tile-m = 16REDUCE_BLOCK_SIZE_M = 16
tile-n = 64REDUCE_BLOCK_SIZE_N = 64

Kernel source

submission.py308 lines
"""
A16WFP4 GEMM with BPRESHUFFLE: BF16 A + MXFP4 B (shuffled) -> BF16 C.
Kernel definitions at module top-level; custom_kernel() prepares data and launches.
"""
from task import input_t, output_t
from aiter.ops.triton._triton_kernels.gemm.basic.gemm_a16wfp4 import _get_config
from aiter.ops.triton.utils._triton.kernel_repr import make_kernel_repr
import torch
import triton
import triton.language as tl
import triton._utils as _tu

# --- ROCm custom dtype monkey-patch ---
_tu.type_canonicalisation_dict["float4_e2m1fn_x2"] = "u8"
_tu.type_canonicalisation_dict["float8_e8m0fnu"] = "u8"

# --- Global helper: pid_grid ---
@triton.jit
def pid_grid(pid, num_pid_m, num_pid_n, GROUP_SIZE_M: tl.constexpr = 1):
    if GROUP_SIZE_M == 1:
        pid_m = pid // num_pid_n
        pid_n = pid % num_pid_n
    else:
        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)
        tl.assume(group_size_m >= 0)
        pid_m = first_pid_m + (pid % group_size_m)
        pid_n = (pid % num_pid_in_group) // group_size_m
    return pid_m, pid_n

# --- Global: _mxfp4_quant_op (from aiter) ---
@triton.jit
def _mxfp4_quant_op(
    x,
    BLOCK_SIZE_N,
    BLOCK_SIZE_M,
    MXFP4_QUANT_BLOCK_SIZE,
):
    EXP_BIAS_FP32: tl.constexpr = 127
    EXP_BIAS_FP4: tl.constexpr = 1
    EBITS_F32: tl.constexpr = 8
    EBITS_FP4: tl.constexpr = 2
    MBITS_FP32: tl.constexpr = 23
    MBITS_FP4: tl.constexpr = 1
    max_normal: tl.constexpr = 6
    min_normal: tl.constexpr = 1
    NUM_QUANT_BLOCKS: tl.constexpr = BLOCK_SIZE_N // MXFP4_QUANT_BLOCK_SIZE
    x = x.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_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)
    bs_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)
    denorm_exp: tl.constexpr = (EXP_BIAS_FP32 - EXP_BIAS_FP4) + (MBITS_FP32 - MBITS_FP4) + 1
    denorm_mask_int: tl.constexpr = denorm_exp << MBITS_FP32
    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
    mant_odd = (normal_x >> (MBITS_FP32 - MBITS_FP4)) & 1
    val_to_add = ((EXP_BIAS_FP4 - EXP_BIAS_FP32) << MBITS_FP32) + (1 << 21) - 1
    normal_x += val_to_add
    normal_x += mant_odd
    normal_x = normal_x >> (MBITS_FP32 - 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_FP32 + 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, MXFP4_QUANT_BLOCK_SIZE // 2, 2])
    evens, odds = tl.split(e2m1_value)
    x_fp4 = evens | (odds << 4)
    x_fp4 = x_fp4.reshape(BLOCK_SIZE_M, BLOCK_SIZE_N // 2)
    return x_fp4, bs_e8m0.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS)

# --- Global: _gemm_a16wfp4_preshuffle_repr & kernel ---
_gemm_a16wfp4_preshuffle_repr = make_kernel_repr(
    "_gemm_a16wfp4_preshuffle_kernel",
    ["BLOCK_SIZE_M","BLOCK_SIZE_N","BLOCK_SIZE_K","GROUP_SIZE_M","num_warps","num_stages","waves_per_eu","matrix_instr_nonkdim","cache_modifier","NUM_KSPLIT"],
)

@triton.heuristics({
    "EVEN_K": lambda args: (args["K"] % (args["BLOCK_SIZE_K"] // 2) == 0)
        and (args["SPLITK_BLOCK_SIZE"] % args["BLOCK_SIZE_K"] == 0)
        and (args["K"] % (args["SPLITK_BLOCK_SIZE"] // 2) == 0),
    "GRID_MN": lambda args: triton.cdiv(args["M"], args["BLOCK_SIZE_M"]) * triton.cdiv(args["N"], args["BLOCK_SIZE_N"]),
})
@triton.jit(repr=_gemm_a16wfp4_preshuffle_repr)
def _gemm_a16wfp4_preshuffle_kernel(
    a_ptr, b_ptr, c_ptr, b_scales_ptr,
    M, N, K,
    stride_am, stride_ak,
    stride_bn, stride_bk,
    stride_ck,
    stride_cm, stride_cn,
    stride_bsn, stride_bsk,
    BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr, BLOCK_SIZE_K: tl.constexpr,
    GROUP_SIZE_M: tl.constexpr, NUM_KSPLIT: tl.constexpr, SPLITK_BLOCK_SIZE: tl.constexpr,
    EVEN_K: tl.constexpr, num_warps: tl.constexpr, num_stages: tl.constexpr,
    waves_per_eu: tl.constexpr, matrix_instr_nonkdim: tl.constexpr, GRID_MN: tl.constexpr,
    PREQUANT: tl.constexpr, cache_modifier: tl.constexpr,
):
    tl.assume(stride_am > 0); tl.assume(stride_ak > 0); tl.assume(stride_bk > 0); tl.assume(stride_bn > 0)
    tl.assume(stride_cm > 0); tl.assume(stride_cn > 0); tl.assume(stride_bsk > 0); tl.assume(stride_bsn > 0)
    pid_unified = tl.program_id(axis=0)
    pid_k = pid_unified % NUM_KSPLIT
    pid = pid_unified // NUM_KSPLIT
    num_pid_m = tl.cdiv(M, BLOCK_SIZE_M)
    num_pid_n = tl.cdiv(N, BLOCK_SIZE_N)
    if NUM_KSPLIT == 1:
        pid_m, pid_n = pid_grid(pid, num_pid_m, num_pid_n, GROUP_SIZE_M=GROUP_SIZE_M)
    else:
        pid_m = pid // num_pid_n
        pid_n = pid % num_pid_n
    tl.assume(pid_m >= 0); tl.assume(pid_n >= 0); tl.assume(pid_k >= 0)
    SCALE_GROUP_SIZE: tl.constexpr = 32
    if (pid_k * SPLITK_BLOCK_SIZE // 2) < K:
        num_k_iter = tl.cdiv(SPLITK_BLOCK_SIZE // 2, BLOCK_SIZE_K // 2)
        offs_k_bf16 = tl.arange(0, BLOCK_SIZE_K)
        offs_k_split_bf16 = pid_k * SPLITK_BLOCK_SIZE + offs_k_bf16
        offs_am = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % M
        a_ptrs = a_ptr + (offs_am[:, None] * stride_am + offs_k_split_bf16[None, :] * stride_ak)
        offs_k_shuffle_arr = tl.arange(0, (BLOCK_SIZE_K // 2) * 16)
        offs_k_shuffle = pid_k * (SPLITK_BLOCK_SIZE // 2) * 16 + offs_k_shuffle_arr
        offs_bn = (pid_n * (BLOCK_SIZE_N // 16) + tl.arange(0, BLOCK_SIZE_N // 16)) % N
        b_ptrs = b_ptr + (offs_bn[:, None] * stride_bn + offs_k_shuffle[None, :] * stride_bk)
        offs_bsn = (pid_n * (BLOCK_SIZE_N // 32) + tl.arange(0, (BLOCK_SIZE_N // 32))) % N
        offs_ks = (pid_k * (SPLITK_BLOCK_SIZE // SCALE_GROUP_SIZE) * 32) + tl.arange(0, BLOCK_SIZE_K // SCALE_GROUP_SIZE * 32)
        b_scale_ptrs = b_scales_ptr + (offs_bsn[:, None] * stride_bsn + offs_ks[None, :] * stride_bsk)
        accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
        for k in range(pid_k * num_k_iter, (pid_k + 1) * num_k_iter):
            b_scales = (tl.load(b_scale_ptrs, cache_modifier=cache_modifier)
                .reshape(BLOCK_SIZE_N // 32, BLOCK_SIZE_K // SCALE_GROUP_SIZE // 8, 4, 16, 2, 2, 1)
                .permute(0,5,3,1,4,2,6)
                .reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // SCALE_GROUP_SIZE))
            if EVEN_K:
                a_bf16 = tl.load(a_ptrs)
                b = tl.load(b_ptrs, cache_modifier=cache_modifier)
            b = (b.reshape(1, BLOCK_SIZE_N // 16, BLOCK_SIZE_K // 64, 2, 16, 16)
                .permute(0,1,4,2,3,5)
                .reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // 2)
                .trans(1,0))
            if PREQUANT:
                a, a_scales = _mxfp4_quant_op(a_bf16, BLOCK_SIZE_K, BLOCK_SIZE_M, 32)
            accumulator += tl.dot_scaled(a, a_scales, "e2m1", b, b_scales, "e2m1")
            a_ptrs += BLOCK_SIZE_K * stride_ak
            b_ptrs += (BLOCK_SIZE_K // 2) * 16 * stride_bk
            b_scale_ptrs += BLOCK_SIZE_K * stride_bsk
        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, :] + pid_k * stride_ck)
        c_mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < N)
        tl.store(c_ptrs, c, mask=c_mask)

# --- Global: get_splitk (Python heuristic) ---
def get_splitk(K: int, BLOCK_SIZE_K: int, NUM_KSPLIT: int):
    NUM_KSPLIT_STEP = 2
    BLOCK_SIZE_K_STEP = 2
    SPLITK_BLOCK_SIZE = (triton.cdiv((2 * triton.cdiv(K, NUM_KSPLIT)), BLOCK_SIZE_K) * BLOCK_SIZE_K)
    while NUM_KSPLIT > 1 and BLOCK_SIZE_K > 16:
        if (K % (SPLITK_BLOCK_SIZE // 2) == 0 and SPLITK_BLOCK_SIZE % BLOCK_SIZE_K == 0 and K % (BLOCK_SIZE_K // 2) == 0):
            break
        elif K % (SPLITK_BLOCK_SIZE // 2) != 0 and NUM_KSPLIT > 1:
            NUM_KSPLIT = NUM_KSPLIT // NUM_KSPLIT_STEP
        elif SPLITK_BLOCK_SIZE % BLOCK_SIZE_K != 0:
            if NUM_KSPLIT > 1:
                NUM_KSPLIT = NUM_KSPLIT // NUM_KSPLIT_STEP
            elif BLOCK_SIZE_K > 16:
                BLOCK_SIZE_K = BLOCK_SIZE_K // BLOCK_SIZE_K_STEP
        elif K % (BLOCK_SIZE_K // 2) != 0 and BLOCK_SIZE_K > 16:
            BLOCK_SIZE_K = BLOCK_SIZE_K // BLOCK_SIZE_K_STEP
        else:
            break
        SPLITK_BLOCK_SIZE = (triton.cdiv((2 * triton.cdiv(K, NUM_KSPLIT)), BLOCK_SIZE_K) * BLOCK_SIZE_K)
    NUM_KSPLIT = triton.cdiv(K, (SPLITK_BLOCK_SIZE // 2))
    return SPLITK_BLOCK_SIZE, BLOCK_SIZE_K, NUM_KSPLIT

# --- Global: _gemm_afp4wfp4_reduce_repr & kernel ---
_gemm_afp4wfp4_reduce_repr = make_kernel_repr("_gemm_afp4wfp4_reduce_kernel", ["BLOCK_SIZE_M","BLOCK_SIZE_N","ACTUAL_KSPLIT","MAX_KSPLIT"])
@triton.heuristics({})
@triton.jit(repr=_gemm_afp4wfp4_reduce_repr)
def _gemm_afp4wfp4_reduce_kernel(
    c_in_ptr, c_out_ptr,
    M, N,
    stride_c_in_k, stride_c_in_m, stride_c_in_n,
    stride_c_out_m, stride_c_out_n,
    BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr,
    ACTUAL_KSPLIT: tl.constexpr, MAX_KSPLIT: tl.constexpr,
):
    pid_m = tl.program_id(axis=0)
    pid_n = tl.program_id(axis=1)
    offs_m = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % M
    offs_n = (pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)) % N
    offs_k = tl.arange(0, MAX_KSPLIT)
    c_in_ptrs = (c_in_ptr + (offs_k[:,None,None] * stride_c_in_k) + (offs_m[None,:,None] * stride_c_in_m) + (offs_n[None,None,:] * stride_c_in_n))
    if ACTUAL_KSPLIT == MAX_KSPLIT:
        c = tl.load(c_in_ptrs)
    else:
        c = tl.load(c_in_ptrs, mask=offs_k[:,None,None] < ACTUAL_KSPLIT)
    c = tl.sum(c, axis=0)
    c = c.to(c_out_ptr.type.element_ty)
    c_out_ptrs = (c_out_ptr + (offs_m[:,None] * stride_c_out_m) + (offs_n[None,:] * stride_c_out_n))
    tl.store(c_out_ptrs, c)

# --- Entry point: custom_kernel ---
def custom_kernel(data: input_t) -> output_t:
    A, B, B_q, B_shuffle, B_scale_sh = data
    A = A.contiguous()
    B_shuffle = B_shuffle.reshape(B_shuffle.shape[0] // 16, B_shuffle.shape[1] * 16).contiguous()
    B_scale_sh = B_scale_sh.reshape(B_scale_sh.shape[0] // 32, B_scale_sh.shape[1] * 32).contiguous()
    M, K_orig = A.shape
    N_packed, K_packed = B_shuffle.shape
    N = N_packed * 16
    K = K_packed // 16

    # Load configuration
    config, _ = _get_config(M, N, K, shuffle=True)

    # Coarse split & BLOCK_SIZE_K adjustment
    if config["BLOCK_SIZE_K"] >= 2 * K:
        config["BLOCK_SIZE_K"] = triton.next_power_of_2(2 * K)
        config["SPLITK_BLOCK_SIZE"] = 2 * K
        config["NUM_KSPLIT"] = 1

    if config["NUM_KSPLIT"] > 1:
        SPLITK_BLOCK_SIZE, BLOCK_SIZE_K, NUM_KSPLIT = get_splitk(
            K, config["BLOCK_SIZE_K"], config["NUM_KSPLIT"]
        )
        config["SPLITK_BLOCK_SIZE"] = SPLITK_BLOCK_SIZE
        config["BLOCK_SIZE_K"] = BLOCK_SIZE_K
        config["NUM_KSPLIT"] = NUM_KSPLIT

    # Post‑get_splitk check (mirrors official)
    if config["BLOCK_SIZE_K"] >= 2 * K:
        config["BLOCK_SIZE_K"] = triton.next_power_of_2(2 * K)
        config["SPLITK_BLOCK_SIZE"] = 2 * K
        config["NUM_KSPLIT"] = 1

    # BLOCK_SIZE_N lower bound
    config["BLOCK_SIZE_N"] = max(config["BLOCK_SIZE_N"], 32)

    # Decide output buffer layout
    use_y_pp = config["NUM_KSPLIT"] > 1
    if use_y_pp:
        y_pp = torch.empty((config["NUM_KSPLIT"], M, N), dtype=torch.float32, device=A.device)
        y = None
    else:
        config["SPLITK_BLOCK_SIZE"] = 2 * K
        y_pp = None
        y = torch.empty((M, N), dtype=torch.bfloat16, device=A.device)

    ck_stride = 0 if not use_y_pp else y_pp.stride(0)

    grid = lambda META: (META["NUM_KSPLIT"] * triton.cdiv(M, META["BLOCK_SIZE_M"]) * triton.cdiv(N, META["BLOCK_SIZE_N"]),)
    _gemm_a16wfp4_preshuffle_kernel[grid](
        A,
        B_shuffle,
        y if not use_y_pp else y_pp,
        B_scale_sh,
        M, N, K,
        A.stride(0), A.stride(1),
        B_shuffle.stride(0), B_shuffle.stride(1),
        ck_stride,
        y.stride(0) if not use_y_pp else y_pp.stride(1),
        y.stride(1) if not use_y_pp else y_pp.stride(2),
        B_scale_sh.stride(0), B_scale_sh.stride(1),
        PREQUANT=True,
        **config,
    )

    if use_y_pp:
        REDUCE_BLOCK_SIZE_M = 16
        REDUCE_BLOCK_SIZE_N = 64
        ACTUAL_KSPLIT = triton.cdiv(K, (config["SPLITK_BLOCK_SIZE"] // 2))
        MAX_KSPLIT = triton.next_power_of_2(config["NUM_KSPLIT"])
        if y is None:
            y = torch.empty((M, N), dtype=torch.bfloat16, device=A.device)
        grid_reduce = (triton.cdiv(M, REDUCE_BLOCK_SIZE_M), triton.cdiv(N, REDUCE_BLOCK_SIZE_N))
        _gemm_afp4wfp4_reduce_kernel[grid_reduce](
            y_pp,
            y,
            M, N,
            y_pp.stride(0), y_pp.stride(1), y_pp.stride(2),
            y.stride(0), y.stride(1),
            REDUCE_BLOCK_SIZE_M, REDUCE_BLOCK_SIZE_N,
            ACTUAL_KSPLIT, MAX_KSPLIT,
        )

    return y[:M, :N]
scrolls · 308 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