Skip to content
KernelIndex
Search⌘K

submission 754951

dc1312 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_v160.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-754951?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
8.03µs
#17 of 1143
2026-04-07

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:dc02f98a78e809d952f31eb83bd58739c23ff3f9aabfbcf6fd7665036f4ec254
license declaredunknown
license concludedunknown
authorsdc1312
imported2026-08-15

Techniques

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

autotuneStrategy: HIP assembly kernel for small-K shapes, autotuned Triton for large-K.
fp4MXFP4 GEMM: hybrid HIP + Triton approach for mixed-precision FP4 matmul.
num-warps = 4num_warps=4, num_stages=2,
shared-memory__shared__ __align__(16) uint8_t Alds[2 * BLOCK_M * LDS_ROW];
split-kSplit-K decomposition for M=16 with high K to maximize CU utilization.
stages = 2num_warps=4, num_stages=2,
tile-m = 16constexpr int MFMA_K=128, DOUBLE_K=256, LDS_ROW=128, HALF_K=64, BLOCK_M=16;
tile-n = 64constexpr int BLOCK_N = 64, K_FIXED = 512;
vector-width = int4int4 ra = reinterpret_cast<const int4*>(s)[0];

Kernel source

submission_v160.py683 lines
"""
MXFP4 GEMM: hybrid HIP + Triton approach for mixed-precision FP4 matmul.
Optimized for MI355X (gfx950/CDNA4) with block-scaled FP4 MFMA.
Strategy: HIP assembly kernel for small-K shapes, autotuned Triton for large-K.
Split-K decomposition for M=16 with high K to maximize CU utilization.
"""
import torch
import triton
import triton.language as tl
from task import input_t, output_t
import os
os.environ["PYTORCH_ROCM_ARCH"] = "gfx950"

# ===================== HIP kernel for M<=32 K=512 =====================

_HIP_CPP = r"""
#include <torch/extension.h>
void launch_sk2_k512(int64_t a, int64_t bq, int64_t bsc, int64_t c, int M, int N);
"""

_HIP_SRC = r"""
#include <hip/hip_runtime.h>
#include <hip/hip_bf16.h>
#include <torch/extension.h>

constexpr int MFMA_K=128, DOUBLE_K=256, LDS_ROW=128, HALF_K=64, BLOCK_M=16;
typedef int __attribute__((ext_vector_type(4))) int4_vec;
typedef float __attribute__((ext_vector_type(4))) float4_vec;
typedef int32_t __attribute__((ext_vector_type(4))) i32x4;
typedef uint32_t __attribute__((address_space(3)))* as3_uint32_ptr;
extern "C" __device__ void llvm_amdgcn_raw_buffer_load_lds(
    i32x4 rsrc, as3_uint32_ptr lds_ptr, int size, int voffset,
    int soffset, int offset, int aux) __asm("llvm.amdgcn.raw.buffer.load.lds");

struct buffer_resource { uint64_t ptr; uint32_t range; uint32_t config; };

__device__ __forceinline__ i32x4 make_srsrc(const void* p, uint32_t r) {
    buffer_resource s = {reinterpret_cast<uint64_t>(p), r, 0x110000};
    return *reinterpret_cast<const i32x4*>(&s);
}

__device__ __forceinline__ float4_vec mfma_fp4(int4_vec A, int4_vec B, float4_vec C, int sA, int sB) {
    float4_vec D;
    asm("v_mfma_scale_f32_16x16x128_f8f6f4 %0,%1,%2,%3,%4,%5 cbsz:4 blgp:4"
        : "=v"(D) : "v"(A), "v"(B), "v"(C), "v"(sA), "v"(sB));
    return D;
}

__device__ __forceinline__ int lds_swz(int o) {
    return o ^ (((o & 2047) >> 8) << 4);
}

__device__ __forceinline__ void compute_scale(float mx, uint8_t& sc, float& sf) {
    if (mx > 0.f) {
        uint32_t b = __float_as_uint(mx);
        b = (b + 0x200000u) & 0xFF800000u;
        int su = ((b >> 23) & 0xFF) - 129;
        su = su < -127 ? -127 : (su > 127 ? 127 : su);
        sc = (uint8_t)(su + 127);
        sf = __uint_as_float((uint32_t)(su + 127) << 23);
    } else {
        sc = 0; sf = 0.f;
    }
}

__device__ __forceinline__ float tree_max8(const float* v) {
    float a0 = fmaxf(fabsf(v[0]), fabsf(v[1]));
    float a1 = fmaxf(fabsf(v[2]), fabsf(v[3]));
    float a2 = fmaxf(fabsf(v[4]), fabsf(v[5]));
    float a3 = fmaxf(fabsf(v[6]), fabsf(v[7]));
    return fmaxf(fmaxf(a0, a1), fmaxf(a2, a3));
}

__global__ __launch_bounds__(512, 3)
void gemm_bn64_sk2_k512(
    const __hip_bfloat16* __restrict__ A,
    const uint8_t* __restrict__ Bq,
    const uint8_t* __restrict__ Bsc,
    __hip_bfloat16* __restrict__ C,
    int M, int N)
{
    constexpr int BLOCK_N = 64, K_FIXED = 512;
    constexpr int BSTRIDE = K_FIXED >> 1, SCS = K_FIXED >> 5;
    const int group = threadIdx.x >> 8;
    const int ltid = threadIdx.x & 255;
    const int wid_local = ltid >> 6;
    const int lm = ltid & 15;
    const int lk = (ltid >> 4) & 3;

    const int bn = blockIdx.x * BLOCK_N, bm = blockIdx.y * BLOCK_M;
    const int wn = bn + (wid_local << 4);

    __shared__ __align__(16) uint8_t Alds[2 * BLOCK_M * LDS_ROW];
    __shared__ uint8_t Asclds[2 * BLOCK_M * 8];
    __shared__ __align__(16) uint8_t Blds[2 * BLOCK_N * LDS_ROW];
    __shared__ uint8_t Bslds[2 * BLOCK_N * 8];

    const int a_base = group * BLOCK_M * LDS_ROW;
    const int as_base = group * BLOCK_M * 8;
    const int b_base = group * BLOCK_N * LDS_ROW;
    const int bs_base = group * BLOCK_N * 8;

    const int ke = group * DOUBLE_K;
    const int kb = group * LDS_ROW;

    const i32x4 srsrc = make_srsrc(Bq, N * BSTRIDE);
    float4_vec acc = {0, 0, 0, 0};

    if (ltid < 128) {
        const int qr = ltid >> 3, qg = ltid & 7;
        const int gr = bm + qr, ko = ke + qg * 32;
        uint32_t p0 = 0, p1 = 0, p2 = 0, p3 = 0;
        uint8_t asc = 0x7f;
        if (gr < M) {
            const __hip_bfloat16* s = A + gr * K_FIXED + ko;
            int4 ra = reinterpret_cast<const int4*>(s)[0];
            int4 rb = reinterpret_cast<const int4*>(s)[1];
            int4 rc = reinterpret_cast<const int4*>(s)[2];
            int4 rd = reinterpret_cast<const int4*>(s)[3];
            const __hip_bfloat16* bfa = reinterpret_cast<const __hip_bfloat16*>(&ra);
            const __hip_bfloat16* bfb = reinterpret_cast<const __hip_bfloat16*>(&rb);
            const __hip_bfloat16* bfc = reinterpret_cast<const __hip_bfloat16*>(&rc);
            const __hip_bfloat16* bfd = reinterpret_cast<const __hip_bfloat16*>(&rd);
            float v0[8], v1[8], v2[8], v3[8];
            for (int i = 0; i < 8; i++) v0[i] = __bfloat162float(bfa[i]);
            for (int i = 0; i < 8; i++) v1[i] = __bfloat162float(bfb[i]);
            for (int i = 0; i < 8; i++) v2[i] = __bfloat162float(bfc[i]);
            for (int i = 0; i < 8; i++) v3[i] = __bfloat162float(bfd[i]);
            float l = fmaxf(fmaxf(tree_max8(v0), tree_max8(v1)),
                            fmaxf(tree_max8(v2), tree_max8(v3)));
            float sf;
            compute_scale(l, asc, sf);
            p0 = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(p0, v0[0], v0[1], sf, 0);
            p0 = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(p0, v0[2], v0[3], sf, 1);
            p0 = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(p0, v0[4], v0[5], sf, 2);
            p0 = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(p0, v0[6], v0[7], sf, 3);
            p1 = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(p1, v1[0], v1[1], sf, 0);
            p1 = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(p1, v1[2], v1[3], sf, 1);
            p1 = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(p1, v1[4], v1[5], sf, 2);
            p1 = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(p1, v1[6], v1[7], sf, 3);
            p2 = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(p2, v2[0], v2[1], sf, 0);
            p2 = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(p2, v2[2], v2[3], sf, 1);
            p2 = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(p2, v2[4], v2[5], sf, 2);
            p2 = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(p2, v2[6], v2[7], sf, 3);
            p3 = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(p3, v3[0], v3[1], sf, 0);
            p3 = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(p3, v3[2], v3[3], sf, 1);
            p3 = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(p3, v3[4], v3[5], sf, 2);
            p3 = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(p3, v3[6], v3[7], sf, 3);
        }
        int4 wr;
        wr.x = (int)p0; wr.y = (int)p1; wr.z = (int)p2; wr.w = (int)p3;
        *reinterpret_cast<int4*>(&Alds[a_base + lds_swz(qr * LDS_ROW + qg * 16)]) = wr;
        Asclds[as_base + qr * 8 + qg] = asc;
    }

    if (ltid >= 128) {
        const int btid = ltid - 128;
        constexpr int BNT = 128;
        constexpr int BL = (BLOCK_N * LDS_ROW / 16 + BNT - 1) / BNT;
        for (int ld = 0; ld < BL; ld++) {
            const int f = (ld * BNT + btid) << 4;
            const int r = f >> 7;
            const int gn = bn + r;
            if (r < BLOCK_N && gn < N) {
                const int sf = lds_swz(f), sc = sf & 127, ac = kb + sc;
                llvm_amdgcn_raw_buffer_load_lds(srsrc,
                    (as3_uint32_ptr)(reinterpret_cast<uintptr_t>(Blds) + b_base + f),
                    16,
                    (gn >> 4) * (BSTRIDE << 4) + (ac >> 5) * 512 + ((ac >> 4) & 1) * 256 + (gn & 15) * 16,
                    0, 0, 0);
            }
        }
    }

    { const int so = group << 3;
    if (ltid >= 128) {
        const int stid = ltid - 128;
        const int row = stid & (BLOCK_N - 1), gb = (stid >> 6) << 2;
        const int gn = bn + row;
        if (gn < N) {
            const int rb = (gn >> 5) * 32 * SCS + (gn & 15) * 4 + ((gn >> 4) & 1);
            for (int g = 0; g < 4; g++) {
                const int grp = gb + g, ac = so + grp;
                Bslds[bs_base + row * 8 + grp] = Bsc[rb + (ac & 3) * 64 + ((ac & 7) >> 2) * 2 + (ac >> 3) * 256];
            }
        } else {
            for (int g = 0; g < 4; g++) Bslds[bs_base + row * 8 + gb + g] = 0x7f;
        }
    }}

    asm volatile("s_waitcnt vmcnt(0)");
    __syncthreads();

    const int br = (wid_local << 4) + lm;
    {
        int4_vec A0, B0; int as0, bs0;
        { const int o = lds_swz(lm * LDS_ROW + (lk << 4));
          const int4 t = *reinterpret_cast<const int4*>(&Alds[a_base + o]);
          A0 = {t.x, t.y, t.z, t.w}; }
        as0 = (int)Asclds[as_base + (lm << 3) + lk];
        { const int o = lds_swz(br * LDS_ROW + (lk << 4));
          const int4 t = *reinterpret_cast<const int4*>(&Blds[b_base + o]);
          B0 = {t.x, t.y, t.z, t.w}; }
        bs0 = (int)Bslds[bs_base + br * 8 + lk];
        acc = mfma_fp4(A0, B0, acc, as0, bs0);
    }
    {
        int4_vec A1, B1; int as1, bs1;
        { const int o = lds_swz(lm * LDS_ROW + HALF_K + (lk << 4));
          const int4 t = *reinterpret_cast<const int4*>(&Alds[a_base + o]);
          A1 = {t.x, t.y, t.z, t.w}; }
        as1 = (int)Asclds[as_base + (lm << 3) + 4 + lk];
        { const int o = lds_swz(br * LDS_ROW + HALF_K + (lk << 4));
          const int4 t = *reinterpret_cast<const int4*>(&Blds[b_base + o]);
          B1 = {t.x, t.y, t.z, t.w}; }
        bs1 = (int)Bslds[bs_base + br * 8 + 4 + lk];
        acc = mfma_fp4(A1, B1, acc, as1, bs1);
    }

    __syncthreads();
    float* reduce_buf = reinterpret_cast<float*>(Alds);

    if (group == 1) {
        const float* ap = reinterpret_cast<const float*>(&acc);
        for (int r = 0; r < 4; r++)
            reduce_buf[ltid * 4 + r] = ap[r];
    }
    __syncthreads();

    if (group == 0) {
        float* mp = reinterpret_cast<float*>(&acc);
        for (int r = 0; r < 4; r++)
            mp[r] += reduce_buf[ltid * 4 + r];

        const int or_ = bm + (lk << 2), oc = wn + lm;
        if (oc < N) {
            for (int r = 0; r < 4; r++) {
                int g = or_ + r;
                if (g < M) C[g * N + oc] = __float2bfloat16(mp[r]);
            }
        }
    }
}

void launch_sk2_k512(int64_t a, int64_t bq, int64_t bsc, int64_t c, int M, int N) {
    dim3 b(512), g((N + 63) / 64, (M + 15) / 16);
    hipLaunchKernelGGL(gemm_bn64_sk2_k512, g, b, 0, 0,
        reinterpret_cast<const __hip_bfloat16*>(a),
        reinterpret_cast<const uint8_t*>(bq),
        reinterpret_cast<const uint8_t*>(bsc),
        reinterpret_cast<__hip_bfloat16*>(c), M, N);
}
"""

from torch.utils.cpp_extension import load_inline
_hip_mod = load_inline(
    name="hip_sk2_k512",
    cpp_sources=_HIP_CPP,
    cuda_sources=_HIP_SRC,
    functions=["launch_sk2_k512"],
    verbose=False,
    extra_cuda_cflags=["-O3", "-std=c++17", "-fno-gpu-rdc", "-ffp-contract=fast",
                       "--offload-arch=gfx950", "-ffast-math",
                       "-funsafe-math-optimizations",
                       "-mllvm", "-amdgpu-max-memory-clause=64",
                       "-mllvm", "-amdgpu-load-store-vectorizer",
                       "-mllvm", "-amdgpu-early-ifcvt",
                       "-mllvm", "-amdgpu-early-inline-all",
                       "-mllvm", "-amdgpu-internalize-symbols",
                       "-mllvm", "-amdgpu-scalarize-global-loads",
                       "-mllvm", "-amdgpu-dpp-combine",
                       "-mllvm", "-amdgpu-enable-pre-ra-optimizations",
                       "-mllvm", "-amdgpu-promote-alloca-to-vector-limit=256"],
)
_hip_dispatch = _hip_mod.launch_sk2_k512


# ---- Triton MXFP4 GEMM kernels ----

@triton.jit
def _xcd_reorder(pid, GRID_SIZE, NUM_XCDS: tl.constexpr = 8):
    pids_per_xcd = (GRID_SIZE + NUM_XCDS - 1) // NUM_XCDS
    tall_xcds = GRID_SIZE % NUM_XCDS
    tall_xcds = tl.where(tall_xcds == 0, NUM_XCDS, tall_xcds)
    xcd = pid % NUM_XCDS
    local_pid = pid // NUM_XCDS
    new_pid = tl.where(
        xcd < tall_xcds,
        xcd * pids_per_xcd + local_pid,
        tall_xcds * pids_per_xcd + (xcd - tall_xcds) * (pids_per_xcd - 1) + local_pid,
    )
    return new_pid


@triton.jit
def _quantize_to_mxfp4(x_f32, BLOCK_M: tl.constexpr, NK_SC: tl.constexpr):
    SZ16: tl.constexpr = NK_SC * 16
    x_g = x_f32.reshape(BLOCK_M, NK_SC, 32)
    amax = tl.max(tl.abs(x_g), axis=2, keep_dims=True)
    amax_u32 = amax.to(tl.int32, bitcast=True)
    amax_u32 = ((amax_u32 + 0x200000).to(tl.uint32, bitcast=True)) & 0xFF800000
    exp_bits = (amax_u32 >> 23)
    raw_exp = tl.maximum(exp_bits.to(tl.int32) - 2, 0)
    a_scale = raw_exp.to(tl.uint8).reshape(BLOCK_M, NK_SC)
    sf_exp = tl.maximum(raw_exp, 1).to(tl.uint32)
    sf = (sf_exp << 23).to(tl.float32, bitcast=True)
    x_pairs = x_g.reshape(BLOCK_M, NK_SC, 16, 2)
    a_elems, b_elems = tl.split(x_pairs)
    sf_bc = tl.broadcast_to(sf, (BLOCK_M, NK_SC, 16))
    raw = tl.inline_asm_elementwise(
        "v_cvt_scalef32_pk_fp4_f32 $0, $1, $2, $3",
        "=v,v,v,v",
        [a_elems, b_elems, sf_bc],
        dtype=tl.uint32, is_pure=True, pack=1,
    )
    a_fp4 = (raw & 0xFF).to(tl.uint8).reshape(BLOCK_M, SZ16)
    return a_fp4, a_scale


# Split-K path: distributes K dimension across CTAs, then reduces

@triton.jit
def _splitk_fp4_gemm(
    a_ptr, b_ptr, c_ptr, b_sc_ptr,
    M, N, K, N16, N32,
    stride_am, stride_ak, stride_bn, stride_bk,
    stride_cm, stride_cn, stride_bsn, stride_bsk,
    BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr,
    BLOCK_K: tl.constexpr, NUM_KSPLIT: tl.constexpr,
):
    SG: tl.constexpr = 32
    NK_SC: tl.constexpr = BLOCK_K // SG
    tl.assume(stride_am > 0)
    tl.assume(stride_ak > 0)
    tl.assume(stride_bn > 0)
    tl.assume(stride_bk > 0)
    tl.assume(stride_cm > 0)
    tl.assume(stride_cn > 0)
    tl.assume(stride_bsn > 0)
    tl.assume(stride_bsk > 0)

    num_pid_n = tl.cdiv(N, BLOCK_N)
    num_pid_m = tl.cdiv(M, BLOCK_M)
    GRID_MN = num_pid_m * num_pid_n
    pid = tl.program_id(0)
    pid = _xcd_reorder(pid, GRID_MN * NUM_KSPLIT)
    pid_k = pid % NUM_KSPLIT
    pid_mn = pid // NUM_KSPLIT
    pid_m = pid_mn // num_pid_n
    pid_n = pid_mn % num_pid_n

    nk_total = K // BLOCK_K
    nk_base = nk_total // NUM_KSPLIT
    nk_rem = nk_total % NUM_KSPLIT
    my_nk = nk_base + tl.where(pid_k < nk_rem, 1, 0)
    k_start_iter = pid_k * nk_base + tl.where(pid_k < nk_rem, pid_k, nk_rem)
    k_offset = k_start_iter * BLOCK_K

    offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
    m_mask = offs_m < M
    offs_k_bf16 = tl.arange(0, BLOCK_K)
    a_ptrs = a_ptr + offs_m[:, None] * stride_am + (k_offset + offs_k_bf16[None, :]) * stride_ak
    offs_bn_sh = (pid_n * (BLOCK_N // 16) + tl.arange(0, BLOCK_N // 16)) % N16
    offs_k_sh = (k_offset // 2) * 16 + tl.arange(0, (BLOCK_K // 2) * 16)
    b_ptrs = b_ptr + offs_bn_sh[:, None].to(tl.int64) * stride_bn + offs_k_sh[None, :].to(tl.int64) * stride_bk
    offs_bn_sc = (pid_n * (BLOCK_N // 32) + tl.arange(0, BLOCK_N // 32)) % N32
    offs_ks_raw = k_offset + tl.arange(0, BLOCK_K)
    b_sc_ptrs = b_sc_ptr + offs_bn_sc[:, None].to(tl.int64) * stride_bsn + offs_ks_raw[None, :].to(tl.int64) * stride_bsk
    offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
    n_mask = offs_n < N

    acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
    for _ in range(my_nk):
        a_bf16 = tl.load(a_ptrs, mask=m_mask[:, None], other=0)
        b_sc_raw = tl.load(b_sc_ptrs, cache_modifier=".cg")
        b_raw = tl.load(b_ptrs, cache_modifier=".cg")
        a_fp4, a_sc = _quantize_to_mxfp4(a_bf16.to(tl.float32), BLOCK_M, NK_SC)
        b_sc = (b_sc_raw
                .reshape(BLOCK_N // 32, BLOCK_K // 256, 4, 16, 2, 2, 1)
                .permute(0, 5, 3, 1, 4, 2, 6)
                .reshape(BLOCK_N, NK_SC))
        b = (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))
        acc = tl.dot_scaled(a_fp4, a_sc, "e2m1", b, b_sc, "e2m1", acc)
        a_ptrs += BLOCK_K * stride_ak
        b_ptrs += (BLOCK_K // 2) * 16 * stride_bk
        b_sc_ptrs += BLOCK_K * stride_bsk

    c_ptrs = c_ptr + pid_k.to(tl.int64) * (M * N) + offs_m[:, None].to(tl.int64) * stride_cm + offs_n[None, :].to(tl.int64) * stride_cn
    c_mask = m_mask[:, None] & n_mask[None, :]
    if my_nk > 0:
        tl.store(c_ptrs, acc, mask=c_mask, cache_modifier=".wt")
    else:
        tl.store(c_ptrs, tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32), mask=c_mask, cache_modifier=".wt")


@triton.jit
def _reduce_partials(
    partial_ptr, c_ptr, total_elems,
    NUM_KSPLIT: tl.constexpr, BLOCK: tl.constexpr,
):
    pid = tl.program_id(0)
    offs = pid * BLOCK + tl.arange(0, BLOCK)
    mask = offs < total_elems
    acc = tl.zeros((BLOCK,), dtype=tl.float32)
    for ks in range(NUM_KSPLIT):
        p = tl.load(partial_ptr + ks * total_elems + offs, mask=mask, other=0.0)
        acc += p
    tl.store(c_ptr + offs, acc.to(tl.bfloat16), mask=mask)


# Autotuned FP4 matmul with XCD-aware tile scheduling

def _cfg(bm, bn, bk, gsm, xcds, w, s):
    return triton.Config(
        {'BLOCK_M': bm, 'BLOCK_N': bn, 'BLOCK_K': bk,
         'GROUP_SIZE_M': gsm, 'NUM_XCDS': xcds},
        num_warps=w, num_stages=s,
    )

_MATMUL_CONFIGS = [
    _cfg(16, 128, 256, 4, 8, 4, 2),
    _cfg(16, 128, 512, 4, 8, 4, 2),
    _cfg(16, 128, 256, 1, 8, 4, 2),
    _cfg(16, 128, 512, 1, 8, 4, 2),
    _cfg(16, 128, 1024, 4, 8, 4, 1),
    _cfg(16, 128, 1024, 4, 8, 4, 2),
    _cfg(16, 128, 1024, 1, 8, 4, 1),
    _cfg(16, 128, 256, 1, 8, 4, 1),
    _cfg(16, 128, 512, 1, 8, 4, 1),
    _cfg(16, 256, 256, 1, 8, 4, 2),
    _cfg(16, 256, 512, 1, 8, 4, 2),
    _cfg(16, 256, 256, 4, 8, 4, 2),
    _cfg(16, 128, 256, 8, 8, 4, 2),
    _cfg(16, 256, 256, 8, 8, 4, 2),
    _cfg(16, 256, 256, 16, 8, 4, 2),
    _cfg(16, 128, 256, 16, 8, 4, 2),
    _cfg(16, 256, 256, 8, 8, 8, 2),
    _cfg(16, 256, 256, 16, 8, 8, 2),
    _cfg(16, 128, 256, 8, 8, 8, 2),
    _cfg(16, 128, 512, 8, 8, 8, 2),
    _cfg(16, 128, 1024, 4, 8, 8, 1),
    _cfg(16, 256, 512, 8, 8, 8, 2),
    _cfg(32, 128, 256, 4, 8, 4, 2),
    _cfg(32, 128, 512, 4, 8, 4, 2),
    _cfg(32, 128, 256, 4, 8, 8, 2),
    _cfg(32, 128, 512, 4, 8, 8, 2),
    _cfg(32, 128, 1024, 4, 8, 8, 1),
    _cfg(32, 256, 256, 4, 8, 8, 2),
    _cfg(32, 256, 512, 4, 8, 8, 2),
    _cfg(32, 128, 256, 8, 8, 4, 2),
    _cfg(32, 128, 256, 8, 8, 8, 2),
    _cfg(32, 256, 256, 8, 8, 8, 2),
    _cfg(32, 256, 512, 8, 8, 8, 2),
    _cfg(64, 128, 256, 4, 8, 8, 2),
    _cfg(64, 128, 512, 4, 8, 8, 2),
    _cfg(64, 256, 256, 4, 8, 8, 2),
    _cfg(64, 128, 256, 8, 8, 8, 2),
    _cfg(16, 128, 512, 8, 8, 4, 2),
    _cfg(16, 256, 512, 4, 8, 4, 2),
    _cfg(16, 256, 512, 8, 8, 4, 2),
    _cfg(16, 128, 512, 8, 1, 4, 2),
    _cfg(16, 128, 512, 4, 1, 4, 2),
    _cfg(16, 256, 512, 8, 1, 8, 2),
    _cfg(16, 128, 256, 8, 1, 4, 2),
    _cfg(16, 128, 256, 4, 1, 4, 2),
    _cfg(32, 128, 512, 4, 1, 4, 2),
    _cfg(32, 128, 256, 8, 1, 8, 2),
    _cfg(64, 128, 256, 4, 1, 8, 2),
    _cfg(64, 128, 512, 4, 1, 8, 2),
    _cfg(16, 128, 512, 8, 4, 4, 2),
    _cfg(16, 256, 512, 8, 4, 8, 2),
    _cfg(32, 128, 256, 8, 4, 8, 2),
    _cfg(16, 128, 512, 8, 8, 4, 3),
    _cfg(16, 128, 512, 4, 8, 4, 3),
    _cfg(16, 256, 512, 8, 8, 8, 3),
    _cfg(16, 128, 256, 8, 8, 4, 3),
    _cfg(32, 128, 512, 4, 8, 4, 3),
    _cfg(32, 128, 256, 8, 8, 8, 3),
    _cfg(64, 128, 256, 4, 8, 8, 3),
    _cfg(16, 128, 512, 8, 1, 4, 3),
    _cfg(16, 256, 512, 8, 1, 8, 3),
    # Wider grid variants for better CU utilization at M=64
    _cfg(16, 64, 512, 8, 8, 4, 2),
    _cfg(16, 64, 512, 4, 8, 4, 2),
    _cfg(16, 64, 256, 8, 8, 4, 2),
    _cfg(16, 64, 256, 4, 8, 4, 2),
    _cfg(16, 64, 1024, 4, 8, 4, 1),
    _cfg(16, 64, 1024, 4, 8, 4, 2),
    _cfg(16, 64, 512, 8, 8, 2, 2),
    _cfg(16, 64, 256, 8, 8, 2, 2),
    _cfg(16, 64, 512, 8, 1, 4, 2),
    _cfg(16, 64, 512, 8, 8, 4, 3),
    # Larger BM tiles with narrow BN for high-M shapes
    _cfg(32, 64, 512, 8, 8, 4, 2),
    _cfg(32, 64, 512, 4, 8, 4, 2),
    _cfg(32, 64, 256, 8, 8, 4, 2),
    _cfg(32, 64, 512, 8, 8, 2, 2),
    _cfg(64, 64, 512, 4, 8, 4, 2),
    _cfg(64, 64, 256, 4, 8, 4, 2),
    _cfg(64, 64, 512, 4, 8, 8, 2),
    _cfg(64, 64, 256, 4, 8, 8, 2),
    _cfg(64, 64, 512, 8, 8, 8, 2),
    _cfg(32, 64, 256, 8, 8, 2, 2),
]


def _filter_valid_configs(configs, named_args, **kwargs):
    K = named_args['K']
    M = named_args['M']
    return [c for c in configs
            if K % c.kwargs['BLOCK_K'] == 0
            and (c.kwargs['BLOCK_M'] == 16 or c.kwargs['BLOCK_M'] <= M)]


@triton.autotune(configs=_MATMUL_CONFIGS, key=['M', 'N', 'K'],
                 prune_configs_by={'early_config_prune': _filter_valid_configs})
@triton.heuristics({
    'EVEN_M': lambda args: args['M'] % args['BLOCK_M'] == 0,
    'EVEN_N': lambda args: args['N'] % args['BLOCK_N'] == 0,
    'NUM_ITERS': lambda args: args['K'] // args['BLOCK_K'] if args['K'] >= 1024 else 0,
})
@triton.jit
def _fp4_matmul_kernel(
    a_ptr, b_ptr, c_ptr, b_sc_ptr,
    M, N, K, N16, N32,
    stride_am, stride_ak, stride_bn, stride_bk,
    stride_cm, stride_cn, stride_bsn, stride_bsk,
    BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr,
    BLOCK_K: tl.constexpr, GROUP_SIZE_M: tl.constexpr,
    NUM_XCDS: tl.constexpr,
    EVEN_M: tl.constexpr, EVEN_N: tl.constexpr,
    NUM_ITERS: tl.constexpr,
):
    SG: tl.constexpr = 32
    NK_SC: tl.constexpr = BLOCK_K // SG
    tl.assume(stride_am > 0)
    tl.assume(stride_ak > 0)
    tl.assume(stride_bn > 0)
    tl.assume(stride_bk > 0)
    tl.assume(stride_cm > 0)
    tl.assume(stride_cn > 0)
    tl.assume(stride_bsn > 0)
    tl.assume(stride_bsk > 0)

    pid = tl.program_id(0)
    num_pid_m = tl.cdiv(M, BLOCK_M)
    num_pid_n = tl.cdiv(N, BLOCK_N)
    GRID_MN = num_pid_m * num_pid_n
    pid = _xcd_reorder(pid, GRID_MN, NUM_XCDS)

    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 = tl.minimum(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)
    m_mask = offs_m < M
    offs_k_bf16 = tl.arange(0, BLOCK_K)
    a_ptrs = a_ptr + offs_m[:, None] * stride_am + offs_k_bf16[None, :] * stride_ak
    offs_bn_sh = (pid_n * (BLOCK_N // 16) + tl.arange(0, BLOCK_N // 16)) % N16
    offs_k_sh = tl.arange(0, (BLOCK_K // 2) * 16)
    b_ptrs = b_ptr + offs_bn_sh[:, None].to(tl.int64) * stride_bn + offs_k_sh[None, :].to(tl.int64) * stride_bk
    offs_bn_sc = (pid_n * (BLOCK_N // 32) + tl.arange(0, BLOCK_N // 32)) % N32
    offs_ks_raw = tl.arange(0, BLOCK_K)
    b_sc_ptrs = b_sc_ptr + offs_bn_sc[:, None].to(tl.int64) * stride_bsn + offs_ks_raw[None, :].to(tl.int64) * stride_bsk
    offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
    n_mask = offs_n < N

    acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
    niters = NUM_ITERS if NUM_ITERS > 0 else K // BLOCK_K
    for _ in range(niters):
        if EVEN_M:
            a_bf16 = tl.load(a_ptrs)
        else:
            a_bf16 = tl.load(a_ptrs, mask=m_mask[:, None], other=0)
        b_sc_raw = tl.load(b_sc_ptrs, cache_modifier=".cg")
        b_raw = tl.load(b_ptrs, cache_modifier=".cg")
        a_fp4, a_sc = _quantize_to_mxfp4(a_bf16.to(tl.float32), BLOCK_M, NK_SC)
        b_sc = (b_sc_raw
                .reshape(BLOCK_N // 32, BLOCK_K // 256, 4, 16, 2, 2, 1)
                .permute(0, 5, 3, 1, 4, 2, 6)
                .reshape(BLOCK_N, NK_SC))
        b = (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))
        acc = tl.dot_scaled(a_fp4, a_sc, "e2m1", b, b_sc, "e2m1", acc)
        a_ptrs += BLOCK_K * stride_ak
        b_ptrs += (BLOCK_K // 2) * 16 * stride_bk
        b_sc_ptrs += BLOCK_K * stride_bsk

    c = acc.to(tl.bfloat16)
    offs_cn = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
    c_ptrs = c_ptr + offs_m[:, None] * stride_cm + offs_cn[None, :] * stride_cn
    if EVEN_M and EVEN_N:
        tl.store(c_ptrs, c, cache_modifier=".wt")
    elif EVEN_N:
        tl.store(c_ptrs, c, mask=m_mask[:, None], cache_modifier=".wt")
    else:
        tl.store(c_ptrs, c, mask=m_mask[:, None] & n_mask[None, :], cache_modifier=".wt")


# ---- Entry point ----

_N_COMPUTE_UNITS = 304


def _run_splitk_path(A_in, B_sh, B_sc_raw, m, n, k, n_padded, dev):
    """Handle M<=16 large-K via split-K decomposition + partial reduction."""
    tile_m, tile_n, tile_k = 16, 128, 512
    num_k_tiles = k // tile_k
    mn_grid = triton.cdiv(m, tile_m) * triton.cdiv(n, tile_n)
    n_splits = min(num_k_tiles, max(1, _N_COMPUTE_UNITS // mn_grid))
    while num_k_tiles % n_splits != 0 and n_splits > 1:
        n_splits -= 1
    if n_splits < 2:
        return None
    total = m * n
    partials = torch.empty(n_splits * total, dtype=torch.float32, device=dev)
    out = torch.empty(m, n, dtype=torch.bfloat16, device=dev)
    _splitk_fp4_gemm[(n_splits * mn_grid,)](
        A_in, B_sh, partials, B_sc_raw,
        m, n, k, n // 16, n_padded // 32,
        A_in.stride(0), A_in.stride(1),
        B_sh.stride(0), B_sh.stride(1),
        n, 1,
        B_sc_raw.stride(0), B_sc_raw.stride(1),
        BLOCK_M=tile_m, BLOCK_N=tile_n, BLOCK_K=tile_k, NUM_KSPLIT=n_splits,
        num_warps=4, num_stages=2,
    )
    merge_blk = 256
    _reduce_partials[(triton.cdiv(total, merge_blk),)](
        partials, out, total,
        NUM_KSPLIT=n_splits, BLOCK=merge_blk,
        num_warps=4, num_stages=2,
    )
    return out


def custom_kernel(data: input_t) -> output_t:
    A_in, B, B_q, B_shuffle, B_scale_sh = data
    m, k = A_in.shape
    n = B.shape[0]
    dev = A_in.device

    # Fast path: small M with K=512 uses hand-tuned HIP assembly kernel
    if m <= 32 and k == 512:
        out = torch.empty(m, n, dtype=torch.bfloat16, device=dev)
        _hip_dispatch(A_in.data_ptr(), B_shuffle.data_ptr(), B_scale_sh.data_ptr(),
                      out.data_ptr(), m, n)
        return out

    # Prepare B operands for Triton path
    B_sh = B_shuffle.contiguous().view(torch.uint8).view(n // 16, (k // 2) * 16)
    n_pad = (n + 255) // 256 * 256
    B_sc = B_scale_sh.contiguous().view(torch.uint8).reshape(n_pad // 32, k)

    # Try split-K for skinny M with large K
    if m <= 16 and k >= 1024:
        result = _run_splitk_path(A_in, B_sh, B_sc, m, n, k, n_pad, dev)
        if result is not None:
            return result

    # General autotuned GEMM path
    out = torch.empty(m, n, dtype=torch.bfloat16, device=dev)
    grid_fn = lambda META: (triton.cdiv(m, META['BLOCK_M']) * triton.cdiv(n, META['BLOCK_N']),)
    _fp4_matmul_kernel[grid_fn](
        A_in, B_sh, out, B_sc,
        m, n, k, n // 16, n_pad // 32,
        A_in.stride(0), A_in.stride(1),
        B_sh.stride(0), B_sh.stride(1),
        out.stride(0), out.stride(1),
        B_sc.stride(0), B_sc.stride(1),
    )
    return out
scrolls · 683 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