Skip to content
KernelIndex
Search⌘K

submission 721028

jiab_85281 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

mm_best.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-721028?include=source"
interfacepython
Compatibility
measured onAMD Instinct MI355X
declared hardwareAMD Instinct MI355X
architecturesgfx950
dtypesbf16, mxfp4

Benchmark evidence

1 measurement across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
AMD MXFP4 GEMMsuite of 6 cases
AMD Instinct MI355X
8.70µs
#81 of 1143
2026-04-04

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:8b2fce1267afcbca99623003a09138d0d75fc084c6ec54b04c77232e250b2699
license declaredunknown
license concludedunknown
authorsjiab_85281
imported2026-08-15

Techniques

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

shared-memory__shared__ Lds lds;
split-knamespace splitk {

Kernel source

mm_best.py704 lines
# Optimized MXFP4 GEMM: shape6 with 12 waves (all active), shape5 original 16 waves
import os
os.environ['PYTORCH_ROCM_ARCH'] = 'gfx950'

from torch.utils.cpp_extension import load_inline
from task import input_t, output_t

CPP_WRAPPER = """
void fp8_mm(torch::Tensor a_full, torch::Tensor b_full, torch::Tensor b_shuffle, torch::Tensor b_scales, torch::Tensor c);
"""

cuda_src = """
#include <torch/extension.h>
#include <cstddef>
#include <cstdint>
#include <hip/amd_detail/amd_hip_fp8.h>
#include <hip/amd_detail/amd_hip_bf16.h>

constexpr int BLOCK = 128;
constexpr int SCALE_GROUP = 32;
constexpr int MMA_M = 16;
constexpr int MMA_N = 16;
constexpr int WAVE = 64;

typedef uint32_t u32x4 __attribute__((ext_vector_type(4)));
typedef uint32_t u32x8 __attribute__((ext_vector_type(8)));
typedef float f32x4 __attribute__((ext_vector_type(4)));
typedef __bf16 bf16x2 __attribute__((ext_vector_type(2)));

union FragRegs { u32x4 v4; u32x8 v8; uint64_t v2[4]; };

__host__ __device__ __forceinline__ int cdiv(int x, int y) { return (x + y - 1) / y; }
__device__ __forceinline__ int lid() { return threadIdx.x & 63; }
__device__ __forceinline__ int wid() { return threadIdx.x / WAVE; }

// ============================================================
// MXFP4 quantization
// ============================================================
__device__ __forceinline__ uint32_t pk_max_u16(uint32_t a, uint32_t b) {
    uint32_t r; asm("v_pk_max_u16 %0, %1, %2" : "=v"(r) : "v"(a), "v"(b)); return r;
}
__device__ __forceinline__ uint32_t reduce_pk_u16(uint32_t v) {
    const uint32_t hi = v >> 16, lo = v & 0xFFFFu; return hi > lo ? hi : lo;
}
__device__ __forceinline__ void mxfp4_scale(uint32_t mx, uint8_t& sb, float& sf) {
    float mf = __builtin_bit_cast(float, mx << 16);
    uint32_t re = (__builtin_bit_cast(uint32_t, mf) + 0x00200000u) >> 23;
    sb = (re > 2u) ? uint8_t(re - 2u) : uint8_t(0u);
    sf = __builtin_bit_cast(float, uint32_t(sb) << 23);
}
__device__ __forceinline__ uint32_t pack_fp4(const uint32_t* bp, float sf) {
    uint32_t p = 0u;
    p = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(p, __builtin_bit_cast(bf16x2, bp[0]), sf, 0);
    p = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(p, __builtin_bit_cast(bf16x2, bp[1]), sf, 1);
    p = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(p, __builtin_bit_cast(bf16x2, bp[2]), sf, 2);
    p = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(p, __builtin_bit_cast(bf16x2, bp[3]), sf, 3);
    return p;
}
__device__ __forceinline__ void quant_bf16x32(const u32x4* src, u32x4& packed, uint8_t& sb) {
    constexpr uint32_t SM = 0x7FFF7FFFu;
    uint32_t m0 = pk_max_u16(src[0][0]&SM, src[0][1]&SM);
    uint32_t m1 = pk_max_u16(src[0][2]&SM, src[0][3]&SM);
    uint32_t m2 = pk_max_u16(src[1][0]&SM, src[1][1]&SM);
    uint32_t m3 = pk_max_u16(src[1][2]&SM, src[1][3]&SM);
    uint32_t m4 = pk_max_u16(src[2][0]&SM, src[2][1]&SM);
    uint32_t m5 = pk_max_u16(src[2][2]&SM, src[2][3]&SM);
    uint32_t m6 = pk_max_u16(src[3][0]&SM, src[3][1]&SM);
    uint32_t m7 = pk_max_u16(src[3][2]&SM, src[3][3]&SM);
    m0 = pk_max_u16(m0,m1); m2 = pk_max_u16(m2,m3);
    m4 = pk_max_u16(m4,m5); m6 = pk_max_u16(m6,m7);
    m0 = pk_max_u16(m0,m2); m4 = pk_max_u16(m4,m6);
    m0 = pk_max_u16(m0,m4);
    float sf; mxfp4_scale(reduce_pk_u16(m0), sb, sf);
    const uint32_t* raw = reinterpret_cast<const uint32_t*>(src);
    packed = { pack_fp4(&raw[0], sf), pack_fp4(&raw[4], sf),
               pack_fp4(&raw[8], sf), pack_fp4(&raw[12], sf) };
}

// ============================================================
// Scale addressing
// ============================================================
__device__ __forceinline__ uint32_t padded_scale_cols(int k) {
    return (uint32_t(cdiv(k, SCALE_GROUP)) + 7u) & ~7u;
}
__device__ __forceinline__ size_t shuffled_scale_offset(uint32_t row, uint32_t kg, uint32_t pkg) {
    uint32_t rb = row >> 5, rh = (row >> 4) & 1u, ri = row & 15u;
    uint32_t gb = kg >> 3, gh = (kg >> 2) & 1u, gi = kg & 3u;
    return (size_t(rb) * (pkg >> 3) + gb) * 256u +
           size_t(gi) * 64u + size_t(ri) * 4u + size_t(gh) * 2u + size_t(rh);
}

// ============================================================
// Helper: write output tile
// ============================================================
__device__ __forceinline__ void write_c(
    __hip_bfloat16* c, f32x4 sum, int m0, int tn, int m, int n, int lane
) {
    int col = tn * MMA_N + (lane & 15);
    int rb = m0 + (lane >> 4) * 4;
    uint32_t pk01, pk23;
    asm("v_cvt_pk_bf16_f32 %0, %1, %2" : "=v"(pk01) : "v"(sum[0]), "v"(sum[1]));
    asm("v_cvt_pk_bf16_f32 %0, %1, %2" : "=v"(pk23) : "v"(sum[2]), "v"(sum[3]));
    const uint16_t* v01 = reinterpret_cast<const uint16_t*>(&pk01);
    const uint16_t* v23 = reinterpret_cast<const uint16_t*>(&pk23);
    #pragma unroll
    for (int i = 0; i < 4; ++i) {
        int r = rb + i;
        if (r < m && col < n) {
            uint16_t val = (i < 2) ? v01[i] : v23[i - 2];
            reinterpret_cast<uint16_t*>(c)[size_t(r) * n + col] = val;
        }
    }
}

// ============================================================
// Shape 2: M=16, N=2112, K=7168  (original 16 waves, 4 iters)
// ============================================================
namespace shape2 {

constexpr int WAVES = 16;
constexpr int KT = 56;

struct Lds {
    alignas(16) float reduce[WAVES][WAVE][4];
};

__global__ __launch_bounds__(WAVES * WAVE)
void kernel(
    const __hip_bfloat16* a, const uint8_t* bsh, const uint8_t* bsc,
    __hip_bfloat16* c, int m, int n, int k
) {
    __shared__ Lds lds;
    const int w = wid();
    const int lane = lid();
    const int tn = int(blockIdx.x);
    const int tm = int(blockIdx.y);
    const int m0 = tm * MMA_M;
    const int row = lane & 15, kgrp = lane >> 4;
    const uint32_t pkg = padded_scale_cols(k);

    f32x4 acc = {0.f, 0.f, 0.f, 0.f};

    for (int kk = 0; kk < 4; ++kk) {
        int ki = w + kk * WAVES;
        if (ki >= KT) break;

        int vrows = m - m0;
        vrows = vrows < 0 ? 0 : (vrows > MMA_M ? MMA_M : vrows);
        u32x4 chunks[4] = {};
        if (row < vrows) {
            const char* base = reinterpret_cast<const char*>(a) +
                (size_t(m0) * k + size_t(ki) * BLOCK) * sizeof(__hip_bfloat16);
            const char* lsrc = base + size_t(row * k + kgrp * SCALE_GROUP) * sizeof(__hip_bfloat16);
            const u32x4* p = reinterpret_cast<const u32x4*>(lsrc);
            chunks[0] = p[0]; chunks[1] = p[1]; chunks[2] = p[2]; chunks[3] = p[3];
        }
        u32x4 a_packed; uint8_t a_sb;
        quant_bf16x32(chunks, a_packed, a_sb);
        uint32_t a_sc = a_sb;

        u32x4 b_data = {};
        uint32_t b_sc = 0;
        int vraw = n - tn * MMA_N;
        int valid = vraw < 0 ? 0 : (vraw > MMA_N ? MMA_N : vraw);
        if (row < uint32_t(valid)) {
            size_t toff = size_t(tn) * MMA_N * (k / 2) + size_t(ki) * MMA_N * (BLOCK / 2);
            uint32_t off16 = row * 16 + kgrp * 256;
            const u32x4* bptr = reinterpret_cast<const u32x4*>(bsh + toff + off16);
            b_data = *bptr;
            uint32_t grow = tn * MMA_N + row;
            uint32_t gkg = ki * (BLOCK / SCALE_GROUP) + kgrp;
            b_sc = reinterpret_cast<const uint8_t*>(bsc)[
                shuffled_scale_offset(grow, gkg, pkg)];
        }

        FragRegs ar, br;
        ar.v4 = a_packed;
        br.v4 = b_data;
        acc = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
            ar.v8, br.v8, acc, 4, 4, 0, a_sc, 0, b_sc);
    }

    lds.reduce[w][lane][0] = acc[0];
    lds.reduce[w][lane][1] = acc[1];
    lds.reduce[w][lane][2] = acc[2];
    lds.reduce[w][lane][3] = acc[3];
    __syncthreads();

    if (w == 0) {
        f32x4 sum = {0.f, 0.f, 0.f, 0.f};
        #pragma unroll
        for (int s = 0; s < WAVES; ++s) {
            sum[0] += lds.reduce[s][lane][0];
            sum[1] += lds.reduce[s][lane][1];
            sum[2] += lds.reduce[s][lane][2];
            sum[3] += lds.reduce[s][lane][3];
        }
        write_c(c, sum, m0, tn, m, n, lane);
    }
}
} // shape2

// ============================================================
// Shape 5: M=64, N=7168, K=2048 (original 16 waves, 4 N-tiles)
// ============================================================
namespace shape5 {

constexpr int WAVES = 16;
constexpr int N_TILES_PER_CTA = 4;
constexpr int S5_N = 7168;
constexpr int S5_K = 2048;
constexpr int S5_PKG = 64;

struct Lds {
    alignas(16) float reduce[N_TILES_PER_CTA][WAVES][WAVE][4];
};

__global__ __launch_bounds__(WAVES * WAVE)
void kernel(
    const __hip_bfloat16* a, const uint8_t* bsh, const uint8_t* bsc,
    __hip_bfloat16* c, int m, int n, int k
) {
    __shared__ Lds lds;
    const int w = wid();
    const int lane = lid();
    const int tm = int(blockIdx.y);
    const int ng = int(blockIdx.x);
    const int m0 = tm * MMA_M;
    const int row = lane & 15, kgrp = lane >> 4;
    const int ki = w;

    const char* base = reinterpret_cast<const char*>(a) +
        (size_t(m0) * S5_K + size_t(ki) * BLOCK) * sizeof(__hip_bfloat16);
    const char* lsrc = base + size_t(row * S5_K + kgrp * SCALE_GROUP) * sizeof(__hip_bfloat16);
    const u32x4* p = reinterpret_cast<const u32x4*>(lsrc);
    u32x4 chunks[4];
    chunks[0] = p[0]; chunks[1] = p[1]; chunks[2] = p[2]; chunks[3] = p[3];

    u32x4 a_packed; uint8_t a_sb;
    quant_bf16x32(chunks, a_packed, a_sb);
    uint32_t a_sc = a_sb;
    FragRegs ar; ar.v4 = a_packed;

    #pragma unroll
    for (int nt = 0; nt < N_TILES_PER_CTA; ++nt) {
        int tn = ng * N_TILES_PER_CTA + nt;
        size_t toff = size_t(tn) * MMA_N * (S5_K / 2) + size_t(ki) * MMA_N * (BLOCK / 2);
        uint32_t off16 = row * 16 + kgrp * 256;
        const u32x4* bptr = reinterpret_cast<const u32x4*>(bsh + toff + off16);
        u32x4 b_data = *bptr;

        uint32_t grow = tn * MMA_N + row;
        uint32_t gkg = ki * (BLOCK / SCALE_GROUP) + kgrp;
        uint32_t b_sc = reinterpret_cast<const uint8_t*>(bsc)[
            shuffled_scale_offset(grow, gkg, S5_PKG)];

        f32x4 acc = {0.f, 0.f, 0.f, 0.f};
        FragRegs br; br.v4 = b_data;
        acc = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
            ar.v8, br.v8, acc, 4, 4, 0, a_sc, 0, b_sc);

        lds.reduce[nt][w][lane][0] = acc[0];
        lds.reduce[nt][w][lane][1] = acc[1];
        lds.reduce[nt][w][lane][2] = acc[2];
        lds.reduce[nt][w][lane][3] = acc[3];
    }

    __syncthreads();

    if (w < N_TILES_PER_CTA) {
        int tn = ng * N_TILES_PER_CTA + w;
        f32x4 sum = {0.f, 0.f, 0.f, 0.f};
        #pragma unroll
        for (int s = 0; s < WAVES; ++s) {
            sum[0] += lds.reduce[w][s][lane][0];
            sum[1] += lds.reduce[w][s][lane][1];
            sum[2] += lds.reduce[w][s][lane][2];
            sum[3] += lds.reduce[w][s][lane][3];
        }
        int col = tn * MMA_N + (lane & 15);
        int rb = m0 + (lane >> 4) * 4;
        uint32_t pk01, pk23;
        asm("v_cvt_pk_bf16_f32 %0, %1, %2" : "=v"(pk01) : "v"(sum[0]), "v"(sum[1]));
        asm("v_cvt_pk_bf16_f32 %0, %1, %2" : "=v"(pk23) : "v"(sum[2]), "v"(sum[3]));
        const uint16_t* v01 = reinterpret_cast<const uint16_t*>(&pk01);
        const uint16_t* v23 = reinterpret_cast<const uint16_t*>(&pk23);
        reinterpret_cast<uint16_t*>(c)[size_t(rb) * S5_N + col] = v01[0];
        reinterpret_cast<uint16_t*>(c)[size_t(rb + 1) * S5_N + col] = v01[1];
        reinterpret_cast<uint16_t*>(c)[size_t(rb + 2) * S5_N + col] = v23[0];
        reinterpret_cast<uint16_t*>(c)[size_t(rb + 3) * S5_N + col] = v23[1];
    }
}
} // shape5

// ============================================================
// Shape 6: M=256, N=3072, K=1536
// 12 waves (all active), 6 n-tiles in 2 batches of 3
// ============================================================
namespace shape6 {

constexpr int WAVES = 12;
constexpr int K_TILES = 12;
constexpr int N_TILES_PER_CTA = 6;
constexpr int BATCH = 3;
constexpr int S6_N = 3072;
constexpr int S6_K = 1536;
constexpr int S6_PKG = 48;

struct Lds {
    alignas(16) float reduce[BATCH][K_TILES][WAVE][4];
};

__global__ __launch_bounds__(WAVES * WAVE)
void kernel(
    const __hip_bfloat16* a, const uint8_t* bsh, const uint8_t* bsc,
    __hip_bfloat16* c, int m, int n, int k
) {
    __shared__ Lds lds;
    const int w = wid();
    const int lane = lid();
    const int tm = int(blockIdx.y);
    const int ng = int(blockIdx.x);
    const int m0 = tm * MMA_M;
    const int row = lane & 15, kgrp = lane >> 4;
    const int ki = w;

    // All 12 waves active
    const char* base = reinterpret_cast<const char*>(a) +
        (size_t(m0) * S6_K + size_t(ki) * BLOCK) * sizeof(__hip_bfloat16);
    const char* lsrc = base + size_t(row * S6_K + kgrp * SCALE_GROUP) * sizeof(__hip_bfloat16);
    const u32x4* p = reinterpret_cast<const u32x4*>(lsrc);
    u32x4 chunks[4];
    chunks[0] = p[0]; chunks[1] = p[1]; chunks[2] = p[2]; chunks[3] = p[3];
    u32x4 a_packed; uint8_t a_sb;
    quant_bf16x32(chunks, a_packed, a_sb);
    uint32_t a_sc = a_sb;
    FragRegs ar; ar.v4 = a_packed;

    #pragma unroll
    for (int batch = 0; batch < 2; ++batch) {
        #pragma unroll
        for (int bi = 0; bi < BATCH; ++bi) {
            int nt = batch * BATCH + bi;
            int tn = ng * N_TILES_PER_CTA + nt;

            size_t toff = size_t(tn) * MMA_N * (S6_K / 2) + size_t(ki) * MMA_N * (BLOCK / 2);
            uint32_t off16 = row * 16 + kgrp * 256;
            const u32x4* bptr = reinterpret_cast<const u32x4*>(bsh + toff + off16);
            u32x4 b_data = *bptr;
            uint32_t grow = tn * MMA_N + row;
            uint32_t gkg = ki * (BLOCK / SCALE_GROUP) + kgrp;
            uint32_t b_sc = reinterpret_cast<const uint8_t*>(bsc)[
                shuffled_scale_offset(grow, gkg, S6_PKG)];
            FragRegs br; br.v4 = b_data;
            f32x4 acc = {0.f, 0.f, 0.f, 0.f};
            acc = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
                ar.v8, br.v8, acc, 4, 4, 0, a_sc, 0, b_sc);

            lds.reduce[bi][w][lane][0] = acc[0];
            lds.reduce[bi][w][lane][1] = acc[1];
            lds.reduce[bi][w][lane][2] = acc[2];
            lds.reduce[bi][w][lane][3] = acc[3];
        }

        __syncthreads();

        if (w < BATCH) {
            int nt = batch * BATCH + w;
            int tn = ng * N_TILES_PER_CTA + nt;
            f32x4 sum = {0.f, 0.f, 0.f, 0.f};
            #pragma unroll
            for (int s = 0; s < K_TILES; ++s) {
                sum[0] += lds.reduce[w][s][lane][0];
                sum[1] += lds.reduce[w][s][lane][1];
                sum[2] += lds.reduce[w][s][lane][2];
                sum[3] += lds.reduce[w][s][lane][3];
            }
            int col = tn * MMA_N + (lane & 15);
            int rb = m0 + (lane >> 4) * 4;
            uint32_t pk01, pk23;
            asm("v_cvt_pk_bf16_f32 %0, %1, %2" : "=v"(pk01) : "v"(sum[0]), "v"(sum[1]));
            asm("v_cvt_pk_bf16_f32 %0, %1, %2" : "=v"(pk23) : "v"(sum[2]), "v"(sum[3]));
            const uint16_t* v01 = reinterpret_cast<const uint16_t*>(&pk01);
            const uint16_t* v23 = reinterpret_cast<const uint16_t*>(&pk23);
            reinterpret_cast<uint16_t*>(c)[size_t(rb) * S6_N + col] = v01[0];
            reinterpret_cast<uint16_t*>(c)[size_t(rb + 1) * S6_N + col] = v01[1];
            reinterpret_cast<uint16_t*>(c)[size_t(rb + 2) * S6_N + col] = v23[0];
            reinterpret_cast<uint16_t*>(c)[size_t(rb + 3) * S6_N + col] = v23[1];
        }
        __syncthreads();
    }
}
} // shape6

// ============================================================
// Generic nosplit (shapes with kt<=4)
// ============================================================
namespace nosplit {

struct Lds {
    alignas(16) float reduce[4][WAVE][4];
};

__global__ void kernel(
    const __hip_bfloat16* a, const uint8_t* bsh, const uint8_t* bsc,
    __hip_bfloat16* c, int m, int n, int k
) {
    __shared__ Lds lds;
    const int w = wid();
    const int lane = lid();
    const int tm = int(blockIdx.y);
    const int tn = int(blockIdx.x);
    if (tm >= cdiv(m, MMA_M) || tn >= cdiv(n, MMA_N)) return;

    const int ki = w;
    const int m0 = tm * MMA_M;
    const int row = lane & 15, kgrp = lane >> 4;

    int vrows = m - m0; vrows = vrows < 0 ? 0 : (vrows > MMA_M ? MMA_M : vrows);
    u32x4 chunks[4] = {};
    if (row < vrows) {
        const char* base = reinterpret_cast<const char*>(a) +
            (size_t(m0) * k + size_t(ki) * BLOCK) * sizeof(__hip_bfloat16);
        const char* lsrc = base + size_t(row * k + kgrp * SCALE_GROUP) * sizeof(__hip_bfloat16);
        const u32x4* p = reinterpret_cast<const u32x4*>(lsrc);
        chunks[0] = p[0]; chunks[1] = p[1]; chunks[2] = p[2]; chunks[3] = p[3];
    }
    u32x4 a_packed; uint8_t a_sb;
    quant_bf16x32(chunks, a_packed, a_sb);
    uint32_t a_sc = a_sb;

    u32x4 b_data = {};
    uint32_t b_sc = 0;
    int vraw = n - tn * MMA_N;
    int valid = vraw < 0 ? 0 : (vraw > MMA_N ? MMA_N : vraw);
    if (row < uint32_t(valid)) {
        size_t toff = size_t(tn) * MMA_N * (k / 2) + size_t(ki) * MMA_N * (BLOCK / 2);
        uint32_t off16 = row * 16 + kgrp * 256;
        const u32x4* bptr = reinterpret_cast<const u32x4*>(bsh + toff + off16);
        b_data = *bptr;
        uint32_t grow = tn * MMA_N + row;
        uint32_t gkg = ki * (BLOCK / SCALE_GROUP) + kgrp;
        b_sc = reinterpret_cast<const uint8_t*>(bsc)[
            shuffled_scale_offset(grow, gkg, padded_scale_cols(k))];
    }

    f32x4 acc = {0.f, 0.f, 0.f, 0.f};
    FragRegs ar, br; ar.v4 = a_packed; br.v4 = b_data;
    acc = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
        ar.v8, br.v8, acc, 4, 4, 0, a_sc, 0, b_sc);

    lds.reduce[w][lane][0] = acc[0];
    lds.reduce[w][lane][1] = acc[1];
    lds.reduce[w][lane][2] = acc[2];
    lds.reduce[w][lane][3] = acc[3];
    __syncthreads();

    if (w == 0) {
        f32x4 sum = {0.f, 0.f, 0.f, 0.f};
        #pragma unroll
        for (int s = 0; s < 4; ++s) {
            sum[0] += lds.reduce[s][lane][0];
            sum[1] += lds.reduce[s][lane][1];
            sum[2] += lds.reduce[s][lane][2];
            sum[3] += lds.reduce[s][lane][3];
        }
        write_c(c, sum, tm * MMA_M, tn, m, n, lane);
    }
}
} // nosplit

// ============================================================
// Generic split-K (fallback)
// ============================================================
namespace splitk {

constexpr int MAX_WAVES = 16;

struct Lds {
    alignas(16) float reduce[MAX_WAVES][WAVE][4];
};

__global__ void kernel(
    const __hip_bfloat16* a, const uint8_t* bsh, const uint8_t* bsc,
    float* workspace, __hip_bfloat16* c_out,
    int m, int n, int k, int k_split
) {
    __shared__ Lds lds;
    const int waves_per_cta = int(blockDim.x) / WAVE;
    const int w = wid();
    const int lane = lid();
    const int tm = int(blockIdx.y);
    const int tn = int(blockIdx.x);
    const int kz = int(blockIdx.z);
    if (tm >= cdiv(m, MMA_M) || tn >= cdiv(n, MMA_N)) return;

    const int ki = kz * waves_per_cta + w;
    const int k_tiles = cdiv(k, BLOCK);
    const int row = lane & 15, kgrp = lane >> 4;

    f32x4 acc = {0.f, 0.f, 0.f, 0.f};
    if (ki < k_tiles) {
        const int m0 = tm * MMA_M;
        int vrows = m - m0; vrows = vrows < 0 ? 0 : (vrows > MMA_M ? MMA_M : vrows);
        u32x4 chunks[4] = {};
        if (row < vrows) {
            const char* base = reinterpret_cast<const char*>(a) +
                (size_t(m0) * k + size_t(ki) * BLOCK) * sizeof(__hip_bfloat16);
            const char* lsrc = base + size_t(row * k + kgrp * SCALE_GROUP) * sizeof(__hip_bfloat16);
            const u32x4* p = reinterpret_cast<const u32x4*>(lsrc);
            chunks[0] = p[0]; chunks[1] = p[1]; chunks[2] = p[2]; chunks[3] = p[3];
        }
        u32x4 a_packed; uint8_t a_sb;
        quant_bf16x32(chunks, a_packed, a_sb);
        uint32_t a_sc = a_sb;

        u32x4 b_data = {};
        uint32_t b_sc = 0;
        int vraw = n - tn * MMA_N;
        int valid = vraw < 0 ? 0 : (vraw > MMA_N ? MMA_N : vraw);
        if (row < uint32_t(valid)) {
            size_t toff = size_t(tn) * MMA_N * (k / 2) + size_t(ki) * MMA_N * (BLOCK / 2);
            uint32_t off16 = row * 16 + kgrp * 256;
            const u32x4* bptr = reinterpret_cast<const u32x4*>(bsh + toff + off16);
            b_data = *bptr;
            uint32_t grow = tn * MMA_N + row;
            uint32_t gkg = ki * (BLOCK / SCALE_GROUP) + kgrp;
            b_sc = reinterpret_cast<const uint8_t*>(bsc)[
                shuffled_scale_offset(grow, gkg, padded_scale_cols(k))];
        }

        FragRegs ar, br; ar.v4 = a_packed; br.v4 = b_data;
        acc = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
            ar.v8, br.v8, acc, 4, 4, 0, a_sc, 0, b_sc);
    }

    lds.reduce[w][lane][0] = acc[0];
    lds.reduce[w][lane][1] = acc[1];
    lds.reduce[w][lane][2] = acc[2];
    lds.reduce[w][lane][3] = acc[3];
    __syncthreads();

    if (w == 0) {
        f32x4 sum = {0.f, 0.f, 0.f, 0.f};
        for (int s = 0; s < waves_per_cta; ++s) {
            sum[0] += lds.reduce[s][lane][0];
            sum[1] += lds.reduce[s][lane][1];
            sum[2] += lds.reduce[s][lane][2];
            sum[3] += lds.reduce[s][lane][3];
        }

        int col = tn * MMA_N + (lane & 15);
        int rb = tm * MMA_M + (lane >> 4) * 4;

        if (k_split == 1) {
            uint32_t pk01, pk23;
            asm("v_cvt_pk_bf16_f32 %0, %1, %2" : "=v"(pk01) : "v"(sum[0]), "v"(sum[1]));
            asm("v_cvt_pk_bf16_f32 %0, %1, %2" : "=v"(pk23) : "v"(sum[2]), "v"(sum[3]));
            const uint16_t* v01 = reinterpret_cast<const uint16_t*>(&pk01);
            const uint16_t* v23 = reinterpret_cast<const uint16_t*>(&pk23);
            #pragma unroll
            for (int i = 0; i < 4; ++i) {
                int r = rb + i;
                if (r < m && col < n) {
                    uint16_t val = (i < 2) ? v01[i] : v23[i - 2];
                    reinterpret_cast<uint16_t*>(c_out)[size_t(r) * n + col] = val;
                }
            }
        } else {
            #pragma unroll
            for (int i = 0; i < 4; ++i) {
                int r = rb + i;
                if (r < m && col < n)
                    atomicAdd(&workspace[size_t(r) * n + col], sum[i]);
            }

            size_t mn = size_t(m) * n;
            int tile_id = tm * cdiv(n, MMA_N) + tn;
            int* counters = reinterpret_cast<int*>(workspace + mn);
            int done = 0;
            if (lane == 0)
                done = atomicAdd(&counters[tile_id], 1);
            done = __builtin_amdgcn_readfirstlane(done);

            if (done == k_split - 1) {
                #pragma unroll
                for (int i = 0; i < 4; ++i) {
                    int r = rb + i;
                    if (r < m && col < n) {
                        float v = reinterpret_cast<volatile float*>(workspace)[size_t(r) * n + col];
                        uint32_t pk;
                        asm("v_cvt_pk_bf16_f32 %0, %1, %2" : "=v"(pk) : "v"(v), "v"(0.f));
                        reinterpret_cast<uint16_t*>(c_out)[size_t(r) * n + col] =
                            uint16_t(pk & 0xFFFFu);
                    }
                }
            }
        }
    }
}
} // splitk

// ============================================================
// Host dispatch
// ============================================================
void fp8_mm(torch::Tensor a_full, torch::Tensor b_full,
            torch::Tensor b_shuffle, torch::Tensor b_scales, torch::Tensor c) {
    int m = a_full.size(0);
    int n = c.size(1);
    int k = a_full.size(1);

    const auto* ap = reinterpret_cast<const __hip_bfloat16*>(a_full.data_ptr());
    const auto* bsh = reinterpret_cast<const uint8_t*>(b_shuffle.data_ptr());
    const auto* bsc = reinterpret_cast<const uint8_t*>(b_scales.data_ptr());
    auto* cp = reinterpret_cast<__hip_bfloat16*>(c.data_ptr());

    int tm = cdiv(m, MMA_M);
    int tn = cdiv(n, MMA_N);
    int kt = cdiv(k, BLOCK);

    // Shape 2: N=2112, K=7168 (any M)
    if (n == 2112 && k == 7168) {
        shape2::kernel<<<dim3(tn, tm), dim3(shape2::WAVES * WAVE)>>>(
            ap, bsh, bsc, cp, m, n, k);
        return;
    }

    // Shape 5: M=64, N=7168, K=2048
    if (m == 64 && n == 7168 && k == 2048) {
        int ng = cdiv(tn, shape5::N_TILES_PER_CTA);
        shape5::kernel<<<dim3(ng, tm), dim3(shape5::WAVES * WAVE)>>>(ap, bsh, bsc, cp, m, n, k);
        return;
    }

    // Shape 6: M=256, N=3072, K=1536
    if (m == 256 && n == 3072 && k == 1536) {
        int ng = cdiv(tn, shape6::N_TILES_PER_CTA);
        shape6::kernel<<<dim3(ng, tm), dim3(shape6::WAVES * WAVE)>>>(ap, bsh, bsc, cp, m, n, k);
        return;
    }

    if (kt <= 4) {
        nosplit::kernel<<<dim3(tn, tm), dim3(4 * WAVE)>>>(ap, bsh, bsc, cp, m, n, k);
    } else {
        int waves_per_cta = kt;
        int k_split = 1;
        while (waves_per_cta > 16) {
            k_split++;
            waves_per_cta = cdiv(kt, k_split);
        }

        float* wp = reinterpret_cast<float*>(
            reinterpret_cast<__hip_bfloat16*>(b_full.data_ptr()) + m * n);

        splitk::kernel<<<dim3(tn, tm, k_split), dim3(waves_per_cta * WAVE)>>>(
            ap, bsh, bsc, wp, cp, m, n, k, k_split);
    }
}
"""

module = load_inline(
    name='fp8_mm_v2',
    cpp_sources=[CPP_WRAPPER],
    cuda_sources=[cuda_src],
    functions=['fp8_mm'],
    verbose=True,
    extra_cuda_cflags=["--offload-arch=gfx950",
    "-O3",
    "-std=c++20"],
)

import torch

def custom_kernel(data: input_t) -> output_t:
    a_full, b_full, b_fp4, b_shuffle, b_scales = data
    m = a_full.size(0)
    n = b_fp4.size(0)
    k = a_full.size(1)

    flat = b_full.view(-1)
    c = flat.narrow(0, 0, m * n).view(m, n)

    kt = (k + 127) // 128
    is_shape2 = (n == 2112 and k == 7168)
    is_shape5 = (m == 64 and n == 7168 and k == 2048)
    is_shape6 = (m == 256 and n == 3072 and k == 1536)

    # Zero workspace for generic splitk path
    if kt > 4 and not is_shape2 and not is_shape5 and not is_shape6:
        waves_per_cta = kt
        k_split = 1
        while waves_per_cta > 16:
            k_split += 1
            waves_per_cta = (kt + k_split - 1) // k_split
        if k_split > 1:
            ws_bf16 = m * n * 2
            tm = (m + 15) // 16
            tn = (n + 15) // 16
            cnt_bf16 = ((tm * tn * 2) + 1) & ~1
            flat.narrow(0, m * n, ws_bf16 + cnt_bf16).zero_()

    module.fp8_mm(a_full, b_full, b_shuffle, b_scales, c)
    return c
scrolls · 704 lines total

Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0

Best evidence level for this revision: reported

JSON