Skip to content
KernelIndex
Search⌘K

submission 737364

divc13 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_v78.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-737364?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
7.88µs
#8 of 1143
2026-04-05

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:86df4e4ffab700bd425d04881481b1d5a6f4130f427c7de32feff462d983850e
license declaredunknown
license concludedunknown
authorsdivc13
imported2026-08-15

Techniques

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

autotunereturn triton.Config(
num-warps = 4num_warps=4, num_stages=2,
shared-memory__shared__ __align__(16) uint8_t Alds[2 * BLOCK_M * LDS_ROW];
split-kNo split-K changes (MY_NK constexpr causes regression).
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 = 64v78: v76 + BN=64 configs for M=64/M=256 (better CU utilization).
vector-width = int4int4 ra = reinterpret_cast<const int4*>(s)[0];

Kernel source

submission_v78.py676 lines
"""
v78: v76 + BN=64 configs for M=64/M=256 (better CU utilization).
  M=64: grid 224→448 (CU 74%→147%). M=256: grid 192→768 (CU 63%→253%).
  No split-K changes (MY_NK constexpr causes regression).
"""
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 kernels (v75) =====================

@triton.jit
def _remap_xcd(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 _quant_fp4_hw(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 GEMM (M=16) ---------------

@triton.jit
def _fused_splitk_gemm_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, 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 = _remap_xcd(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 = _quant_fp4_hw(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 _merge_splitk_kernel(
    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 GEMM (v75 configs) ---------------

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,
    )

_FUSED_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),
    # --- v78: BN=64 for M=64 (grid 224→448) ---
    _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),
    # --- v78: BN=64 w/ BM=32/64 for M=256 ---
    _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 _prune_fused_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=_FUSED_CONFIGS, key=['M', 'N', 'K'],
                 prune_configs_by={'early_config_prune': _prune_fused_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 _fused_gemm_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 = _remap_xcd(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 = _quant_fp4_hw(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")


# ===================== Host dispatch =====================

NUM_CUS = 304

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

    # HIP path for M<=32 K=512 (M=4, M=32 benchmark shapes)
    if m <= 32 and k == 512:
        C = torch.empty(m, n, dtype=torch.bfloat16, device=dev)
        _hip_dispatch(A_in.data_ptr(), B_shuffle.data_ptr(), B_scale_sh.data_ptr(),
                      C.data_ptr(), m, n)
        return C

    # Triton path for everything else (M=16 K=7168, M=64 K=2048, M=256 K=1536)
    B_sh = B_shuffle.contiguous().view(torch.uint8).view(n // 16, (k // 2) * 16)
    n_padded = (n + 255) // 256 * 256
    B_sc_raw = B_scale_sh.contiguous().view(torch.uint8).reshape(n_padded // 32, k)

    use_splitk = False
    if m <= 16 and k >= 1024:
        BM_sk, BN_sk, BK_sk = 16, 128, 512
        nk = k // BK_sk
        grid_mn = triton.cdiv(m, BM_sk) * triton.cdiv(n, BN_sk)
        if nk > 1:
            NUM_KSPLIT = min(nk, max(1, NUM_CUS // grid_mn))
            while nk % NUM_KSPLIT != 0 and NUM_KSPLIT > 1:
                NUM_KSPLIT -= 1
            if NUM_KSPLIT >= 2:
                use_splitk = True

    if use_splitk:
        total_elems = m * n
        C_parts = torch.empty(NUM_KSPLIT * total_elems, dtype=torch.float32, device=dev)
        C = torch.empty(m, n, dtype=torch.bfloat16, device=dev)
        _fused_splitk_gemm_kernel[(NUM_KSPLIT * grid_mn,)](
            A_in, B_sh, C_parts, 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=BM_sk, BLOCK_N=BN_sk, BLOCK_K=BK_sk, NUM_KSPLIT=NUM_KSPLIT,
            num_warps=4, num_stages=2,
        )
        MERGE_BLOCK = 256
        _merge_splitk_kernel[(triton.cdiv(total_elems, MERGE_BLOCK),)](
            C_parts, C, total_elems,
            NUM_KSPLIT=NUM_KSPLIT, BLOCK=MERGE_BLOCK,
            num_warps=4, num_stages=2,
        )
    else:
        C = torch.empty(m, n, dtype=torch.bfloat16, device=dev)
        grid = lambda META: (triton.cdiv(m, META['BLOCK_M']) * triton.cdiv(n, META['BLOCK_N']),)
        _fused_gemm_kernel[grid](
            A_in, B_sh, C, 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),
            C.stride(0), C.stride(1),
            B_sc_raw.stride(0), B_sc_raw.stride(1),
        )

    return C
scrolls · 676 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 572802.

"""
- Phase 60: Hybrid dispatch — sub59 B_scale + sub57 BLOCK_N=128 for large-M.
- Benchmarks 5/6 (M>=64) use BLOCK_N=128; benchmarks 1-4 (M<=32) use BLOCK_N=64.
- Both paths use sub59's 2*BLOCK_N-thread B_scale loading (4 loads each).
+ v78: v76 + BN=64 configs for M=64/M=256 (better CU utilization).
+ M=64: grid 224→448 (CU 74%→147%). M=256: grid 192→768 (CU 63%→253%).
+ No split-K changes (MY_NK constexpr causes regression).
"""
- from task import input_t, output_t
import torch
+ import triton
+ import triton.language as tl
+ from task import input_t, output_t
import os
os.environ["PYTORCH_ROCM_ARCH"] = "gfx950"
- CPP_SOURCE = r"""
+ # ===================== HIP kernel for M<=32 K=512 =====================
+
+ _HIP_CPP = r"""
#include <torch/extension.h>
- // BLOCK_N=64 (sub53) path
- void launch_n64_nosplit(
- torch::Tensor A_bf16, torch::Tensor B_q, torch::Tensor B_scale,
- torch::Tensor C, int M, int N, int K);
- void launch_n64_splitk(
- torch::Tensor A_bf16, torch::Tensor B_q, torch::Tensor B_scale,
- torch::Tensor workspace, int M, int N, int K, int split_k);
- // BLOCK_N=128 (sub57) path
- void launch_n128_nosplit(
- torch::Tensor A_bf16, torch::Tensor B_q, torch::Tensor B_scale,
- torch::Tensor C, int M, int N, int K);
- void launch_n128_splitk(
- torch::Tensor A_bf16, torch::Tensor B_q, torch::Tensor B_scale,
- torch::Tensor workspace, int M, int N, int K, int split_k);
- // Reduce
- void launch_reduce(torch::Tensor workspace, torch::Tensor C, int M, int N, int split_k);
+ void launch_sk2_k512(int64_t a, int64_t bq, int64_t bsc, int64_t c, int M, int N);
"""
- HIP_SOURCE = r"""
+ _HIP_SRC = r"""
#include <hip/hip_runtime.h>
#include <hip/hip_bf16.h>
#include <torch/extension.h>
- constexpr int WARP_SIZE = 64;
- constexpr int MFMA_K = 128;
- constexpr int DOUBLE_K = MFMA_K * 2;
- constexpr int LDS_ROW = DOUBLE_K >> 1;
- constexpr int HALF_K = MFMA_K >> 1;
- constexpr int SCALE_GROUP = 32;
- constexpr int BLOCK_M = 16;
-
+ 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");
+ 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* ptr, uint32_t range_bytes) {
- buffer_resource rsrc = {reinterpret_cast<uint64_t>(ptr), range_bytes, 0x110000};
- return *reinterpret_cast<const i32x4*>(&rsrc);
+ __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_scaled(
- int4_vec A, int4_vec B, float4_vec C, int sA, int sB
- ) {
+ __device__ __forceinline__ float4_vec mfma_fp4(int4_vec A, int4_vec B, float4_vec C, int sA, int sB) {
float4_vec D;
- asm volatile(
- "v_mfma_scale_f32_16x16x128_f8f6f4 %0, %1, %2, %3, %4, %5 cbsz:4 blgp:4"
+ 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 offset) {
- return offset ^ (((offset & 2047) >> 8) << 4);
+ __device__ __forceinline__ int lds_swz(int o) {
+ return o ^ (((o & 2047) >> 8) << 4);
}
- __device__ __forceinline__ void compute_scale(float max_abs, uint8_t& sc, float& scale_f) {
- if (max_abs > 0.0f) {
- uint32_t b = __float_as_uint(max_abs);
+ __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);
- scale_f = __uint_as_float((uint32_t)(su + 127) << 23);
- } else { sc = 0; scale_f = 0.0f; }
+ sf = __uint_as_float((uint32_t)(su + 127) << 23);
+ } else {
+ sc = 0; sf = 0.f;
+ }
}
- // ============================================================
- // Templated kernel — BLOCK_N and NUM_WARPS as template params
- // ============================================================
- template <int BLOCK_N, int NUM_WARPS>
- __global__ __launch_bounds__(NUM_WARPS * 64, (NUM_WARPS == 4 ? 3 : 2))
- void gemm_kernel(
- const __hip_bfloat16* __restrict__ A_bf16,
- const uint8_t* __restrict__ B_q,
- const uint8_t* __restrict__ B_scale,
- float* __restrict__ workspace,
- __hip_bfloat16* __restrict__ C_out,
- const int M, const int N, const int K,
- const int k_steps_per_split
- ) {
- constexpr int NUM_THREADS = NUM_WARPS * 64;
- const int warp_id = threadIdx.x >> 6;
- const int lane_id = threadIdx.x & 63;
- const int lane_m = lane_id & 15;
- const int lane_k = lane_id >> 4;
- const int tid = threadIdx.x;
+ __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));
+ }
- const int block_m = blockIdx.y * BLOCK_M;
- const int block_n = blockIdx.x * BLOCK_N;
- const int warp_n = block_n + (warp_id << 4);
- const int split_id = blockIdx.z;
+ __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 b_stride = K >> 1;
- const int sc_stride = K >> 5;
+ const int bn = blockIdx.x * BLOCK_N, bm = blockIdx.y * BLOCK_M;
+ const int wn = bn + (wid_local << 4);
- __shared__ __align__(16) uint8_t A_lds[BLOCK_M * LDS_ROW];
- __shared__ __align__(16) uint8_t B_lds[BLOCK_N * LDS_ROW];
- __shared__ uint8_t A_scale_lds[BLOCK_M * 8];
- __shared__ uint8_t B_scale_lds[BLOCK_N * 8];
+ __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 i32x4 b_srsrc = make_srsrc(B_q, N * b_stride);
- float4_vec acc = {0.0f, 0.0f, 0.0f, 0.0f};
+ 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 ks_start = split_id * k_steps_per_split;
- const int ks_end = ks_start + k_steps_per_split;
+ const int ke = group * DOUBLE_K;
+ const int kb = group * LDS_ROW;
- for (int ks = ks_start; ks < ks_end; ks++) {
- const int k_elem = ks * DOUBLE_K;
- const int k_byte = ks * LDS_ROW;
+ const i32x4 srsrc = make_srsrc(Bq, N * BSTRIDE);
+ float4_vec acc = {0, 0, 0, 0};
- // A quant: HW FP4 conversion (only 256 threads needed)
- if (tid < 256) {
- const int group_id = tid >> 1;
- const int half = tid & 1;
- const int q_row = group_id >> 3;
- const int q_grp = group_id & 7;
- const int g_row = block_m + q_row;
- const int k_off = k_elem + q_grp * SCALE_GROUP + half * 16;
-
- uint32_t pk_lo = 0, pk_hi = 0;
- uint8_t a_scale_val = 0x7f;
-
- if (g_row < M) {
- const __hip_bfloat16* src = A_bf16 + g_row * K + k_off;
- int4 raw[2];
- #pragma unroll
- for (int j = 0; j < 2; j++)
- raw[j] = reinterpret_cast<const int4*>(src)[j];
- const __hip_bfloat16* bf = reinterpret_cast<const __hip_bfloat16*>(raw);
- float vals[16];
- float local_max = 0.0f;
- #pragma unroll
- for (int i = 0; i < 16; i++) {
- vals[i] = __bfloat162float(bf[i]);
- local_max = fmaxf(local_max, fabsf(vals[i]));
- }
- float global_max = fmaxf(local_max, __shfl_xor(local_max, 1));
- float scale_f;
- compute_scale(global_max, a_scale_val, scale_f);
- pk_lo = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk_lo, vals[0], vals[1], scale_f, 0);
- pk_lo = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk_lo, vals[2], vals[3], scale_f, 1);
- pk_lo = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk_lo, vals[4], vals[5], scale_f, 2);
- pk_lo = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk_lo, vals[6], vals[7], scale_f, 3);
- pk_hi = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk_hi, vals[8], vals[9], scale_f, 0);
- pk_hi = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk_hi, vals[10], vals[11], scale_f, 1);
- pk_hi = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk_hi, vals[12], vals[13], scale_f, 2);
- pk_hi = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk_hi, vals[14], vals[15], scale_f, 3);
- }
- const int a_lds_off = lds_swz(q_row * LDS_ROW + q_grp * 16 + half * 8);
- int2 packed; packed.x = (int)pk_lo; packed.y = (int)pk_hi;
- *reinterpret_cast<int2*>(&A_lds[a_lds_off]) = packed;
- if (half == 0) A_scale_lds[q_row * 8 + q_grp] = a_scale_val;
+ 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;
+ }
- // B: buffer_load_lds from B_shuffle with swizzled global source
- {
- constexpr int B_LOADS = (BLOCK_N * LDS_ROW / 16 + NUM_THREADS - 1) / NUM_THREADS;
- #pragma unroll
- for (int ld = 0; ld < B_LOADS; ld++) {
- const int flat = (ld * NUM_THREADS + tid) << 4;
- const int row = flat >> 7;
- const int g_row = block_n + row;
- if (row < BLOCK_N && g_row < N) {
- const int swz_flat = lds_swz(flat);
- const int swz_col = swz_flat & 127;
- const int abs_col = k_byte + swz_col;
- const int tile_n = g_row >> 4;
- const int inner_n = g_row & 15;
- const int tile_k = abs_col >> 5;
- const int inner_k_hi = (abs_col >> 4) & 1;
- const int src_off = tile_n * (b_stride << 4) + tile_k * 512 + inner_k_hi * 256 + inner_n * 16;
-
- llvm_amdgcn_raw_buffer_load_lds(b_srsrc,
- (as3_uint32_ptr)(reinterpret_cast<uintptr_t>(B_lds) + flat),
- 16, src_off, 0, 0, 2); // aux=2: SLC (non-temporal)
- }
+ 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);
}
}
+ }
- // B_scale from B_scale_sh — sub59 style: 2*BLOCK_N threads, 4 loads each
- // For BLOCK_N=64: 128 threads active, for BLOCK_N=128: 256 threads active
- {
- const int sc_off = ks << 3;
- if (tid < BLOCK_N * 2) {
- const int row = tid & (BLOCK_N - 1); // 0..BLOCK_N-1
- constexpr int BN_SHIFT = (BLOCK_N == 64) ? 6 : 7;
- const int grp_base = (tid >> BN_SHIFT) << 2; // 0 or 4
- const int g_row = block_n + row;
-
- if (g_row < N) {
- const int row_base = (g_row >> 5) * 32 * sc_stride
- + (g_row & 15) * 4
- + ((g_row >> 4) & 1);
- #pragma unroll
- for (int g = 0; g < 4; g++) {
- const int grp = grp_base + g;
- const int abs_col = sc_off + grp;
- const int col_off = (abs_col & 3) * 64
- + ((abs_col & 7) >> 2) * 2
- + (abs_col >> 3) * 256;
- B_scale_lds[row * 8 + grp] = B_scale[row_base + col_off];
- }
- } else {
- #pragma unroll
- for (int g = 0; g < 4; g++)
- B_scale_lds[row * 8 + grp_base + g] = 0x7f;
- }
+ { 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();
+ asm volatile("s_waitcnt vmcnt(0)");
+ __syncthreads();
- // MFMA
- #pragma unroll
- for (int half = 0; half < 2; half++) {
- const int kh = half * HALF_K;
- int4_vec A_reg;
- {
- const int a_off = lds_swz(lane_m * LDS_ROW + kh + (lane_k << 4));
- const int4 tmp = *reinterpret_cast<const int4*>(&A_lds[a_off]);
- A_reg.s0 = tmp.x; A_reg.s1 = tmp.y; A_reg.s2 = tmp.z; A_reg.s3 = tmp.w;
- }
- int4_vec B_reg;
- {
- const int b_row = (warp_id << 4) + lane_m;
- const int b_off = lds_swz(b_row * LDS_ROW + kh + (lane_k << 4));
- const int4 tmp = *reinterpret_cast<const int4*>(&B_lds[b_off]);
- B_reg.s0 = tmp.x; B_reg.s1 = tmp.y; B_reg.s2 = tmp.z; B_reg.s3 = tmp.w;
- }
- const int a_sc = (int)A_scale_lds[(lane_m << 3) + (half << 2) + lane_k];
- const int b_sc = (int)B_scale_lds[((warp_id << 4) + lane_m) * 8 + (half << 2) + lane_k];
- acc = mfma_fp4_scaled(A_reg, B_reg, acc, a_sc, b_sc);
- }
- __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);
+ }
- // Store
- const int out_row = block_m + (lane_k << 2);
- const int out_col = warp_n + lane_m;
- if (out_col < N) {
+ __syncthreads();
+ float* reduce_buf = reinterpret_cast<float*>(Alds);
+
+ if (group == 1) {
const float* ap = reinterpret_cast<const float*>(&acc);
- if (C_out) {
- #pragma unroll
+ 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++) {
- const int gm = out_row + r;
- if (gm < M) C_out[gm * N + out_col] = __float2bfloat16(ap[r]);
+ int g = or_ + r;
+ if (g < M) C[g * N + oc] = __float2bfloat16(mp[r]);
}
- } else {
- float* ws = workspace + split_id * M * N;
- #pragma unroll
- for (int r = 0; r < 4; r++) {
- const int gm = out_row + r;
- if (gm < M) ws[gm * N + out_col] = ap[r];
- }
}
}
}
- // Reduction kernel
- __global__ void reduce_kernel(
- const float* __restrict__ workspace,
- __hip_bfloat16* __restrict__ C,
- const int M, const int N, const int split_k
- ) {
- const int idx = blockIdx.x * blockDim.x + threadIdx.x;
- if (idx >= M * N) return;
- float sum = 0.0f;
- for (int s = 0; s < split_k; s++)
- sum += workspace[s * M * N + idx];
- C[idx] = __float2bfloat16(sum);
+ 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);
}
+ """
- // ---- Launch functions for BLOCK_N=64 (4 warps, 256 threads) ----
- void launch_n64_nosplit(
- torch::Tensor A_bf16, torch::Tensor B_q, torch::Tensor B_scale,
- torch::Tensor C, int M, int N, int K
- ) {
- const int k_steps = K / (MFMA_K * 2);
- dim3 block(256);
- dim3 grid((N + 63) / 64, (M + BLOCK_M - 1) / BLOCK_M, 1);
- hipLaunchKernelGGL((gemm_kernel<64, 4>), grid, block, 0, 0,
- reinterpret_cast<const __hip_bfloat16*>(A_bf16.data_ptr()),
- reinterpret_cast<const uint8_t*>(B_q.data_ptr()),
- reinterpret_cast<const uint8_t*>(B_scale.data_ptr()),
- (float*)nullptr,
- reinterpret_cast<__hip_bfloat16*>(C.data_ptr()), M, N, K, k_steps);
- }
+ 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
- void launch_n64_splitk(
- torch::Tensor A_bf16, torch::Tensor B_q, torch::Tensor B_scale,
- torch::Tensor workspace, int M, int N, int K, int split_k
- ) {
- const int k_steps = K / (MFMA_K * 2);
- dim3 block(256);
- dim3 grid((N + 63) / 64, (M + BLOCK_M - 1) / BLOCK_M, split_k);
- hipLaunchKernelGGL((gemm_kernel<64, 4>), grid, block, 0, 0,
- reinterpret_cast<const __hip_bfloat16*>(A_bf16.data_ptr()),
- reinterpret_cast<const uint8_t*>(B_q.data_ptr()),
- reinterpret_cast<const uint8_t*>(B_scale.data_ptr()),
- reinterpret_cast<float*>(workspace.data_ptr()),
- (__hip_bfloat16*)nullptr, M, N, K, k_steps / split_k);
- }
- // ---- Launch functions for BLOCK_N=128 (8 warps, 512 threads) ----
- void launch_n128_nosplit(
- torch::Tensor A_bf16, torch::Tensor B_q, torch::Tensor B_scale,
- torch::Tensor C, int M, int N, int K
- ) {
- const int k_steps = K / (MFMA_K * 2);
- dim3 block(512);
- dim3 grid((N + 127) / 128, (M + BLOCK_M - 1) / BLOCK_M, 1);
- hipLaunchKernelGGL((gemm_kernel<128, 8>), grid, block, 0, 0,
- reinterpret_cast<const __hip_bfloat16*>(A_bf16.data_ptr()),
- reinterpret_cast<const uint8_t*>(B_q.data_ptr()),
- reinterpret_cast<const uint8_t*>(B_scale.data_ptr()),
- (float*)nullptr,
- reinterpret_cast<__hip_bfloat16*>(C.data_ptr()), M, N, K, k_steps);
- }
+ # ===================== Triton kernels (v75) =====================
- void launch_n128_splitk(
- torch::Tensor A_bf16, torch::Tensor B_q, torch::Tensor B_scale,
- torch::Tensor workspace, int M, int N, int K, int split_k
- ) {
- const int k_steps = K / (MFMA_K * 2);
- dim3 block(512);
- dim3 grid((N + 127) / 128, (M + BLOCK_M - 1) / BLOCK_M, split_k);
- hipLaunchKernelGGL((gemm_kernel<128, 8>), grid, block, 0, 0,
- reinterpret_cast<const __hip_bfloat16*>(A_bf16.data_ptr()),
- reinterpret_cast<const uint8_t*>(B_q.data_ptr()),
- reinterpret_cast<const uint8_t*>(B_scale.data_ptr()),
- reinterpret_cast<float*>(workspace.data_ptr()),
- (__hip_bfloat16*)nullptr, M, N, K, k_steps / split_k);
- }
+ @triton.jit
+ def _remap_xcd(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
- void launch_reduce(torch::Tensor workspace, torch::Tensor C, int M, int N, int split_k) {
- const int num = M * N;
- hipLaunchKernelGGL(reduce_kernel, dim3((num+255)/256), dim3(256), 0, 0,
- reinterpret_cast<const float*>(workspace.data_ptr()),
- reinterpret_cast<__hip_bfloat16*>(C.data_ptr()), M, N, split_k);
- }
- """
- from torch.utils.cpp_extension import load_inline
+ @triton.jit
+ def _quant_fp4_hw(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
- _module = None
- def _get_module():
- global _module
- if _module is None:
- _module = load_inline(
- name="hybrid_v74",
- cpp_sources=CPP_SOURCE,
- cuda_sources=HIP_SOURCE,
- functions=[
- "launch_n64_nosplit", "launch_n64_splitk",
- "launch_n128_nosplit", "launch_n128_splitk",
- "launch_reduce",
- ],
- verbose=False,
- extra_cuda_cflags=["-O3", "-fno-gpu-rdc", "-ffp-contract=fast"],
- )
- return _module
+ # --------------- Split-K GEMM (M=16) ---------------
- def _pick_split_k(m, n, k, block_n):
- k_steps = k // 256
- blocks_mn = ((n + block_n - 1) // block_n) * ((m + 15) // 16)
- if blocks_mn >= 304:
- return 1
- # Target ~912 total blocks (304 CUs × 3 blocks/CU)
- target_split = max(1, (912 + blocks_mn - 1) // blocks_mn)
- best = 1
- for s in range(1, k_steps + 1):
- if k_steps % s == 0 and s <= target_split:
- best = s
- while best > 1 and k_steps // best < 2:
- best //= 2
- return max(1, best)
+ @triton.jit
+ def _fused_splitk_gemm_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, 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 = _remap_xcd(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 = _quant_fp4_hw(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 _merge_splitk_kernel(
+ 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 GEMM (v75 configs) ---------------
+
+ 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,
+ )
+
+ _FUSED_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),
+ # --- v78: BN=64 for M=64 (grid 224→448) ---
+ _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),
+ # --- v78: BN=64 w/ BM=32/64 for M=256 ---
+ _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 _prune_fused_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=_FUSED_CONFIGS, key=['M', 'N', 'K'],
+ prune_configs_by={'early_config_prune': _prune_fused_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 _fused_gemm_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 = _remap_xcd(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 = _quant_fp4_hw(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")
+
+
+ # ===================== Host dispatch =====================
+
+ NUM_CUS = 304
+
def custom_kernel(data: input_t) -> output_t:
- A, B, B_q, B_shuffle, B_scale_sh = data
- A = A.contiguous()
- m, k = A.shape
+ A_in, B, B_q, B_shuffle, B_scale_sh = data
+ m, k = A_in.shape
n = B.shape[0]
+ dev = A_in.device
- mod = _get_module()
+ # HIP path for M<=32 K=512 (M=4, M=32 benchmark shapes)
+ if m <= 32 and k == 512:
+ C = torch.empty(m, n, dtype=torch.bfloat16, device=dev)
+ _hip_dispatch(A_in.data_ptr(), B_shuffle.data_ptr(), B_scale_sh.data_ptr(),
+ C.data_ptr(), m, n)
+ return C
- B_sh_u8 = B_shuffle.contiguous().view(torch.uint8)
- B_sc = B_scale_sh.contiguous().view(torch.uint8)
+ # Triton path for everything else (M=16 K=7168, M=64 K=2048, M=256 K=1536)
+ B_sh = B_shuffle.contiguous().view(torch.uint8).view(n // 16, (k // 2) * 16)
+ n_padded = (n + 255) // 256 * 256
+ B_sc_raw = B_scale_sh.contiguous().view(torch.uint8).reshape(n_padded // 32, k)
- # Dispatch: use BLOCK_N=128 for large-M benchmarks (M>=64)
- use_n128 = (m >= 64)
+ use_splitk = False
+ if m <= 16 and k >= 1024:
+ BM_sk, BN_sk, BK_sk = 16, 128, 512
+ nk = k // BK_sk
+ grid_mn = triton.cdiv(m, BM_sk) * triton.cdiv(n, BN_sk)
+ if nk > 1:
+ NUM_KSPLIT = min(nk, max(1, NUM_CUS // grid_mn))
+ while nk % NUM_KSPLIT != 0 and NUM_KSPLIT > 1:
+ NUM_KSPLIT -= 1
+ if NUM_KSPLIT >= 2:
+ use_splitk = True
- if use_n128:
- block_n = 128
- split_k = _pick_split_k(m, n, k, block_n)
- if split_k == 1:
- C = torch.empty((m, n), dtype=torch.bfloat16, device="cuda")
- mod.launch_n128_nosplit(A, B_sh_u8, B_sc, C, m, n, k)
- else:
- workspace = torch.empty((split_k, m, n), dtype=torch.float32, device="cuda")
- mod.launch_n128_splitk(A, B_sh_u8, B_sc, workspace, m, n, k, split_k)
- C = torch.empty((m, n), dtype=torch.bfloat16, device="cuda")
- mod.launch_reduce(workspace, C, m, n, split_k)
+ if use_splitk:
+ total_elems = m * n
+ C_parts = torch.empty(NUM_KSPLIT * total_elems, dtype=torch.float32, device=dev)
+ C = torch.empty(m, n, dtype=torch.bfloat16, device=dev)
+ _fused_splitk_gemm_kernel[(NUM_KSPLIT * grid_mn,)](
+ A_in, B_sh, C_parts, 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=BM_sk, BLOCK_N=BN_sk, BLOCK_K=BK_sk, NUM_KSPLIT=NUM_KSPLIT,
+ num_warps=4, num_stages=2,
+ )
+ MERGE_BLOCK = 256
+ _merge_splitk_kernel[(triton.cdiv(total_elems, MERGE_BLOCK),)](
+ C_parts, C, total_elems,
+ NUM_KSPLIT=NUM_KSPLIT, BLOCK=MERGE_BLOCK,
+ num_warps=4, num_stages=2,
+ )
else:
- block_n = 64
- split_k = _pick_split_k(m, n, k, block_n)
- if split_k == 1:
- C = torch.empty((m, n), dtype=torch.bfloat16, device="cuda")
- mod.launch_n64_nosplit(A, B_sh_u8, B_sc, C, m, n, k)
- else:
- workspace = torch.empty((split_k, m, n), dtype=torch.float32, device="cuda")
- mod.launch_n64_splitk(A, B_sh_u8, B_sc, workspace, m, n, k, split_k)
- C = torch.empty((m, n), dtype=torch.bfloat16, device="cuda")
- mod.launch_reduce(workspace, C, m, n, split_k)
+ C = torch.empty(m, n, dtype=torch.bfloat16, device=dev)
+ grid = lambda META: (triton.cdiv(m, META['BLOCK_M']) * triton.cdiv(n, META['BLOCK_N']),)
+ _fused_gemm_kernel[grid](
+ A_in, B_sh, C, 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),
+ C.stride(0), C.stride(1),
+ B_sc_raw.stride(0), B_sc_raw.stride(1),
+ )
return C
scrolls · 1030 diff lines total

Best evidence level for this revision: reported

JSON