Skip to content
KernelIndex
Search⌘K

submission 570151

div22 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

solution_new_25d_v4_i1.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-570151?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
12.2µs
#371 of 1143
2026-03-16

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:6c5d1546dfa31a473b593780fef6655f77bd3be8bf87abcbfcbddb21f70ed409
license declaredunknown
license concludedunknown
authorsdiv22
imported2026-08-15

Techniques

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

fp4MXFP4 GEMM v25d_v4_i1 — gfx950 (MI355X) with LDS + software pipelining.
shared-memoryextern __shared__ uint8_t smem_raw[];
split-ktemplate<int BM, int BN, int NWARPS, bool SPLITK, bool A_VALID, bool B_VALID,
tile-m = 0constexpr auto full_m = (BM >= 16) ? C_M / BM : 0;
tile-n = 32constexpr auto BN = 32;

Kernel source

solution_new_25d_v4_i1.py905 lines
"""
MXFP4 GEMM v25d_v4_i1 — gfx950 (MI355X) with LDS + software pipelining.

Based on v4_i0 (12.264μs). Changes:
  v4_i1: Quant intrinsic dest_sel packing (no shift/OR/mask).
         LdsLayout forceinline constexpr const auto.
         launch_bounds tuning for quant kernel.
"""
import os
os.environ["PYTORCH_ROCM_ARCH"] = "gfx950"

from typing import Tuple
import torch
from torch.utils.cpp_extension import load_inline
import uuid


HIP_KERNEL = r"""
#include <hip/hip_runtime.h>
#include <stdint.h>

using int4_v   = int   __attribute__((ext_vector_type(4)));
using float4_v = float __attribute__((ext_vector_type(4)));
using bf16x2   = __bf16 __attribute__((ext_vector_type(2)));

static constexpr auto FP4_E2M1 = 4;

__device__ float4_v __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
    int4_v a, int4_v b, float4_v c,
    int cbsz, int blgp, int op_sel_a, int scale_a, int op_sel_b, int scale_b
) __asm("llvm.amdgcn.mfma.scale.f32.16x16x128.f8f6f4.v4i32.v4i32");

__device__ __forceinline__ auto load16(const uint8_t* __restrict__ p) {
    return *reinterpret_cast<const int4_v*>(p);
}

// Non-temporal load: bypass L1, keep in L2.
__device__ __forceinline__ auto load16_nt(const uint8_t* __restrict__ p) {
    return __builtin_nontemporal_load(reinterpret_cast<const int4_v*>(p));
}

__device__ __forceinline__ auto float_to_bf16(float f) {
    bf16x2 v;
    v[0] = static_cast<__bf16>(f);
    auto r = uint16_t{};
    __builtin_memcpy(&r, &v, sizeof(r));
    return r;
}

// ── Quant kernel — fully templated on M, K ──────────────────────────────────

template<int C_M, int C_K>
__global__ void __launch_bounds__(128, (C_M * (C_K / 32) <= 256) ? 2 : 4)
mxfp4_quant(
    const __bf16* __restrict__ A_bf16,
    uint8_t*      __restrict__ A_fp4,
    uint8_t*      __restrict__ A_scale)
{
    constexpr auto KS = C_K / 32;
    constexpr auto K2 = C_K / 2;
    const auto group = blockIdx.x * 128 + threadIdx.x;
    const auto row   = group / KS;
    const auto kg    = group % KS;

    if (row >= C_M) return;

    const auto row_k = (long)row * C_K;
    const auto* src = A_bf16 + row_k + kg * 32;

    // 4 wide loads: 64 bytes in 4 × global_load_dwordx4
    const auto w0 = *reinterpret_cast<const int4_v*>(src);
    const auto w1 = *reinterpret_cast<const int4_v*>(src + 8);
    const auto w2 = *reinterpret_cast<const int4_v*>(src + 16);
    const auto w3 = *reinterpret_cast<const int4_v*>(src + 24);

    // Reinterpret as uint32 pairs (zero-cost, same registers)
    const auto* p0 = reinterpret_cast<const uint32_t*>(&w0);
    const auto* p1 = reinterpret_cast<const uint32_t*>(&w1);
    const auto* p2 = reinterpret_cast<const uint32_t*>(&w2);
    const auto* p3 = reinterpret_cast<const uint32_t*>(&w3);

    // absmax across all 32 bf16 values (16 uint32 pairs, each = 2 bf16)
    auto absMax = 1e-10f;
    #define AMAX_PAIR(pair) { \
        const auto lo = __builtin_bit_cast(__bf16, static_cast<uint16_t>(pair)); \
        const auto hi = __builtin_bit_cast(__bf16, static_cast<uint16_t>((pair) >> 16)); \
        const auto flo = __builtin_elementwise_abs(static_cast<float>(lo)); \
        const auto fhi = __builtin_elementwise_abs(static_cast<float>(hi)); \
        absMax = (flo > absMax) ? flo : absMax; \
        absMax = (fhi > absMax) ? fhi : absMax; \
    }
    AMAX_PAIR(p0[0]) AMAX_PAIR(p0[1]) AMAX_PAIR(p0[2]) AMAX_PAIR(p0[3])
    AMAX_PAIR(p1[0]) AMAX_PAIR(p1[1]) AMAX_PAIR(p1[2]) AMAX_PAIR(p1[3])
    AMAX_PAIR(p2[0]) AMAX_PAIR(p2[1]) AMAX_PAIR(p2[2]) AMAX_PAIR(p2[3])
    AMAX_PAIR(p3[0]) AMAX_PAIR(p3[1]) AMAX_PAIR(p3[2]) AMAX_PAIR(p3[3])
    #undef AMAX_PAIR

    const auto u32 = __builtin_bit_cast(uint32_t, absMax);
    const auto amax_exp = ((u32 + 0x200000u) >> 23) & 0xFFu;
    const auto inv_exp  = (amax_exp >= 2u) ? (amax_exp - 2u) : 0u;
    A_scale[row_k / 32 + kg] = static_cast<uint8_t>(inv_exp);

    const auto hw_scale = __builtin_bit_cast(float, static_cast<uint32_t>(inv_exp) << 23);

    // Pack 16 FP4 pairs into 4 dwords using intrinsic dest_sel (no shift/OR/mask)
    // Each __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(old, v2bf16, scale, byte_sel)
    // writes 1 byte at position byte_sel in the accumulator dword.
    #define CVT(d, pair, sel) \
        d = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16( \
            d, __builtin_bit_cast(bf16x2, (pair)), hw_scale, sel)

    auto d0 = 0u, d1 = 0u, d2 = 0u, d3 = 0u;
    CVT(d0, p0[0], 0); CVT(d0, p0[1], 1); CVT(d0, p0[2], 2); CVT(d0, p0[3], 3);
    CVT(d1, p1[0], 0); CVT(d1, p1[1], 1); CVT(d1, p1[2], 2); CVT(d1, p1[3], 3);
    CVT(d2, p2[0], 0); CVT(d2, p2[1], 1); CVT(d2, p2[2], 2); CVT(d2, p2[3], 3);
    CVT(d3, p3[0], 0); CVT(d3, p3[1], 1); CVT(d3, p3[2], 2); CVT(d3, p3[3], 3);
    #undef CVT

    auto* dst = reinterpret_cast<int4_v*>(A_fp4 + row_k / 2 + kg * 16);
    *dst = int4_v{static_cast<int>(d0), static_cast<int>(d1),
                  static_cast<int>(d2), static_cast<int>(d3)};
}

// ── Boundary-checked loads ──────────────────────────────────────────────────

template<bool ALWAYS_VALID>
__device__ __forceinline__ auto load_or_zero(bool rt_valid, const uint8_t* p) {
    if constexpr (ALWAYS_VALID) return load16(p);
    else                        return rt_valid ? load16(p) : int4_v{0,0,0,0};
}

// ── LDS layout for software-pipelined path (CKT > 4) ───────────────────────
//
// Double-buffered. Per buffer:
//   A: [WAVES_M][16 rows][CHUNK_K * 4 kgrps * 16 bytes + 16 pad]
//   B: [WAVES_N][CHUNK_K][4 kgrps + 1 pad][16 lrows][16 bytes]
//
// Bank conflict strategy:
//   A: XOR swizzle on row index — lrow ^ (k_idx & 7)
//   B: XOR swizzle on kgrp — kgrp ^ (lrow >> 2)

template<int WAVES_M, int WAVES_N, int CHUNK_K>
struct LdsLayout {
    static constexpr auto A_ROW   = CHUNK_K * 64 + 16;   // +16B pad per row
    static constexpr auto A_TILE  = 16 * A_ROW;
    static constexpr auto A_SIZE  = WAVES_M * A_TILE;

    // B: 5 kgrp slots (4 real + 1 pad) to space kgrps 4-banks apart
    static constexpr auto B_KGRP_STRIDE = 16 * 16;        // 16 lrows × 16B = 256B
    static constexpr auto B_KT_STRIDE   = 5 * B_KGRP_STRIDE;  // 5 slots (4+1 pad)
    static constexpr auto B_TILE  = CHUNK_K * B_KT_STRIDE;
    static constexpr auto B_SIZE  = WAVES_N * B_TILE;

    static constexpr auto BUF_SIZE  = A_SIZE + B_SIZE;
    static constexpr auto TOTAL_LDS = 2 * BUF_SIZE;

    __device__ __forceinline__ static constexpr auto a_off(uint8_t* const buf) { return buf; }
    __device__ __forceinline__ static constexpr auto b_off(uint8_t* const buf) { return buf + A_SIZE; }

    __device__ __forceinline__ static constexpr auto a_idx(const auto wm, const auto row, const auto k_idx) {
        return wm * A_TILE + (row ^ (k_idx & 7)) * A_ROW + k_idx * 16;
    }

    __device__ __forceinline__ static constexpr auto b_idx(const auto wn, const auto kt, const auto kgrp, const auto lrow) {
        return wn * B_TILE + kt * B_KT_STRIDE
             + (kgrp ^ (lrow >> 2)) * B_KGRP_STRIDE + lrow * 16;
    }
};

// ── Per-thread load item counts (constexpr) ─────────────────────────────────

template<int WAVES_M, int WAVES_N, int CHUNK_K, int NWARPS>
struct LoadCounts {
    static constexpr auto NTHREADS = NWARPS * 64;
    static constexpr auto A_TOTAL = WAVES_M * 16 * CHUNK_K * 4;
    static constexpr auto B_TOTAL = WAVES_N * CHUNK_K * 4 * 16;
    static constexpr auto A_PER_THREAD = (A_TOTAL + NTHREADS - 1) / NTHREADS;
    static constexpr auto B_PER_THREAD = (B_TOTAL + NTHREADS - 1) / NTHREADS;
};

// ── Main GEMM kernel ────────────────────────────────────────────────────────

template<int BM, int BN, int NWARPS, bool SPLITK, bool A_VALID, bool B_VALID,
         int CKT, int C_N, int C_K, int C_SCALEN, int C_TOTAL_KT, int C_KPS,
         int C_M, int C_TILE_OFF_X, int C_TILE_OFF_Y>
__global__ void __launch_bounds__(NWARPS * 64, (NWARPS <= 2) ? 4 : 2)
mxfp4_gemm(
    const uint8_t* __restrict__ A,
    const uint8_t* __restrict__ As,
    const uint8_t* __restrict__ Bsh,
    const uint8_t* __restrict__ Bssh,
    float*         __restrict__ C_partial,
    uint16_t*      __restrict__ C_final)
{
    static_assert(BN % 16 == 0);
    constexpr auto WAVES_M = (BM + 15) / 16;
    constexpr auto WAVES_N = BN / 16;
    static_assert(WAVES_M * WAVES_N == NWARPS);

    constexpr auto M = C_M;
    constexpr auto N = C_N;
    constexpr auto K = C_K;
    constexpr auto scaleN = C_SCALEN;
    constexpr auto K2 = K / 2;
    constexpr auto KS = K / 32;
    constexpr auto bsh_n_stride = (long)(K / 64) * 512;

    const auto ks_idx = static_cast<int>(blockIdx.z);
    const auto lane   = static_cast<int>(threadIdx.x % 64);
    const auto wave   = static_cast<int>(threadIdx.x / 64);
    const auto wave_m = wave / WAVES_N;
    const auto wave_n = wave % WAVES_N;
    const auto tid    = static_cast<int>(threadIdx.x);

    const auto tile_m_base = (static_cast<int>(blockIdx.y) + C_TILE_OFF_Y) * BM;
    const auto tile_n_base = (static_cast<int>(blockIdx.x) + C_TILE_OFF_X) * BN;
    const auto tile_m = tile_m_base + wave_m * 16;
    const auto tile_n = tile_n_base + wave_n * 16;

    constexpr auto ktiles_per_split = C_KPS;
    const auto ks_start = ks_idx * ktiles_per_split;

    const auto lrow = lane % 16;
    const auto kgrp = lane / 16;

    const auto gm = tile_m + lrow;
    const auto gn = tile_n + lrow;

    const auto a_rt = A_VALID | (gm < M);
    const auto b_rt = B_VALID | (gn < N);

    // Scale pointers (loaded from global, 1 byte, L1 cached)
    const uint8_t* as_row = nullptr;
    if constexpr (A_VALID) {
        as_row = As + (long)gm * KS;
    } else {
        if (a_rt) as_row = As + (long)gm * KS;
    }

    auto bssh_base = 0;
    if constexpr (B_VALID) {
        bssh_base = ((gn >> 4) & 1) + (gn & 15) * 4 + kgrp * 64 + (gn >> 5) * (32 * scaleN);
    } else {
        if (b_rt) bssh_base = ((gn >> 4) & 1) + (gn & 15) * 4 + kgrp * 64 + (gn >> 5) * (32 * scaleN);
    }

    float4_v acc{0.f, 0.f, 0.f, 0.f};

    // ════════════════════════════════════════════════════════════════════════
    // PATH A: CKT <= 4 — Direct global loads, no LDS
    // ════════════════════════════════════════════════════════════════════════
    if constexpr (CKT <= 4) {
        if (tile_m >= M || tile_n >= N) return;

        const auto n_tile = tile_n / 16;
        const auto bsh_lane_base = Bsh + (long)n_tile * bsh_n_stride + (long)lrow * 16;

        const auto k_half_off = (kgrp & 1) * 256;
        const auto k_blk_base = kgrp >> 1;

        constexpr auto a_kt_stride   = 64L;
        constexpr auto bsh_kt_stride = 1024L;

        const uint8_t* a_ptr    = nullptr;
        const uint8_t* bsh_ptr  = nullptr;
        const uint8_t* bssh_ptr = nullptr;

        if constexpr (A_VALID) {
            a_ptr = A + (long)gm * K2 + (long)ks_start * a_kt_stride + kgrp * 16;
        } else {
            if (a_rt) a_ptr = A + (long)gm * K2 + (long)ks_start * a_kt_stride + kgrp * 16;
        }
        if constexpr (B_VALID) {
            bsh_ptr  = bsh_lane_base + (long)(ks_start * 2 + k_blk_base) * 512 + k_half_off;
            bssh_ptr = Bssh + bssh_base + (ks_start & 1) * 2 + (ks_start >> 1) * 256;
        } else {
            if (b_rt) {
                bsh_ptr  = bsh_lane_base + (long)(ks_start * 2 + k_blk_base) * 512 + k_half_off;
                bssh_ptr = Bssh + bssh_base + (ks_start & 1) * 2 + (ks_start >> 1) * 256;
            }
        }

        const auto bssh_step0 = (ks_start & 1) ? 254 : 2;
        const auto bssh_step1 = 256 - bssh_step0;

        // PATH A: A from regular load (small, L1 reuse), B from .cg (large, L1 bypass)
        #define DO_MFMA(a_off, b_off, bssh_off, ks_val) \
        { \
            const auto av = load_or_zero<A_VALID>(a_rt, a_ptr + (a_off) * a_kt_stride); \
            const auto bv = load_or_zero<B_VALID>(b_rt, bsh_ptr + (b_off) * bsh_kt_stride); \
            const auto ks = (ks_val) * 4 + kgrp; \
            auto sa = 0, sb = 0; \
            if constexpr (A_VALID) sa = static_cast<int>(as_row[ks]); \
            else                   sa = (a_rt & (ks < KS)) ? static_cast<int>(as_row[ks]) : 127; \
            if constexpr (B_VALID) sb = static_cast<int>(*(bssh_ptr + (bssh_off))); \
            else                   sb = (b_rt & (ks < KS)) ? static_cast<int>(*(bssh_ptr + (bssh_off))) : 127; \
            acc = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(av,bv,acc,FP4_E2M1,FP4_E2M1,0,sa,0,sb); \
        }

        static_assert(CKT > 0);
        #pragma unroll
        for (auto q = 0; q < (CKT / 4); ++q) {
            DO_MFMA(0, 0, 0, ks_start + q*4)
            DO_MFMA(1, 1, bssh_step0, ks_start + q*4 + 1)
            DO_MFMA(2, 2, bssh_step0 + bssh_step1, ks_start + q*4 + 2)
            DO_MFMA(3, 3, bssh_step0 + bssh_step1 + bssh_step0, ks_start + q*4 + 3)
            a_ptr    += 4 * a_kt_stride;
            bsh_ptr  += 4 * bsh_kt_stride;
            bssh_ptr += 512;
        }
        if constexpr ((CKT % 4) >= 2) {
            DO_MFMA(0, 0, 0, ks_start + (CKT/4)*4)
            DO_MFMA(1, 1, bssh_step0, ks_start + (CKT/4)*4 + 1)
            a_ptr    += 2 * a_kt_stride;
            bsh_ptr  += 2 * bsh_kt_stride;
            bssh_ptr += 256;
        }
        if constexpr ((CKT % 2) == 1) {
            DO_MFMA(0, 0, 0, ks_start + CKT - 1)
        }
        #undef DO_MFMA

    // ════════════════════════════════════════════════════════════════════════
    // PATH B: CKT > 4 — LDS double-buffered with register-staged pipelining
    // ════════════════════════════════════════════════════════════════════════
    } else {
        constexpr auto CHUNK_K = 4;
        constexpr auto NUM_CHUNKS = CKT / CHUNK_K;
        constexpr auto TAIL_KT = CKT % CHUNK_K;

        using Lds = LdsLayout<WAVES_M, WAVES_N, CHUNK_K>;
        using LC  = LoadCounts<WAVES_M, WAVES_N, CHUNK_K, NWARPS>;

        extern __shared__ uint8_t smem_raw[];
        auto* buf0 = smem_raw;
        auto* buf1 = smem_raw + Lds::BUF_SIZE;

        const auto tile_valid = (tile_m < M) && (tile_n < N);

        // ── Precompute per-item coordinates (hoisted out of hot loop) ────────
        // A items: row, k_idx, lds_off — independent of ks_base
        int  a_row[LC::A_PER_THREAD];
        int  a_k_idx[LC::A_PER_THREAD];
        int  a_lds_offs[LC::A_PER_THREAD];
        bool a_valid[LC::A_PER_THREAD];

        #pragma unroll
        for (auto i = 0; i < LC::A_PER_THREAD; ++i) {
            auto linear = i * LC::NTHREADS + tid;
            if (linear >= LC::A_TOTAL) {
                a_valid[i] = false;
                a_lds_offs[i] = -1;
            } else {
                auto k_idx      = linear % (CHUNK_K * 4);
                auto m_local    = (linear / (CHUNK_K * 4)) % 16;
                auto wave_m_idx = linear / (16 * CHUNK_K * 4);
                a_row[i]      = tile_m_base + wave_m_idx * 16 + m_local;
                a_k_idx[i]    = k_idx;
                a_lds_offs[i] = Lds::a_idx(wave_m_idx, m_local, k_idx);
                if constexpr (A_VALID) a_valid[i] = tile_valid;
                else                   a_valid[i] = tile_valid && (a_row[i] < M);
            }
        }

        // B items: global base offset (without ks-dependent k_blk), lds coords
        int  b_kt_local[LC::B_PER_THREAD];   // kt within chunk (0..CHUNK_K-1)
        long b_base_off[LC::B_PER_THREAD];    // global offset without k_blk term
        int  b_kgrp_half[LC::B_PER_THREAD];   // (b_kgrp / 2) for k_blk calc
        int  b_lds_kgrp[LC::B_PER_THREAD];    // b_kgrp for LDS index
        int  b_lds_lrow[LC::B_PER_THREAD];    // b_lrow for LDS index
        int  b_lds_wn[LC::B_PER_THREAD];      // wave_n_idx for LDS index
        int  b_lds_offs[LC::B_PER_THREAD];
        bool b_valid[LC::B_PER_THREAD];

        #pragma unroll
        for (auto i = 0; i < LC::B_PER_THREAD; ++i) {
            auto linear = i * LC::NTHREADS + tid;
            if (linear >= LC::B_TOTAL) {
                b_valid[i] = false;
                b_lds_offs[i] = -1;
            } else {
                auto b_lrow     = linear % 16;
                auto b_kgrp     = (linear / 16) % 4;
                auto kt         = (linear / 64) % CHUNK_K;
                auto wave_n_idx = linear / (CHUNK_K * 64);
                auto b_tile_n   = tile_n_base + wave_n_idx * 16;
                auto n_tile     = b_tile_n / 16;
                auto k_half     = (b_kgrp & 1) * 256;

                b_kt_local[i]  = kt;
                b_base_off[i]  = (long)n_tile * bsh_n_stride + k_half + (long)b_lrow * 16;
                b_kgrp_half[i] = b_kgrp / 2;
                b_lds_kgrp[i]  = b_kgrp;
                b_lds_lrow[i]  = b_lrow;
                b_lds_wn[i]    = wave_n_idx;
                b_lds_offs[i]  = Lds::b_idx(wave_n_idx, kt, b_kgrp, b_lrow);

                if constexpr (B_VALID) b_valid[i] = tile_valid;
                else                   b_valid[i] = tile_valid && (b_tile_n + b_lrow < N);
            }
        }

        // ── Register buffers for staged loads ────────────────────────────────
        int4_v a_regs[LC::A_PER_THREAD];
        int4_v b_regs[LC::B_PER_THREAD];

        // ── Macros for inlined issue/store/compute (no lambdas) ──────────────

        #define ISSUE_LOADS(ks_base) \
        { \
            _Pragma("unroll") \
            for (auto i = 0; i < LC::A_PER_THREAD; ++i) { \
                a_regs[i] = int4_v{0,0,0,0}; \
                if (a_valid[i]) { \
                    auto k_byte = ((ks_base) * 4 + a_k_idx[i]) * 16; \
                    if constexpr (A_VALID) \
                        a_regs[i] = load16(A + (long)a_row[i] * K2 + k_byte); \
                    else if (k_byte + 16 <= K2) \
                        a_regs[i] = load16(A + (long)a_row[i] * K2 + k_byte); \
                } \
            } \
            _Pragma("unroll") \
            for (auto i = 0; i < LC::B_PER_THREAD; ++i) { \
                b_regs[i] = int4_v{0,0,0,0}; \
                if (b_valid[i]) { \
                    auto ks = (ks_base) + b_kt_local[i]; \
                    auto k_blk = ks * 2 + b_kgrp_half[i]; \
                    auto global_off = b_base_off[i] + (long)k_blk * 512; \
                    b_regs[i] = load16_nt(Bsh + global_off); \
                } \
            } \
        }

        #define STORE_TO_LDS(buf) \
        { \
            auto* _sa = Lds::a_off(buf); \
            auto* _sb = Lds::b_off(buf); \
            _Pragma("unroll") \
            for (auto i = 0; i < LC::A_PER_THREAD; ++i) { \
                if (a_lds_offs[i] >= 0) \
                    *reinterpret_cast<int4_v*>(_sa + a_lds_offs[i]) = a_regs[i]; \
            } \
            _Pragma("unroll") \
            for (auto i = 0; i < LC::B_PER_THREAD; ++i) { \
                if (b_lds_offs[i] >= 0) \
                    *reinterpret_cast<int4_v*>(_sb + b_lds_offs[i]) = b_regs[i]; \
            } \
        }

        #define COMPUTE_CHUNK(buf, chunk_ks) \
        { \
            auto* _ca = Lds::a_off(buf); \
            auto* _cb = Lds::b_off(buf); \
            _Pragma("unroll") \
            for (auto _kt = 0; _kt < CHUNK_K; ++_kt) { \
                auto _av = *reinterpret_cast<const int4_v*>( \
                    _ca + Lds::a_idx(wave_m, lrow, _kt * 4 + kgrp)); \
                auto _bv = *reinterpret_cast<const int4_v*>( \
                    _cb + Lds::b_idx(wave_n, _kt, kgrp, lrow)); \
                auto _ks_val = (chunk_ks) + _kt; \
                auto _ks = _ks_val * 4 + kgrp; \
                auto _sa = 0, _sb = 0; \
                if constexpr (A_VALID) _sa = static_cast<int>(as_row[_ks]); \
                else                   _sa = (a_rt & (_ks < KS)) ? static_cast<int>(as_row[_ks]) : 127; \
                auto* _bssh_p = Bssh + bssh_base + (_ks_val & 1) * 2 + (_ks_val >> 1) * 256; \
                if constexpr (B_VALID) _sb = static_cast<int>(*_bssh_p); \
                else                   _sb = (b_rt & (_ks < KS)) ? static_cast<int>(*_bssh_p) : 127; \
                acc = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4( \
                    _av, _bv, acc, FP4_E2M1, FP4_E2M1, 0, _sa, 0, _sb); \
            } \
        }

        // ── Prologue: load first chunk ───────────────────────────────────────
        auto cur_ks = ks_start;
        ISSUE_LOADS(cur_ks);
        asm volatile("s_waitcnt vmcnt(0) lgkmcnt(0)" ::: "memory");
        STORE_TO_LDS(buf0);
        __syncthreads();

        auto* cur_buf = buf0;
        auto* nxt_buf = buf1;

        // ── Main pipelined loop ──────────────────────────────────────────────
        #pragma unroll
        for (auto c = 0; c < NUM_CHUNKS - 1; ++c) {
            ISSUE_LOADS(cur_ks + CHUNK_K);
            COMPUTE_CHUNK(cur_buf, cur_ks);
            asm volatile("s_waitcnt vmcnt(0)" ::: "memory");
            STORE_TO_LDS(nxt_buf);
            __syncthreads();

            auto* tmp = cur_buf;
            cur_buf = nxt_buf;
            nxt_buf = tmp;
            cur_ks += CHUNK_K;
        }

        // ── Epilogue: compute last full chunk ────────────────────────────────
        COMPUTE_CHUNK(cur_buf, cur_ks);

        // ── Tail: remaining K-tiles if CKT not divisible by CHUNK_K ─────────
        if constexpr (TAIL_KT > 0) {
            cur_ks += CHUNK_K;
            ISSUE_LOADS(cur_ks);
            asm volatile("s_waitcnt vmcnt(0) lgkmcnt(0)" ::: "memory");
            STORE_TO_LDS(buf0);
            __syncthreads();

            // Tail compute — only TAIL_KT tiles, not full CHUNK_K
            auto* _ca = Lds::a_off(buf0);
            auto* _cb = Lds::b_off(buf0);
            #pragma unroll
            for (auto kt = 0; kt < TAIL_KT; ++kt) {
                auto av = *reinterpret_cast<const int4_v*>(
                    _ca + Lds::a_idx(wave_m, lrow, kt * 4 + kgrp));
                auto bv = *reinterpret_cast<const int4_v*>(
                    _cb + Lds::b_idx(wave_n, kt, kgrp, lrow));
                auto ks_val = cur_ks + kt;
                auto ks = ks_val * 4 + kgrp;
                auto sa = 0, sb = 0;
                if constexpr (A_VALID) sa = static_cast<int>(as_row[ks]);
                else                   sa = (a_rt & (ks < KS)) ? static_cast<int>(as_row[ks]) : 127;
                auto* bssh_p = Bssh + bssh_base + (ks_val & 1) * 2 + (ks_val >> 1) * 256;
                if constexpr (B_VALID) sb = static_cast<int>(*bssh_p);
                else                   sb = (b_rt & (ks < KS)) ? static_cast<int>(*bssh_p) : 127;
                acc = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
                    av, bv, acc, FP4_E2M1, FP4_E2M1, 0, sa, 0, sb);
            }
        }

        #undef ISSUE_LOADS
        #undef STORE_TO_LDS
        #undef COMPUTE_CHUNK

        if (!tile_valid) return;
    }

    // ── Store output (shared by both paths) ─────────────────────────────────

    const auto out_col      = tile_n + lrow;
    const auto out_row_base = tile_m + kgrp * 4;

    if constexpr (!B_VALID) { if (out_col >= N) return; }

    constexpr auto out_rows_always_valid = A_VALID && (BM >= 16);

    if constexpr (SPLITK) {
        auto* c_out = C_partial + (long)ks_idx * M * N + (long)out_row_base * N + out_col;
        #pragma unroll
        for (auto i = 0; i < 4; ++i, c_out += N) {
            if constexpr (out_rows_always_valid) *c_out = acc[i];
            else if (out_row_base + i < M)       *c_out = acc[i];
        }
    } else {
        auto* c_out = C_final + (long)out_row_base * N + out_col;
        #pragma unroll
        for (auto i = 0; i < 4; ++i, c_out += N) {
            if constexpr (out_rows_always_valid) *c_out = float_to_bf16(acc[i]);
            else if (out_row_base + i < M)       *c_out = float_to_bf16(acc[i]);
        }
    }
}

// ── Reduce kernel — fully templated ─────────────────────────────────────────

template<int C_N, int C_NUM_KSPLIT, int C_M>
__global__ void mxfp4_reduce(
    const float*  __restrict__ C_partial,
    uint16_t*     __restrict__ C_out)
{
    constexpr auto M = C_M;
    constexpr auto N = C_N;
    constexpr auto mn_stride = (long)M * N;
    const auto col = blockIdx.x * 32 + threadIdx.x;
    const auto row = blockIdx.y * 16 + threadIdx.y;
    if (row >= M || col >= N) return;

    const auto mn = (long)row * N + col;
    auto sum = 0.f;
    const auto* ptr = C_partial + mn;
    #pragma unroll
    for (auto k = 0; k < C_NUM_KSPLIT; ++k)
        sum += ptr[k * mn_stride];

    bf16x2 v;
    v[0] = static_cast<__bf16>(sum);
    auto r = uint16_t{};
    __builtin_memcpy(&r, &v, sizeof(r));
    C_out[mn] = r;
}

// ── Launch helpers — fully templated on M ───────────────────────────────────

template<int C_M, int C_K>
void launch_quant(
    const __bf16* A_bf16, uint8_t* A_fp4, uint8_t* A_scale)
{
    constexpr auto KS       = C_K / 32;
    constexpr auto n_groups = C_M * KS;
    constexpr dim3 block{128};
    constexpr dim3 grid{static_cast<uint32_t>((n_groups + 127) / 128)};
    mxfp4_quant<C_M, C_K><<<grid, block>>>(A_bf16, A_fp4, A_scale);
}

template<int C_N, int C_K, int C_SCALEN, int C_TOTAL_KT, int C_KPS, int C_NUM_KSPLIT, int C_M>
void launch_gemm_nk(
    const uint8_t* A, const uint8_t* As,
    const uint8_t* Bsh, const uint8_t* Bssh,
    float* C_partial, uint16_t* C_final)
{
    constexpr auto do_splitk = C_NUM_KSPLIT > 1;

    // M→BM dispatch resolved at compile time
    constexpr auto BM     = (C_M <=  8) ?  8 : (C_M <= 32) ? 16 : (C_M <= 128) ? 32 : 64;
    constexpr auto BN     = 32;
    constexpr auto NWARPS = ((BM + 15) / 16) * (BN / 16);
    static_assert(((BM + 15) / 16) * (BN / 16) == NWARPS);

    constexpr auto WAVES_M = (BM + 15) / 16;
    constexpr auto WAVES_N = BN / 16;
    constexpr auto CHUNK_K = (C_KPS >= 4) ? 4 : C_KPS;
    constexpr auto smem_size = (C_KPS > 4)
        ? LdsLayout<WAVES_M, WAVES_N, CHUNK_K>::TOTAL_LDS
        : 0;

    constexpr auto full_m  = (BM >= 16) ? C_M / BM : 0;
    constexpr auto full_n  = C_N / BN;
    constexpr auto total_m = (C_M + BM - 1) / BM;
    constexpr auto total_n = (C_N + BN - 1) / BN;
    constexpr auto edge_m  = total_m - full_m;
    constexpr auto edge_n  = total_n - full_n;

    constexpr dim3 block{static_cast<uint32_t>(NWARPS * 64)};

    // Helper: set smem + launch for a specific (AV, BV, ox, oy) combo
    auto sub = [&]<bool AV, bool BV>() {
        constexpr auto gx = BV ? full_n : edge_n;
        constexpr auto gy = AV ? full_m : edge_m;
        constexpr auto ox = BV ? 0 : full_n;
        constexpr auto oy = AV ? 0 : full_m;
        if constexpr (gx > 0 && gy > 0) {
            constexpr auto SK = do_splitk;
            if constexpr (smem_size > 0)
                (void)hipFuncSetAttribute(
                    (const void*)mxfp4_gemm<BM,BN,NWARPS,SK,AV,BV,C_KPS,C_N,C_K,C_SCALEN,C_TOTAL_KT,C_KPS,C_M,ox,oy>,
                    hipFuncAttributeMaxDynamicSharedMemorySize, smem_size);
            constexpr dim3 grid{
                static_cast<uint32_t>(gx),
                static_cast<uint32_t>(gy),
                static_cast<uint32_t>(C_NUM_KSPLIT)
            };
            if constexpr (SK)
                mxfp4_gemm<BM,BN,NWARPS,true,AV,BV,C_KPS,C_N,C_K,C_SCALEN,C_TOTAL_KT,C_KPS,C_M,ox,oy>
                    <<<grid,block,smem_size>>>(A,As,Bsh,Bssh,C_partial,nullptr);
            else
                mxfp4_gemm<BM,BN,NWARPS,false,AV,BV,C_KPS,C_N,C_K,C_SCALEN,C_TOTAL_KT,C_KPS,C_M,ox,oy>
                    <<<grid,block,smem_size>>>(A,As,Bsh,Bssh,nullptr,C_final);
        }
    };

    sub.template operator()<true,  true >();
    sub.template operator()<true,  false>();
    sub.template operator()<false, true >();
    sub.template operator()<false, false>();

    if constexpr (do_splitk) {
        constexpr dim3 rblock{32, 16};
        constexpr dim3 rgrid{
            static_cast<uint32_t>((C_N + 31) / 32),
            static_cast<uint32_t>((C_M + 15) / 16)
        };
        mxfp4_reduce<C_N, C_NUM_KSPLIT, C_M><<<rgrid, rblock>>>(C_partial, C_final);
    }
}

template<int C_N, int C_K, int C_SCALEN, int C_TOTAL_KT, int C_M>
void launch_gemm_nk_7168(
    const uint8_t* A, const uint8_t* As,
    const uint8_t* Bsh, const uint8_t* Bssh,
    float* C_partial, uint16_t* C_final)
{
    if constexpr (C_M <= 8)
        launch_gemm_nk<C_N, C_K, C_SCALEN, C_TOTAL_KT, 8, 7, C_M>(A, As, Bsh, Bssh, C_partial, C_final);
    else if constexpr (C_M <= 16)
        launch_gemm_nk<C_N, C_K, C_SCALEN, C_TOTAL_KT, 4, 14, C_M>(A, As, Bsh, Bssh, C_partial, C_final);
    else
        launch_gemm_nk<C_N, C_K, C_SCALEN, C_TOTAL_KT, 56, 1, C_M>(A, As, Bsh, Bssh, C_partial, C_final);
}

// Combined quant + GEMM dispatch — fully compile-time
template<int C_M, int C_K, int C_N, int C_SCALEN, int C_TOTAL_KT, int C_KPS, int C_NUM_KSPLIT>
void launch_shape(
    const __bf16* A_bf16, uint8_t* A_fp4, uint8_t* A_scale,
    const uint8_t* Bsh, const uint8_t* Bssh,
    float* C_partial, uint16_t* C_final)
{
    launch_quant<C_M, C_K>(A_bf16, A_fp4, A_scale);
    launch_gemm_nk<C_N, C_K, C_SCALEN, C_TOTAL_KT, C_KPS, C_NUM_KSPLIT, C_M>(
        A_fp4, A_scale, Bsh, Bssh, C_partial, C_final);
}

template<int C_M, int C_K, int C_N, int C_SCALEN, int C_TOTAL_KT>
void launch_shape_7168(
    const __bf16* A_bf16, uint8_t* A_fp4, uint8_t* A_scale,
    const uint8_t* Bsh, const uint8_t* Bssh,
    float* C_partial, uint16_t* C_final)
{
    launch_quant<C_M, C_K>(A_bf16, A_fp4, A_scale);
    launch_gemm_nk_7168<C_N, C_K, C_SCALEN, C_TOTAL_KT, C_M>(
        A_fp4, A_scale, Bsh, Bssh, C_partial, C_final);
}

// Runtime M dispatch
template<int C_K, int C_N, int C_SCALEN, int C_TOTAL_KT, int C_KPS, int C_NUM_KSPLIT>
void dispatch_shape(
    const __bf16* A_bf16, uint8_t* A_fp4, uint8_t* A_scale,
    const uint8_t* Bsh, const uint8_t* Bssh,
    float* C_partial, uint16_t* C_final, int M)
{
    #define CASE_M(val) case val: launch_shape<val,C_K,C_N,C_SCALEN,C_TOTAL_KT,C_KPS,C_NUM_KSPLIT>(A_bf16,A_fp4,A_scale,Bsh,Bssh,C_partial,C_final); return
    switch (M) {
        CASE_M(4); CASE_M(8); CASE_M(16); CASE_M(32); CASE_M(64); CASE_M(256);
    }
    #undef CASE_M
}

template<int C_K, int C_N, int C_SCALEN, int C_TOTAL_KT>
void dispatch_shape_7168(
    const __bf16* A_bf16, uint8_t* A_fp4, uint8_t* A_scale,
    const uint8_t* Bsh, const uint8_t* Bssh,
    float* C_partial, uint16_t* C_final, int M)
{
    #define CASE_M(val) case val: launch_shape_7168<val,C_K,C_N,C_SCALEN,C_TOTAL_KT>(A_bf16,A_fp4,A_scale,Bsh,Bssh,C_partial,C_final); return
    switch (M) {
        CASE_M(4); CASE_M(8); CASE_M(16); CASE_M(32); CASE_M(64); CASE_M(256);
    }
    #undef CASE_M
}

extern "C" void launch_all(
    const __bf16* A_bf16, uint8_t* A_fp4, uint8_t* A_scale,
    const uint8_t* Bsh, const uint8_t* Bssh,
    float* C_partial, uint16_t* C_final,
    int M, int N, int K)
{
    if (N == 2880 && K == 512)
        dispatch_shape<512, 2880, 16, 4, 4, 1>(A_bf16, A_fp4, A_scale, Bsh, Bssh, C_partial, C_final, M);
    else if (N == 2112 && K == 7168)
        dispatch_shape_7168<7168, 2112, 224, 56>(A_bf16, A_fp4, A_scale, Bsh, Bssh, C_partial, C_final, M);
    else if (N == 4096 && K == 512)
        dispatch_shape<512, 4096, 16, 4, 4, 1>(A_bf16, A_fp4, A_scale, Bsh, Bssh, C_partial, C_final, M);
    else if (N == 7168 && K == 2048)
        dispatch_shape<2048, 7168, 64, 16, 16, 1>(A_bf16, A_fp4, A_scale, Bsh, Bssh, C_partial, C_final, M);
    else if (N == 3072 && K == 1536)
        dispatch_shape<1536, 3072, 48, 12, 12, 1>(A_bf16, A_fp4, A_scale, Bsh, Bssh, C_partial, C_final, M);
}
"""


CPP = r"""
#include <torch/extension.h>
#include <c10/core/DeviceGuard.h>

extern "C" void launch_all(const __bf16*, uint8_t*, uint8_t*,
                            const uint8_t*, const uint8_t*,
                            float*, uint16_t*, int, int, int);

static int get_num_ksplit(int M, int K) {
    auto total_ktiles = K / 128;
    if (total_ktiles < 28) return 1;
    if (M <= 8)  return 7;
    if (M <= 16) return 14;
    return 1;
}

struct ShapeWorkspace {
    at::Tensor A_fp4;
    at::Tensor A_scale;
    at::Tensor C_partial;
    at::Tensor C;
    uint8_t* a_fp4_ptr = nullptr;
    uint8_t* a_scale_ptr = nullptr;
    float* c_partial_ptr = nullptr;
    uint16_t* c_final_ptr = nullptr;
    int M = 0, N = 0, K = 0;
    int num_ksplit = 0;
};

struct BCache {
    const uint8_t* bsh_ptr = nullptr;
    const uint8_t* bssh_ptr = nullptr;
    int64_t bsh_data_ptr = 0;
    int64_t bssh_data_ptr = 0;
};

static ShapeWorkspace g_ws[10];
static auto g_ws_count = 0;
static BCache g_bcache;

static auto* find_or_create_ws(int M, int N, int K, int num_ksplit,
                                const at::TensorOptions& opts) {
    for (auto i = 0; i < g_ws_count; ++i) {
        if (g_ws[i].M == M && g_ws[i].N == N && g_ws[i].K == K)
            return &g_ws[i];
    }
    auto& ws = g_ws[g_ws_count++];
    ws.M = M; ws.N = N; ws.K = K;
    ws.num_ksplit = num_ksplit;
    auto KS = K / 32;
    ws.A_fp4   = at::empty({(int64_t)M, (int64_t)(K / 2)}, opts.dtype(at::kByte));
    ws.A_scale = at::empty({(int64_t)M, (int64_t)KS},       opts.dtype(at::kByte));
    ws.C       = at::empty({(int64_t)M, (int64_t)N},         opts.dtype(at::kBFloat16));
    if (num_ksplit > 1)
        ws.C_partial = at::empty({(int64_t)num_ksplit, (int64_t)M, (int64_t)N}, opts.dtype(at::kFloat));
    ws.a_fp4_ptr   = ws.A_fp4.data_ptr<uint8_t>();
    ws.a_scale_ptr = ws.A_scale.data_ptr<uint8_t>();
    ws.c_partial_ptr = (num_ksplit > 1) ? ws.C_partial.data_ptr<float>() : nullptr;
    ws.c_final_ptr = reinterpret_cast<uint16_t*>(ws.C.data_ptr<at::BFloat16>());
    return &ws;
}

at::Tensor fwd(const at::Tensor& A,
               const at::Tensor& B_q,
               const at::Tensor& B_shuffle,
               const at::Tensor& B_scale_sh) {
    auto guard = at::DeviceGuard(A.device());

    const auto M = static_cast<int>(A.size(0));
    const auto K = static_cast<int>(A.size(1));
    const auto N = static_cast<int>(B_q.size(0));

    const __bf16* a_bf16_ptr;
    at::Tensor A_bf16;
    if (A.scalar_type() == at::kBFloat16 && A.is_contiguous()) {
        a_bf16_ptr = reinterpret_cast<const __bf16*>(A.data_ptr<at::BFloat16>());
    } else {
        A_bf16 = A.to(A.device(), at::kBFloat16, false, false, at::MemoryFormat::Contiguous);
        a_bf16_ptr = reinterpret_cast<const __bf16*>(A_bf16.data_ptr<at::BFloat16>());
    }

    auto bsh_dp  = reinterpret_cast<int64_t>(B_shuffle.data_ptr());
    auto bssh_dp = reinterpret_cast<int64_t>(B_scale_sh.data_ptr());
    if (bsh_dp != g_bcache.bsh_data_ptr || bssh_dp != g_bcache.bssh_data_ptr) {
        auto Bsh = B_shuffle.view(at::kByte);
        if (!Bsh.is_contiguous()) Bsh = Bsh.contiguous();
        auto Bssh = B_scale_sh.view(at::kByte);
        if (!Bssh.is_contiguous()) Bssh = Bssh.contiguous();
        g_bcache.bsh_ptr  = Bsh.data_ptr<uint8_t>();
        g_bcache.bssh_ptr = Bssh.data_ptr<uint8_t>();
        g_bcache.bsh_data_ptr  = bsh_dp;
        g_bcache.bssh_data_ptr = bssh_dp;
    }

    const auto num_ksplit = get_num_ksplit(M, K);
    auto* ws = find_or_create_ws(M, N, K, num_ksplit, A.options());

    launch_all(
        a_bf16_ptr, ws->a_fp4_ptr, ws->a_scale_ptr,
        g_bcache.bsh_ptr, g_bcache.bssh_ptr,
        ws->c_partial_ptr, ws->c_final_ptr,
        M, N, K);

    return ws->C;
}
"""

_ext = load_inline(
    name=f"g_{uuid.uuid4().hex[:8]}",
    cpp_sources=[CPP],
    cuda_sources=[HIP_KERNEL],
    functions=["fwd"],
    with_cuda=True,
    extra_cflags=["-O3", "-std=c++20"],
    extra_cuda_cflags=[
        "-O3",
        "--offload-arch=gfx950",
        "-ffast-math",
        "-ffinite-math-only",
        "-munsafe-fp-atomics",
        "-std=c++20",
        "-mllvm", "-amdgpu-early-inline-all=true",
        "-mllvm", "-amdgpu-function-calls=false",
        "-mwavefrontsize64",
        "-mcumode",
        "-mllvm", "--amdgpu-kernarg-preload-count=16",
        "-mllvm", "-enable-post-misched=0",
        "-mllvm", "--lsr-drop-solution=1",
        "-mllvm", "-amdgpu-coerce-illegal-types=1",
        "-fgpu-flush-denormals-to-zero",
        "-fno-offload-uniform-block",
        "-mllvm", "-amdgpu-loop-prefetch=true",
        "-mllvm", "-enable-unroll-and-jam=true",
        "-mllvm", "-amdgpu-set-wave-priority=true",
        "-mllvm", "-unroll-threshold=1000",
        "-mllvm", "-amdgpu-internalize-symbols=true",
    ],
    extra_ldflags=["-lamdhip64"],
)


def custom_kernel(data: Tuple[torch.Tensor, ...]) -> torch.Tensor:
    """MXFP4 GEMM v25d_v4_i1: v4_i0 + quant intrinsic dest_sel + LdsLayout + launch_bounds."""
    A, _, B_q, B_shuffle, B_scale_sh = data
    return _ext.fwd(A.cuda(), B_q, B_shuffle, B_scale_sh)
scrolls · 905 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 562906.

"""
- MXFP4 GEMM v25d_v5 — gfx950 (MI355X) with LDS + software pipelining.
+ MXFP4 GEMM v25d_v4_i1 — gfx950 (MI355X) with LDS + software pipelining.
- Based on v4 (12.770μs). Changes:
- v4: nt (non-temporal) on PATH B B loads
- v5: nt on PATH B A loads too (both go to LDS, no L1 reuse)
+ Based on v4_i0 (12.264μs). Changes:
+ v4_i1: Quant intrinsic dest_sel packing (no shift/OR/mask).
+ LdsLayout forceinline constexpr const auto.
+ launch_bounds tuning for quant kernel.
"""
import os
os.environ["PYTORCH_ROCM_ARCH"] = "gfx950"
⋯ 12 unchanged lines
using float4_v = float __attribute__((ext_vector_type(4)));
using bf16x2 = __bf16 __attribute__((ext_vector_type(2)));
- static constexpr int FP4_E2M1 = 4;
+ static constexpr auto FP4_E2M1 = 4;
__device__ float4_v __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
int4_v a, int4_v b, float4_v c,
⋯ 5 unchanged lines
}
// Non-temporal load: bypass L1, keep in L2.
- // Optimal for B data in PATH B: large working set, goes to LDS, no L1 reuse.
__device__ __forceinline__ auto load16_nt(const uint8_t* __restrict__ p) {
return __builtin_nontemporal_load(reinterpret_cast<const int4_v*>(p));
}
⋯ 1 unchanged lines
__device__ __forceinline__ auto float_to_bf16(float f) {
bf16x2 v;
v[0] = static_cast<__bf16>(f);
- uint16_t r;
+ auto r = uint16_t{};
__builtin_memcpy(&r, &v, sizeof(r));
return r;
}
- __device__ __forceinline__ auto hw_bf16x2_to_fp4x2(uint32_t bf16_pair, float scale) {
- uint32_t result;
- asm volatile("v_cvt_scalef32_pk_fp4_bf16 %0, %1, %2"
- : "=v"(result) : "v"(bf16_pair), "v"(scale));
- return static_cast<uint8_t>(result & 0xFFu);
- }
+ // ── Quant kernel — fully templated on M, K ──────────────────────────────────
- // ── Quant kernel (unchanged from v25d) ──────────────────────────────────────
-
- __global__ void __launch_bounds__(128, 4)
+ template<int C_M, int C_K>
+ __global__ void __launch_bounds__(128, (C_M * (C_K / 32) <= 256) ? 2 : 4)
mxfp4_quant(
const __bf16* __restrict__ A_bf16,
uint8_t* __restrict__ A_fp4,
- uint8_t* __restrict__ A_scale,
- int M, int K)
+ uint8_t* __restrict__ A_scale)
{
- const auto KS = K / 32;
- const auto K2 = K / 2;
+ constexpr auto KS = C_K / 32;
+ constexpr auto K2 = C_K / 2;
const auto group = blockIdx.x * 128 + threadIdx.x;
const auto row = group / KS;
const auto kg = group % KS;
- if (row >= M) return;
+ if (row >= C_M) return;
- const auto* src = A_bf16 + (long)row * K + kg * 32;
+ const auto row_k = (long)row * C_K;
+ const auto* src = A_bf16 + row_k + kg * 32;
+ // 4 wide loads: 64 bytes in 4 × global_load_dwordx4
+ const auto w0 = *reinterpret_cast<const int4_v*>(src);
+ const auto w1 = *reinterpret_cast<const int4_v*>(src + 8);
+ const auto w2 = *reinterpret_cast<const int4_v*>(src + 16);
+ const auto w3 = *reinterpret_cast<const int4_v*>(src + 24);
+
+ // Reinterpret as uint32 pairs (zero-cost, same registers)
+ const auto* p0 = reinterpret_cast<const uint32_t*>(&w0);
+ const auto* p1 = reinterpret_cast<const uint32_t*>(&w1);
+ const auto* p2 = reinterpret_cast<const uint32_t*>(&w2);
+ const auto* p3 = reinterpret_cast<const uint32_t*>(&w3);
+
+ // absmax across all 32 bf16 values (16 uint32 pairs, each = 2 bf16)
auto absMax = 1e-10f;
- #pragma unroll
- for (int i = 0; i < 32; ++i) {
- auto v = __builtin_elementwise_abs(static_cast<float>(src[i]));
- absMax = (v > absMax) ? v : absMax;
+ #define AMAX_PAIR(pair) { \
+ const auto lo = __builtin_bit_cast(__bf16, static_cast<uint16_t>(pair)); \
+ const auto hi = __builtin_bit_cast(__bf16, static_cast<uint16_t>((pair) >> 16)); \
+ const auto flo = __builtin_elementwise_abs(static_cast<float>(lo)); \
+ const auto fhi = __builtin_elementwise_abs(static_cast<float>(hi)); \
+ absMax = (flo > absMax) ? flo : absMax; \
+ absMax = (fhi > absMax) ? fhi : absMax; \
}
+ AMAX_PAIR(p0[0]) AMAX_PAIR(p0[1]) AMAX_PAIR(p0[2]) AMAX_PAIR(p0[3])
+ AMAX_PAIR(p1[0]) AMAX_PAIR(p1[1]) AMAX_PAIR(p1[2]) AMAX_PAIR(p1[3])
+ AMAX_PAIR(p2[0]) AMAX_PAIR(p2[1]) AMAX_PAIR(p2[2]) AMAX_PAIR(p2[3])
+ AMAX_PAIR(p3[0]) AMAX_PAIR(p3[1]) AMAX_PAIR(p3[2]) AMAX_PAIR(p3[3])
+ #undef AMAX_PAIR
- auto u32 = __builtin_bit_cast(uint32_t, absMax);
+ const auto u32 = __builtin_bit_cast(uint32_t, absMax);
const auto amax_exp = ((u32 + 0x200000u) >> 23) & 0xFFu;
const auto inv_exp = (amax_exp >= 2u) ? (amax_exp - 2u) : 0u;
- A_scale[(long)row * KS + kg] = static_cast<uint8_t>(inv_exp);
+ A_scale[row_k / 32 + kg] = static_cast<uint8_t>(inv_exp);
const auto hw_scale = __builtin_bit_cast(float, static_cast<uint32_t>(inv_exp) << 23);
- const auto* src_u32 = reinterpret_cast<const uint32_t*>(src);
- auto* dst = reinterpret_cast<uint8_t*>(A_fp4 + (long)row * K2 + kg * 16);
- #pragma unroll
- for (int i = 0; i < 16; ++i)
- dst[i] = hw_bf16x2_to_fp4x2(src_u32[i], hw_scale);
+ // Pack 16 FP4 pairs into 4 dwords using intrinsic dest_sel (no shift/OR/mask)
+ // Each __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(old, v2bf16, scale, byte_sel)
+ // writes 1 byte at position byte_sel in the accumulator dword.
+ #define CVT(d, pair, sel) \
+ d = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16( \
+ d, __builtin_bit_cast(bf16x2, (pair)), hw_scale, sel)
+
+ auto d0 = 0u, d1 = 0u, d2 = 0u, d3 = 0u;
+ CVT(d0, p0[0], 0); CVT(d0, p0[1], 1); CVT(d0, p0[2], 2); CVT(d0, p0[3], 3);
+ CVT(d1, p1[0], 0); CVT(d1, p1[1], 1); CVT(d1, p1[2], 2); CVT(d1, p1[3], 3);
+ CVT(d2, p2[0], 0); CVT(d2, p2[1], 1); CVT(d2, p2[2], 2); CVT(d2, p2[3], 3);
+ CVT(d3, p3[0], 0); CVT(d3, p3[1], 1); CVT(d3, p3[2], 2); CVT(d3, p3[3], 3);
+ #undef CVT
+
+ auto* dst = reinterpret_cast<int4_v*>(A_fp4 + row_k / 2 + kg * 16);
+ *dst = int4_v{static_cast<int>(d0), static_cast<int>(d1),
+ static_cast<int>(d2), static_cast<int>(d3)};
}
- // ── Boundary-checked load ───────────────────────────────────────────────────
+ // ── Boundary-checked loads ──────────────────────────────────────────────────
template<bool ALWAYS_VALID>
__device__ __forceinline__ auto load_or_zero(bool rt_valid, const uint8_t* p) {
⋯ 1 unchanged lines
else return rt_valid ? load16(p) : int4_v{0,0,0,0};
}
- template<bool ALWAYS_VALID>
- __device__ __forceinline__ auto load_or_zero_nt(bool rt_valid, const uint8_t* p) {
- if constexpr (ALWAYS_VALID) return load16_nt(p);
- else return rt_valid ? load16_nt(p) : int4_v{0,0,0,0};
- }
-
-
// ── LDS layout for software-pipelined path (CKT > 4) ───────────────────────
//
// Double-buffered. Per buffer:
⋯ 2 unchanged lines
//
// Bank conflict strategy:
// A: XOR swizzle on row index — lrow ^ (k_idx & 7)
- // Makes consecutive rows 4-bank-apart (pad ensures ≤2-way conflict)
// B: XOR swizzle on kgrp — kgrp ^ (lrow >> 2)
- // Prevents kgrp 0 and kgrp 2 (same half-wave) from hitting same banks
template<int WAVES_M, int WAVES_N, int CHUNK_K>
struct LdsLayout {
- static constexpr int A_ROW = CHUNK_K * 64 + 16; // +16B pad per row
- static constexpr int A_TILE = 16 * A_ROW;
- static constexpr int A_SIZE = WAVES_M * A_TILE;
+ static constexpr auto A_ROW = CHUNK_K * 64 + 16; // +16B pad per row
+ static constexpr auto A_TILE = 16 * A_ROW;
+ static constexpr auto A_SIZE = WAVES_M * A_TILE;
// B: 5 kgrp slots (4 real + 1 pad) to space kgrps 4-banks apart
- static constexpr int B_KGRP_STRIDE = 16 * 16; // 16 lrows × 16B = 256B
- static constexpr int B_KT_STRIDE = 5 * B_KGRP_STRIDE; // 5 slots (4+1 pad)
- static constexpr int B_TILE = CHUNK_K * B_KT_STRIDE;
- static constexpr int B_SIZE = WAVES_N * B_TILE;
+ static constexpr auto B_KGRP_STRIDE = 16 * 16; // 16 lrows × 16B = 256B
+ static constexpr auto B_KT_STRIDE = 5 * B_KGRP_STRIDE; // 5 slots (4+1 pad)
+ static constexpr auto B_TILE = CHUNK_K * B_KT_STRIDE;
+ static constexpr auto B_SIZE = WAVES_N * B_TILE;
- static constexpr int BUF_SIZE = A_SIZE + B_SIZE;
- static constexpr int TOTAL_LDS = 2 * BUF_SIZE;
+ static constexpr auto BUF_SIZE = A_SIZE + B_SIZE;
+ static constexpr auto TOTAL_LDS = 2 * BUF_SIZE;
- __device__ static constexpr auto a_off(uint8_t* buf) { return buf; }
- __device__ static constexpr auto b_off(uint8_t* buf) { return buf + A_SIZE; }
+ __device__ __forceinline__ static constexpr auto a_off(uint8_t* const buf) { return buf; }
+ __device__ __forceinline__ static constexpr auto b_off(uint8_t* const buf) { return buf + A_SIZE; }
- // A: swizzled store/load offset
- __device__ static auto a_idx(int wm, int row, int k_idx) {
+ __device__ __forceinline__ static constexpr auto a_idx(const auto wm, const auto row, const auto k_idx) {
return wm * A_TILE + (row ^ (k_idx & 7)) * A_ROW + k_idx * 16;
}
- // B: swizzled store/load offset
- __device__ static auto b_idx(int wn, int kt, int kgrp, int lrow) {
+ __device__ __forceinline__ static constexpr auto b_idx(const auto wn, const auto kt, const auto kgrp, const auto lrow) {
return wn * B_TILE + kt * B_KT_STRIDE
+ (kgrp ^ (lrow >> 2)) * B_KGRP_STRIDE + lrow * 16;
}
⋯ 3 unchanged lines
template<int WAVES_M, int WAVES_N, int CHUNK_K, int NWARPS>
struct LoadCounts {
- static constexpr int NTHREADS = NWARPS * 64;
- static constexpr int A_TOTAL = WAVES_M * 16 * CHUNK_K * 4;
- static constexpr int B_TOTAL = WAVES_N * CHUNK_K * 4 * 16;
- static constexpr int A_PER_THREAD = (A_TOTAL + NTHREADS - 1) / NTHREADS;
- static constexpr int B_PER_THREAD = (B_TOTAL + NTHREADS - 1) / NTHREADS;
+ static constexpr auto NTHREADS = NWARPS * 64;
+ static constexpr auto A_TOTAL = WAVES_M * 16 * CHUNK_K * 4;
+ static constexpr auto B_TOTAL = WAVES_N * CHUNK_K * 4 * 16;
+ static constexpr auto A_PER_THREAD = (A_TOTAL + NTHREADS - 1) / NTHREADS;
+ static constexpr auto B_PER_THREAD = (B_TOTAL + NTHREADS - 1) / NTHREADS;
};
// ── Main GEMM kernel ────────────────────────────────────────────────────────
template<int BM, int BN, int NWARPS, bool SPLITK, bool A_VALID, bool B_VALID,
- int CKT, int C_N, int C_K, int C_SCALEN, int C_TOTAL_KT, int C_KPS>
+ int CKT, int C_N, int C_K, int C_SCALEN, int C_TOTAL_KT, int C_KPS,
+ int C_M, int C_TILE_OFF_X, int C_TILE_OFF_Y>
__global__ void __launch_bounds__(NWARPS * 64, (NWARPS <= 2) ? 4 : 2)
mxfp4_gemm(
const uint8_t* __restrict__ A,
⋯ 1 unchanged lines
const uint8_t* __restrict__ Bsh,
const uint8_t* __restrict__ Bssh,
float* __restrict__ C_partial,
- uint16_t* __restrict__ C_final,
- int M,
- int tile_off_x, int tile_off_y)
+ uint16_t* __restrict__ C_final)
{
static_assert(BN % 16 == 0);
- constexpr int WAVES_M = (BM + 15) / 16;
- constexpr int WAVES_N = BN / 16;
+ constexpr auto WAVES_M = (BM + 15) / 16;
+ constexpr auto WAVES_N = BN / 16;
static_assert(WAVES_M * WAVES_N == NWARPS);
- constexpr int N = C_N;
- constexpr int K = C_K;
- constexpr int scaleN = C_SCALEN;
- constexpr int K2 = K / 2;
- constexpr int KS = K / 32;
- constexpr long bsh_n_stride = (long)(K / 64) * 512;
+ constexpr auto M = C_M;
+ constexpr auto N = C_N;
+ constexpr auto K = C_K;
+ constexpr auto scaleN = C_SCALEN;
+ constexpr auto K2 = K / 2;
+ constexpr auto KS = K / 32;
+ constexpr auto bsh_n_stride = (long)(K / 64) * 512;
const auto ks_idx = static_cast<int>(blockIdx.z);
const auto lane = static_cast<int>(threadIdx.x % 64);
⋯ 2 unchanged lines
const auto wave_n = wave % WAVES_N;
const auto tid = static_cast<int>(threadIdx.x);
- const auto tile_m_base = (static_cast<int>(blockIdx.y) + tile_off_y) * BM;
- const auto tile_n_base = (static_cast<int>(blockIdx.x) + tile_off_x) * BN;
+ const auto tile_m_base = (static_cast<int>(blockIdx.y) + C_TILE_OFF_Y) * BM;
+ const auto tile_n_base = (static_cast<int>(blockIdx.x) + C_TILE_OFF_X) * BN;
const auto tile_m = tile_m_base + wave_m * 16;
const auto tile_n = tile_n_base + wave_n * 16;
- constexpr int ktiles_per_split = C_KPS;
+ constexpr auto ktiles_per_split = C_KPS;
const auto ks_start = ks_idx * ktiles_per_split;
const auto lrow = lane % 16;
⋯ 20 unchanged lines
if (b_rt) bssh_base = ((gn >> 4) & 1) + (gn & 15) * 4 + kgrp * 64 + (gn >> 5) * (32 * scaleN);
}
- // B_scale pointer for a given absolute K-tile index
- auto bssh_for = [&](int ks_val) -> const uint8_t* {
- return Bssh + bssh_base + (ks_val & 1) * 2 + (ks_val >> 1) * 256;
- };
-
float4_v acc{0.f, 0.f, 0.f, 0.f};
// ════════════════════════════════════════════════════════════════════════
- // PATH A: CKT <= 4 — Direct global loads, no LDS (same as v25d)
+ // PATH A: CKT <= 4 — Direct global loads, no LDS
// ════════════════════════════════════════════════════════════════════════
if constexpr (CKT <= 4) {
if (tile_m >= M || tile_n >= N) return;
⋯ 4 unchanged lines
const auto k_half_off = (kgrp & 1) * 256;
const auto k_blk_base = kgrp >> 1;
- constexpr long a_kt_stride = 64L;
- constexpr long bsh_kt_stride = 1024L;
+ constexpr auto a_kt_stride = 64L;
+ constexpr auto bsh_kt_stride = 1024L;
const uint8_t* a_ptr = nullptr;
const uint8_t* bsh_ptr = nullptr;
⋯ 17 unchanged lines
const auto bssh_step0 = (ks_start & 1) ? 254 : 2;
const auto bssh_step1 = 256 - bssh_step0;
+ // PATH A: A from regular load (small, L1 reuse), B from .cg (large, L1 bypass)
#define DO_MFMA(a_off, b_off, bssh_off, ks_val) \
{ \
const auto av = load_or_zero<A_VALID>(a_rt, a_ptr + (a_off) * a_kt_stride); \
const auto bv = load_or_zero<B_VALID>(b_rt, bsh_ptr + (b_off) * bsh_kt_stride); \
const auto ks = (ks_val) * 4 + kgrp; \
- int sa, sb; \
+ auto sa = 0, sb = 0; \
if constexpr (A_VALID) sa = static_cast<int>(as_row[ks]); \
else sa = (a_rt & (ks < KS)) ? static_cast<int>(as_row[ks]) : 127; \
if constexpr (B_VALID) sb = static_cast<int>(*(bssh_ptr + (bssh_off))); \
⋯ 3 unchanged lines
static_assert(CKT > 0);
#pragma unroll
- for (int q = 0; q < (CKT / 4); ++q) {
+ for (auto q = 0; q < (CKT / 4); ++q) {
DO_MFMA(0, 0, 0, ks_start + q*4)
DO_MFMA(1, 1, bssh_step0, ks_start + q*4 + 1)
DO_MFMA(2, 2, bssh_step0 + bssh_step1, ks_start + q*4 + 2)
⋯ 18 unchanged lines
// PATH B: CKT > 4 — LDS double-buffered with register-staged pipelining
// ════════════════════════════════════════════════════════════════════════
} else {
- constexpr int CHUNK_K = 4;
- constexpr int NUM_CHUNKS = CKT / CHUNK_K;
- constexpr int TAIL_KT = CKT % CHUNK_K;
+ constexpr auto CHUNK_K = 4;
+ constexpr auto NUM_CHUNKS = CKT / CHUNK_K;
+ constexpr auto TAIL_KT = CKT % CHUNK_K;
using Lds = LdsLayout<WAVES_M, WAVES_N, CHUNK_K>;
using LC = LoadCounts<WAVES_M, WAVES_N, CHUNK_K, NWARPS>;
⋯ 2 unchanged lines
auto* buf0 = smem_raw;
auto* buf1 = smem_raw + Lds::BUF_SIZE;
- // NOTE: no early return before this point — all threads must participate
- // in __syncthreads(). Invalid tiles produce zeros (handled via A_VALID/B_VALID).
- const bool tile_valid = (tile_m < M) && (tile_n < N);
+ const auto tile_valid = (tile_m < M) && (tile_n < N);
- // ── Helpers: decompose thread linear index to load coordinates ───────
+ // ── Precompute per-item coordinates (hoisted out of hot loop) ────────
+ // A items: row, k_idx, lds_off — independent of ks_base
+ int a_row[LC::A_PER_THREAD];
+ int a_k_idx[LC::A_PER_THREAD];
+ int a_lds_offs[LC::A_PER_THREAD];
+ bool a_valid[LC::A_PER_THREAD];
- // Compute A global address + LDS offset for a given linear item index
- auto a_item = [&](int item_idx, int ks_base) {
- struct { int4_v data; int lds_off; } result;
- auto linear = item_idx * LC::NTHREADS + tid;
- if (linear >= LC::A_TOTAL) { result.data = int4_v{0,0,0,0}; result.lds_off = -1; return result; }
-
- auto k_idx = linear % (CHUNK_K * 4);
- auto m_local = (linear / (CHUNK_K * 4)) % 16;
- auto wave_m_idx = linear / (16 * CHUNK_K * 4);
-
- auto row = tile_m_base + wave_m_idx * 16 + m_local;
- auto k_byte = (ks_base * 4 + k_idx) * 16;
-
- result.data = int4_v{0,0,0,0};
- if (tile_valid) {
- if constexpr (A_VALID) {
- result.data = load16_nt(A + (long)row * K2 + k_byte);
- } else {
- if (row < M && k_byte + 16 <= K2)
- result.data = load16_nt(A + (long)row * K2 + k_byte);
- }
+ #pragma unroll
+ for (auto i = 0; i < LC::A_PER_THREAD; ++i) {
+ auto linear = i * LC::NTHREADS + tid;
+ if (linear >= LC::A_TOTAL) {
+ a_valid[i] = false;
+ a_lds_offs[i] = -1;
+ } else {
+ auto k_idx = linear % (CHUNK_K * 4);
+ auto m_local = (linear / (CHUNK_K * 4)) % 16;
+ auto wave_m_idx = linear / (16 * CHUNK_K * 4);
+ a_row[i] = tile_m_base + wave_m_idx * 16 + m_local;
+ a_k_idx[i] = k_idx;
+ a_lds_offs[i] = Lds::a_idx(wave_m_idx, m_local, k_idx);
+ if constexpr (A_VALID) a_valid[i] = tile_valid;
+ else a_valid[i] = tile_valid && (a_row[i] < M);
}
- result.lds_off = Lds::a_idx(wave_m_idx, m_local, k_idx);
- return result;
- };
+ }
- // Compute B global address + LDS offset for a given linear item index
- auto b_item = [&](int item_idx, int ks_base) {
- struct { int4_v data; int lds_off; } result;
- auto linear = item_idx * LC::NTHREADS + tid;
- if (linear >= LC::B_TOTAL) { result.data = int4_v{0,0,0,0}; result.lds_off = -1; return result; }
+ // B items: global base offset (without ks-dependent k_blk), lds coords
+ int b_kt_local[LC::B_PER_THREAD]; // kt within chunk (0..CHUNK_K-1)
+ long b_base_off[LC::B_PER_THREAD]; // global offset without k_blk term
+ int b_kgrp_half[LC::B_PER_THREAD]; // (b_kgrp / 2) for k_blk calc
+ int b_lds_kgrp[LC::B_PER_THREAD]; // b_kgrp for LDS index
+ int b_lds_lrow[LC::B_PER_THREAD]; // b_lrow for LDS index
+ int b_lds_wn[LC::B_PER_THREAD]; // wave_n_idx for LDS index
+ int b_lds_offs[LC::B_PER_THREAD];
+ bool b_valid[LC::B_PER_THREAD];
- auto b_lrow = linear % 16;
- auto b_kgrp = (linear / 16) % 4;
- auto kt = (linear / 64) % CHUNK_K;
- auto wave_n_idx = linear / (CHUNK_K * 64);
+ #pragma unroll
+ for (auto i = 0; i < LC::B_PER_THREAD; ++i) {
+ auto linear = i * LC::NTHREADS + tid;
+ if (linear >= LC::B_TOTAL) {
+ b_valid[i] = false;
+ b_lds_offs[i] = -1;
+ } else {
+ auto b_lrow = linear % 16;
+ auto b_kgrp = (linear / 16) % 4;
+ auto kt = (linear / 64) % CHUNK_K;
+ auto wave_n_idx = linear / (CHUNK_K * 64);
+ auto b_tile_n = tile_n_base + wave_n_idx * 16;
+ auto n_tile = b_tile_n / 16;
+ auto k_half = (b_kgrp & 1) * 256;
- auto b_tile_n = tile_n_base + wave_n_idx * 16;
- auto ks = ks_base + kt;
+ b_kt_local[i] = kt;
+ b_base_off[i] = (long)n_tile * bsh_n_stride + k_half + (long)b_lrow * 16;
+ b_kgrp_half[i] = b_kgrp / 2;
+ b_lds_kgrp[i] = b_kgrp;
+ b_lds_lrow[i] = b_lrow;
+ b_lds_wn[i] = wave_n_idx;
+ b_lds_offs[i] = Lds::b_idx(wave_n_idx, kt, b_kgrp, b_lrow);
- // B_shuffle address
- auto n_tile = b_tile_n / 16;
- auto k_blk = ks * 2 + b_kgrp / 2;
- auto k_half = (b_kgrp & 1) * 256;
- auto global_off = (long)n_tile * bsh_n_stride + (long)k_blk * 512 + k_half + (long)b_lrow * 16;
-
- result.data = int4_v{0,0,0,0};
- if (tile_valid) {
- if constexpr (B_VALID) {
- result.data = load16_nt(Bsh + global_off);
- } else {
- if (b_tile_n + b_lrow < N)
- result.data = load16_nt(Bsh + global_off);
- }
+ if constexpr (B_VALID) b_valid[i] = tile_valid;
+ else b_valid[i] = tile_valid && (b_tile_n + b_lrow < N);
}
- result.lds_off = Lds::b_idx(wave_n_idx, kt, b_kgrp, b_lrow);
- return result;
- };
+ }
- // ── Phase: issue global loads into register arrays ───────────────────
- // Returns register arrays holding prefetched data + LDS offsets
-
- // Register buffers for staged loads
+ // ── Register buffers for staged loads ────────────────────────────────
int4_v a_regs[LC::A_PER_THREAD];
- int a_lds_offs[LC::A_PER_THREAD];
int4_v b_regs[LC::B_PER_THREAD];
- int b_lds_offs[LC::B_PER_THREAD];
- auto issue_loads = [&](int ks_base) {
- #pragma unroll
- for (int i = 0; i < LC::A_PER_THREAD; ++i) {
- auto [data, off] = a_item(i, ks_base);
- a_regs[i] = data;
- a_lds_offs[i] = off;
- }
- #pragma unroll
- for (int i = 0; i < LC::B_PER_THREAD; ++i) {
- auto [data, off] = b_item(i, ks_base);
- b_regs[i] = data;
- b_lds_offs[i] = off;
- }
- };
+ // ── Macros for inlined issue/store/compute (no lambdas) ──────────────
- auto store_to_lds = [&](uint8_t* buf) {
- auto* sa = Lds::a_off(buf);
- auto* sb = Lds::b_off(buf);
- #pragma unroll
- for (int i = 0; i < LC::A_PER_THREAD; ++i) {
- if (a_lds_offs[i] >= 0)
- *reinterpret_cast<int4_v*>(sa + a_lds_offs[i]) = a_regs[i];
- }
- #pragma unroll
- for (int i = 0; i < LC::B_PER_THREAD; ++i) {
- if (b_lds_offs[i] >= 0)
- *reinterpret_cast<int4_v*>(sb + b_lds_offs[i]) = b_regs[i];
- }
- };
+ #define ISSUE_LOADS(ks_base) \
+ { \
+ _Pragma("unroll") \
+ for (auto i = 0; i < LC::A_PER_THREAD; ++i) { \
+ a_regs[i] = int4_v{0,0,0,0}; \
+ if (a_valid[i]) { \
+ auto k_byte = ((ks_base) * 4 + a_k_idx[i]) * 16; \
+ if constexpr (A_VALID) \
+ a_regs[i] = load16(A + (long)a_row[i] * K2 + k_byte); \
+ else if (k_byte + 16 <= K2) \
+ a_regs[i] = load16(A + (long)a_row[i] * K2 + k_byte); \
+ } \
+ } \
+ _Pragma("unroll") \
+ for (auto i = 0; i < LC::B_PER_THREAD; ++i) { \
+ b_regs[i] = int4_v{0,0,0,0}; \
+ if (b_valid[i]) { \
+ auto ks = (ks_base) + b_kt_local[i]; \
+ auto k_blk = ks * 2 + b_kgrp_half[i]; \
+ auto global_off = b_base_off[i] + (long)k_blk * 512; \
+ b_regs[i] = load16_nt(Bsh + global_off); \
+ } \
+ } \
+ }
- // ── Compute CHUNK_K MFMAs from LDS buffer ───────────────────────────
+ #define STORE_TO_LDS(buf) \
+ { \
+ auto* _sa = Lds::a_off(buf); \
+ auto* _sb = Lds::b_off(buf); \
+ _Pragma("unroll") \
+ for (auto i = 0; i < LC::A_PER_THREAD; ++i) { \
+ if (a_lds_offs[i] >= 0) \
+ *reinterpret_cast<int4_v*>(_sa + a_lds_offs[i]) = a_regs[i]; \
+ } \
+ _Pragma("unroll") \
+ for (auto i = 0; i < LC::B_PER_THREAD; ++i) { \
+ if (b_lds_offs[i] >= 0) \
+ *reinterpret_cast<int4_v*>(_sb + b_lds_offs[i]) = b_regs[i]; \
+ } \
+ }
- auto compute_chunk_from_lds = [&](uint8_t* buf, int chunk_ks) {
- auto* smem_a = Lds::a_off(buf);
- auto* smem_b = Lds::b_off(buf);
+ #define COMPUTE_CHUNK(buf, chunk_ks) \
+ { \
+ auto* _ca = Lds::a_off(buf); \
+ auto* _cb = Lds::b_off(buf); \
+ _Pragma("unroll") \
+ for (auto _kt = 0; _kt < CHUNK_K; ++_kt) { \
+ auto _av = *reinterpret_cast<const int4_v*>( \
+ _ca + Lds::a_idx(wave_m, lrow, _kt * 4 + kgrp)); \
+ auto _bv = *reinterpret_cast<const int4_v*>( \
+ _cb + Lds::b_idx(wave_n, _kt, kgrp, lrow)); \
+ auto _ks_val = (chunk_ks) + _kt; \
+ auto _ks = _ks_val * 4 + kgrp; \
+ auto _sa = 0, _sb = 0; \
+ if constexpr (A_VALID) _sa = static_cast<int>(as_row[_ks]); \
+ else _sa = (a_rt & (_ks < KS)) ? static_cast<int>(as_row[_ks]) : 127; \
+ auto* _bssh_p = Bssh + bssh_base + (_ks_val & 1) * 2 + (_ks_val >> 1) * 256; \
+ if constexpr (B_VALID) _sb = static_cast<int>(*_bssh_p); \
+ else _sb = (b_rt & (_ks < KS)) ? static_cast<int>(*_bssh_p) : 127; \
+ acc = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4( \
+ _av, _bv, acc, FP4_E2M1, FP4_E2M1, 0, _sa, 0, _sb); \
+ } \
+ }
- #pragma unroll
- for (int kt = 0; kt < CHUNK_K; ++kt) {
- auto av = *reinterpret_cast<const int4_v*>(
- smem_a + Lds::a_idx(wave_m, lrow, kt * 4 + kgrp));
- auto bv = *reinterpret_cast<const int4_v*>(
- smem_b + Lds::b_idx(wave_n, kt, kgrp, lrow));
-
- auto ks_val = chunk_ks + kt;
- auto ks = ks_val * 4 + kgrp;
-
- int sa, sb;
- if constexpr (A_VALID) sa = static_cast<int>(as_row[ks]);
- else sa = (a_rt & (ks < KS)) ? static_cast<int>(as_row[ks]) : 127;
-
- auto* bssh_p = bssh_for(ks_val);
- if constexpr (B_VALID) sb = static_cast<int>(*bssh_p);
- else sb = (b_rt & (ks < KS)) ? static_cast<int>(*bssh_p) : 127;
-
- acc = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
- av, bv, acc, FP4_E2M1, FP4_E2M1, 0, sa, 0, sb);
- }
- };
-
// ── Prologue: load first chunk ───────────────────────────────────────
auto cur_ks = ks_start;
- issue_loads(cur_ks);
- // No overlap opportunity yet, so just wait and store
+ ISSUE_LOADS(cur_ks);
asm volatile("s_waitcnt vmcnt(0) lgkmcnt(0)" ::: "memory");
- store_to_lds(buf0);
+ STORE_TO_LDS(buf0);
__syncthreads();
auto* cur_buf = buf0;
auto* nxt_buf = buf1;
// ── Main pipelined loop ──────────────────────────────────────────────
- // For each chunk except the last:
- // 1. Issue global loads for chunk[c+1] into register buffers
- // 2. Compute chunk[c] from LDS (MFMA pipe, overlaps with VMEM loads)
- // 3. Wait for global loads
- // 4. Store register buffers to LDS[nxt_buf]
- // 5. Barrier, swap buffers
-
#pragma unroll
- for (int c = 0; c < NUM_CHUNKS - 1; ++c) {
- // 1. Issue loads for NEXT chunk (non-blocking VMEM)
- issue_loads(cur_ks + CHUNK_K);
-
- // 2. Compute CURRENT chunk from LDS (overlaps with VMEM loads)
- compute_chunk_from_lds(cur_buf, cur_ks);
-
- // 3. Wait for next chunk's global loads to complete
+ for (auto c = 0; c < NUM_CHUNKS - 1; ++c) {
+ ISSUE_LOADS(cur_ks + CHUNK_K);
+ COMPUTE_CHUNK(cur_buf, cur_ks);
asm volatile("s_waitcnt vmcnt(0)" ::: "memory");
-
- // 4. Store to next buffer's LDS
- store_to_lds(nxt_buf);
-
- // 5. Barrier — all threads done storing to nxt_buf
+ STORE_TO_LDS(nxt_buf);
__syncthreads();
- // Swap
auto* tmp = cur_buf;
cur_buf = nxt_buf;
nxt_buf = tmp;
⋯ 1 unchanged lines
}
// ── Epilogue: compute last full chunk ────────────────────────────────
- compute_chunk_from_lds(cur_buf, cur_ks);
+ COMPUTE_CHUNK(cur_buf, cur_ks);
// ── Tail: remaining K-tiles if CKT not divisible by CHUNK_K ─────────
if constexpr (TAIL_KT > 0) {
cur_ks += CHUNK_K;
-
- // Issue loads for tail tiles
- // Note: a_item/b_item handle out-of-bounds with zero data via K2 check
- issue_loads(cur_ks);
+ ISSUE_LOADS(cur_ks);
asm volatile("s_waitcnt vmcnt(0) lgkmcnt(0)" ::: "memory");
- store_to_lds(buf0);
+ STORE_TO_LDS(buf0);
__syncthreads();
- // Compute only TAIL_KT tiles (not full CHUNK_K)
- auto* smem_a = Lds::a_off(buf0);
- auto* smem_b = Lds::b_off(buf0);
-
+ // Tail compute — only TAIL_KT tiles, not full CHUNK_K
+ auto* _ca = Lds::a_off(buf0);
+ auto* _cb = Lds::b_off(buf0);
#pragma unroll
- for (int kt = 0; kt < TAIL_KT; ++kt) {
+ for (auto kt = 0; kt < TAIL_KT; ++kt) {
auto av = *reinterpret_cast<const int4_v*>(
- smem_a + Lds::a_idx(wave_m, lrow, kt * 4 + kgrp));
+ _ca + Lds::a_idx(wave_m, lrow, kt * 4 + kgrp));
auto bv = *reinterpret_cast<const int4_v*>(
- smem_b + Lds::b_idx(wave_n, kt, kgrp, lrow));
-
+ _cb + Lds::b_idx(wave_n, kt, kgrp, lrow));
auto ks_val = cur_ks + kt;
auto ks = ks_val * 4 + kgrp;
-
- int sa, sb;
+ auto sa = 0, sb = 0;
if constexpr (A_VALID) sa = static_cast<int>(as_row[ks]);
else sa = (a_rt & (ks < KS)) ? static_cast<int>(as_row[ks]) : 127;
-
- auto* bssh_p = bssh_for(ks_val);
+ auto* bssh_p = Bssh + bssh_base + (ks_val & 1) * 2 + (ks_val >> 1) * 256;
if constexpr (B_VALID) sb = static_cast<int>(*bssh_p);
else sb = (b_rt & (ks < KS)) ? static_cast<int>(*bssh_p) : 127;
-
acc = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
av, bv, acc, FP4_E2M1, FP4_E2M1, 0, sa, 0, sb);
}
}
- // If tile was invalid, don't store
+ #undef ISSUE_LOADS
+ #undef STORE_TO_LDS
+ #undef COMPUTE_CHUNK
+
if (!tile_valid) return;
}
⋯ 4 unchanged lines
if constexpr (!B_VALID) { if (out_col >= N) return; }
- constexpr bool out_rows_always_valid = A_VALID && (BM >= 16);
+ constexpr auto out_rows_always_valid = A_VALID && (BM >= 16);
if constexpr (SPLITK) {
- auto c_out = C_partial + (long)ks_idx * M * N + (long)out_row_base * N + out_col;
+ auto* c_out = C_partial + (long)ks_idx * M * N + (long)out_row_base * N + out_col;
#pragma unroll
- for (int i = 0; i < 4; ++i) {
- if constexpr (out_rows_always_valid) c_out[i * N] = acc[i];
- else if (out_row_base + i < M) c_out[i * N] = acc[i];
+ for (auto i = 0; i < 4; ++i, c_out += N) {
+ if constexpr (out_rows_always_valid) *c_out = acc[i];
+ else if (out_row_base + i < M) *c_out = acc[i];
}
} else {
- auto c_out = C_final + (long)out_row_base * N + out_col;
+ auto* c_out = C_final + (long)out_row_base * N + out_col;
#pragma unroll
- for (int i = 0; i < 4; ++i) {
- if constexpr (out_rows_always_valid) c_out[i * N] = float_to_bf16(acc[i]);
- else if (out_row_base + i < M) c_out[i * N] = float_to_bf16(acc[i]);
+ for (auto i = 0; i < 4; ++i, c_out += N) {
+ if constexpr (out_rows_always_valid) *c_out = float_to_bf16(acc[i]);
+ else if (out_row_base + i < M) *c_out = float_to_bf16(acc[i]);
}
}
}
- // ── Reduce kernel (unchanged) ───────────────────────────────────────────────
+ // ── Reduce kernel — fully templated ─────────────────────────────────────────
- template<int C_N, int C_NUM_KSPLIT>
+ template<int C_N, int C_NUM_KSPLIT, int C_M>
__global__ void mxfp4_reduce(
const float* __restrict__ C_partial,
- uint16_t* __restrict__ C_out,
- int M)
+ uint16_t* __restrict__ C_out)
{
+ constexpr auto M = C_M;
constexpr auto N = C_N;
+ constexpr auto mn_stride = (long)M * N;
const auto col = blockIdx.x * 32 + threadIdx.x;
const auto row = blockIdx.y * 16 + threadIdx.y;
if (row >= M || col >= N) return;
const auto mn = (long)row * N + col;
- const auto mn_stride = (long)M * N;
auto sum = 0.f;
- auto* ptr = C_partial + mn;
+ const auto* ptr = C_partial + mn;
#pragma unroll
- for (auto k = 0; k < C_NUM_KSPLIT; ++k, ptr += mn_stride)
- sum += *ptr;
+ for (auto k = 0; k < C_NUM_KSPLIT; ++k)
+ sum += ptr[k * mn_stride];
bf16x2 v;
v[0] = static_cast<__bf16>(sum);
- uint16_t r;
+ auto r = uint16_t{};
__builtin_memcpy(&r, &v, sizeof(r));
C_out[mn] = r;
}
- // ── Launch helpers ──────────────────────────────────────────────────────────
+ // ── Launch helpers — fully templated on M ───────────────────────────────────
- extern "C" void launch_quant(
- const __bf16* A_bf16, uint8_t* A_fp4, uint8_t* A_scale, int M, int K)
+ template<int C_M, int C_K>
+ void launch_quant(
+ const __bf16* A_bf16, uint8_t* A_fp4, uint8_t* A_scale)
{
- const auto KS = K / 32;
- const auto n_groups = M * KS;
- const dim3 block{128};
- const dim3 grid{static_cast<uint32_t>((n_groups + 127) / 128)};
- mxfp4_quant<<<grid, block>>>(A_bf16, A_fp4, A_scale, M, K);
+ constexpr auto KS = C_K / 32;
+ constexpr auto n_groups = C_M * KS;
+ constexpr dim3 block{128};
+ constexpr dim3 grid{static_cast<uint32_t>((n_groups + 127) / 128)};
+ mxfp4_quant<C_M, C_K><<<grid, block>>>(A_bf16, A_fp4, A_scale);
}
- template<int C_N, int C_K, int C_SCALEN, int C_TOTAL_KT, int C_KPS, int C_NUM_KSPLIT>
+ template<int C_N, int C_K, int C_SCALEN, int C_TOTAL_KT, int C_KPS, int C_NUM_KSPLIT, int C_M>
void launch_gemm_nk(
const uint8_t* A, const uint8_t* As,
const uint8_t* Bsh, const uint8_t* Bssh,
- float* C_partial, uint16_t* C_final, int M)
+ float* C_partial, uint16_t* C_final)
{
- constexpr bool do_splitk = C_NUM_KSPLIT > 1;
+ constexpr auto do_splitk = C_NUM_KSPLIT > 1;
- auto launch = [&]<int BM, int BN, int NWARPS>() {
- static_assert(((BM + 15) / 16) * (BN / 16) == NWARPS);
+ // M→BM dispatch resolved at compile time
+ constexpr auto BM = (C_M <= 8) ? 8 : (C_M <= 32) ? 16 : (C_M <= 128) ? 32 : 64;
+ constexpr auto BN = 32;
+ constexpr auto NWARPS = ((BM + 15) / 16) * (BN / 16);
+ static_assert(((BM + 15) / 16) * (BN / 16) == NWARPS);
- // Compute LDS size for pipelined path
- constexpr int WAVES_M = (BM + 15) / 16;
- constexpr int WAVES_N = BN / 16;
- constexpr int CHUNK_K = (C_KPS >= 4) ? 4 : C_KPS;
- constexpr int smem_size = (C_KPS > 4)
- ? LdsLayout<WAVES_M, WAVES_N, CHUNK_K>::TOTAL_LDS
- : 0;
+ constexpr auto WAVES_M = (BM + 15) / 16;
+ constexpr auto WAVES_N = BN / 16;
+ constexpr auto CHUNK_K = (C_KPS >= 4) ? 4 : C_KPS;
+ constexpr auto smem_size = (C_KPS > 4)
+ ? LdsLayout<WAVES_M, WAVES_N, CHUNK_K>::TOTAL_LDS
+ : 0;
- const auto full_m = (BM >= 16) ? M / BM : 0;
- constexpr int full_n = C_N / BN;
- const auto total_m = (M + BM - 1) / BM;
- constexpr int total_n = (C_N + BN - 1) / BN;
- const auto edge_m = total_m - full_m;
- constexpr int edge_n = total_n - full_n;
+ constexpr auto full_m = (BM >= 16) ? C_M / BM : 0;
+ constexpr auto full_n = C_N / BN;
+ constexpr auto total_m = (C_M + BM - 1) / BM;
+ constexpr auto total_n = (C_N + BN - 1) / BN;
+ constexpr auto edge_m = total_m - full_m;
+ constexpr auto edge_n = total_n - full_n;
- const dim3 block{static_cast<uint32_t>(NWARPS * 64)};
+ constexpr dim3 block{static_cast<uint32_t>(NWARPS * 64)};
- // Hoist hipFuncSetAttribute outside the sub lambda (avoids constexpr capture issues)
- if (smem_size > 0) {
- // Set max dynamic shared memory for ALL template instantiations we'll launch
- // hipFuncSetAttribute is safe to call even for sizes <= 48KB
- if constexpr (do_splitk) {
- (void)hipFuncSetAttribute((const void*)mxfp4_gemm<BM,BN,NWARPS,true,true,true,C_KPS,C_N,C_K,C_SCALEN,C_TOTAL_KT,C_KPS>,
+ // Helper: set smem + launch for a specific (AV, BV, ox, oy) combo
+ auto sub = [&]<bool AV, bool BV>() {
+ constexpr auto gx = BV ? full_n : edge_n;
+ constexpr auto gy = AV ? full_m : edge_m;
+ constexpr auto ox = BV ? 0 : full_n;
+ constexpr auto oy = AV ? 0 : full_m;
+ if constexpr (gx > 0 && gy > 0) {
+ constexpr auto SK = do_splitk;
+ if constexpr (smem_size > 0)
+ (void)hipFuncSetAttribute(
+ (const void*)mxfp4_gemm<BM,BN,NWARPS,SK,AV,BV,C_KPS,C_N,C_K,C_SCALEN,C_TOTAL_KT,C_KPS,C_M,ox,oy>,
hipFuncAttributeMaxDynamicSharedMemorySize, smem_size);
- (void)hipFuncSetAttribute((const void*)mxfp4_gemm<BM,BN,NWARPS,true,true,false,C_KPS,C_N,C_K,C_SCALEN,C_TOTAL_KT,C_KPS>,
- hipFuncAttributeMaxDynamicSharedMemorySize, smem_size);
- (void)hipFuncSetAttribute((const void*)mxfp4_gemm<BM,BN,NWARPS,true,false,true,C_KPS,C_N,C_K,C_SCALEN,C_TOTAL_KT,C_KPS>,
- hipFuncAttributeMaxDynamicSharedMemorySize, smem_size);
- (void)hipFuncSetAttribute((const void*)mxfp4_gemm<BM,BN,NWARPS,true,false,false,C_KPS,C_N,C_K,C_SCALEN,C_TOTAL_KT,C_KPS>,
- hipFuncAttributeMaxDynamicSharedMemorySize, smem_size);
- } else {
- (void)hipFuncSetAttribute((const void*)mxfp4_gemm<BM,BN,NWARPS,false,true,true,C_KPS,C_N,C_K,C_SCALEN,C_TOTAL_KT,C_KPS>,
- hipFuncAttributeMaxDynamicSharedMemorySize, smem_size);
- (void)hipFuncSetAttribute((const void*)mxfp4_gemm<BM,BN,NWARPS,false,true,false,C_KPS,C_N,C_K,C_SCALEN,C_TOTAL_KT,C_KPS>,
- hipFuncAttributeMaxDynamicSharedMemorySize, smem_size);
- (void)hipFuncSetAttribute((const void*)mxfp4_gemm<BM,BN,NWARPS,false,false,true,C_KPS,C_N,C_K,C_SCALEN,C_TOTAL_KT,C_KPS>,
- hipFuncAttributeMaxDynamicSharedMemorySize, smem_size);
- (void)hipFuncSetAttribute((const void*)mxfp4_gemm<BM,BN,NWARPS,false,false,false,C_KPS,C_N,C_K,C_SCALEN,C_TOTAL_KT,C_KPS>,
- hipFuncAttributeMaxDynamicSharedMemorySize, smem_size);
- }
- }
-
- auto sub = [&]<bool AV, bool BV>(int gx, int gy, int ox, int oy) {
- if (gx <= 0 || gy <= 0) return;
- const dim3 grid{
+ constexpr dim3 grid{
static_cast<uint32_t>(gx),
static_cast<uint32_t>(gy),
static_cast<uint32_t>(C_NUM_KSPLIT)
};
-
- if constexpr (do_splitk)
- mxfp4_gemm<BM,BN,NWARPS,true,AV,BV,C_KPS,C_N,C_K,C_SCALEN,C_TOTAL_KT,C_KPS>
- <<<grid,block,smem_size>>>(A,As,Bsh,Bssh,C_partial,nullptr,M,ox,oy);
+ if constexpr (SK)
+ mxfp4_gemm<BM,BN,NWARPS,true,AV,BV,C_KPS,C_N,C_K,C_SCALEN,C_TOTAL_KT,C_KPS,C_M,ox,oy>
+ <<<grid,block,smem_size>>>(A,As,Bsh,Bssh,C_partial,nullptr);
else
- mxfp4_gemm<BM,BN,NWARPS,false,AV,BV,C_KPS,C_N,C_K,C_SCALEN,C_TOTAL_KT,C_KPS>
- <<<grid,block,smem_size>>>(A,As,Bsh,Bssh,nullptr,C_final,M,ox,oy);
- };
-
- sub.template operator()<true, true >(full_n, full_m, 0, 0);
- sub.template operator()<true, false>(edge_n, full_m, full_n, 0);
- sub.template operator()<false, true >(full_n, edge_m, 0, full_m);
- sub.template operator()<false, false>(edge_n, edge_m, full_n, full_m);
-
- if constexpr (do_splitk) {
- const dim3 rblock{32, 16};
- const dim3 rgrid{
- static_cast<uint32_t>((C_N + 31) / 32),
- static_cast<uint32_t>((M + 15) / 16)
- };
- mxfp4_reduce<C_N, C_NUM_KSPLIT><<<rgrid, rblock>>>(C_partial, C_final, M);
+ mxfp4_gemm<BM,BN,NWARPS,false,AV,BV,C_KPS,C_N,C_K,C_SCALEN,C_TOTAL_KT,C_KPS,C_M,ox,oy>
+ <<<grid,block,smem_size>>>(A,As,Bsh,Bssh,nullptr,C_final);
}
};
- if (M <= 8) launch.template operator()< 8, 32, 2>();
- else if (M <= 16) launch.template operator()< 16, 32, 2>();
- else if (M <= 32) launch.template operator()< 16, 32, 2>();
- else if (M <= 64) launch.template operator()< 32, 32, 4>();
- else if (M <=128) launch.template operator()< 32, 32, 4>();
- else launch.template operator()< 64, 32, 8>();
+ sub.template operator()<true, true >();
+ sub.template operator()<true, false>();
+ sub.template operator()<false, true >();
+ sub.template operator()<false, false>();
+
+ if constexpr (do_splitk) {
+ constexpr dim3 rblock{32, 16};
+ constexpr dim3 rgrid{
+ static_cast<uint32_t>((C_N + 31) / 32),
+ static_cast<uint32_t>((C_M + 15) / 16)
+ };
+ mxfp4_reduce<C_N, C_NUM_KSPLIT, C_M><<<rgrid, rblock>>>(C_partial, C_final);
+ }
}
- template<int C_N, int C_K, int C_SCALEN, int C_TOTAL_KT>
+ template<int C_N, int C_K, int C_SCALEN, int C_TOTAL_KT, int C_M>
void launch_gemm_nk_7168(
const uint8_t* A, const uint8_t* As,
const uint8_t* Bsh, const uint8_t* Bssh,
- float* C_partial, uint16_t* C_final, int M)
+ float* C_partial, uint16_t* C_final)
{
- if (M <= 8)
- launch_gemm_nk<C_N, C_K, C_SCALEN, C_TOTAL_KT, 8, 7>(A, As, Bsh, Bssh, C_partial, C_final, M);
- else if (M <= 16)
- launch_gemm_nk<C_N, C_K, C_SCALEN, C_TOTAL_KT, 4, 14>(A, As, Bsh, Bssh, C_partial, C_final, M);
+ if constexpr (C_M <= 8)
+ launch_gemm_nk<C_N, C_K, C_SCALEN, C_TOTAL_KT, 8, 7, C_M>(A, As, Bsh, Bssh, C_partial, C_final);
+ else if constexpr (C_M <= 16)
+ launch_gemm_nk<C_N, C_K, C_SCALEN, C_TOTAL_KT, 4, 14, C_M>(A, As, Bsh, Bssh, C_partial, C_final);
else
- launch_gemm_nk<C_N, C_K, C_SCALEN, C_TOTAL_KT, 56, 1>(A, As, Bsh, Bssh, C_partial, C_final, M);
+ launch_gemm_nk<C_N, C_K, C_SCALEN, C_TOTAL_KT, 56, 1, C_M>(A, As, Bsh, Bssh, C_partial, C_final);
}
- extern "C" void launch_gemm_raw(
- const uint8_t* A_fp4, const uint8_t* A_scale,
+ // Combined quant + GEMM dispatch — fully compile-time
+ template<int C_M, int C_K, int C_N, int C_SCALEN, int C_TOTAL_KT, int C_KPS, int C_NUM_KSPLIT>
+ void launch_shape(
+ const __bf16* A_bf16, uint8_t* A_fp4, uint8_t* A_scale,
const uint8_t* Bsh, const uint8_t* Bssh,
+ float* C_partial, uint16_t* C_final)
+ {
+ launch_quant<C_M, C_K>(A_bf16, A_fp4, A_scale);
+ launch_gemm_nk<C_N, C_K, C_SCALEN, C_TOTAL_KT, C_KPS, C_NUM_KSPLIT, C_M>(
+ A_fp4, A_scale, Bsh, Bssh, C_partial, C_final);
+ }
+
+ template<int C_M, int C_K, int C_N, int C_SCALEN, int C_TOTAL_KT>
+ void launch_shape_7168(
+ const __bf16* A_bf16, uint8_t* A_fp4, uint8_t* A_scale,
+ const uint8_t* Bsh, const uint8_t* Bssh,
+ float* C_partial, uint16_t* C_final)
+ {
+ launch_quant<C_M, C_K>(A_bf16, A_fp4, A_scale);
+ launch_gemm_nk_7168<C_N, C_K, C_SCALEN, C_TOTAL_KT, C_M>(
+ A_fp4, A_scale, Bsh, Bssh, C_partial, C_final);
+ }
+
+ // Runtime M dispatch
+ template<int C_K, int C_N, int C_SCALEN, int C_TOTAL_KT, int C_KPS, int C_NUM_KSPLIT>
+ void dispatch_shape(
+ const __bf16* A_bf16, uint8_t* A_fp4, uint8_t* A_scale,
+ const uint8_t* Bsh, const uint8_t* Bssh,
+ float* C_partial, uint16_t* C_final, int M)
+ {
+ #define CASE_M(val) case val: launch_shape<val,C_K,C_N,C_SCALEN,C_TOTAL_KT,C_KPS,C_NUM_KSPLIT>(A_bf16,A_fp4,A_scale,Bsh,Bssh,C_partial,C_final); return
+ switch (M) {
+ CASE_M(4); CASE_M(8); CASE_M(16); CASE_M(32); CASE_M(64); CASE_M(256);
+ }
+ #undef CASE_M
+ }
+
+ template<int C_K, int C_N, int C_SCALEN, int C_TOTAL_KT>
+ void dispatch_shape_7168(
+ const __bf16* A_bf16, uint8_t* A_fp4, uint8_t* A_scale,
+ const uint8_t* Bsh, const uint8_t* Bssh,
+ float* C_partial, uint16_t* C_final, int M)
+ {
+ #define CASE_M(val) case val: launch_shape_7168<val,C_K,C_N,C_SCALEN,C_TOTAL_KT>(A_bf16,A_fp4,A_scale,Bsh,Bssh,C_partial,C_final); return
+ switch (M) {
+ CASE_M(4); CASE_M(8); CASE_M(16); CASE_M(32); CASE_M(64); CASE_M(256);
+ }
+ #undef CASE_M
+ }
+
+ extern "C" void launch_all(
+ const __bf16* A_bf16, uint8_t* A_fp4, uint8_t* A_scale,
+ const uint8_t* Bsh, const uint8_t* Bssh,
float* C_partial, uint16_t* C_final,
int M, int N, int K)
{
if (N == 2880 && K == 512)
- launch_gemm_nk<2880, 512, 16, 4, 4, 1>(A_fp4, A_scale, Bsh, Bssh, C_partial, C_final, M);
+ dispatch_shape<512, 2880, 16, 4, 4, 1>(A_bf16, A_fp4, A_scale, Bsh, Bssh, C_partial, C_final, M);
else if (N == 2112 && K == 7168)
- launch_gemm_nk_7168<2112, 7168, 224, 56>(A_fp4, A_scale, Bsh, Bssh, C_partial, C_final, M);
+ dispatch_shape_7168<7168, 2112, 224, 56>(A_bf16, A_fp4, A_scale, Bsh, Bssh, C_partial, C_final, M);
else if (N == 4096 && K == 512)
- launch_gemm_nk<4096, 512, 16, 4, 4, 1>(A_fp4, A_scale, Bsh, Bssh, C_partial, C_final, M);
+ dispatch_shape<512, 4096, 16, 4, 4, 1>(A_bf16, A_fp4, A_scale, Bsh, Bssh, C_partial, C_final, M);
else if (N == 7168 && K == 2048)
- launch_gemm_nk<7168, 2048, 64, 16, 16, 1>(A_fp4, A_scale, Bsh, Bssh, C_partial, C_final, M);
+ dispatch_shape<2048, 7168, 64, 16, 16, 1>(A_bf16, A_fp4, A_scale, Bsh, Bssh, C_partial, C_final, M);
else if (N == 3072 && K == 1536)
- launch_gemm_nk<3072, 1536, 48, 12, 12, 1>(A_fp4, A_scale, Bsh, Bssh, C_partial, C_final, M);
+ dispatch_shape<1536, 3072, 48, 12, 12, 1>(A_bf16, A_fp4, A_scale, Bsh, Bssh, C_partial, C_final, M);
}
"""
⋯ 2 unchanged lines
#include <torch/extension.h>
#include <c10/core/DeviceGuard.h>
- extern "C" void launch_quant(const __bf16*, uint8_t*, uint8_t*, int, int);
- extern "C" void launch_gemm_raw(const uint8_t*, const uint8_t*, const uint8_t*, const uint8_t*,
- float*, uint16_t*, int, int, int);
+ extern "C" void launch_all(const __bf16*, uint8_t*, uint8_t*,
+ const uint8_t*, const uint8_t*,
+ float*, uint16_t*, int, int, int);
static int get_num_ksplit(int M, int K) {
auto total_ktiles = K / 128;
⋯ 24 unchanged lines
};
static ShapeWorkspace g_ws[10];
- static int g_ws_count = 0;
+ static auto g_ws_count = 0;
static BCache g_bcache;
static auto* find_or_create_ws(int M, int N, int K, int num_ksplit,
const at::TensorOptions& opts) {
- for (int i = 0; i < g_ws_count; ++i) {
+ for (auto i = 0; i < g_ws_count; ++i) {
if (g_ws[i].M == M && g_ws[i].N == N && g_ws[i].K == K)
return &g_ws[i];
}
⋯ 48 unchanged lines
const auto num_ksplit = get_num_ksplit(M, K);
auto* ws = find_or_create_ws(M, N, K, num_ksplit, A.options());
- launch_quant(a_bf16_ptr, ws->a_fp4_ptr, ws->a_scale_ptr, M, K);
-
- launch_gemm_raw(
- ws->a_fp4_ptr, ws->a_scale_ptr,
+ launch_all(
+ a_bf16_ptr, ws->a_fp4_ptr, ws->a_scale_ptr,
g_bcache.bsh_ptr, g_bcache.bssh_ptr,
ws->c_partial_ptr, ws->c_final_ptr,
M, N, K);
⋯ 37 unchanged lines
def custom_kernel(data: Tuple[torch.Tensor, ...]) -> torch.Tensor:
- """MXFP4 GEMM v25d_v5: v4 + nt on A loads in PATH B."""
+ """MXFP4 GEMM v25d_v4_i1: v4_i0 + quant intrinsic dest_sel + LdsLayout + launch_bounds."""
A, _, B_q, B_shuffle, B_scale_sh = data
return _ext.fwd(A.cuda(), B_q, B_shuffle, B_scale_sh)
scrolls · 1091 diff lines total

Best evidence level for this revision: reported

JSON