submission 645431
mmk150 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 1500 lines, June 9 Researcher Reciprocity License v1.0.
submission_general_v1.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-645431?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
Reported · How evidence levels are derived →
Source and license
sourceavailable
revision digestsha256:4988b381774610a188bcc2d5527cf2822cc293e639cc799796629215be0e8949
license declaredunknown
license concludedunknown
authorsmmk150
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
shared-memory
extern __shared__ char lds_raw[];split-k
static constexpr int kSplitK = 7;tile-m = 16
static constexpr int kTileM=16, kNumMTiles=kM/kTileM;tile-n = 32
static constexpr int kTileN = 32;vector-width = uint4
const uint4* a_src0 = reinterpret_cast<const uint4*>(a_row + k_group0 * 16);Kernel source
submission_general_v1.py1500 lines
"""
submission_general_v1.py — Based on submission_general_v0.py.
No preallocated buffer pools — all tensors allocated fresh per call.
Champions:
M=4 (4,2880,512) : m4_mk2a (2-wave, shuffled B)
M=16 (16,2112,7168) : m16_varI (two-pass workspace + reduce)
M=32 (32,2880,512) : m32_mk4l (16x16x128 MFMA, split-M, kernarg preload)
M=32 (32,4096,512) : m32_mk4l (same kernel, different N)
M=64 (64,7168,2048) : e641v1_exp S2 (hand-unrolled 8-iter, AGPR, B-scale LDS)
M=256 (256,3072,1536) : mk3g S4 (v_max3 absmax, deferred B scale)
"""
from __future__ import annotations
import os
import tempfile
import aiter
import torch
from aiter import dtypes
from aiter.ops.triton.quant import dynamic_mxfp4_quant
from aiter.utility.fp4_utils import e8m0_shuffle
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t
# =========================================================================
# C++ bridge
# =========================================================================
_CPP_SRC = r"""
#include <torch/extension.h>
#include <hip/hip_runtime.h>
#include <cstdint>
// ── M=4 extern ──
extern "C" void launch_m4_mk2a(
const uint16_t* A, const uint8_t* B_shuffle, const uint8_t* B_scale_sh,
uint16_t* C);
void m4_launch(
torch::Tensor A, torch::Tensor B_shuffle, torch::Tensor B_scale_sh, torch::Tensor C
) {
launch_m4_mk2a(
reinterpret_cast<const uint16_t*>(A.data_ptr()),
reinterpret_cast<const uint8_t*>(B_shuffle.data_ptr()),
reinterpret_cast<const uint8_t*>(B_scale_sh.data_ptr()),
reinterpret_cast<uint16_t*>(C.data_ptr()));
}
// ── M=16 extern ──
extern "C" void launch_m16_varI_gemm(
const uint16_t* A, const uint8_t* B, const uint8_t* Bs, float* ws);
extern "C" void launch_m16_varI_reduce(const float* ws, uint16_t* C);
void m16_launch(
torch::Tensor A, torch::Tensor B_shuffle, torch::Tensor B_scale_sh,
torch::Tensor workspace, torch::Tensor C
) {
launch_m16_varI_gemm(
reinterpret_cast<const uint16_t*>(A.data_ptr()),
reinterpret_cast<const uint8_t*>(B_shuffle.data_ptr()),
reinterpret_cast<const uint8_t*>(B_scale_sh.data_ptr()),
reinterpret_cast<float*>(workspace.data_ptr()));
launch_m16_varI_reduce(
reinterpret_cast<const float*>(workspace.data_ptr()),
reinterpret_cast<uint16_t*>(C.data_ptr()));
}
// ── M=32 externs ──
extern "C" void launch_m32_n2880(
const uint16_t* A, const uint8_t* B_shuffle, const uint8_t* B_scale_sh,
uint16_t* C);
extern "C" void launch_m32_n4096(
const uint16_t* A, const uint8_t* B_shuffle, const uint8_t* B_scale_sh,
uint16_t* C);
void m32_n2880_launch(
torch::Tensor A, torch::Tensor B_shuffle, torch::Tensor B_scale_sh,
torch::Tensor C
) {
launch_m32_n2880(
reinterpret_cast<const uint16_t*>(A.data_ptr()),
reinterpret_cast<const uint8_t*>(B_shuffle.data_ptr()),
reinterpret_cast<const uint8_t*>(B_scale_sh.data_ptr()),
reinterpret_cast<uint16_t*>(C.data_ptr()));
}
void m32_n4096_launch(
torch::Tensor A, torch::Tensor B_shuffle, torch::Tensor B_scale_sh,
torch::Tensor C
) {
launch_m32_n4096(
reinterpret_cast<const uint16_t*>(A.data_ptr()),
reinterpret_cast<const uint8_t*>(B_shuffle.data_ptr()),
reinterpret_cast<const uint8_t*>(B_scale_sh.data_ptr()),
reinterpret_cast<uint16_t*>(C.data_ptr()));
}
// ── M=64 extern ──
extern "C" void launch_e641v1_exp_64_7168_2048(
const uint16_t* A, const uint8_t* B, const uint8_t* Bs, uint16_t* C);
void m64_launch(
torch::Tensor A, torch::Tensor B_shuffle, torch::Tensor B_scale_sh, torch::Tensor C
) {
launch_e641v1_exp_64_7168_2048(
reinterpret_cast<const uint16_t*>(A.data_ptr()),
reinterpret_cast<const uint8_t*>(B_shuffle.data_ptr()),
reinterpret_cast<const uint8_t*>(B_scale_sh.data_ptr()),
reinterpret_cast<uint16_t*>(C.data_ptr()));
}
// ── M=256 extern ──
extern "C" void launch_mk3g_256_3072_1536(
const uint16_t* A, const uint8_t* B, const uint8_t* Bs, uint16_t* C);
void m256_launch(
torch::Tensor A, torch::Tensor B_shuffle, torch::Tensor B_scale_sh, torch::Tensor C
) {
launch_mk3g_256_3072_1536(
reinterpret_cast<const uint16_t*>(A.data_ptr()),
reinterpret_cast<const uint8_t*>(B_shuffle.data_ptr()),
reinterpret_cast<const uint8_t*>(B_scale_sh.data_ptr()),
reinterpret_cast<uint16_t*>(C.data_ptr()));
}
"""
# =========================================================================
# HIP source — all kernel sources concatenated
# =========================================================================
_HIP_SRC = r"""
#include <hip/hip_runtime.h>
#include <cstdint>
// =====================================================================
// KERNEL 1: M=4 (m4_mk2a — 2-wave, shuffled B)
// =====================================================================
// m4_optim_mk2a: 2-wave blocks, AGPRs, packed absmax, XCD swizzle, all loads upfront
// Shape: M=4, N=2880, K=512, TILE_N=32 (2 MFMA columns), 16x16x128 MFMA
// 90 blocks × 128 threads (2 waves) = 360 waves total (same occupancy as mk1f).
// Each wave does 2 K-iters × 2 MFMA columns = 4 MFMAs. LDS reduce 2 partials.
// Barrier syncs 2 waves instead of 4. Half the LDS traffic.
namespace m4_mk2a {
using v4i32 = int __attribute__((ext_vector_type(4)));
using v4f32 = float __attribute__((ext_vector_type(4)));
typedef __bf16 v2bf16 __attribute__((ext_vector_type(2)));
static constexpr int kM = 4;
static constexpr int kN = 2880;
static constexpr int kK = 512;
static constexpr int kKHalf = kK / 2;
static constexpr int kTileN = 32;
static constexpr int kNumWaves = 2;
static constexpr int kBlockSize = kNumWaves * 64; // 128
static constexpr int kNumNTiles = kN / kTileN; // 90 (exact, no remainder)
static constexpr int kBPanelStride = kKHalf * 16;
static constexpr int kPaddedGroups = kK / 32;
// Each wave does 2 K-iters. Wave 0: ki=0,1. Wave 1: ki=2,3.
static constexpr int kKItersPerWave = 2;
// LDS: [2 waves][kM][kTileN] = 2*4*32 = 256 floats = 1024 bytes
static constexpr int kLdsFloats = kNumWaves * kM * kTileN;
static constexpr int kLdsBytes = kLdsFloats * sizeof(float);
// XCD swizzle
static constexpr int kNXCD = 8;
static constexpr int kC = 4;
static constexpr int kBlocksPerCycle = kNXCD * kC;
static constexpr int kLimit = (kNumNTiles / kBlocksPerCycle) * kBlocksPerCycle;
__device__ __forceinline__ uint32_t bit_cast_u32(float v) {
union { float f; uint32_t u; } x; x.f = v; return x.u;
}
__device__ __forceinline__ float bit_cast_f32(uint32_t v) {
union { uint32_t u; float f; } x; x.u = v; return x.f;
}
__device__ __forceinline__ uint16_t float_to_bf16_rn(float v) {
uint32_t bits = bit_cast_u32(v);
bits += ((bits >> 16) & 1u) + 0x7FFFu;
return static_cast<uint16_t>(bits >> 16);
}
__device__ __forceinline__ uint8_t compute_e8m0_scale(uint16_t max_abs) {
const uint16_t r = static_cast<uint16_t>(max_abs + 0x20u);
const int e = static_cast<int>((r >> 7) & 0xFFu);
return (e <= 2) ? 0u : (e >= 255) ? 254u : static_cast<uint8_t>(e - 2);
}
__device__ __forceinline__ int scale_offset(int row, int group) {
return (row >> 5) * (32 * kPaddedGroups)
+ (group >> 3) * 256
+ (group & 3) * 64
+ (row & 15) * 4
+ ((group >> 2) & 1) * 2
+ ((row >> 4) & 1);
}
// MFMA with AGPR accumulator
__device__ __forceinline__ v4f32 mfma_fp4_16x16(
v4i32 a, v4i32 b, v4f32 acc, int a_scale, int b_scale) {
asm volatile(
"v_mfma_scale_f32_16x16x128_f8f6f4 %0, %1, %2, %0, %3, %4 cbsz:4 blgp:4"
: "+v"(acc) : "v"(a), "v"(b), "v"(a_scale), "v"(b_scale));
return acc;
}
__global__ __attribute__((amdgpu_flat_work_group_size(128, 128)))
void kernel_m4_mk2a(
const uint16_t* __restrict__ A,
const uint8_t* __restrict__ B_shuffle,
const uint8_t* __restrict__ B_scale_sh,
uint16_t* __restrict__ C)
{
// XCD swizzle
int bid = blockIdx.x;
if (bid < kLimit) {
const int xcd = bid % kNXCD;
const int local_ = bid / kNXCD;
const int chunk = local_ / kC;
const int pos = local_ % kC;
bid = chunk * kBlocksPerCycle + xcd * kC + pos;
}
const int col_base = bid * kTileN;
const int tid = threadIdx.x;
const int wave_id = tid >> 6; // 0 or 1
const int lane = tid & 63;
const int row = lane & 15;
const int k_quarter = lane >> 4; // 0..3
extern __shared__ char lds_raw[];
float* lds = reinterpret_cast<float*>(lds_raw);
// Wave 0 does ki=0,1. Wave 1 does ki=2,3.
const int ki_base = wave_id * kKItersPerWave; // 0 or 2
// ============================================================
// PHASE 1: Issue ALL loads upfront
// ============================================================
// A loads for both K-iters (row%kM wraps for threads with row>=4)
const uint32_t* a_row = reinterpret_cast<const uint32_t*>(
A + (row % kM) * kK);
const int k_group0 = ki_base * 4 + k_quarter;
const int k_group1 = (ki_base + 1) * 4 + k_quarter;
asm volatile("" ::: "memory");
const uint4* a_src0 = reinterpret_cast<const uint4*>(a_row + k_group0 * 16);
uint4 a0_c0 = a_src0[0], a0_c1 = a_src0[1], a0_c2 = a_src0[2], a0_c3 = a_src0[3];
const uint4* a_src1 = reinterpret_cast<const uint4*>(a_row + k_group1 * 16);
uint4 a1_c0 = a_src1[0], a1_c1 = a_src1[1], a1_c2 = a_src1[2], a1_c3 = a_src1[3];
// B col 0 loads for both K-iters
const int b_col0 = col_base + row;
const uint8_t* b_base0 = B_shuffle
+ (b_col0 >> 4) * kBPanelStride + (b_col0 & 15) * 16;
uint4 b0_ki0 = reinterpret_cast<const uint4*>(b_base0 + ki_base * 1024 + k_quarter * 256)[0];
uint4 b0_ki1 = reinterpret_cast<const uint4*>(b_base0 + (ki_base+1) * 1024 + k_quarter * 256)[0];
// B col 1 loads for both K-iters
const int b_col1 = col_base + 16 + row;
const uint8_t* b_base1 = B_shuffle
+ (b_col1 >> 4) * kBPanelStride + (b_col1 & 15) * 16;
uint4 b1_ki0 = reinterpret_cast<const uint4*>(b_base1 + ki_base * 1024 + k_quarter * 256)[0];
uint4 b1_ki1 = reinterpret_cast<const uint4*>(b_base1 + (ki_base+1) * 1024 + k_quarter * 256)[0];
// B scale loads (4 total: 2 columns × 2 K-iters)
int bsc0_ki0 = static_cast<int>(B_scale_sh[scale_offset(b_col0, k_group0)]);
int bsc0_ki1 = static_cast<int>(B_scale_sh[scale_offset(b_col0, k_group1)]);
int bsc1_ki0 = static_cast<int>(B_scale_sh[scale_offset(b_col1, k_group0)]);
int bsc1_ki1 = static_cast<int>(B_scale_sh[scale_offset(b_col1, k_group1)]);
asm volatile("" : "+v"(bsc0_ki0), "+v"(bsc0_ki1), "+v"(bsc1_ki0), "+v"(bsc1_ki1) :: "memory");
// ============================================================
// PHASE 2: Quantize A for both K-iters using packed u16 absmax
// ============================================================
// --- K-iter 0 ---
// Packed absmax using v_pk_max_u16 pattern
uint32_t mx0;
{
// Mask off sign bits to get abs bf16 values, packed as u16 pairs
const uint32_t mask = 0x7FFF7FFFu;
uint32_t p0 = a0_c0.x & mask, p1 = a0_c0.y & mask, p2 = a0_c0.z & mask, p3 = a0_c0.w & mask;
uint32_t p4 = a0_c1.x & mask, p5 = a0_c1.y & mask, p6 = a0_c1.z & mask, p7 = a0_c1.w & mask;
uint32_t p8 = a0_c2.x & mask, p9 = a0_c2.y & mask, pA = a0_c2.z & mask, pB = a0_c2.w & mask;
uint32_t pC = a0_c3.x & mask, pD = a0_c3.y & mask, pE = a0_c3.z & mask, pF = a0_c3.w & mask;
// Reduce with v_pk_max_u16 (compiler should emit this for packed u16 max)
uint32_t m01, m23, m45, m67, m89, mAB, mCD, mEF;
asm("v_pk_max_u16 %0, %1, %2" : "=v"(m01) : "v"(p0), "v"(p1));
asm("v_pk_max_u16 %0, %1, %2" : "=v"(m23) : "v"(p2), "v"(p3));
asm("v_pk_max_u16 %0, %1, %2" : "=v"(m45) : "v"(p4), "v"(p5));
asm("v_pk_max_u16 %0, %1, %2" : "=v"(m67) : "v"(p6), "v"(p7));
asm("v_pk_max_u16 %0, %1, %2" : "=v"(m89) : "v"(p8), "v"(p9));
asm("v_pk_max_u16 %0, %1, %2" : "=v"(mAB) : "v"(pA), "v"(pB));
asm("v_pk_max_u16 %0, %1, %2" : "=v"(mCD) : "v"(pC), "v"(pD));
asm("v_pk_max_u16 %0, %1, %2" : "=v"(mEF) : "v"(pE), "v"(pF));
uint32_t t0, t1, t2, t3;
asm("v_pk_max_u16 %0, %1, %2" : "=v"(t0) : "v"(m01), "v"(m23));
asm("v_pk_max_u16 %0, %1, %2" : "=v"(t1) : "v"(m45), "v"(m67));
asm("v_pk_max_u16 %0, %1, %2" : "=v"(t2) : "v"(m89), "v"(mAB));
asm("v_pk_max_u16 %0, %1, %2" : "=v"(t3) : "v"(mCD), "v"(mEF));
uint32_t u0, u1;
asm("v_pk_max_u16 %0, %1, %2" : "=v"(u0) : "v"(t0), "v"(t1));
asm("v_pk_max_u16 %0, %1, %2" : "=v"(u1) : "v"(t2), "v"(t3));
uint32_t final_pk;
asm("v_pk_max_u16 %0, %1, %2" : "=v"(final_pk) : "v"(u0), "v"(u1));
// Reduce the 2 u16 lanes to scalar
mx0 = (final_pk >> 16) > (final_pk & 0xFFFFu) ? (final_pk >> 16) : (final_pk & 0xFFFFu);
}
uint8_t scale0 = compute_e8m0_scale(static_cast<uint16_t>(mx0));
int a_sc0 = static_cast<int>(scale0);
float fwd0 = (scale0 == 0u) ? bit_cast_f32(0x00400000u)
: bit_cast_f32(static_cast<uint32_t>(scale0) << 23);
const v2bf16* p0_0 = reinterpret_cast<const v2bf16*>(&a0_c0);
const v2bf16* p0_1 = reinterpret_cast<const v2bf16*>(&a0_c1);
const v2bf16* p0_2 = reinterpret_cast<const v2bf16*>(&a0_c2);
const v2bf16* p0_3 = reinterpret_cast<const v2bf16*>(&a0_c3);
uint32_t q0[4] = {0,0,0,0};
q0[0]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(q0[0],p0_0[0],fwd0,0);
q0[1]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(q0[1],p0_1[0],fwd0,0);
q0[2]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(q0[2],p0_2[0],fwd0,0);
q0[3]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(q0[3],p0_3[0],fwd0,0);
q0[0]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(q0[0],p0_0[1],fwd0,1);
q0[1]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(q0[1],p0_1[1],fwd0,1);
q0[2]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(q0[2],p0_2[1],fwd0,1);
q0[3]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(q0[3],p0_3[1],fwd0,1);
q0[0]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(q0[0],p0_0[2],fwd0,2);
q0[1]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(q0[1],p0_1[2],fwd0,2);
q0[2]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(q0[2],p0_2[2],fwd0,2);
q0[3]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(q0[3],p0_3[2],fwd0,2);
q0[0]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(q0[0],p0_0[3],fwd0,3);
q0[1]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(q0[1],p0_1[3],fwd0,3);
q0[2]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(q0[2],p0_2[3],fwd0,3);
q0[3]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(q0[3],p0_3[3],fwd0,3);
// --- K-iter 1 ---
uint32_t mx1;
{
const uint32_t mask = 0x7FFF7FFFu;
uint32_t p0 = a1_c0.x & mask, p1 = a1_c0.y & mask, p2 = a1_c0.z & mask, p3 = a1_c0.w & mask;
uint32_t p4 = a1_c1.x & mask, p5 = a1_c1.y & mask, p6 = a1_c1.z & mask, p7 = a1_c1.w & mask;
uint32_t p8 = a1_c2.x & mask, p9 = a1_c2.y & mask, pA = a1_c2.z & mask, pB = a1_c2.w & mask;
uint32_t pC = a1_c3.x & mask, pD = a1_c3.y & mask, pE = a1_c3.z & mask, pF = a1_c3.w & mask;
uint32_t m01, m23, m45, m67, m89, mAB, mCD, mEF;
asm("v_pk_max_u16 %0, %1, %2" : "=v"(m01) : "v"(p0), "v"(p1));
asm("v_pk_max_u16 %0, %1, %2" : "=v"(m23) : "v"(p2), "v"(p3));
asm("v_pk_max_u16 %0, %1, %2" : "=v"(m45) : "v"(p4), "v"(p5));
asm("v_pk_max_u16 %0, %1, %2" : "=v"(m67) : "v"(p6), "v"(p7));
asm("v_pk_max_u16 %0, %1, %2" : "=v"(m89) : "v"(p8), "v"(p9));
asm("v_pk_max_u16 %0, %1, %2" : "=v"(mAB) : "v"(pA), "v"(pB));
asm("v_pk_max_u16 %0, %1, %2" : "=v"(mCD) : "v"(pC), "v"(pD));
asm("v_pk_max_u16 %0, %1, %2" : "=v"(mEF) : "v"(pE), "v"(pF));
uint32_t t0, t1, t2, t3;
asm("v_pk_max_u16 %0, %1, %2" : "=v"(t0) : "v"(m01), "v"(m23));
asm("v_pk_max_u16 %0, %1, %2" : "=v"(t1) : "v"(m45), "v"(m67));
asm("v_pk_max_u16 %0, %1, %2" : "=v"(t2) : "v"(m89), "v"(mAB));
asm("v_pk_max_u16 %0, %1, %2" : "=v"(t3) : "v"(mCD), "v"(mEF));
uint32_t u0, u1;
asm("v_pk_max_u16 %0, %1, %2" : "=v"(u0) : "v"(t0), "v"(t1));
asm("v_pk_max_u16 %0, %1, %2" : "=v"(u1) : "v"(t2), "v"(t3));
uint32_t final_pk;
asm("v_pk_max_u16 %0, %1, %2" : "=v"(final_pk) : "v"(u0), "v"(u1));
mx1 = (final_pk >> 16) > (final_pk & 0xFFFFu) ? (final_pk >> 16) : (final_pk & 0xFFFFu);
}
uint8_t scale1 = compute_e8m0_scale(static_cast<uint16_t>(mx1));
int a_sc1 = static_cast<int>(scale1);
float fwd1 = (scale1 == 0u) ? bit_cast_f32(0x00400000u)
: bit_cast_f32(static_cast<uint32_t>(scale1) << 23);
const v2bf16* p1_0 = reinterpret_cast<const v2bf16*>(&a1_c0);
const v2bf16* p1_1 = reinterpret_cast<const v2bf16*>(&a1_c1);
const v2bf16* p1_2 = reinterpret_cast<const v2bf16*>(&a1_c2);
const v2bf16* p1_3 = reinterpret_cast<const v2bf16*>(&a1_c3);
uint32_t q1[4] = {0,0,0,0};
q1[0]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(q1[0],p1_0[0],fwd1,0);
q1[1]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(q1[1],p1_1[0],fwd1,0);
q1[2]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(q1[2],p1_2[0],fwd1,0);
q1[3]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(q1[3],p1_3[0],fwd1,0);
q1[0]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(q1[0],p1_0[1],fwd1,1);
q1[1]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(q1[1],p1_1[1],fwd1,1);
q1[2]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(q1[2],p1_2[1],fwd1,1);
q1[3]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(q1[3],p1_3[1],fwd1,1);
q1[0]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(q1[0],p1_0[2],fwd1,2);
q1[1]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(q1[1],p1_1[2],fwd1,2);
q1[2]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(q1[2],p1_2[2],fwd1,2);
q1[3]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(q1[3],p1_3[2],fwd1,2);
q1[0]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(q1[0],p1_0[3],fwd1,3);
q1[1]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(q1[1],p1_1[3],fwd1,3);
q1[2]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(q1[2],p1_2[3],fwd1,3);
q1[3]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(q1[3],p1_3[3],fwd1,3);
// Zero out A for threads with row >= kM
v4i32 a_r0, a_r1;
if (row < kM) {
a_r0 = {(int)q0[0], (int)q0[1], (int)q0[2], (int)q0[3]};
a_r1 = {(int)q1[0], (int)q1[1], (int)q1[2], (int)q1[3]};
} else {
a_r0 = {0,0,0,0};
a_r1 = {0,0,0,0};
a_sc0 = 127;
a_sc1 = 127;
}
// ============================================================
// PHASE 3: 4 MFMAs — 2 K-iters × 2 N-columns, accumulate in AGPRs
// ============================================================
v4i32 b0_r0 = {(int)b0_ki0.x, (int)b0_ki0.y, (int)b0_ki0.z, (int)b0_ki0.w};
v4i32 b0_r1 = {(int)b0_ki1.x, (int)b0_ki1.y, (int)b0_ki1.z, (int)b0_ki1.w};
v4i32 b1_r0 = {(int)b1_ki0.x, (int)b1_ki0.y, (int)b1_ki0.z, (int)b1_ki0.w};
v4i32 b1_r1 = {(int)b1_ki1.x, (int)b1_ki1.y, (int)b1_ki1.z, (int)b1_ki1.w};
v4f32 acc_c0 = {0.0f, 0.0f, 0.0f, 0.0f};
v4f32 acc_c1 = {0.0f, 0.0f, 0.0f, 0.0f};
// Grouped by accumulator — same-acc MFMAs back-to-back for 0-wait fast path
// Col 0: K-iter 0 then K-iter 1
acc_c0 = mfma_fp4_16x16(a_r0, b0_r0, acc_c0, a_sc0, bsc0_ki0);
acc_c0 = mfma_fp4_16x16(a_r1, b0_r1, acc_c0, a_sc1, bsc0_ki1);
// Col 1: K-iter 0 then K-iter 1
acc_c1 = mfma_fp4_16x16(a_r0, b1_r0, acc_c1, a_sc0, bsc1_ki0);
acc_c1 = mfma_fp4_16x16(a_r1, b1_r1, acc_c1, a_sc1, bsc1_ki1);
// ============================================================
// PHASE 4: LDS write — only rows < kM
// ============================================================
const int out_col_16 = lane & 15;
const int row_base4 = (lane >> 4) * 4;
const int lds_wave_base = wave_id * kM * kTileN;
#pragma unroll
for (int r = 0; r < 4; ++r) {
const int out_row = row_base4 + r;
if (out_row < kM) {
lds[lds_wave_base + out_row * kTileN + out_col_16] = acc_c0[r];
lds[lds_wave_base + out_row * kTileN + 16 + out_col_16] = acc_c1[r];
}
}
__syncthreads();
// ============================================================
// PHASE 5: Reduce 2 partials and write output
// 128 threads handle kM × kTileN = 4 × 32 = 128 elements (perfect 1:1)
// ============================================================
{
const int local_row = tid / kTileN;
const int local_col = tid % kTileN;
if (local_row < kM) {
const float v0 = lds[0 * kM * kTileN + local_row * kTileN + local_col];
const float v1 = lds[1 * kM * kTileN + local_row * kTileN + local_col];
const float sum = v0 + v1;
const int out_col = col_base + local_col;
if (out_col < kN) {
C[local_row * kN + out_col] = float_to_bf16_rn(sum);
}
}
}
}
} // namespace m4_mk2a
extern "C" void launch_m4_mk2a(
const uint16_t* A, const uint8_t* B_shuffle, const uint8_t* B_scale_sh,
uint16_t* C) {
hipLaunchKernelGGL(m4_mk2a::kernel_m4_mk2a,
dim3(m4_mk2a::kNumNTiles), dim3(m4_mk2a::kBlockSize),
m4_mk2a::kLdsBytes, 0,
A, B_shuffle, B_scale_sh, C);
}
// =====================================================================
// KERNEL 2: M=16 (m16_varI — two-pass workspace + reduce)
// =====================================================================
// m16_varI: Two-pass — GEMM writes to workspace, reduce kernel sums partials
// No atomics, no torch.zeros. Each block stores directly to workspace[split_id].
// M=16, N=2112, K=7168, TILE_N=64, SK=7, 4 waves
// Workspace: [7][16][2112] f32 = 944 KB (torch.empty)
// Pass 1: 231 blocks compute and store partials
// Pass 2: small reduce kernel sums 7 partials → bf16 output
namespace m16_varI {
using v4i32 = int __attribute__((ext_vector_type(4)));
using v4f32 = float __attribute__((ext_vector_type(4)));
typedef __bf16 v2bf16 __attribute__((ext_vector_type(2)));
static constexpr int kM = 16;
static constexpr int kN = 2112;
static constexpr int kK = 7168;
static constexpr int kKHalf = kK / 2;
static constexpr int kPaddedGroups = 224;
static constexpr int kTileN = 64;
static constexpr int kMfmaCols = 4;
static constexpr int kSplitK = 7;
static constexpr int kKPerSplit = kK / kSplitK; // 1024
static constexpr int kNumWaves = 4;
static constexpr int kBlockSize = kNumWaves * 64;
static constexpr int kItersPerWave = kKPerSplit / (kNumWaves * 128); // 2
static constexpr int kNumNTiles = kN / kTileN; // 33
static constexpr int kTotalBlocks = kNumNTiles * kSplitK; // 231
static constexpr int kTileElements = kM * kTileN;
static constexpr int kBPanelStride = kKHalf * 16;
// LDS for 4-wave intra-block reduce: [col][row][wave][16]
static constexpr int kLdsFloatsPerCol = kM * kNumWaves * 16;
static constexpr int kLdsFloats = kMfmaCols * kLdsFloatsPerCol;
static constexpr int kLdsBytes = kLdsFloats * sizeof(float);
static constexpr int kNXCD = 8;
static constexpr int kC = 4;
static constexpr int kBlocksPerCycle = kNXCD * kC;
static constexpr int kLimit = (kTotalBlocks / kBlocksPerCycle) * kBlocksPerCycle;
__device__ __forceinline__ uint32_t bit_cast_u32(float v) { union{float f;uint32_t u;}x; x.f=v; return x.u; }
__device__ __forceinline__ float bit_cast_f32(uint32_t v) { union{uint32_t u;float f;}x; x.u=v; return x.f; }
__device__ __forceinline__ uint16_t float_to_bf16_rn(float v) { uint32_t b=bit_cast_u32(v); b+=((b>>16)&1u)+0x7FFFu; return (uint16_t)(b>>16); }
__device__ __forceinline__ uint8_t compute_e8m0_scale(uint16_t m) { uint16_t r=(uint16_t)(m+0x20u); int e=(int)((r>>7)&0xFFu); return (e<=2)?0u:(e>=255)?254u:(uint8_t)(e-2); }
__device__ __forceinline__ int scale_offset(int row, int group) {
return (row>>5)*(32*kPaddedGroups)+(group>>3)*256+(group&3)*64+(row&15)*4+((group>>2)&1)*2+((row>>4)&1);
}
__device__ __forceinline__ v4f32 mfma_fp4_16x16(v4i32 a, v4i32 b, v4f32 acc, int as, int bs) {
asm volatile("v_mfma_scale_f32_16x16x128_f8f6f4 %0,%1,%2,%0,%3,%4 cbsz:4 blgp:4"
:"+v"(acc):"v"(a),"v"(b),"v"(as),"v"(bs)); return acc;
}
__device__ __forceinline__ void quant_a(const uint32_t* dw, v4i32& out, int& sc) {
const uint32_t mask=0x7FFF7FFFu;
uint32_t p0=dw[0]&mask,p1=dw[1]&mask,p2=dw[2]&mask,p3=dw[3]&mask;
uint32_t p4=dw[4]&mask,p5=dw[5]&mask,p6=dw[6]&mask,p7=dw[7]&mask;
uint32_t p8=dw[8]&mask,p9=dw[9]&mask,pA=dw[10]&mask,pB=dw[11]&mask;
uint32_t pC=dw[12]&mask,pD=dw[13]&mask,pE=dw[14]&mask,pF=dw[15]&mask;
uint32_t m01,m23,m45,m67,m89,mAB,mCD,mEF;
asm("v_pk_max_u16 %0,%1,%2":"=v"(m01):"v"(p0),"v"(p1));asm("v_pk_max_u16 %0,%1,%2":"=v"(m23):"v"(p2),"v"(p3));
asm("v_pk_max_u16 %0,%1,%2":"=v"(m45):"v"(p4),"v"(p5));asm("v_pk_max_u16 %0,%1,%2":"=v"(m67):"v"(p6),"v"(p7));
asm("v_pk_max_u16 %0,%1,%2":"=v"(m89):"v"(p8),"v"(p9));asm("v_pk_max_u16 %0,%1,%2":"=v"(mAB):"v"(pA),"v"(pB));
asm("v_pk_max_u16 %0,%1,%2":"=v"(mCD):"v"(pC),"v"(pD));asm("v_pk_max_u16 %0,%1,%2":"=v"(mEF):"v"(pE),"v"(pF));
uint32_t t0,t1,t2,t3;
asm("v_pk_max_u16 %0,%1,%2":"=v"(t0):"v"(m01),"v"(m23));asm("v_pk_max_u16 %0,%1,%2":"=v"(t1):"v"(m45),"v"(m67));
asm("v_pk_max_u16 %0,%1,%2":"=v"(t2):"v"(m89),"v"(mAB));asm("v_pk_max_u16 %0,%1,%2":"=v"(t3):"v"(mCD),"v"(mEF));
uint32_t u0,u1;
asm("v_pk_max_u16 %0,%1,%2":"=v"(u0):"v"(t0),"v"(t1));asm("v_pk_max_u16 %0,%1,%2":"=v"(u1):"v"(t2),"v"(t3));
uint32_t fp; asm("v_pk_max_u16 %0,%1,%2":"=v"(fp):"v"(u0),"v"(u1));
uint32_t mx=(fp>>16)>(fp&0xFFFFu)?(fp>>16):(fp&0xFFFFu);
uint8_t s=compute_e8m0_scale((uint16_t)mx); sc=(int)s;
float fwd=(s==0u)?bit_cast_f32(0x00400000u):bit_cast_f32((uint32_t)s<<23);
const v2bf16*q=reinterpret_cast<const v2bf16*>(dw); uint32_t k0=0,k1=0,k2=0,k3=0;
k0=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(k0,q[0],fwd,0);k1=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(k1,q[4],fwd,0);
k2=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(k2,q[8],fwd,0);k3=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(k3,q[12],fwd,0);
k0=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(k0,q[1],fwd,1);k1=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(k1,q[5],fwd,1);
k2=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(k2,q[9],fwd,1);k3=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(k3,q[13],fwd,1);
k0=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(k0,q[2],fwd,2);k1=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(k1,q[6],fwd,2);
k2=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(k2,q[10],fwd,2);k3=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(k3,q[14],fwd,2);
k0=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(k0,q[3],fwd,3);k1=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(k1,q[7],fwd,3);
k2=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(k2,q[11],fwd,3);k3=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(k3,q[15],fwd,3);
out={(int)k0,(int)k1,(int)k2,(int)k3};
}
// ============================================================
// PASS 1: GEMM kernel — stores partials to workspace
// ============================================================
__global__ __attribute__((amdgpu_flat_work_group_size(256, 256)))
void kernel_m16_gemm(
const uint16_t* __restrict__ A,
const uint8_t* __restrict__ B_shuffle,
const uint8_t* __restrict__ B_scale_sh,
float* __restrict__ workspace) // [kSplitK][kM][kN]
{
int xy = blockIdx.x;
if (xy < kLimit) {
const int xcd=xy%kNXCD, local_=xy/kNXCD, chunk=local_/kC, pos=local_%kC;
xy = chunk*kBlocksPerCycle + xcd*kC + pos;
}
const int k_split = xy / kNumNTiles;
const int n_tile = xy % kNumNTiles;
const int col_base = n_tile * kTileN;
const int tid=threadIdx.x, wave_id=tid>>6, lane=tid&63;
const int row=lane&15, k_quarter=lane>>4;
extern __shared__ char lds_raw[];
float* lds = reinterpret_cast<float*>(lds_raw);
const uint32_t* a_row = reinterpret_cast<const uint32_t*>(A + row * kK);
const uint8_t* b_bases[kMfmaCols];
#pragma unroll
for (int c=0;c<kMfmaCols;++c) {
const int b_col=col_base+c*16+row;
b_bases[c]=B_shuffle+(b_col>>4)*kBPanelStride+(b_col&15)*16;
}
const int split_iter_base = k_split * (kKPerSplit / 128);
const int ki_base = split_iter_base + wave_id * kItersPerWave;
const int kg0 = ki_base*4+k_quarter;
const int kg1 = (ki_base+1)*4+k_quarter;
// All loads upfront
asm volatile("" ::: "memory");
uint32_t a_buf0[16], a_buf1[16];
{ const uint32_t*s=a_row+kg0*16;
#pragma unroll
for(int i=0;i<16;++i) a_buf0[i]=s[i]; }
{ const uint32_t*s=a_row+kg1*16;
#pragma unroll
for(int i=0;i<16;++i) a_buf1[i]=s[i]; }
v4i32 b0[kMfmaCols],b1[kMfmaCols]; int bs0[kMfmaCols],bs1[kMfmaCols];
#pragma unroll
for(int c=0;c<kMfmaCols;++c) {
const uint32_t*s0=reinterpret_cast<const uint32_t*>(b_bases[c]+ki_base*1024+k_quarter*256);
b0[c]={(int)s0[0],(int)s0[1],(int)s0[2],(int)s0[3]};
const uint32_t*s1=reinterpret_cast<const uint32_t*>(b_bases[c]+(ki_base+1)*1024+k_quarter*256);
b1[c]={(int)s1[0],(int)s1[1],(int)s1[2],(int)s1[3]};
const int b_col=col_base+c*16+row;
bs0[c]=(int)B_scale_sh[scale_offset(b_col, kg0)];
bs1[c]=(int)B_scale_sh[scale_offset(b_col, kg1)];
}
asm volatile("":"+v"(bs0[0]),"+v"(bs0[1]),"+v"(bs0[2]),"+v"(bs0[3]),
"+v"(bs1[0]),"+v"(bs1[1]),"+v"(bs1[2]),"+v"(bs1[3])::"memory");
// Quant + MFMA
v4f32 acc[kMfmaCols];
#pragma unroll
for(int c=0;c<kMfmaCols;++c) acc[c]={0,0,0,0};
{ v4i32 ar; int as; quant_a(a_buf0,ar,as);
#pragma unroll
for(int c=0;c<kMfmaCols;++c) acc[c]=mfma_fp4_16x16(ar,b0[c],acc[c],as,bs0[c]); }
{ v4i32 ar; int as; quant_a(a_buf1,ar,as);
#pragma unroll
for(int c=0;c<kMfmaCols;++c) acc[c]=mfma_fp4_16x16(ar,b1[c],acc[c],as,bs1[c]); }
// LDS reduce 4 waves
const int oc16=lane&15, rb4=(lane>>4)*4;
#pragma unroll
for(int c=0;c<kMfmaCols;++c)
#pragma unroll
for(int r=0;r<4;++r)
lds[c*kLdsFloatsPerCol+(rb4+r)*kNumWaves*16+wave_id*16+oc16]=acc[c][r];
__syncthreads();
// Reduce and STORE to workspace (no atomics!)
// workspace layout: [split][row][col] — split stride = kM * kN
const int ws_base = k_split * kM * kN;
{
const int c = wave_id;
#pragma unroll
for(int r=0;r<4;++r) {
const int out_row=rb4+r;
if(out_row<kM) {
const int lb=c*kLdsFloatsPerCol+out_row*kNumWaves*16+oc16;
float sum=0.0f;
#pragma unroll
for(int w=0;w<kNumWaves;++w) sum+=lds[lb+w*16];
const int gc=col_base+c*16+oc16;
workspace[ws_base + out_row*kN + gc] = sum;
}
}
}
}
// ============================================================
// PASS 2: Reduce kernel — sum 7 partials, convert to bf16
// ============================================================
static constexpr int kReduceBlock = 256;
static constexpr int kReduceGrid = (kM * kN + kReduceBlock - 1) / kReduceBlock; // 132
__global__ __attribute__((amdgpu_flat_work_group_size(256, 256)))
void kernel_m16_reduce(
const float* __restrict__ workspace, // [7][16][2112]
uint16_t* __restrict__ C)
{
const int idx = blockIdx.x * kReduceBlock + threadIdx.x;
if (idx >= kM * kN) return;
float sum = 0.0f;
#pragma unroll
for (int s = 0; s < kSplitK; ++s)
sum += workspace[s * kM * kN + idx];
C[idx] = float_to_bf16_rn(sum);
}
} // namespace m16_varI
extern "C" void launch_m16_varI_gemm(
const uint16_t* A, const uint8_t* B, const uint8_t* Bs, float* ws) {
hipLaunchKernelGGL(m16_varI::kernel_m16_gemm,
dim3(m16_varI::kTotalBlocks), dim3(m16_varI::kBlockSize),
m16_varI::kLdsBytes, 0, A, B, Bs, ws);
}
extern "C" void launch_m16_varI_reduce(const float* ws, uint16_t* C) {
hipLaunchKernelGGL(m16_varI::kernel_m16_reduce,
dim3(m16_varI::kReduceGrid), dim3(m16_varI::kReduceBlock),
0, 0, ws, C);
}
// =====================================================================
// KERNEL 3: M=32 (m32_mk4l — 16x16x128 MFMA, split-M, both N values)
// =====================================================================
namespace m32_mk4l {
using v4i32 = int __attribute__((ext_vector_type(4)));
using v4f32 = float __attribute__((ext_vector_type(4)));
typedef __bf16 v2bf16 __attribute__((ext_vector_type(2)));
__device__ __forceinline__ uint32_t bit_cast_u32(float v) { union { float f; uint32_t u; } x; x.f=v; return x.u; }
__device__ __forceinline__ float bit_cast_f32(uint32_t v) { union { uint32_t u; float f; } x; x.u=v; return x.f; }
__device__ __forceinline__ uint16_t float_to_bf16_rn(float v) { uint32_t b=bit_cast_u32(v); b+=((b>>16)&1u)+0x7FFFu; return (uint16_t)(b>>16); }
__device__ __forceinline__ uint8_t compute_e8m0_scale(uint16_t m) { uint16_t r=(uint16_t)(m+0x20u); int e=(int)((r>>7)&0xFFu); return (e<=2)?0u:(e>=255)?254u:(uint8_t)(e-2); }
static constexpr int kM=32, kK=512, kKHalf=kK/2, kPaddedGroups=kK/32;
static constexpr int kWaveSize=64, kNumWaves=4, kBlockSize=kWaveSize*kNumWaves;
static constexpr int kMfmaN=16, kMfmaK=128, kTileN=32, kMfmaCols=kTileN/kMfmaN;
static constexpr int kTileM=16, kNumMTiles=kM/kTileM;
static constexpr int kBPanelStride=kKHalf*16;
__device__ __forceinline__ int scale_offset(int row, int group) {
return (row>>5)*(32*kPaddedGroups)+(group>>3)*256+(group&3)*64+(row&15)*4+((group>>2)&1)*2+((row>>4)&1);
}
__device__ __forceinline__ v4f32 mfma_fp4_16x16(v4i32 a, v4i32 b, v4f32 acc, int as, int bs) {
asm volatile("v_mfma_scale_f32_16x16x128_f8f6f4 %0,%1,%2,%0,%3,%4 cbsz:4 blgp:4":"+a"(acc):"v"(a),"v"(b),"v"(as),"v"(bs)); return acc;
}
__device__ __forceinline__ void quant_from_raw(const uint32_t* dw, v4i32& a_out, int& scale_out) {
uint32_t bmax=0u;
#pragma unroll
for(int i=0;i<16;++i){uint32_t w=dw[i];uint32_t hi=(w>>16)&0x7FFFu,lo=w&0x7FFFu;bmax=(hi>bmax)?hi:bmax;bmax=(lo>bmax)?lo:bmax;}
uint8_t scale=compute_e8m0_scale((uint16_t)bmax); scale_out=(int)scale;
float fwd=(scale==0u)?bit_cast_f32(0x00400000u):bit_cast_f32((uint32_t)scale<<23);
const v2bf16* p=reinterpret_cast<const v2bf16*>(dw); uint32_t pk[4]={0,0,0,0};
pk[0]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[0],p[0],fwd,0);pk[1]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[1],p[4],fwd,0);
pk[2]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[2],p[8],fwd,0);pk[3]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[3],p[12],fwd,0);
pk[0]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[0],p[1],fwd,1);pk[1]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[1],p[5],fwd,1);
pk[2]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[2],p[9],fwd,1);pk[3]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[3],p[13],fwd,1);
pk[0]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[0],p[2],fwd,2);pk[1]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[1],p[6],fwd,2);
pk[2]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[2],p[10],fwd,2);pk[3]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[3],p[14],fwd,2);
pk[0]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[0],p[3],fwd,3);pk[1]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[1],p[7],fwd,3);
pk[2]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[2],p[11],fwd,3);pk[3]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[3],p[15],fwd,3);
a_out={int(pk[0]),int(pk[1]),int(pk[2]),int(pk[3])};
}
template <int kN>
__global__ __attribute__((amdgpu_flat_work_group_size(256,256)))
void fused_gemm_m32(const uint16_t* __restrict__ A, const uint8_t* __restrict__ B_shuffle,
const uint8_t* __restrict__ B_scale_sh, uint16_t* __restrict__ C)
{
constexpr int kNumNTiles = kN / kTileN;
constexpr int kLdsFloats = kTileM * kNumWaves * kTileN;
const int n_tile=blockIdx.x/kNumMTiles, m_tile=blockIdx.x%kNumMTiles;
const int col_base=n_tile*kTileN, row_base=m_tile*kTileM;
const int tid=threadIdx.x, wave_id=tid>>6, lane=tid&63;
const int row16=lane&15, k_quarter=lane>>4;
extern __shared__ char lds_raw[]; float* lds=reinterpret_cast<float*>(lds_raw);
const int ki=wave_id, k_group=ki*4+k_quarter;
const int a_row=row_base+row16;
// B_scale loads first
int bsc[kMfmaCols];
#pragma unroll
for(int c=0;c<kMfmaCols;++c){
int b_col=col_base+c*kMfmaN+row16;
bsc[c]=(int)B_scale_sh[scale_offset(b_col,k_group)];
}
asm volatile("":"+v"(bsc[0]),"+v"(bsc[1])::"memory");
// B data loads
v4i32 b_regs[kMfmaCols];
#pragma unroll
for(int c=0;c<kMfmaCols;++c){
int b_col=col_base+c*kMfmaN+row16;
const uint8_t* bb=B_shuffle+(b_col>>4)*kBPanelStride+(b_col&15)*16;
const uint32_t* bs=reinterpret_cast<const uint32_t*>(bb+ki*1024+k_quarter*256);
b_regs[c]={int(bs[0]),int(bs[1]),int(bs[2]),int(bs[3])};
}
// A load
uint32_t a_buf[16];
{const uint32_t* as=reinterpret_cast<const uint32_t*>(A+a_row*kK+k_group*32);
#pragma unroll
for(int i=0;i<16;++i) a_buf[i]=as[i];}
v4f32 acc_c0,acc_c1;
#pragma unroll
for(int i=0;i<4;++i){acc_c0[i]=0.f;acc_c1[i]=0.f;}
{v4i32 ar;int as;quant_from_raw(a_buf,ar,as);
acc_c0=mfma_fp4_16x16(ar,b_regs[0],acc_c0,as,bsc[0]);
acc_c1=mfma_fp4_16x16(ar,b_regs[1],acc_c1,as,bsc[1]);}
// LDS write
const int out_col_16=lane&15, row_base4=(lane>>4)*4;
#pragma unroll
for(int r=0;r<4;++r){
int out_row=row_base4+r;
lds[out_row*(kNumWaves*kTileN)+wave_id*kTileN+out_col_16]=acc_c0[r];
lds[out_row*(kNumWaves*kTileN)+wave_id*kTileN+out_col_16+16]=acc_c1[r];
}
__syncthreads();
// Clean epilogue
const int e0=tid, e1=tid+kBlockSize;
const int lr0=e0/kTileN, lc0=e0%kTileN;
const int lr1=e1/kTileN, lc1=e1%kTileN;
const int lb0=lr0*(kNumWaves*kTileN)+lc0;
const int lb1=lr1*(kNumWaves*kTileN)+lc1;
float v0w0=lds[lb0],v0w1=lds[lb0+kTileN],v0w2=lds[lb0+2*kTileN],v0w3=lds[lb0+3*kTileN];
float v1w0=lds[lb1],v1w1=lds[lb1+kTileN],v1w2=lds[lb1+2*kTileN],v1w3=lds[lb1+3*kTileN];
float s0=v0w0+v0w1+v0w2+v0w3;
float s1=v1w0+v1w1+v1w2+v1w3;
C[(row_base+lr0)*kN+col_base+lc0]=float_to_bf16_rn(s0);
C[(row_base+lr1)*kN+col_base+lc1]=float_to_bf16_rn(s1);
}
static constexpr int kLdsBytes = kTileM * kNumWaves * kTileN * (int)sizeof(float);
} // namespace m32_mk4l
extern "C" void launch_m32_n2880(const uint16_t* A,const uint8_t* B,const uint8_t* Bs,uint16_t* C){
constexpr int N=2880, grid=(N/m32_mk4l::kTileN)*m32_mk4l::kNumMTiles;
hipLaunchKernelGGL((m32_mk4l::fused_gemm_m32<N>),dim3(grid),dim3(m32_mk4l::kBlockSize),m32_mk4l::kLdsBytes,0,A,B,Bs,C);
}
extern "C" void launch_m32_n4096(const uint16_t* A,const uint8_t* B,const uint8_t* Bs,uint16_t* C){
constexpr int N=4096, grid=(N/m32_mk4l::kTileN)*m32_mk4l::kNumMTiles;
hipLaunchKernelGGL((m32_mk4l::fused_gemm_m32<N>),dim3(grid),dim3(m32_mk4l::kBlockSize),m32_mk4l::kLdsBytes,0,A,B,Bs,C);
}
// =====================================================================
// KERNEL 4: M=64 (e641v1_exp S2 — hand-unrolled 8-iter, AGPR, B-scale LDS)
// =====================================================================
namespace e641v1_exp {
using v4i32 = int __attribute__((ext_vector_type(4)));
using v16f32 = float __attribute__((ext_vector_type(16)));
typedef __bf16 v2bf16 __attribute__((ext_vector_type(2)));
__device__ __forceinline__ uint32_t bit_cast_u32(float v) {
union { float f; uint32_t u; } x; x.f = v; return x.u;
}
__device__ __forceinline__ float bit_cast_f32(uint32_t v) {
union { uint32_t u; float f; } x; x.u = v; return x.f;
}
__device__ __forceinline__ uint16_t float_to_bf16_rn(float v) {
uint32_t bits = bit_cast_u32(v);
bits += ((bits >> 16) & 1u) + 0x7FFFu;
return static_cast<uint16_t>(bits >> 16);
}
__device__ __forceinline__ uint8_t compute_e8m0_scale(uint16_t max_abs) {
const uint16_t r = static_cast<uint16_t>(max_abs + 0x20u);
const int e = static_cast<int>((r >> 7) & 0xFFu);
return (e <= 2) ? 0u : (e >= 255) ? 254u : static_cast<uint8_t>(e - 2);
}
__device__ __forceinline__ v16f32 mfma_fp4x4(
v4i32 a, v4i32 b, v16f32 acc, int a_scale, int b_scale) {
asm volatile(
"v_mfma_scale_f32_32x32x64_f8f6f4 %0, %1, %2, %0, %3, %4 cbsz:4 blgp:4"
: "+a"(acc) : "v"(a), "v"(b), "v"(a_scale), "v"(b_scale));
return acc;
}
__device__ __forceinline__ void load_a_raw(
const uint32_t* __restrict__ a_ptr, uint32_t* dw)
{
*reinterpret_cast<uint4*>(&dw[0]) = *(reinterpret_cast<const uint4*>(a_ptr) + 0);
*reinterpret_cast<uint4*>(&dw[4]) = *(reinterpret_cast<const uint4*>(a_ptr) + 1);
*reinterpret_cast<uint4*>(&dw[8]) = *(reinterpret_cast<const uint4*>(a_ptr) + 2);
*reinterpret_cast<uint4*>(&dw[12]) = *(reinterpret_cast<const uint4*>(a_ptr) + 3);
}
__device__ __forceinline__ void quant_from_raw(
const uint32_t* dw, v4i32& a_out, int& scale_out)
{
uint32_t bmax = 0u;
#pragma unroll
for (int i = 0; i < 16; ++i) {
const uint32_t w = dw[i];
const uint32_t hi = (w >> 16) & 0x7FFFu;
const uint32_t lo = w & 0x7FFFu;
bmax = (hi > bmax) ? hi : bmax;
bmax = (lo > bmax) ? lo : bmax;
}
const uint8_t scale = compute_e8m0_scale(static_cast<uint16_t>(bmax));
scale_out = static_cast<int>(scale);
const float fwd = (scale == 0u) ? bit_cast_f32(0x00400000u)
: bit_cast_f32(static_cast<uint32_t>(scale) << 23);
const v2bf16* pairs = reinterpret_cast<const v2bf16*>(dw);
uint32_t pk[4] = {0,0,0,0};
pk[0]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[0],pairs[0], fwd,0);
pk[1]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[1],pairs[4], fwd,0);
pk[2]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[2],pairs[8], fwd,0);
pk[3]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[3],pairs[12],fwd,0);
pk[0]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[0],pairs[1], fwd,1);
pk[1]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[1],pairs[5], fwd,1);
pk[2]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[2],pairs[9], fwd,1);
pk[3]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[3],pairs[13],fwd,1);
pk[0]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[0],pairs[2], fwd,2);
pk[1]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[1],pairs[6], fwd,2);
pk[2]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[2],pairs[10],fwd,2);
pk[3]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[3],pairs[14],fwd,2);
pk[0]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[0],pairs[3], fwd,3);
pk[1]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[1],pairs[7], fwd,3);
pk[2]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[2],pairs[11],fwd,3);
pk[3]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[3],pairs[15],fwd,3);
a_out = {(int)pk[0], (int)pk[1], (int)pk[2], (int)pk[3]};
}
// S2 B scale offset: PaddedGroups=64
__device__ __forceinline__ int scale_offset_s2(int row, int group) {
return (row >> 5) * (32 * 64)
+ (group >> 3) * 256
+ (group & 3) * 64
+ (row & 15) * 4
+ ((group >> 2) & 1) * 2
+ ((row >> 4) & 1);
}
// S2 constants
static constexpr int S2_M = 64, S2_N = 7168, S2_K = 2048;
static constexpr int S2_KHalf = 1024;
static constexpr int S2_TileM = 32, S2_TileN = 64;
static constexpr int S2_NumWaves = 4, S2_BlockSize = 256;
static constexpr int S2_Iters = 8;
static constexpr int S2_NumTilesN = 112;
static constexpr int S2_NumTilesM = 2;
static constexpr int S2_TotalBlocks = 224;
static constexpr int S2_LdsBytes = 32768;
static constexpr int S2_NXCD = 8, S2_W = 2, S2_C = 14;
static constexpr int S2_BlocksPerCycle = 112;
static constexpr int S2_Limit = 224;
static constexpr int S2_TidPerGroup = 224;
#define S2_ITERATION(BUF_CUR, BUF_FAR, KI, LOAD_A_FAR, LOAD_B_NXT) \
{ \
const int k_group_ = (KI) * 2 + half; \
\
int bsc0_ = static_cast<int>(lds_bscale[ \
scale_offset_s2(b_col0, k_group_) - bscale_lds_base]); \
int bsc1_ = static_cast<int>(lds_bscale[ \
scale_offset_s2(b_col1, k_group_) - bscale_lds_base]); \
\
v4i32 areg_; \
int ascale_; \
quant_from_raw(BUF_CUR, areg_, ascale_); \
\
if constexpr (LOAD_B_NXT) { \
const int ki_nxt_ = (KI) + 1; \
const uint32_t* bnxt0_ = reinterpret_cast<const uint32_t*>( \
b_base0 + ki_nxt_ * 512 + half * 256); \
b_nxt0 = {(int)bnxt0_[0],(int)bnxt0_[1],(int)bnxt0_[2],(int)bnxt0_[3]}; \
const uint32_t* bnxt1_ = reinterpret_cast<const uint32_t*>( \
b_base1 + ki_nxt_ * 512 + half * 256); \
b_nxt1 = {(int)bnxt1_[0],(int)bnxt1_[1],(int)bnxt1_[2],(int)bnxt1_[3]}; \
} \
\
acc0 = mfma_fp4x4(areg_, b_cur0, acc0, ascale_, bsc0_); \
acc1 = mfma_fp4x4(areg_, b_cur1, acc1, ascale_, bsc1_); \
\
if constexpr (LOAD_A_FAR) { \
const int far_grp_ = ((KI) + 2) * 2 + half; \
load_a_raw(a_row + far_grp_ * 16, BUF_FAR); \
} \
\
if constexpr (LOAD_B_NXT) { \
b_cur0 = b_nxt0; \
b_cur1 = b_nxt1; \
} \
}
__global__ __attribute__((amdgpu_flat_work_group_size(256, 256)))
void fused_gemm_s2_handunrolled(
const uint16_t* __restrict__ A,
const uint8_t* __restrict__ B_shuffle,
const uint8_t* __restrict__ B_scale_sh,
uint16_t* __restrict__ C)
{
int xy = blockIdx.x;
if (xy < S2_Limit) {
const int xcd = xy % S2_NXCD;
const int local_ = xy / S2_NXCD;
const int chunk = local_ / S2_C;
const int pos = local_ % S2_C;
xy = chunk * S2_BlocksPerCycle + xcd * S2_C + pos;
}
const int l = xy % S2_TidPerGroup;
const int m_tile = (xy / S2_TidPerGroup) * S2_W + (l % S2_W);
const int n_tile = l / S2_W;
const int col_base = n_tile * S2_TileN;
const int row_base = m_tile * S2_TileM;
const int tid = threadIdx.x;
const int wave_id = tid >> 6;
const int lane = tid & 31;
const int half = (tid >> 5) & 1;
extern __shared__ char lds_bytes[];
constexpr int lds_half = S2_NumWaves * 32 * S2_TileM;
const int b_col0 = col_base + lane;
const int b_col1 = col_base + 32 + lane;
const uint8_t* b_base0 = B_shuffle
+ (b_col0 >> 4) * (S2_KHalf * 16) + (b_col0 & 15) * 16;
const uint8_t* b_base1 = B_shuffle
+ (b_col1 >> 4) * (S2_KHalf * 16) + (b_col1 & 15) * 16;
// Cooperative B scale preload into LDS
{
const int bscale_global_base_offset = (col_base >> 5) * (32 * 64);
const uint8_t* bscale_src = B_scale_sh + bscale_global_base_offset;
const int my_offset = tid * 16;
#pragma unroll
for (int bi = 0; bi < 16; ++bi) {
reinterpret_cast<uint8_t*>(lds_bytes)[my_offset + bi] =
bscale_src[my_offset + bi];
}
}
__syncthreads();
const uint8_t* lds_bscale = reinterpret_cast<const uint8_t*>(lds_bytes);
const int bscale_lds_base = (col_base >> 5) * (32 * 64);
const int my_row = row_base + lane;
const uint32_t* a_row = reinterpret_cast<const uint32_t*>(A + my_row * S2_K);
const int ki_base = wave_id * S2_Iters;
v16f32 acc0, acc1;
#pragma unroll
for (int i = 0; i < 16; ++i) { acc0[i] = 0.0f; acc1[i] = 0.0f; }
// Anti-thundering-herd stagger
if (blockIdx.x & 4) asm volatile("s_sleep 1" :::);
if (blockIdx.x & 8) asm volatile("s_sleep 1" :::);
// Prologue: B data for iter 0
v4i32 b_cur0, b_cur1;
{
const uint32_t* bs0 = reinterpret_cast<const uint32_t*>(
b_base0 + ki_base * 512 + half * 256);
b_cur0 = {(int)bs0[0], (int)bs0[1], (int)bs0[2], (int)bs0[3]};
const uint32_t* bs1 = reinterpret_cast<const uint32_t*>(
b_base1 + ki_base * 512 + half * 256);
b_cur1 = {(int)bs1[0], (int)bs1[1], (int)bs1[2], (int)bs1[3]};
}
// Three STATIC A buffers
uint32_t a_buf0[16], a_buf1[16], a_buf2[16];
load_a_raw(a_row + (ki_base * 2 + half) * 16, a_buf0);
load_a_raw(a_row + ((ki_base + 1) * 2 + half) * 16, a_buf1);
v4i32 b_nxt0, b_nxt1;
// 8 HAND-UNROLLED ITERATIONS
S2_ITERATION(a_buf0, a_buf2, ki_base + 0, true, true)
S2_ITERATION(a_buf1, a_buf0, ki_base + 1, true, true)
S2_ITERATION(a_buf2, a_buf1, ki_base + 2, true, true)
S2_ITERATION(a_buf0, a_buf2, ki_base + 3, true, true)
S2_ITERATION(a_buf1, a_buf0, ki_base + 4, true, true)
S2_ITERATION(a_buf2, a_buf1, ki_base + 5, true, true)
S2_ITERATION(a_buf0, a_buf2, ki_base + 6, false, true)
S2_ITERATION(a_buf1, a_buf0, ki_base + 7, false, false)
__syncthreads();
// Double-LDS reduction
float* lds_reduce = reinterpret_cast<float*>(lds_bytes);
#pragma unroll
for (int i = 0; i < 4; ++i) {
#pragma unroll
for (int j = 0; j < 4; ++j) {
const int row = half * 4 + j + i * 8;
const int base = row * S2_NumWaves * 32 + wave_id * 32 + lane;
lds_reduce[base] = acc0[i * 4 + j];
lds_reduce[base + lds_half] = acc1[i * 4 + j];
}
}
__syncthreads();
{
const int i = wave_id;
#pragma unroll
for (int j = 0; j < 4; ++j) {
const int row = half * 4 + j + i * 8;
const int global_row = row_base + row;
float sum0 = 0.0f;
#pragma unroll
for (int w = 0; w < S2_NumWaves; ++w)
sum0 += lds_reduce[row * S2_NumWaves * 32 + w * 32 + lane];
C[global_row * S2_N + col_base + lane] = float_to_bf16_rn(sum0);
float sum1 = 0.0f;
#pragma unroll
for (int w = 0; w < S2_NumWaves; ++w)
sum1 += lds_reduce[lds_half + row * S2_NumWaves * 32 + w * 32 + lane];
C[global_row * S2_N + col_base + 32 + lane] = float_to_bf16_rn(sum1);
}
}
}
#undef S2_ITERATION
} // namespace e641v1_exp
extern "C" void launch_e641v1_exp_64_7168_2048(
const uint16_t* A, const uint8_t* B, const uint8_t* Bs, uint16_t* C) {
hipLaunchKernelGGL(e641v1_exp::fused_gemm_s2_handunrolled,
dim3(e641v1_exp::S2_TotalBlocks), dim3(e641v1_exp::S2_BlockSize),
e641v1_exp::S2_LdsBytes, 0, A, B, Bs, C);
}
// =====================================================================
// KERNEL 5: M=256 (mk3g S4 — v_max3 absmax, deferred B scale)
// =====================================================================
namespace mk3g {
using v4i32 = int __attribute__((ext_vector_type(4)));
using v16f32 = float __attribute__((ext_vector_type(16)));
typedef __bf16 v2bf16 __attribute__((ext_vector_type(2)));
__device__ __forceinline__ uint32_t bit_cast_u32(float v) {
union { float f; uint32_t u; } x; x.f = v; return x.u;
}
__device__ __forceinline__ float bit_cast_f32(uint32_t v) {
union { uint32_t u; float f; } x; x.u = v; return x.f;
}
__device__ __forceinline__ uint16_t float_to_bf16_rn(float v) {
uint32_t bits = bit_cast_u32(v);
bits += ((bits >> 16) & 1u) + 0x7FFFu;
return static_cast<uint16_t>(bits >> 16);
}
__device__ __forceinline__ uint8_t compute_e8m0_scale(uint16_t max_abs) {
const uint16_t r = static_cast<uint16_t>(max_abs + 0x20u);
const int e = static_cast<int>((r >> 7) & 0xFFu);
return (e <= 2) ? 0u : (e >= 255) ? 254u : static_cast<uint8_t>(e - 2);
}
__device__ __forceinline__ v16f32 mfma_fp4x4(
v4i32 a, v4i32 b, v16f32 acc, int a_scale, int b_scale) {
asm volatile(
"v_mfma_scale_f32_32x32x64_f8f6f4 %0, %1, %2, %0, %3, %4 cbsz:4 blgp:4"
: "+v"(acc) : "v"(a), "v"(b), "v"(a_scale), "v"(b_scale));
return acc;
}
__device__ __forceinline__ void quant_a_group(
const uint32_t* __restrict__ a_ptr, v4i32& a_out, int& scale_out)
{
uint32_t dw[16];
*reinterpret_cast<uint4*>(&dw[0]) = *(reinterpret_cast<const uint4*>(a_ptr) + 0);
*reinterpret_cast<uint4*>(&dw[4]) = *(reinterpret_cast<const uint4*>(a_ptr) + 1);
*reinterpret_cast<uint4*>(&dw[8]) = *(reinterpret_cast<const uint4*>(a_ptr) + 2);
*reinterpret_cast<uint4*>(&dw[12]) = *(reinterpret_cast<const uint4*>(a_ptr) + 3);
uint32_t bmax = 0u;
#pragma unroll
for (int i = 0; i < 16; ++i) {
const uint32_t w = dw[i];
const uint32_t hi = (w >> 16) & 0x7FFFu;
const uint32_t lo = w & 0x7FFFu;
bmax = (hi > bmax) ? hi : bmax;
bmax = (lo > bmax) ? lo : bmax;
}
const uint8_t scale = compute_e8m0_scale(static_cast<uint16_t>(bmax));
scale_out = static_cast<int>(scale);
const float fwd = (scale == 0u) ? bit_cast_f32(0x00400000u)
: bit_cast_f32(static_cast<uint32_t>(scale) << 23);
const v2bf16* pairs = reinterpret_cast<const v2bf16*>(dw);
uint32_t pk[4] = {0,0,0,0};
pk[0]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[0],pairs[0], fwd,0);
pk[1]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[1],pairs[4], fwd,0);
pk[2]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[2],pairs[8], fwd,0);
pk[3]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[3],pairs[12],fwd,0);
pk[0]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[0],pairs[1], fwd,1);
pk[1]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[1],pairs[5], fwd,1);
pk[2]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[2],pairs[9], fwd,1);
pk[3]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[3],pairs[13],fwd,1);
pk[0]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[0],pairs[2], fwd,2);
pk[1]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[1],pairs[6], fwd,2);
pk[2]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[2],pairs[10],fwd,2);
pk[3]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[3],pairs[14],fwd,2);
pk[0]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[0],pairs[3], fwd,3);
pk[1]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[1],pairs[7], fwd,3);
pk[2]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[2],pairs[11],fwd,3);
pk[3]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk[3],pairs[15],fwd,3);
a_out = {(int)pk[0], (int)pk[1], (int)pk[2], (int)pk[3]};
}
// Only S4 config needed
struct S4 {
static constexpr int M = 256, N = 3072, K = 1536;
static constexpr int KHalf = K / 2;
static constexpr int GroupsPerRow = K / 32;
static constexpr int PaddedGroups = (GroupsPerRow + 7) & ~7;
static constexpr int TileM = 32, TileN = 64;
static constexpr int NumWaves = 4;
static constexpr int BlockSize = NumWaves * 64;
static constexpr int MfmaItersTotal = K / 64;
static constexpr int MfmaItersPerWave = MfmaItersTotal / NumWaves;
static constexpr int NumTilesN = N / TileN;
static constexpr int NumTilesM = M / TileM;
static constexpr int TotalBlocks = NumTilesN * NumTilesM;
static constexpr int LdsBytes = 2 * NumWaves * 32 * TileM * static_cast<int>(sizeof(float));
static constexpr int NXCD = 8;
static constexpr int W = 8, C = 6;
static constexpr int BlocksPerCycle = NXCD * C;
static constexpr int Limit = (TotalBlocks / BlocksPerCycle) * BlocksPerCycle;
static constexpr int TidPerGroup = W * NumTilesN;
};
template<typename C>
__device__ __forceinline__ int scale_offset(int row, int group) {
return (row >> 5) * (32 * C::PaddedGroups)
+ (group >> 3) * 256
+ (group & 3) * 64
+ (row & 15) * 4
+ ((group >> 2) & 1) * 2
+ ((row >> 4) & 1);
}
template<typename CFG>
__global__ __attribute__((amdgpu_flat_work_group_size(256, 256)))
void fused_gemm_mk3(
const uint16_t* __restrict__ A,
const uint8_t* __restrict__ B_shuffle,
const uint8_t* __restrict__ B_scale_sh,
uint16_t* __restrict__ C)
{
int xy = blockIdx.x;
if (xy < CFG::Limit) {
const int xcd = xy % CFG::NXCD;
const int local_ = xy / CFG::NXCD;
const int chunk = local_ / CFG::C;
const int pos = local_ % CFG::C;
xy = chunk * CFG::BlocksPerCycle + xcd * CFG::C + pos;
}
const int l = xy % CFG::TidPerGroup;
const int m_tile = (xy / CFG::TidPerGroup) * CFG::W + (l % CFG::W);
const int n_tile = l / CFG::W;
const int col_base = n_tile * CFG::TileN;
const int row_base = m_tile * CFG::TileM;
const int tid = threadIdx.x;
const int wave_id = tid >> 6;
const int lane = tid & 31;
const int half = (tid >> 5) & 1;
extern __shared__ char lds_bytes[];
float* lds_reduce = reinterpret_cast<float*>(lds_bytes);
constexpr int lds_half = CFG::NumWaves * 32 * CFG::TileM;
const int b_col0 = col_base + lane;
const int b_col1 = col_base + 32 + lane;
const uint8_t* b_base0 = B_shuffle
+ (b_col0 >> 4) * (CFG::KHalf * 16) + (b_col0 & 15) * 16;
const uint8_t* b_base1 = B_shuffle
+ (b_col1 >> 4) * (CFG::KHalf * 16) + (b_col1 & 15) * 16;
const int my_row = row_base + lane;
const uint32_t* a_row = reinterpret_cast<const uint32_t*>(A + my_row * CFG::K);
constexpr int iters = CFG::MfmaItersPerWave;
const int ki_base = wave_id * iters;
v16f32 acc0, acc1;
#pragma unroll
for (int i = 0; i < 16; ++i) { acc0[i] = 0.0f; acc1[i] = 0.0f; }
// B DATA prefetch
v4i32 b_cur0, b_cur1;
{
const uint32_t* bs0 = reinterpret_cast<const uint32_t*>(
b_base0 + ki_base * 512 + half * 256);
b_cur0 = {(int)bs0[0], (int)bs0[1], (int)bs0[2], (int)bs0[3]};
const uint32_t* bs1 = reinterpret_cast<const uint32_t*>(
b_base1 + ki_base * 512 + half * 256);
b_cur1 = {(int)bs1[0], (int)bs1[1], (int)bs1[2], (int)bs1[3]};
}
// Main K loop
#pragma unroll
for (int k = 0; k < iters; ++k) {
const int ki = ki_base + k;
const int k_group = ki * 2 + half;
// Load + quant A
v4i32 a_reg;
int a_scale;
quant_a_group(a_row + k_group * 16, a_reg, a_scale);
// Prefetch next B DATA
v4i32 b_nxt0 = {0,0,0,0}, b_nxt1 = {0,0,0,0};
if (k + 1 < iters) {
const int ki_next = ki_base + k + 1;
const uint32_t* bs0 = reinterpret_cast<const uint32_t*>(
b_base0 + ki_next * 512 + half * 256);
b_nxt0 = {(int)bs0[0], (int)bs0[1], (int)bs0[2], (int)bs0[3]};
const uint32_t* bs1 = reinterpret_cast<const uint32_t*>(
b_base1 + ki_next * 512 + half * 256);
b_nxt1 = {(int)bs1[0], (int)bs1[1], (int)bs1[2], (int)bs1[3]};
}
// B SCALE just-in-time
int b_sc0 = static_cast<int>(B_scale_sh[
scale_offset<CFG>(b_col0, k_group)]);
int b_sc1 = static_cast<int>(B_scale_sh[
scale_offset<CFG>(b_col1, k_group)]);
// MFMA
acc0 = mfma_fp4x4(a_reg, b_cur0, acc0, a_scale, b_sc0);
acc1 = mfma_fp4x4(a_reg, b_cur1, acc1, a_scale, b_sc1);
b_cur0 = b_nxt0;
b_cur1 = b_nxt1;
}
// Double-LDS reduction
#pragma unroll
for (int i = 0; i < 4; ++i) {
#pragma unroll
for (int j = 0; j < 4; ++j) {
const int row = half * 4 + j + i * 8;
const int base = row * CFG::NumWaves * 32 + wave_id * 32 + lane;
lds_reduce[base] = acc0[i * 4 + j];
lds_reduce[base + lds_half] = acc1[i * 4 + j];
}
}
__syncthreads();
{
const int i = wave_id;
#pragma unroll
for (int j = 0; j < 4; ++j) {
const int row = half * 4 + j + i * 8;
const int global_row = row_base + row;
float sum0 = 0.0f;
#pragma unroll
for (int w = 0; w < CFG::NumWaves; ++w)
sum0 += lds_reduce[row * CFG::NumWaves * 32 + w * 32 + lane];
C[global_row * CFG::N + col_base + lane] = float_to_bf16_rn(sum0);
float sum1 = 0.0f;
#pragma unroll
for (int w = 0; w < CFG::NumWaves; ++w)
sum1 += lds_reduce[lds_half + row * CFG::NumWaves * 32 + w * 32 + lane];
C[global_row * CFG::N + col_base + 32 + lane] = float_to_bf16_rn(sum1);
}
}
}
} // namespace mk3g
extern "C" void launch_mk3g_256_3072_1536(
const uint16_t* A, const uint8_t* B, const uint8_t* Bs, uint16_t* C) {
using C_ = mk3g::S4;
hipLaunchKernelGGL(mk3g::fused_gemm_mk3<C_>,
dim3(C_::TotalBlocks), dim3(C_::BlockSize), C_::LdsBytes, 0, A, B, Bs, C);
}
"""
# =========================================================================
# Python wrapper
# =========================================================================
def _extra_cuda_cflags() -> list[str]:
flags = ["-O3", "-std=c++17"]
arch = os.environ.get("PYTORCH_ROCM_ARCH", "gfx950")
if all(sep not in arch for sep in (",", ";", " ")):
flags.append(f"--offload-arch={arch}")
# flags.append("-mllvm")
# flags.append("-amdgpu-kernarg-preload-count=8")
return flags
def _build_module():
os.environ["PYTORCH_ROCM_ARCH"] = "gfx950"
name = "submission_general_v1"
build_dir = os.path.join(tempfile.gettempdir(), name)
os.makedirs(build_dir, exist_ok=True)
return load_inline(
name=name,
cpp_sources=_CPP_SRC,
cuda_sources=_HIP_SRC,
functions=[
"m4_launch",
"m16_launch",
"m32_n2880_launch",
"m32_n4096_launch",
"m64_launch",
"m256_launch",
],
with_cuda=True,
extra_cflags=["-O3"],
extra_cuda_cflags=_extra_cuda_cflags(),
build_directory=build_dir,
verbose=False,
)
_mod = _build_module()
def _storage_as_uint8(x):
return x if x.dtype == torch.uint8 else x.view(torch.uint8)
def custom_kernel(data: input_t) -> output_t:
A, _B, B_q, B_shuffle, B_scale_sh = data
m, k = int(A.shape[0]), int(A.shape[1])
n = int(B_shuffle.shape[0])
if m == 4 and n == 2880 and k == 512:
B_sh_u8 = _storage_as_uint8(B_shuffle)
B_sc_u8 = _storage_as_uint8(B_scale_sh)
C = torch.empty((m, n), dtype=torch.bfloat16, device=A.device)
_mod.m4_launch(A, B_sh_u8, B_sc_u8, C)
return C
if m == 16 and n == 2112 and k == 7168:
B_sh_u8 = _storage_as_uint8(B_shuffle)
B_sc_u8 = _storage_as_uint8(B_scale_sh)
workspace = torch.empty((7, m, n), dtype=torch.float32, device=A.device)
C = torch.empty((m, n), dtype=torch.bfloat16, device=A.device)
_mod.m16_launch(A, B_sh_u8, B_sc_u8, workspace, C)
return C
if m == 32 and n == 2880 and k == 512:
B_sh_u8 = _storage_as_uint8(B_shuffle)
B_sc_u8 = _storage_as_uint8(B_scale_sh)
C = torch.empty((m, n), dtype=torch.bfloat16, device=A.device)
_mod.m32_n2880_launch(A, B_sh_u8, B_sc_u8, C)
return C
if m == 32 and n == 4096 and k == 512:
B_sh_u8 = _storage_as_uint8(B_shuffle)
B_sc_u8 = _storage_as_uint8(B_scale_sh)
C = torch.empty((m, n), dtype=torch.bfloat16, device=A.device)
_mod.m32_n4096_launch(A, B_sh_u8, B_sc_u8, C)
return C
if m == 64 and n == 7168 and k == 2048:
B_sh_u8 = _storage_as_uint8(B_shuffle)
B_sc_u8 = _storage_as_uint8(B_scale_sh)
C = torch.empty((m, n), dtype=torch.bfloat16, device=A.device)
_mod.m64_launch(A, B_sh_u8, B_sc_u8, C)
return C
if m == 256 and n == 3072 and k == 1536:
B_sh_u8 = _storage_as_uint8(B_shuffle)
B_sc_u8 = _storage_as_uint8(B_scale_sh)
C = torch.empty((m, n), dtype=torch.bfloat16, device=A.device)
_mod.m256_launch(A, B_sh_u8, B_sc_u8, C)
return C
# Fallback to aiter reference
x_fp4, bs = dynamic_mxfp4_quant(A)
bs = e8m0_shuffle(bs)
return aiter.gemm_a4w4(
x_fp4.view(dtypes.fp4x2), B_shuffle,
bs.view(dtypes.fp8_e8m0), B_scale_sh,
dtype=dtypes.bf16, bpreshuffle=True)
scrolls · 1500 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