Skip to content
KernelIndex
Search⌘K

submission 685211

jiahuizz · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submit_cc_v7.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-685211?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
11.4µs
#330 of 1143
2026-04-01

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:73ca812d1bc9d9df897741d293a6885bdf7bd7fd17bf5e7a4ea6331aca8695fd
license declaredunknown
license concludedunknown
authorsjiahuizz
imported2026-08-26

Techniques

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

fp4submit_cc_v7.py - Hybrid MXFP4 GEMM with shape-specific large-M K tiles.
split-kfrom aiter.ops.triton.gemm.basic.gemm_afp4wfp4 import get_splitk

Kernel source

submit_cc_v7.py634 lines
"""
submit_cc_v7.py - Hybrid MXFP4 GEMM with shape-specific large-M K tiles.

Changes vs submit_cc_v4.py:
- Keep the low-overhead fast paths from v4.
- Use a 512-wide bf16 K tile only for the 64x7168x2048 benchmark shape, where
  it helps materially.
- Keep the 256x3072x1536 shape on the safer 256-wide K tile.
"""
from task import input_t, output_t
import torch
import triton
import triton.language as tl

from aiter.ops.triton.utils._triton.pid_preprocessing import pid_grid, remap_xcd
from aiter.ops.triton._triton_kernels.gemm.basic.gemm_afp4wfp4 import (
    _gemm_afp4wfp4_reduce_kernel,
    _get_config,
)
from aiter.ops.triton.gemm.basic.gemm_afp4wfp4 import get_splitk


# ---------------------------------------------------------------------------
# Inline MXFP4 quantization (shared by fused kernel and quant-only kernel)
# ---------------------------------------------------------------------------
@triton.jit
def _mxfp4_quant_op(x, BK: tl.constexpr, BM: tl.constexpr, QBS: tl.constexpr):
    EXP_BIAS_FP32: tl.constexpr = 127
    EXP_BIAS_FP4: tl.constexpr = 1
    EBITS_F32: tl.constexpr = 8
    EBITS_FP4: tl.constexpr = 2
    MBITS_F32: tl.constexpr = 23
    MBITS_FP4: tl.constexpr = 1
    max_normal: tl.constexpr = 6
    min_normal: tl.constexpr = 1
    NQB: tl.constexpr = BK // QBS
    x = x.reshape(BM, NQB, QBS)
    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_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
    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) & 0xFFFFFFFF
    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, [BM, NQB, QBS // 2, 2])
    evens, odds = tl.split(e2m1_value)
    x_fp4 = evens | (odds << 4)
    return x_fp4.reshape(BM, BK // 2), bs_e8m0.reshape(BM, NQB)


# ---------------------------------------------------------------------------
# Standalone A quantization kernel (for large M path)
# ---------------------------------------------------------------------------
@triton.heuristics({
    "EVEN_M": lambda args: (args["M"] % args["BM_Q"]) == 0,
    "EVEN_KQ": lambda args: (args["K_bf16"] % args["BK_Q"]) == 0,
})
@triton.jit
def _quant_a_kernel(
    a_ptr, a_fp4_ptr, a_scale_ptr,
    M, K_bf16,
    stride_am, stride_ak,
    stride_afm, stride_afk,
    stride_asm, stride_ask,
    BM_Q: tl.constexpr, BK_Q: tl.constexpr,
    EVEN_M: tl.constexpr, EVEN_KQ: tl.constexpr,
):
    SCALE_GROUP_SIZE: tl.constexpr = 32
    NQB: tl.constexpr = BK_Q // SCALE_GROUP_SIZE
    pid_m = tl.program_id(0)
    pid_k = tl.program_id(1)
    offs_m = pid_m * BM_Q + tl.arange(0, BM_Q)
    offs_k = pid_k * BK_Q + tl.arange(0, BK_Q)
    a_ptrs = a_ptr + offs_m[:, None] * stride_am + offs_k[None, :] * stride_ak
    if EVEN_M and EVEN_KQ:
        a = tl.load(a_ptrs)
    else:
        mask = (offs_m[:, None] < M) & (offs_k[None, :] < K_bf16)
        a = tl.load(a_ptrs, mask=mask, other=0.0)
    a_fp4, a_scales = _mxfp4_quant_op(a.to(tl.float32), BK_Q, BM_Q, SCALE_GROUP_SIZE)
    offs_kp = pid_k * (BK_Q // 2) + tl.arange(0, BK_Q // 2)
    a_fp4_ptrs = a_fp4_ptr + offs_m[:, None] * stride_afm + offs_kp[None, :] * stride_afk
    if EVEN_M and EVEN_KQ:
        tl.store(a_fp4_ptrs, a_fp4)
    else:
        fp4_mask = (offs_m[:, None] < M) & (offs_kp[None, :] < K_bf16 // 2)
        tl.store(a_fp4_ptrs, a_fp4, mask=fp4_mask)
    offs_ks = pid_k * NQB + tl.arange(0, NQB)
    a_scale_ptrs = a_scale_ptr + offs_m[:, None] * stride_asm + offs_ks[None, :] * stride_ask
    if EVEN_M and EVEN_KQ:
        tl.store(a_scale_ptrs, a_scales.reshape(BM_Q, NQB))
    else:
        scale_mask = (offs_m[:, None] < M) & (offs_ks[None, :] < K_bf16 // SCALE_GROUP_SIZE)
        tl.store(a_scale_ptrs, a_scales.reshape(BM_Q, NQB), mask=scale_mask)


@triton.jit
def _load_b_scales(
    b_scales_ptr,
    stride_bsn,
    stride_bsk,
    n_group,
    n2,
    n16,
    k_iter,
    BK_SCALES: tl.constexpr,
    K_SCALE_PAD: tl.constexpr,
):
    if BK_SCALES == 8:
        k_local = tl.arange(0, 8)
        k4 = k_local % 4
        k2 = k_local // 4
        base = n16[:, None] * 4 + k2[None, :] * 2 + n2[:, None]
        if K_SCALE_PAD == 64:
            bs_row = n_group[:, None] * 32 + k_iter * 4 + k4[None, :]
            bs_col = base
            return tl.load(b_scales_ptr + bs_row * stride_bsn + bs_col * stride_bsk)
        if K_SCALE_PAD == 48:
            k_mod3 = k_iter % 3
            k_div3 = k_iter // 3
            mixed = base + k_mod3 * 16 + k4[None, :] * 16
            carry = mixed >= 48
            bs_row = n_group[:, None] * 32 + k_iter * 5 + k_div3 + k4[None, :] + carry
            bs_col = mixed - carry * 48
            return tl.load(b_scales_ptr + bs_row * stride_bsn + bs_col * stride_bsk)
        if K_SCALE_PAD == 16:
            bs_row = n_group[:, None] * 32 + k_iter * 16 + k4[None, :] * 4 + (base >> 4)
            bs_col = base & 0xF
            return tl.load(b_scales_ptr + bs_row * stride_bsn + bs_col * stride_bsk)
        flat_local = k4 * 64
        flat_bs = k_iter * 256 + base + flat_local[None, :]
    else:
        offs_ks = k_iter * BK_SCALES + tl.arange(0, BK_SCALES)
        ks8 = offs_ks // 8
        k2 = (offs_ks // 4) % 2
        k4 = offs_ks % 4
        flat_bs = ks8[None, :] * 256 + k4[None, :] * 64 + n16[:, None] * 4 + k2[None, :] * 2 + n2[:, None]
    bs_row = n_group[:, None] * 32 + flat_bs // K_SCALE_PAD
    bs_col = flat_bs % K_SCALE_PAD
    return tl.load(b_scales_ptr + bs_row * stride_bsn + bs_col * stride_bsk)


# ---------------------------------------------------------------------------
# GEMM kernel without inline quant (reads pre-quantized A fp4 + A scales)
# ---------------------------------------------------------------------------
@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),
    "EVEN_M": lambda args: (args["M"] % args["BLOCK_SIZE_M"]) == 0,
    "EVEN_N": lambda args: (args["N"] % args["BLOCK_SIZE_N"]) == 0,
})
@triton.jit
def _gemm_noquant_kernel(
    a_fp4_ptr, a_scale_ptr, b_ptr, c_ptr, b_scales_ptr,
    M, N, K,
    stride_afm, stride_afk,
    stride_asm, stride_ask,
    stride_bk, stride_bn,
    stride_cm, stride_cn, stride_bsn, stride_bsk,
    K_SCALE_PAD: tl.constexpr,
    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,
    EVEN_M: tl.constexpr, EVEN_N: tl.constexpr,
    num_warps: tl.constexpr, num_stages: tl.constexpr,
    waves_per_eu: tl.constexpr, matrix_instr_nonkdim: tl.constexpr,
    cache_modifier: tl.constexpr,
):
    SCALE_GROUP_SIZE: tl.constexpr = 32
    BK_PACKED: tl.constexpr = BLOCK_SIZE_K // 2
    BK_SCALES: tl.constexpr = BLOCK_SIZE_K // SCALE_GROUP_SIZE

    GRID_MN = tl.cdiv(M, BLOCK_SIZE_M) * tl.cdiv(N, BLOCK_SIZE_N)
    pid_unified = tl.program_id(axis=0)
    pid_unified = remap_xcd(pid_unified, GRID_MN * NUM_KSPLIT, NUM_XCDS=8)
    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(stride_afm > 0)
    tl.assume(stride_afk > 0)
    tl.assume(stride_bk > 0)
    tl.assume(stride_bn > 0)
    tl.assume(stride_cm > 0)
    tl.assume(stride_cn > 0)
    tl.assume(stride_bsn > 0)
    tl.assume(stride_bsk > 0)

    if (pid_k * SPLITK_BLOCK_SIZE // 2) < K:
        num_k_iter = tl.cdiv(SPLITK_BLOCK_SIZE // 2, BK_PACKED)
        if EVEN_M:
            offs_am = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
        else:
            offs_am = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % M

        # A fp4 pointers [M, K_packed]
        offs_afk = tl.arange(0, BK_PACKED)
        offs_afk_split = pid_k * (SPLITK_BLOCK_SIZE // 2) + offs_afk
        a_fp4_ptrs = a_fp4_ptr + offs_am[:, None] * stride_afm + offs_afk_split[None, :] * stride_afk

        # A scale pointers [M, K_bf16//32]
        offs_ask = tl.arange(0, BK_SCALES)
        offs_ask_split = pid_k * (SPLITK_BLOCK_SIZE // SCALE_GROUP_SIZE) + offs_ask
        a_scale_ptrs = a_scale_ptr + offs_am[:, None] * stride_asm + offs_ask_split[None, :] * stride_ask

        # B pointers [N, K_packed] loaded as [BN, BK_packed] coalesced
        offs_k = tl.arange(0, BK_PACKED)
        offs_k_split = pid_k * (SPLITK_BLOCK_SIZE // 2) + offs_k
        if EVEN_N:
            offs_bn = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)
        else:
            offs_bn = (pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)) % N
        b_ptrs = b_ptr + offs_bn[:, None] * stride_bn + offs_k_split[None, :] * stride_bk

        # B scale N-side indices (invariant across K iterations)
        n_group = offs_bn // 32
        n2 = (offs_bn % 32) // 16
        n16 = offs_bn % 16

        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):
            if EVEN_K:
                a_fp4 = tl.load(a_fp4_ptrs)
                a_scales = tl.load(a_scale_ptrs)
            else:
                a_fp4 = tl.load(a_fp4_ptrs, mask=offs_afk[None, :] < K - k * BK_PACKED, other=0)
                a_scales = tl.load(a_scale_ptrs, mask=offs_ask[None, :] < tl.cdiv(2 * K, SCALE_GROUP_SIZE) - k * BK_SCALES, other=0)

            b_scales = _load_b_scales(
                b_scales_ptr,
                stride_bsn,
                stride_bsk,
                n_group,
                n2,
                n16,
                k,
                BK_SCALES,
                K_SCALE_PAD,
            )

            if EVEN_K:
                b = tl.load(b_ptrs, cache_modifier=cache_modifier).trans(1, 0)
            else:
                b = tl.load(b_ptrs, mask=offs_k[None, :] < K - k * BK_PACKED, other=0, cache_modifier=cache_modifier).trans(1, 0)

            accumulator = tl.dot_scaled(a_fp4, a_scales, "e2m1", b, b_scales, "e2m1", accumulator)
            a_fp4_ptrs += BK_PACKED * stride_afk
            a_scale_ptrs += BK_SCALES * stride_ask
            b_ptrs += BK_PACKED * stride_bk

        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, :]
        if EVEN_M and EVEN_N:
            tl.store(c_ptrs, c)
        else:
            c_mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < N)
            tl.store(c_ptrs, c, mask=c_mask)


# ---------------------------------------------------------------------------
# Fused kernel (inline A quant + GEMM, for small M shapes)
# ---------------------------------------------------------------------------
@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),
    "EVEN_M": lambda args: (args["M"] % args["BLOCK_SIZE_M"]) == 0,
    "EVEN_N": lambda args: (args["N"] % args["BLOCK_SIZE_N"]) == 0,
})
@triton.jit
def _fused_kernel(
    a_ptr, b_ptr, c_ptr, b_scales_ptr,
    M, N, K,
    stride_am, stride_ak, stride_bk, stride_bn,
    stride_ck, stride_cm, stride_cn, stride_bsn, stride_bsk,
    K_SCALE_PAD: tl.constexpr,
    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,
    EVEN_M: tl.constexpr, EVEN_N: tl.constexpr,
    num_warps: tl.constexpr, num_stages: tl.constexpr,
    waves_per_eu: tl.constexpr, matrix_instr_nonkdim: tl.constexpr,
    cache_modifier: tl.constexpr,
):
    SCALE_GROUP_SIZE: tl.constexpr = 32
    GRID_MN = tl.cdiv(M, BLOCK_SIZE_M) * tl.cdiv(N, BLOCK_SIZE_N)
    pid_unified = tl.program_id(axis=0)
    pid_unified = remap_xcd(pid_unified, GRID_MN * NUM_KSPLIT, NUM_XCDS=8)
    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(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_bsn > 0)
    tl.assume(stride_bsk > 0)

    if (pid_k * SPLITK_BLOCK_SIZE // 2) < K:
        num_k_iter = tl.cdiv(SPLITK_BLOCK_SIZE // 2, BLOCK_SIZE_K // 2)
        if EVEN_M:
            offs_am = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
        else:
            offs_am = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % M
        offs_ak = tl.arange(0, BLOCK_SIZE_K)
        offs_ak_split = pid_k * SPLITK_BLOCK_SIZE + offs_ak
        a_ptrs = a_ptr + offs_am[:, None] * stride_am + offs_ak_split[None, :] * stride_ak

        offs_k = tl.arange(0, BLOCK_SIZE_K // 2)
        offs_k_split = pid_k * (SPLITK_BLOCK_SIZE // 2) + offs_k
        if EVEN_N:
            offs_bn = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)
        else:
            offs_bn = (pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)) % N
        b_ptrs = b_ptr + offs_bn[:, None] * stride_bn + offs_k_split[None, :] * stride_bk

        n_group = offs_bn // 32
        n2 = (offs_bn % 32) // 16
        n16 = offs_bn % 16

        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):
            if EVEN_K:
                a_bf16 = tl.load(a_ptrs)
            else:
                a_bf16 = tl.load(a_ptrs, mask=offs_ak[None, :] < 2 * K - k * BLOCK_SIZE_K, other=0.0)
            a_fp4, a_scales = _mxfp4_quant_op(a_bf16.to(tl.float32), BLOCK_SIZE_K, BLOCK_SIZE_M, SCALE_GROUP_SIZE)

            b_scales = _load_b_scales(
                b_scales_ptr,
                stride_bsn,
                stride_bsk,
                n_group,
                n2,
                n16,
                k,
                BLOCK_SIZE_K // SCALE_GROUP_SIZE,
                K_SCALE_PAD,
            )

            if EVEN_K:
                b = tl.load(b_ptrs, cache_modifier=cache_modifier).trans(1, 0)
            else:
                b = tl.load(b_ptrs, mask=offs_k[None, :] < K - k * (BLOCK_SIZE_K // 2), other=0, cache_modifier=cache_modifier).trans(1, 0)

            accumulator = tl.dot_scaled(a_fp4, a_scales, "e2m1", b, b_scales, "e2m1", accumulator)
            a_ptrs += BLOCK_SIZE_K * stride_ak
            b_ptrs += (BLOCK_SIZE_K // 2) * stride_bk

        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
        if EVEN_M and EVEN_N:
            tl.store(c_ptrs, c)
        else:
            c_mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < N)
            tl.store(c_ptrs, c, mask=c_mask)


# ---------------------------------------------------------------------------
# Per-shape config
# ---------------------------------------------------------------------------
def _make_cfg(bm, bn, bk, gm=4, nks=1, nw=4, ns=2, wpe=0, mind=16, cm=".cg"):
    return {
        "BLOCK_SIZE_M": bm, "BLOCK_SIZE_N": bn, "BLOCK_SIZE_K": bk,
        "GROUP_SIZE_M": gm, "NUM_KSPLIT": nks, "SPLITK_BLOCK_SIZE": bk,
        "num_warps": nw, "num_stages": ns, "waves_per_eu": wpe,
        "matrix_instr_nonkdim": mind, "cache_modifier": cm,
    }

# Fused kernel configs (small M shapes)
_SHAPE_CONFIGS = {
    (4, 2880, 256):    _make_cfg(16, 64, 256, nw=4, ns=2),
    (16, 2112, 3584):  _make_cfg(16, 128, 256, nks=7, nw=4, ns=2),
    (32, 4096, 256):   _make_cfg(32, 64, 256, nw=4, ns=2),
    (32, 2880, 256):   _make_cfg(32, 64, 256, nw=4, ns=2),
    # Fallback for large M (used if noquant path disabled)
    (64, 7168, 1024):  _make_cfg(16, 64, 256, nw=4, ns=2),
    (256, 3072, 768):  _make_cfg(64, 64, 256, nw=8, ns=2, wpe=2),
    # Test shapes
    (8, 2112, 3584):   _make_cfg(16, 128, 256, nw=4, ns=2),
    (16, 3072, 768):   _make_cfg(16, 128, 256, nw=4, ns=2),
    (64, 3072, 768):   _make_cfg(64, 64, 256, nw=8, ns=2),
    (256, 2880, 256):  _make_cfg(64, 64, 256, nw=8, ns=2),
}

# Noquant GEMM configs (large M shapes, lower VGPR → higher occupancy)
_NOQUANT_GEMM_CONFIGS = {
    (64, 7168, 1024):  _make_cfg(32, 64, 512, gm=4, nw=8, ns=2),
    (256, 3072, 768):  _make_cfg(64, 64, 256, nw=8, ns=2, wpe=2),
    (64, 3072, 768):   _make_cfg(64, 64, 256, nw=8, ns=2),
    (256, 2880, 256):  _make_cfg(64, 64, 256, nw=8, ns=2),
}

_NOQUANT_M_THRESHOLD = 64

_config_cache = {}
_noquant_config_cache = {}
_ksplit_cache = {}
_out_cache = {}
_a_fp4_cache = {}
_a_scale_cache = {}


def _get_quant_launch_params(m):
    if m == 64:
        return 32, 512, 8, 1
    return 32, 256, 4, 1


def _get_shape_config(m, n, k_packed):
    key = (m, n, k_packed)
    cached = _config_cache.get(key)
    if cached is not None:
        return dict(cached)
    cfg = _SHAPE_CONFIGS.get(key)
    if cfg is None:
        cfg, _ = _get_config(m, n, k_packed)
    _config_cache[key] = cfg
    return dict(cfg)


def _get_noquant_config(m, n, k_packed):
    key = (m, n, k_packed)
    cached = _noquant_config_cache.get(key)
    if cached is not None:
        return dict(cached)
    cfg = _NOQUANT_GEMM_CONFIGS.get(key)
    if cfg is None:
        cfg, _ = _get_config(m, n, k_packed)
    _noquant_config_cache[key] = cfg
    return dict(cfg)


# ---------------------------------------------------------------------------
# Wrapper
# ---------------------------------------------------------------------------
def custom_kernel(data: input_t) -> output_t:
    A, _, B_q, _, B_scale_sh = data
    m, k_bf16 = A.shape
    n = B_q.shape[0]
    k_packed = k_bf16 // 2

    if m >= _NOQUANT_M_THRESHOLD:
        return _noquant_path(A, B_q, B_scale_sh, m, n, k_bf16, k_packed)
    return _fused_path(A, B_q, B_scale_sh, m, n, k_packed)


def _fused_path(A, B_q, B_scale_sh, m, n, k_packed):
    w = B_q.view(torch.uint8)
    b_scale = B_scale_sh.view(torch.uint8)
    k_scale_pad = b_scale.shape[1]

    config = _get_shape_config(m, n, k_packed)
    key = (m, n, k_packed)

    sk_cached = _ksplit_cache.get(key)
    if sk_cached is not None:
        config["SPLITK_BLOCK_SIZE"], config["BLOCK_SIZE_K"], config["NUM_KSPLIT"] = sk_cached
    else:
        if config["NUM_KSPLIT"] > 1:
            sbs, bk, nks = get_splitk(k_packed, config["BLOCK_SIZE_K"], config["NUM_KSPLIT"])
            config["SPLITK_BLOCK_SIZE"] = sbs
            config["BLOCK_SIZE_K"] = bk
            config["NUM_KSPLIT"] = nks
        else:
            config["SPLITK_BLOCK_SIZE"] = 2 * k_packed
            config["NUM_KSPLIT"] = 1
        if config["BLOCK_SIZE_K"] >= 2 * k_packed:
            config["BLOCK_SIZE_K"] = triton.next_power_of_2(2 * k_packed)
            config["SPLITK_BLOCK_SIZE"] = 2 * k_packed
            config["NUM_KSPLIT"] = 1
        config["BLOCK_SIZE_K"] = max(config["BLOCK_SIZE_K"], 128)
        _ksplit_cache[key] = (config["SPLITK_BLOCK_SIZE"], config["BLOCK_SIZE_K"], config["NUM_KSPLIT"])

    NUM_KSPLIT = config["NUM_KSPLIT"]
    config["K_SCALE_PAD"] = k_scale_pad

    y_key = (m, n, A.device)
    y = _out_cache.get(y_key)
    if y is None:
        y = torch.empty((m, n), dtype=torch.bfloat16, device=A.device)
        _out_cache[y_key] = y

    if NUM_KSPLIT > 1:
        y_pp = torch.empty((NUM_KSPLIT, m, n), dtype=torch.float32, device=A.device)
    else:
        y_pp = None

    grid = lambda META: (
        META["NUM_KSPLIT"] * triton.cdiv(m, META["BLOCK_SIZE_M"]) * triton.cdiv(n, META["BLOCK_SIZE_N"]),
    )

    _fused_kernel[grid](
        A, w, y if NUM_KSPLIT == 1 else y_pp, b_scale,
        m, n, k_packed,
        A.stride(0), A.stride(1), w.stride(1), w.stride(0),
        0 if NUM_KSPLIT == 1 else y_pp.stride(0),
        y.stride(0) if NUM_KSPLIT == 1 else y_pp.stride(1),
        y.stride(1) if NUM_KSPLIT == 1 else y_pp.stride(2),
        b_scale.stride(0), b_scale.stride(1),
        **config,
    )

    if NUM_KSPLIT > 1:
        ACTUAL_KSPLIT = triton.cdiv(k_packed, (config["SPLITK_BLOCK_SIZE"] // 2))
        grid_r = (triton.cdiv(m, 16), triton.cdiv(n, 64))
        _gemm_afp4wfp4_reduce_kernel[grid_r](
            y_pp, y, m, n,
            y_pp.stride(0), y_pp.stride(1), y_pp.stride(2),
            y.stride(0), y.stride(1),
            16, 64, ACTUAL_KSPLIT,
            triton.next_power_of_2(NUM_KSPLIT),
        )

    return y


def _noquant_path(A, B_q, B_scale_sh, m, n, k_bf16, k_packed):
    w = B_q.view(torch.uint8)
    b_scale = B_scale_sh.view(torch.uint8)
    k_scale_pad = b_scale.shape[1]
    k_scales = k_bf16 // 32

    # Cached intermediate buffers
    buf_key = (m, k_packed, A.device)
    a_fp4 = _a_fp4_cache.get(buf_key)
    if a_fp4 is None:
        a_fp4 = torch.empty((m, k_packed), dtype=torch.uint8, device=A.device)
        _a_fp4_cache[buf_key] = a_fp4
    scale_key = (m, k_scales, A.device)
    a_scale = _a_scale_cache.get(scale_key)
    if a_scale is None:
        a_scale = torch.empty((m, k_scales), dtype=torch.uint8, device=A.device)
        _a_scale_cache[scale_key] = a_scale

    # Step 1: Quantize A (separate kernel)
    BM_Q, BK_Q, Q_NUM_WARPS, Q_NUM_STAGES = _get_quant_launch_params(m)
    grid_q = (triton.cdiv(m, BM_Q), triton.cdiv(k_bf16, BK_Q))
    _quant_a_kernel[grid_q](
        A, a_fp4, a_scale,
        m, k_bf16,
        A.stride(0), A.stride(1),
        a_fp4.stride(0), a_fp4.stride(1),
        a_scale.stride(0), a_scale.stride(1),
        BM_Q=BM_Q, BK_Q=BK_Q,
        num_warps=Q_NUM_WARPS, num_stages=Q_NUM_STAGES,
    )

    # Step 2: GEMM with pre-quantized A (no inline quant → lower VGPR → higher occupancy)
    config = _get_noquant_config(m, n, k_packed)
    config["SPLITK_BLOCK_SIZE"] = 2 * k_packed
    config["NUM_KSPLIT"] = 1
    if config["BLOCK_SIZE_K"] >= 2 * k_packed:
        config["BLOCK_SIZE_K"] = triton.next_power_of_2(2 * k_packed)
        config["SPLITK_BLOCK_SIZE"] = 2 * k_packed
    config["BLOCK_SIZE_K"] = max(config["BLOCK_SIZE_K"], 128)
    config["K_SCALE_PAD"] = k_scale_pad

    y_key = (m, n, A.device)
    y = _out_cache.get(y_key)
    if y is None:
        y = torch.empty((m, n), dtype=torch.bfloat16, device=A.device)
        _out_cache[y_key] = y

    grid = lambda META: (
        triton.cdiv(m, META["BLOCK_SIZE_M"]) * triton.cdiv(n, META["BLOCK_SIZE_N"]),
    )

    _gemm_noquant_kernel[grid](
        a_fp4, a_scale, w, y, b_scale,
        m, n, k_packed,
        a_fp4.stride(0), a_fp4.stride(1),
        a_scale.stride(0), a_scale.stride(1),
        w.stride(1), w.stride(0),
        y.stride(0), y.stride(1),
        b_scale.stride(0), b_scale.stride(1),
        **config,
    )

    return y
scrolls · 634 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