Skip to content
KernelIndex
Search⌘K

submission 720553

lgc0338 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-720553?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
9.33µs
#159 of 1143
2026-04-04

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:e2b725dfe68317386495100512fa1de5ef4cce288901a5dd88578bdd953e3846
license declaredunknown
license concludedunknown
authorslgc0338
imported2026-08-15

Techniques

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

shared-memory__shared__ float red[NW * 64 * 16];
split-kvoid fused_mfma_splitk_kernel(
vector-width = uint4uint4 b_pf = {};

Kernel source

submission.py788 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X

import os
os.environ.setdefault('PYTORCH_ROCM_ARCH', 'gfx950')
os.environ.setdefault('CXX', 'clang++')

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

# =============================================================================
# v35: Based on v22 code style (proven to compile), incremental additions:
#   - 16×16 kernel for M=4
#   - 32×32 8wf splitK for M=16
#   - 32×96 wide for M=256
# =============================================================================

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

typedef uint8_t fp4x2_t;
typedef fp4x2_t fp4x64_t __attribute__((ext_vector_type(32)));
typedef float fp32x16_t __attribute__((ext_vector_type(16)));
typedef float fp32x4_t __attribute__((ext_vector_type(4)));

__device__ __forceinline__ uint32_t f2u(float f) {
    uint32_t u; __builtin_memcpy(&u, &f, 4); return u;
}
__device__ __forceinline__ float u2f(uint32_t u) {
    float f; __builtin_memcpy(&f, &u, 4); return f;
}

__device__ __forceinline__ uint8_t read_shuffled_scale(
    const uint8_t* s, int n, int ks, int sn
) {
    int br = n / 32, rh = (n & 31) / 16, rl = (n & 31) & 15;
    int bc = ks / 8, ch = (ks & 7) / 4, cl = (ks & 7) & 3;
    return s[br * (sn * 32) + bc * 256 + cl * 64 + rl * 4 + ch * 2 + rh];
}

__device__ __forceinline__ int b_shuffle_addr(int n, int k_byte, int K_half) {
    int n_block = n / 16;
    int n_local = n & 15;
    int k_block = k_byte / 32;
    int k_group = (k_byte & 31) / 16;
    return n_block * (K_half * 16) + k_block * 512 + k_group * 256 + n_local * 16;
}

// ============================================================
// Kernel 0: 32×32 tile (template NW) — for M=32, M=64
// ============================================================
template <int NW, bool EXACT_M = false>
__attribute__((amdgpu_waves_per_eu(2)))
__global__ __launch_bounds__(64 * NW)
void fused_mfma_kernel(
    const uint16_t* __restrict__ A,
    const uint8_t* __restrict__ B_sh,
    const uint8_t* __restrict__ B_scale_sh,
    uint16_t* __restrict__ C,
    const int M, const int N, const int K,
    const int K_half, const int sn_pad
) {
    __shared__ float red[NW * 64 * 16];
    const int warp_id = threadIdx.x / 64;
    const int tid = threadIdx.x & 63;
    const int wf_row = tid & 31;
    const int k_half = tid >> 5;
    const int m_base = blockIdx.y * 32;
    const int n_base = blockIdx.x * 32;
    const int my_m = m_base + wf_row;
    const int my_n = n_base + wf_row;
    const int K_per_wf = K / NW;
    const int k_start = warp_id * K_per_wf;
    const int k_end = k_start + K_per_wf;

    fp32x16_t c_reg = {};
    uint4 b_pf = {};
    uint8_t bs_pf = 127;
    if (my_n < N) {
        int bk = k_start / 2 + k_half * 16;
        b_pf = *(const uint4*)(B_sh + b_shuffle_addr(my_n, bk, K_half));
        bs_pf = read_shuffled_scale(B_scale_sh, my_n, k_start / 32 + k_half, sn_pad);
    }
    uint4 a_pf[4] = {};
    if (EXACT_M || my_m < M) {
        const int a_k0 = k_start + k_half * 32;
        if (a_k0 + 32 <= K) {
            #pragma unroll
            for (int i = 0; i < 4; i++)
                a_pf[i] = reinterpret_cast<const uint4*>(A + my_m * K + a_k0)[i];
        }
    }

    for (int k = k_start; k < k_end; k += 64) {
        // Schedule: loads first → VALU (A quant) → MFMA
        __builtin_amdgcn_sched_group_barrier(0x020, 6, 0);
        __builtin_amdgcn_sched_group_barrier(0x002, 120, 0);
        __builtin_amdgcn_sched_group_barrier(0x008, 1, 0);
        fp4x64_t b_reg = {};
        __builtin_memcpy(&b_reg, &b_pf, 16);
        uint8_t scale_b = bs_pf;

        uint4 a_data[4];
        __builtin_memcpy(a_data, a_pf, sizeof(uint4) * 4);

        int nk = k + 64;
        if (nk < k_end) {
            if (my_n < N) {
                int bk = nk / 2 + k_half * 16;
                b_pf = *(const uint4*)(B_sh + b_shuffle_addr(my_n, bk, K_half));
                bs_pf = read_shuffled_scale(B_scale_sh, my_n, nk / 32 + k_half, sn_pad);
            }
            if (EXACT_M || my_m < M) {
                const int a_k_next = nk + k_half * 32;
                if (a_k_next + 32 <= K) {
                    #pragma unroll
                    for (int i = 0; i < 4; i++)
                        a_pf[i] = reinterpret_cast<const uint4*>(A + my_m * K + a_k_next)[i];
                }
            }
        }

        fp4x64_t a_reg = {};
        uint8_t scale_a = 127;
        if (EXACT_M || my_m < M) {
            const int a_k = k + k_half * 32;
            if (a_k + 32 <= K) {
                const uint16_t* a_u16 = reinterpret_cast<const uint16_t*>(a_data);
                float vals[32];
                float amax = 0.0f;
                #pragma unroll
                for (int i = 0; i < 32; i++) {
                    vals[i] = u2f((uint32_t)a_u16[i] << 16);
                    amax = fmaxf(amax, fabsf(vals[i]));
                }
                uint32_t au = (f2u(amax) + 0x200000u) & 0xFF800000u;
                int exp_biased = (int)((au >> 23) & 0xFFu);
                int sub_i = (exp_biased == 0) ? -127 : max(min(exp_biased - 129, 127), -127);
                scale_a = (uint8_t)(sub_i + 127);
                float qs = u2f((uint32_t)(sub_i + 127) << 23);
                {
                    uint32_t pk[4] = {0, 0, 0, 0};
                    pk[0] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[0], vals[0],  vals[1],  qs, 0);
                    pk[1] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[1], vals[8],  vals[9],  qs, 0);
                    pk[2] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[2], vals[16], vals[17], qs, 0);
                    pk[3] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[3], vals[24], vals[25], qs, 0);
                    pk[0] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[0], vals[2],  vals[3],  qs, 1);
                    pk[1] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[1], vals[10], vals[11], qs, 1);
                    pk[2] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[2], vals[18], vals[19], qs, 1);
                    pk[3] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[3], vals[26], vals[27], qs, 1);
                    pk[0] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[0], vals[4],  vals[5],  qs, 2);
                    pk[1] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[1], vals[12], vals[13], qs, 2);
                    pk[2] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[2], vals[20], vals[21], qs, 2);
                    pk[3] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[3], vals[28], vals[29], qs, 2);
                    pk[0] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[0], vals[6],  vals[7],  qs, 3);
                    pk[1] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[1], vals[14], vals[15], qs, 3);
                    pk[2] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[2], vals[22], vals[23], qs, 3);
                    pk[3] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[3], vals[30], vals[31], qs, 3);
                    __builtin_memcpy(&a_reg, pk, 16);
                }
            }
        }

        c_reg = __builtin_amdgcn_mfma_scale_f32_32x32x64_f8f6f4(
            a_reg, b_reg, c_reg, 4, 4,
            0, (uint32_t)scale_a, 0, (uint32_t)scale_b);
    }

    const int lb = warp_id * 64 * 16 + tid * 16;
    #pragma unroll
    for (int i = 0; i < 16; i++) red[lb + i] = c_reg[i];
    __syncthreads();

    if (warp_id == 0) {
        int c_col = n_base + wf_row;
        if (c_col < N) {
            #pragma unroll
            for (int i = 0; i < 4; i++) {
                #pragma unroll
                for (int j = 0; j < 4; j++) {
                    int c_row = m_base + k_half * 4 + i * 8 + j;
                    if (EXACT_M || c_row < M) {
                        float sum = 0.0f;
                        #pragma unroll
                        for (int w = 0; w < NW; w++)
                            sum += red[w * 64 * 16 + tid * 16 + i * 4 + j];
                        uint32_t bits = f2u(sum);
                        bits += (0x7FFFu + ((bits >> 16) & 1u));
                        C[c_row * N + c_col] = (uint16_t)(bits >> 16);
                    }
                }
            }
        }
    }
}

// ============================================================
// Kernel 1: 32×32 8wf splitK → float32 partials (for M=16)
// ============================================================
__attribute__((amdgpu_waves_per_eu(2)))
__global__ __launch_bounds__(512)
void fused_mfma_splitk_kernel(
    const uint16_t* __restrict__ A,
    const uint8_t* __restrict__ B_sh,
    const uint8_t* __restrict__ B_scale_sh,
    float* __restrict__ partial,
    const int M, const int N, const int K,
    const int K_half, const int sn_pad,
    const int splitK
) {
    const int NW = 8;
    __shared__ float red[NW * 64 * 16];
    const int warp_id = threadIdx.x / 64;
    const int tid = threadIdx.x & 63;
    const int wf_row = tid & 31;
    const int k_half = tid >> 5;
    const int m_base = blockIdx.y * 32;
    const int n_base = blockIdx.x * 32;
    const int split_idx = blockIdx.z;
    const int my_m = m_base + wf_row;
    const int my_n = n_base + wf_row;
    const int K_per_split = K / splitK;
    const int K_split_start = split_idx * K_per_split;
    const int K_per_wf = K_per_split / NW;
    const int k_start = K_split_start + warp_id * K_per_wf;
    const int k_end = k_start + K_per_wf;

    fp32x16_t c_reg = {};
    uint4 b_pf = {};
    uint8_t bs_pf = 127;
    if (my_n < N) {
        int bk = k_start / 2 + k_half * 16;
        b_pf = *(const uint4*)(B_sh + b_shuffle_addr(my_n, bk, K_half));
        bs_pf = read_shuffled_scale(B_scale_sh, my_n, k_start / 32 + k_half, sn_pad);
    }
    uint4 a_pf[4] = {};
    if (my_m < M) {
        const int a_k0 = k_start + k_half * 32;
        if (a_k0 + 32 <= K) {
            #pragma unroll
            for (int i = 0; i < 4; i++)
                a_pf[i] = reinterpret_cast<const uint4*>(A + my_m * K + a_k0)[i];
        }
    }

    for (int k = k_start; k < k_end; k += 64) {
        // Schedule: loads first → VALU (A quant) → MFMA
        __builtin_amdgcn_sched_group_barrier(0x020, 6, 0);
        __builtin_amdgcn_sched_group_barrier(0x002, 120, 0);
        __builtin_amdgcn_sched_group_barrier(0x008, 1, 0);
        fp4x64_t b_reg = {};
        __builtin_memcpy(&b_reg, &b_pf, 16);
        uint8_t scale_b = bs_pf;
        uint4 a_data[4];
        __builtin_memcpy(a_data, a_pf, sizeof(uint4) * 4);
        int nk = k + 64;
        if (nk < k_end) {
            if (my_n < N) {
                int bk = nk / 2 + k_half * 16;
                b_pf = *(const uint4*)(B_sh + b_shuffle_addr(my_n, bk, K_half));
                bs_pf = read_shuffled_scale(B_scale_sh, my_n, nk / 32 + k_half, sn_pad);
            }
            if (my_m < M) {
                const int a_k_next = nk + k_half * 32;
                if (a_k_next + 32 <= K) {
                    #pragma unroll
                    for (int i = 0; i < 4; i++)
                        a_pf[i] = reinterpret_cast<const uint4*>(A + my_m * K + a_k_next)[i];
                }
            }
        }
        fp4x64_t a_reg = {};
        uint8_t scale_a = 127;
        if (my_m < M) {
            const int a_k = k + k_half * 32;
            if (a_k + 32 <= K) {
                const uint16_t* a_u16 = reinterpret_cast<const uint16_t*>(a_data);
                float vals[32];
                float amax = 0.0f;
                #pragma unroll
                for (int i = 0; i < 32; i++) {
                    vals[i] = u2f((uint32_t)a_u16[i] << 16);
                    amax = fmaxf(amax, fabsf(vals[i]));
                }
                uint32_t au = (f2u(amax) + 0x200000u) & 0xFF800000u;
                int exp_biased = (int)((au >> 23) & 0xFFu);
                int sub_i = (exp_biased == 0) ? -127 : max(min(exp_biased - 129, 127), -127);
                scale_a = (uint8_t)(sub_i + 127);
                float qs = u2f((uint32_t)(sub_i + 127) << 23);
                {
                    uint32_t pk[4] = {0, 0, 0, 0};
                    pk[0] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[0], vals[0],  vals[1],  qs, 0);
                    pk[1] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[1], vals[8],  vals[9],  qs, 0);
                    pk[2] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[2], vals[16], vals[17], qs, 0);
                    pk[3] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[3], vals[24], vals[25], qs, 0);
                    pk[0] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[0], vals[2],  vals[3],  qs, 1);
                    pk[1] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[1], vals[10], vals[11], qs, 1);
                    pk[2] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[2], vals[18], vals[19], qs, 1);
                    pk[3] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[3], vals[26], vals[27], qs, 1);
                    pk[0] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[0], vals[4],  vals[5],  qs, 2);
                    pk[1] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[1], vals[12], vals[13], qs, 2);
                    pk[2] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[2], vals[20], vals[21], qs, 2);
                    pk[3] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[3], vals[28], vals[29], qs, 2);
                    pk[0] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[0], vals[6],  vals[7],  qs, 3);
                    pk[1] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[1], vals[14], vals[15], qs, 3);
                    pk[2] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[2], vals[22], vals[23], qs, 3);
                    pk[3] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[3], vals[30], vals[31], qs, 3);
                    __builtin_memcpy(&a_reg, pk, 16);
                }
            }
        }
        c_reg = __builtin_amdgcn_mfma_scale_f32_32x32x64_f8f6f4(
            a_reg, b_reg, c_reg, 4, 4,
            0, (uint32_t)scale_a, 0, (uint32_t)scale_b);
    }

    const int lb = warp_id * 64 * 16 + tid * 16;
    #pragma unroll
    for (int i = 0; i < 16; i++) red[lb + i] = c_reg[i];
    __syncthreads();

    if (warp_id == 0) {
        int c_col = n_base + wf_row;
        if (c_col < N) {
            #pragma unroll
            for (int i = 0; i < 4; i++) {
                #pragma unroll
                for (int j = 0; j < 4; j++) {
                    int c_row = m_base + k_half * 4 + i * 8 + j;
                    if (c_row < M) {
                        float sum = 0.0f;
                        #pragma unroll
                        for (int w = 0; w < 8; w++)
                            sum += red[w * 64 * 16 + tid * 16 + i * 4 + j];
                        partial[(long long)split_idx * M * N + c_row * N + c_col] = sum;
                    }
                }
            }
        }
    }
}

// ============================================================
// Kernel 2: 16×16 tile (for M=4)
// ============================================================
#define SMALL_NW 4
template <bool EXACT_M = false>
__attribute__((amdgpu_waves_per_eu(2)))
__global__ __launch_bounds__(64 * SMALL_NW)
void fused_mfma_16x16_kernel(
    const uint16_t* __restrict__ A,
    const uint8_t* __restrict__ B_sh,
    const uint8_t* __restrict__ B_scale_sh,
    uint16_t* __restrict__ C,
    const int M, const int N, const int K,
    const int K_half, const int sn_pad
) {
    __shared__ float red[SMALL_NW * 64 * 4];
    const int warp_id = threadIdx.x / 64;
    const int tid = threadIdx.x & 63;
    const int lane16 = tid & 15;
    const int k_quarter = tid >> 4;
    const int m_base = blockIdx.y * 16;
    const int n_base = blockIdx.x * 16;
    const int my_m = m_base + lane16;
    const int my_n = n_base + lane16;
    const int K_per_wf = K / SMALL_NW;
    const int k_start = warp_id * K_per_wf;
    const int k_end = k_start + K_per_wf;

    fp32x4_t c_reg = {};
    uint4 b_pf = {};
    uint8_t bs_pf = 127;
    if (my_n < N) {
        int bk = k_start / 2 + k_quarter * 16;
        b_pf = *(const uint4*)(B_sh + b_shuffle_addr(my_n, bk, K_half));
        bs_pf = read_shuffled_scale(B_scale_sh, my_n, k_start / 32 + k_quarter, sn_pad);
    }
    uint4 a_pf[4] = {};
    if (EXACT_M || my_m < M) {
        const int a_k0 = k_start + k_quarter * 32;
        if (a_k0 + 32 <= K) {
            #pragma unroll
            for (int i = 0; i < 4; i++)
                a_pf[i] = reinterpret_cast<const uint4*>(A + my_m * K + a_k0)[i];
        }
    }

    for (int k = k_start; k < k_end; k += 128) {
        __builtin_amdgcn_sched_group_barrier(0x020, 6, 0);
        __builtin_amdgcn_sched_group_barrier(0x002, 120, 0);
        __builtin_amdgcn_sched_group_barrier(0x008, 1, 0);
        fp4x64_t b_reg = {};
        __builtin_memcpy(&b_reg, &b_pf, 16);
        uint8_t scale_b = bs_pf;
        uint4 a_data[4];
        __builtin_memcpy(a_data, a_pf, sizeof(uint4) * 4);
        int nk = k + 128;
        if (nk < k_end) {
            if (my_n < N) {
                int bk = nk / 2 + k_quarter * 16;
                b_pf = *(const uint4*)(B_sh + b_shuffle_addr(my_n, bk, K_half));
                bs_pf = read_shuffled_scale(B_scale_sh, my_n, nk / 32 + k_quarter, sn_pad);
            }
            if (EXACT_M || my_m < M) {
                const int a_k_next = nk + k_quarter * 32;
                if (a_k_next + 32 <= K) {
                    #pragma unroll
                    for (int i = 0; i < 4; i++)
                        a_pf[i] = reinterpret_cast<const uint4*>(A + my_m * K + a_k_next)[i];
                }
            }
        }
        fp4x64_t a_reg = {};
        uint8_t scale_a = 127;
        if (EXACT_M || my_m < M) {
            const int a_k = k + k_quarter * 32;
            if (a_k + 32 <= K) {
                const uint16_t* a_u16 = reinterpret_cast<const uint16_t*>(a_data);
                float vals[32];
                float amax = 0.0f;
                #pragma unroll
                for (int i = 0; i < 32; i++) {
                    vals[i] = u2f((uint32_t)a_u16[i] << 16);
                    amax = fmaxf(amax, fabsf(vals[i]));
                }
                uint32_t au = (f2u(amax) + 0x200000u) & 0xFF800000u;
                int exp_biased = (int)((au >> 23) & 0xFFu);
                int sub_i = (exp_biased == 0) ? -127 : max(min(exp_biased - 129, 127), -127);
                scale_a = (uint8_t)(sub_i + 127);
                float qs = u2f((uint32_t)(sub_i + 127) << 23);
                {
                    uint32_t pk[4] = {0, 0, 0, 0};
                    pk[0] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[0], vals[0],  vals[1],  qs, 0);
                    pk[1] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[1], vals[8],  vals[9],  qs, 0);
                    pk[2] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[2], vals[16], vals[17], qs, 0);
                    pk[3] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[3], vals[24], vals[25], qs, 0);
                    pk[0] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[0], vals[2],  vals[3],  qs, 1);
                    pk[1] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[1], vals[10], vals[11], qs, 1);
                    pk[2] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[2], vals[18], vals[19], qs, 1);
                    pk[3] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[3], vals[26], vals[27], qs, 1);
                    pk[0] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[0], vals[4],  vals[5],  qs, 2);
                    pk[1] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[1], vals[12], vals[13], qs, 2);
                    pk[2] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[2], vals[20], vals[21], qs, 2);
                    pk[3] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[3], vals[28], vals[29], qs, 2);
                    pk[0] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[0], vals[6],  vals[7],  qs, 3);
                    pk[1] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[1], vals[14], vals[15], qs, 3);
                    pk[2] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[2], vals[22], vals[23], qs, 3);
                    pk[3] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[3], vals[30], vals[31], qs, 3);
                    __builtin_memcpy(&a_reg, pk, 16);
                }
            }
        }
        c_reg = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
            a_reg, b_reg, c_reg, 4, 4,
            0, (uint32_t)scale_a, 0, (uint32_t)scale_b);
    }

    const int lb = warp_id * 64 * 4 + tid * 4;
    #pragma unroll
    for (int i = 0; i < 4; i++) red[lb + i] = c_reg[i];
    __syncthreads();

    if (warp_id == 0) {
        int c_col = n_base + lane16;
        if (c_col < N) {
            int group = k_quarter;
            #pragma unroll
            for (int i = 0; i < 4; i++) {
                int c_row = m_base + group * 4 + i;
                if (EXACT_M || c_row < M) {
                    float sum = 0.0f;
                    #pragma unroll
                    for (int w = 0; w < SMALL_NW; w++)
                        sum += red[w * 64 * 4 + tid * 4 + i];
                    uint32_t bits = f2u(sum);
                    bits += (0x7FFFu + ((bits >> 16) & 1u));
                    C[c_row * N + c_col] = (uint16_t)(bits >> 16);
                }
            }
        }
    }
}

// ============================================================
// Kernel 3: 32×96 wide tile (3 MFMA, for M=256)
// ============================================================
#define WIDE96_NW 4
template <bool EXACT_M = false>
__attribute__((amdgpu_waves_per_eu(2)))
__global__ __launch_bounds__(64 * WIDE96_NW)
void fused_mfma_wide96_kernel(
    const uint16_t* __restrict__ A,
    const uint8_t* __restrict__ B_sh,
    const uint8_t* __restrict__ B_scale_sh,
    uint16_t* __restrict__ C,
    const int M, const int N, const int K,
    const int K_half, const int sn_pad
) {
    __shared__ float red[WIDE96_NW * 64 * 48];
    const int warp_id = threadIdx.x / 64;
    const int tid = threadIdx.x & 63;
    const int wf_row = tid & 31;
    const int k_half = tid >> 5;
    const int m_base = blockIdx.y * 32;
    const int n_base = blockIdx.x * 96;
    const int my_m = m_base + wf_row;
    const int K_per_wf = K / WIDE96_NW;
    const int k_start = warp_id * K_per_wf;
    const int k_end = k_start + K_per_wf;

    fp32x16_t c0 = {}, c1 = {}, c2 = {};

    // Prefetch B for 3 N-subtiles
    uint4 b_pf0 = {}, b_pf1 = {}, b_pf2 = {};
    uint8_t bs_pf0 = 127, bs_pf1 = 127, bs_pf2 = 127;
    {
        int bk = k_start / 2 + k_half * 16;
        int n0 = n_base + wf_row, n1 = n_base + 32 + wf_row, n2 = n_base + 64 + wf_row;
        if (n0 < N) { b_pf0 = *(const uint4*)(B_sh + b_shuffle_addr(n0, bk, K_half)); bs_pf0 = read_shuffled_scale(B_scale_sh, n0, k_start/32+k_half, sn_pad); }
        if (n1 < N) { b_pf1 = *(const uint4*)(B_sh + b_shuffle_addr(n1, bk, K_half)); bs_pf1 = read_shuffled_scale(B_scale_sh, n1, k_start/32+k_half, sn_pad); }
        if (n2 < N) { b_pf2 = *(const uint4*)(B_sh + b_shuffle_addr(n2, bk, K_half)); bs_pf2 = read_shuffled_scale(B_scale_sh, n2, k_start/32+k_half, sn_pad); }
    }
    uint4 a_pf[4] = {};
    if (EXACT_M || my_m < M) {
        const int a_k0 = k_start + k_half * 32;
        if (a_k0 + 32 <= K) {
            #pragma unroll
            for (int i = 0; i < 4; i++)
                a_pf[i] = reinterpret_cast<const uint4*>(A + my_m * K + a_k0)[i];
        }
    }

    for (int k = k_start; k < k_end; k += 64) {
        uint4 bc0 = b_pf0, bc1 = b_pf1, bc2 = b_pf2;
        uint8_t bsc0 = bs_pf0, bsc1 = bs_pf1, bsc2 = bs_pf2;
        uint4 a_data[4];
        __builtin_memcpy(a_data, a_pf, sizeof(uint4) * 4);

        int nk = k + 64;
        if (nk < k_end) {
            int bk = nk / 2 + k_half * 16;
            int n0 = n_base + wf_row, n1 = n_base + 32 + wf_row, n2 = n_base + 64 + wf_row;
            if (n0 < N) { b_pf0 = *(const uint4*)(B_sh + b_shuffle_addr(n0, bk, K_half)); bs_pf0 = read_shuffled_scale(B_scale_sh, n0, nk/32+k_half, sn_pad); }
            if (n1 < N) { b_pf1 = *(const uint4*)(B_sh + b_shuffle_addr(n1, bk, K_half)); bs_pf1 = read_shuffled_scale(B_scale_sh, n1, nk/32+k_half, sn_pad); }
            if (n2 < N) { b_pf2 = *(const uint4*)(B_sh + b_shuffle_addr(n2, bk, K_half)); bs_pf2 = read_shuffled_scale(B_scale_sh, n2, nk/32+k_half, sn_pad); }
            if (EXACT_M || my_m < M) {
                const int a_k_next = nk + k_half * 32;
                if (a_k_next + 32 <= K) {
                    #pragma unroll
                    for (int i = 0; i < 4; i++)
                        a_pf[i] = reinterpret_cast<const uint4*>(A + my_m * K + a_k_next)[i];
                }
            }
        }

        fp4x64_t a_reg = {};
        uint8_t scale_a = 127;
        if (EXACT_M || my_m < M) {
            const int a_k = k + k_half * 32;
            if (a_k + 32 <= K) {
                const uint16_t* a_u16 = reinterpret_cast<const uint16_t*>(a_data);
                float vals[32];
                float amax = 0.0f;
                #pragma unroll
                for (int i = 0; i < 32; i++) {
                    vals[i] = u2f((uint32_t)a_u16[i] << 16);
                    amax = fmaxf(amax, fabsf(vals[i]));
                }
                uint32_t au = (f2u(amax) + 0x200000u) & 0xFF800000u;
                int exp_biased = (int)((au >> 23) & 0xFFu);
                int sub_i = (exp_biased == 0) ? -127 : max(min(exp_biased - 129, 127), -127);
                scale_a = (uint8_t)(sub_i + 127);
                float qs = u2f((uint32_t)(sub_i + 127) << 23);
                {
                    uint32_t pk[4] = {0, 0, 0, 0};
                    pk[0] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[0], vals[0],  vals[1],  qs, 0);
                    pk[1] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[1], vals[8],  vals[9],  qs, 0);
                    pk[2] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[2], vals[16], vals[17], qs, 0);
                    pk[3] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[3], vals[24], vals[25], qs, 0);
                    pk[0] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[0], vals[2],  vals[3],  qs, 1);
                    pk[1] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[1], vals[10], vals[11], qs, 1);
                    pk[2] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[2], vals[18], vals[19], qs, 1);
                    pk[3] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[3], vals[26], vals[27], qs, 1);
                    pk[0] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[0], vals[4],  vals[5],  qs, 2);
                    pk[1] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[1], vals[12], vals[13], qs, 2);
                    pk[2] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[2], vals[20], vals[21], qs, 2);
                    pk[3] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[3], vals[28], vals[29], qs, 2);
                    pk[0] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[0], vals[6],  vals[7],  qs, 3);
                    pk[1] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[1], vals[14], vals[15], qs, 3);
                    pk[2] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[2], vals[22], vals[23], qs, 3);
                    pk[3] = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk[3], vals[30], vals[31], qs, 3);
                    __builtin_memcpy(&a_reg, pk, 16);
                }
            }
        }

        // 3 MFMA sharing same A
        { fp4x64_t b_reg = {}; __builtin_memcpy(&b_reg, &bc0, 16);
          c0 = __builtin_amdgcn_mfma_scale_f32_32x32x64_f8f6f4(a_reg, b_reg, c0, 4, 4, 0, (uint32_t)scale_a, 0, (uint32_t)bsc0); }
        { fp4x64_t b_reg = {}; __builtin_memcpy(&b_reg, &bc1, 16);
          c1 = __builtin_amdgcn_mfma_scale_f32_32x32x64_f8f6f4(a_reg, b_reg, c1, 4, 4, 0, (uint32_t)scale_a, 0, (uint32_t)bsc1); }
        { fp4x64_t b_reg = {}; __builtin_memcpy(&b_reg, &bc2, 16);
          c2 = __builtin_amdgcn_mfma_scale_f32_32x32x64_f8f6f4(a_reg, b_reg, c2, 4, 4, 0, (uint32_t)scale_a, 0, (uint32_t)bsc2); }
    }

    // LDS reduce + output all 3 subtiles
    const int lb = warp_id * 64 * 48 + tid * 48;
    #pragma unroll
    for (int i = 0; i < 16; i++) { red[lb + i] = c0[i]; red[lb + 16 + i] = c1[i]; red[lb + 32 + i] = c2[i]; }
    __syncthreads();

    if (warp_id == 0) {
        for (int s = 0; s < 3; s++) {
            int c_col = n_base + s * 32 + wf_row;
            if (c_col < N) {
                #pragma unroll
                for (int i = 0; i < 4; i++) {
                    #pragma unroll
                    for (int j = 0; j < 4; j++) {
                        int c_row = m_base + k_half * 4 + i * 8 + j;
                        if (EXACT_M || c_row < M) {
                            float sum = 0.0f;
                            #pragma unroll
                            for (int w = 0; w < WIDE96_NW; w++)
                                sum += red[w * 64 * 48 + tid * 48 + s * 16 + i * 4 + j];
                            uint32_t bits = f2u(sum);
                            bits += (0x7FFFu + ((bits >> 16) & 1u));
                            C[c_row * N + c_col] = (uint16_t)(bits >> 16);
                        }
                    }
                }
            }
        }
    }
}

// ============================================================
// Kernel 4: Reduction (sum splitK partials → bf16)
// ============================================================
__global__ void reduce_splitk_kernel(
    const float* __restrict__ partial,
    uint16_t* __restrict__ C,
    const int M, const int N, const int splitK
) {
    const int idx = blockIdx.x * blockDim.x + threadIdx.x;
    if (idx >= M * N) return;
    const int m = idx / N, n = idx % N;
    float sum = 0.0f;
    for (int s = 0; s < splitK; s++)
        sum += partial[(long long)s * M * N + m * N + n];
    uint32_t bits = f2u(sum);
    bits += (0x7FFFu + ((bits >> 16) & 1u));
    C[m * N + n] = (uint16_t)(bits >> 16);
}

// ============================================================
// Dispatch
// ============================================================
void launch_v35(
    torch::Tensor A, torch::Tensor B_shuffle, torch::Tensor B_scale_sh,
    torch::Tensor C, int sn_pad, int kernel_id, int splitK,
    torch::Tensor partial
) {
    const int M = A.size(0), K = A.size(1), N = B_shuffle.size(0);
    const uint16_t* Ap = (const uint16_t*)A.data_ptr();
    const uint8_t* Bp = (const uint8_t*)B_shuffle.data_ptr();
    const uint8_t* Sp = (const uint8_t*)B_scale_sh.data_ptr();
    uint16_t* Cp = (uint16_t*)C.data_ptr();

    if (kernel_id == 5) {
        // 32×32 8wf splitK
        dim3 grid_sk((N + 31) / 32, (M + 31) / 32, splitK);
        dim3 block_sk(64 * 8);
        fused_mfma_splitk_kernel<<<grid_sk, block_sk>>>(
            Ap, Bp, Sp, partial.data_ptr<float>(), M, N, K, K/2, sn_pad, splitK);
        int total = M * N;
        dim3 grid_r((total + 255) / 256);
        dim3 block_r(256);
        reduce_splitk_kernel<<<grid_r, block_r>>>(partial.data_ptr<float>(), Cp, M, N, splitK);
    } else if (kernel_id == 9) {
        // 32×96 wide EXACT_M
        dim3 grid_w((N + 95) / 96, (M + 31) / 32);
        dim3 block_w(64 * WIDE96_NW);
        fused_mfma_wide96_kernel<true><<<grid_w, block_w>>>(Ap, Bp, Sp, Cp, M, N, K, K/2, sn_pad);
    } else if (kernel_id == 4) {
        // 32×96 wide
        dim3 grid_w((N + 95) / 96, (M + 31) / 32);
        dim3 block_w(64 * WIDE96_NW);
        fused_mfma_wide96_kernel<<<grid_w, block_w>>>(Ap, Bp, Sp, Cp, M, N, K, K/2, sn_pad);
    } else if (kernel_id == 10) {
        // 16×16 EXACT_M
        dim3 grid_16((N + 15) / 16, (M + 15) / 16);
        dim3 block_16(64 * SMALL_NW);
        fused_mfma_16x16_kernel<true><<<grid_16, block_16>>>(Ap, Bp, Sp, Cp, M, N, K, K/2, sn_pad);
    } else if (kernel_id == 3) {
        // 16×16
        dim3 grid_16((N + 15) / 16, (M + 15) / 16);
        dim3 block_16(64 * SMALL_NW);
        fused_mfma_16x16_kernel<<<grid_16, block_16>>>(Ap, Bp, Sp, Cp, M, N, K, K/2, sn_pad);
    } else if (kernel_id == 1) {
        // 32×32 8wf
        dim3 grid_8((N + 31) / 32, (M + 31) / 32);
        dim3 block_8(64 * 8);
        fused_mfma_kernel<8><<<grid_8, block_8>>>(Ap, Bp, Sp, Cp, M, N, K, K/2, sn_pad);
    } else if (kernel_id == 8) {
        // 32×32 4wf EXACT_M (no M bounds check)
        dim3 grid_4((N + 31) / 32, (M + 31) / 32);
        dim3 block_4(64 * 4);
        fused_mfma_kernel<4, true><<<grid_4, block_4>>>(Ap, Bp, Sp, Cp, M, N, K, K/2, sn_pad);
    } else {
        // 32×32 4wf
        dim3 grid_4((N + 31) / 32, (M + 31) / 32);
        dim3 block_4(64 * 4);
        fused_mfma_kernel<4><<<grid_4, block_4>>>(Ap, Bp, Sp, Cp, M, N, K, K/2, sn_pad);
    }
}
"""

CPP_SRC = "void launch_v35(torch::Tensor A, torch::Tensor B_shuffle, torch::Tensor B_scale_sh, torch::Tensor C, int sn_pad, int kernel_id, int splitK, torch::Tensor partial);"

_CFLAGS = ["--offload-arch=gfx950", "-std=c++20", "-O3", "-ffast-math", "-funsafe-math-optimizations", "-mno-wavefrontsize64"]
import hashlib as _hl
_hip = load_inline(name='fm_'+_hl.md5(HIP_SRC.encode()).hexdigest()[:10], cpp_sources=[CPP_SRC], cuda_sources=[HIP_SRC],
    functions=['launch_v35'], verbose=True, extra_cuda_cflags=_CFLAGS)

# Separate wide96 EXACT_MN module (isolated compilation → no icache/regalloc interference)
_W96_HELPERS = HIP_SRC[HIP_SRC.index('#include'):HIP_SRC.index('// Kernel 0')]
_W96_BODY = HIP_SRC[HIP_SRC.index('// Kernel 3:'):HIP_SRC.index('// Kernel 4:')]
# Replace all M/N bounds with true
_W96_BODY = _W96_BODY.replace('EXACT_M || my_m < M', 'true').replace('EXACT_M || c_row < M', 'true')
_W96_BODY = _W96_BODY.replace('if (n0 < N)', 'if (true)').replace('if (n1 < N)', 'if (true)').replace('if (n2 < N)', 'if (true)')
_W96_BODY = _W96_BODY.replace('if (c_col < N)', 'if (true)')
_W96_BODY = _W96_BODY.replace('template <bool EXACT_M = false>\n', '')
_W96_SRC = _W96_HELPERS + _W96_BODY + r"""
void launch_w96(torch::Tensor A, torch::Tensor B_sh, torch::Tensor B_sc, torch::Tensor C, int sn) {
    const int M=A.size(0),K=A.size(1),N=B_sh.size(0);
    dim3 g((N+95)/96,(M+31)/32);
    fused_mfma_wide96_kernel<<<g,64*WIDE96_NW>>>((const uint16_t*)A.data_ptr(),(const uint8_t*)B_sh.data_ptr(),(const uint8_t*)B_sc.data_ptr(),(uint16_t*)C.data_ptr(),M,N,K,K/2,sn);
}
"""
_W96_CPP = "void launch_w96(torch::Tensor A, torch::Tensor B_sh, torch::Tensor B_sc, torch::Tensor C, int sn);"
_hip_w96 = load_inline(name='w9_'+_hl.md5(_W96_SRC.encode()).hexdigest()[:10], cpp_sources=[_W96_CPP], cuda_sources=[_W96_SRC],
    functions=['launch_w96'], verbose=True, extra_cuda_cflags=_CFLAGS)

@torch.inference_mode()
def custom_kernel(data: input_t) -> output_t:
    A, B, B_q, B_shuffle, B_scale_sh = data
    A = A.contiguous()
    m, k = A.shape
    n = B_shuffle.shape[0]
    C = torch.empty(m, n, dtype=torch.bfloat16, device=A.device)

    if m >= 128:
        blocks_96 = ((n + 95) // 96) * ((m + 31) // 32)
        blocks_32 = ((n + 31) // 32) * ((m + 31) // 32)
        if (blocks_96 + 255) // 256 < (blocks_32 + 255) // 256 and m % 32 == 0:
            # Use separate wide96 EXACT_MN module
            _hip_w96.launch_w96(A, B_shuffle, B_scale_sh, C, B_scale_sh.shape[1])
            return C
        elif (blocks_96 + 255) // 256 < (blocks_32 + 255) // 256:
            kernel_id, splitK = 4, 0
        else:
            kernel_id, splitK = (8 if m % 32 == 0 else 0), 0
    elif m <= 8:
        kernel_id, splitK = 3, 0
    elif m <= 16 and k >= 2048:
        splitK = 7
        if (k // splitK) % (8 * 64) == 0:
            kernel_id = 5
        else:
            kernel_id, splitK = 1, 0
    elif m <= 16:
        kernel_id, splitK = 3, 0
    elif m <= 32 and k <= 512:
        kernel_id, splitK = (10 if m % 16 == 0 else 3), 0
    else:
        # M=64: use EXACT_M version if M is multiple of 32
        if m % 32 == 0:
            kernel_id, splitK = 8, 0  # no bounds check
        else:
            kernel_id, splitK = 0, 0

    partial = torch.empty(splitK * m * n, dtype=torch.float32, device=A.device) if splitK > 0 else torch.empty(0, device=A.device)
    _hip.launch_v35(A, B_shuffle, B_scale_sh, C, B_scale_sh.shape[1], kernel_id, splitK, partial)
    return C
scrolls · 788 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