submission 744947
ak65432 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 1127 lines, June 9 Researcher Reciprocity License v1.0.
v274_var.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-744947?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:6e4b17de8ed03b893862fed118945d1faa4eb9567bbc1e2fa5d466295c2b885e
license declaredunknown
license concludedunknown
authorsak65432
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
autotune
@triton.autotune(shared-memory
__shared__ uint8_t B_lds[4096];split-k
SPLITK_BLOCK = 512tile-k = 16
BLOCK_M, BLOCK_K = 16, 512tile-n = 128
REDUCE_BN = 128vector-width = float4
union { bf16x2_t bf[16]; float4 f4[4]; } a0_u, a1_u;Kernel source
v274_var.py1127 lines
# v274: Triton S2 no XCD remap
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
# v234: Route S3/S4 through 2-wave split-K instead of single-wave
# Hypothesis: 2 waves/WG improves latency hiding for under-parallelized shapes
# S3: 128 WGs, S4: 90 WGs on 304 CUs → 2-wave gives 2x more wavefronts
import os, sys
os.environ['PYTORCH_ROCM_ARCH'] = 'gfx950'
os.environ['HIP_FORCE_DEV_KERNARG'] = '1'
os.environ['GPU_FORCE_BLIT_COPY_SIZE'] = '64'
os.environ['CXX'] = 'clang++'
import torch
import triton
import triton.language as tl
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t
from aiter.ops.triton._triton_kernels.quant.quant import _mxfp4_quant_op
SHAPES = [(4,2880,512),(16,2112,7168),(32,4096,512),(32,2880,512),
(64,7168,2048),(256,3072,1536),(8,2112,7168),(16,3072,1536),
(64,3072,1536),(256,2880,512)]
NARROW_MFMA_SHAPES = {(4, 2880, 512), (16, 3072, 1536)}
WIDE_MFMA_SHAPES = {(64, 3072, 1536), (256, 2880, 512)}
TRITON_SHAPES = {(16, 2112, 7168), (8, 2112, 7168)}
# S5/S6 now use the new fused ASM-scheduled kernel
# S3/S4 moved from single-wave to 2-wave split-K
FUSED_ASM_SHAPES = {(64, 7168, 2048), (256, 3072, 1536)}
TWOWAVE_SHAPES = {(32, 4096, 512), (32, 2880, 512)}
NUM_KSPLIT = 14
SPLITK_BLOCK = 512
BLOCK_M, BLOCK_K = 16, 512
REDUCE_BN = 128
# ============================================================
# Triton XCD remap
# ============================================================
@triton.jit
def _remap_xcd(pid, GRID_TOTAL, NUM_XCDS: tl.constexpr = 8):
return pid # v274: disabled XCD remap
# ============================================================
# Triton split-K with SHUFFLED B_scale indexing (proven from v772)
# ============================================================
@triton.autotune(
configs=[
triton.Config({'BLOCK_N': bn}, num_warps=w, num_stages=s)
for bn in [64, 128, 256]
for w in [4, 8]
for s in [2, 3]
],
key=['M', 'N', 'K'],
)
@triton.jit
def _triton_fused_quant_gemm_splitk(
a_bf16_ptr, b_ptr, ws_ptr, b_scale_sh_ptr,
M, N, K,
stride_am, stride_ak, stride_bk, stride_bn,
stride_wk, stride_wm, stride_wn,
BSD0_STRIDE,
BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr,
NUM_KSPLIT: tl.constexpr, SPLITK_SIZE: tl.constexpr,
):
SCALE_GROUP: tl.constexpr = 32
pid_raw = tl.program_id(0)
grid_mn = tl.cdiv(M, BLOCK_M) * tl.cdiv(N, BLOCK_N)
pid_raw = _remap_xcd(pid_raw, grid_mn * NUM_KSPLIT)
pid_k = pid_raw % NUM_KSPLIT
pid = pid_raw // NUM_KSPLIT
num_pid_n = tl.cdiv(N, BLOCK_N)
pid_m = pid // num_pid_n
pid_n = pid % num_pid_n
num_k_iter = SPLITK_SIZE // BLOCK_K
offs_am = (pid_m * BLOCK_M + tl.arange(0, BLOCK_M)) % M
offs_k_bf16 = tl.arange(0, BLOCK_K)
k_start_bf16 = pid_k * SPLITK_SIZE
a_bf16_ptrs = a_bf16_ptr + offs_am[:, None] * stride_am + (k_start_bf16 + offs_k_bf16[None, :]) * stride_ak
offs_k_packed = tl.arange(0, BLOCK_K // 2)
k_start_packed = pid_k * (SPLITK_SIZE // 2)
offs_bn = (pid_n * BLOCK_N + tl.arange(0, BLOCK_N)) % N
b_ptrs = b_ptr + (k_start_packed + offs_k_packed[:, None]) * stride_bk + offs_bn[None, :] * stride_bn
# B_scale SHUFFLED index precompute (N-dependent parts)
bs_d0 = offs_bn // 32
bs_d1 = (offs_bn & 31) >> 4
bs_d2 = offs_bn & 15
bs_n_part = bs_d0 * BSD0_STRIDE + bs_d2 * 4 + bs_d1
acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
k_start_scale = pid_k * (SPLITK_SIZE // SCALE_GROUP)
offs_ks = tl.arange(0, BLOCK_K // SCALE_GROUP)
for ki in range(num_k_iter):
a_bf16 = tl.load(a_bf16_ptrs)
a_fp4, a_scale = _mxfp4_quant_op(a_bf16, BLOCK_K, BLOCK_M, SCALE_GROUP)
gs = k_start_scale + ki * (BLOCK_K // SCALE_GROUP) + offs_ks
bs_s3 = gs >> 3
bs_s4 = (gs & 7) >> 2
bs_s5 = gs & 3
bs_g_part = bs_s3 * 256 + bs_s5 * 64 + bs_s4 * 2
b_scale_ptrs = b_scale_sh_ptr + bs_n_part[:, None] + bs_g_part[None, :]
b_scales = tl.load(b_scale_ptrs, cache_modifier=".cg")
b = tl.load(b_ptrs, cache_modifier=".cg")
acc = tl.dot_scaled(a_fp4, a_scale, "e2m1", b, b_scales, "e2m1", acc)
a_bf16_ptrs += BLOCK_K * stride_ak
b_ptrs += (BLOCK_K // 2) * stride_bk
offs_cm = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
offs_cn = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
ws_ptrs = ws_ptr + pid_k * stride_wk + offs_cm[:, None] * stride_wm + offs_cn[None, :] * stride_wn
mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < N)
tl.store(ws_ptrs, acc, mask=mask)
@triton.jit
def _splitk_reduce(
ws_ptr, out_ptr, M, N,
stride_wk, stride_wm, stride_wn, stride_om, stride_on,
NUM_KSPLIT: tl.constexpr, BLOCK_RN: tl.constexpr,
):
pid_m = tl.program_id(0)
pid_n = tl.program_id(1)
offs_n = pid_n * BLOCK_RN + tl.arange(0, BLOCK_RN)
mask = offs_n < N
acc = tl.zeros((BLOCK_RN,), dtype=tl.float32)
base = ws_ptr + pid_m * stride_wm
for k in range(NUM_KSPLIT):
val = tl.load(base + k * stride_wk + offs_n * stride_wn, mask=mask, other=0.0)
acc += val
tl.store(out_ptr + pid_m * stride_om + offs_n * stride_on, acc.to(tl.bfloat16), mask=mask)
# ============================================================
# HIP C++ source
# ============================================================
HIP_SRC = r"""
#include <hip/hip_runtime.h>
#include <hip/amd_detail/amd_hip_bf16.h>
#include <torch/types.h>
#include <unordered_map>
#include <vector>
#include <cstdio>
#include <cstring>
using v8i32 = int32_t __attribute__((ext_vector_type(8)));
using v4f32 = float __attribute__((ext_vector_type(4)));
using i32x4 = int32_t __attribute__((ext_vector_type(4)));
typedef __attribute__((ext_vector_type(2))) __bf16 bf16x2_t;
using as3_ptr = uint32_t __attribute__((address_space(3)))*;
#define SPTR(_p_) reinterpret_cast<as3_ptr>(reinterpret_cast<uintptr_t>(_p_))
extern "C" __device__ void llvm_amdgcn_raw_buffer_load_lds(
i32x4 rsrc, as3_ptr lds_ptr, int size,
int voffset, int soffset, int offset, int aux
) __asm("llvm.amdgcn.raw.buffer.load.lds");
__device__ __forceinline__ i32x4 make_buffer_srd(const void* ptr, uint32_t nbytes) {
i32x4 r;
uint64_t a = reinterpret_cast<uint64_t>(ptr);
r[0] = (int32_t)(a); r[1] = (int32_t)(a >> 32);
r[2] = (int32_t)nbytes; r[3] = 0x00020000;
return r;
}
__device__ __forceinline__ void quant_32bf16_to_fp4(
const bf16x2_t vals[16], bool valid, uint32_t ap[4], int& a_scale_out)
{
uint32_t max_packed = 0;
if (valid) {
#pragma unroll
for (int i = 0; i < 16; i++) {
uint32_t packed; __builtin_memcpy(&packed, &vals[i], 4);
uint32_t abs_packed = packed & 0x7FFF7FFF;
asm volatile("v_pk_max_u16 %0, %0, %1" : "+v"(max_packed) : "v"(abs_packed));
}
}
uint16_t lo = max_packed & 0xFFFF, hi = max_packed >> 16;
uint16_t umax = lo > hi ? lo : hi;
uint32_t amax_bits = (uint32_t)umax << 16;
float amax; __builtin_memcpy(&amax, &amax_bits, 4);
uint32_t ab; __builtin_memcpy(&ab, &amax, 4);
ab = (ab + 0x200000u) & 0xFF800000u;
int su = (ab == 0) ? -127 : ((int)((ab >> 23) & 0xFF) - 129);
su = su < -127 ? -127 : (su > 127 ? 127 : su);
a_scale_out = su + 127;
float sf; { uint32_t qb = (uint32_t)a_scale_out << 23; __builtin_memcpy(&sf, &qb, 4); }
if (valid) {
#pragma unroll
for (int d = 0; d < 4; d++) {
uint32_t pk = 0;
pk = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk, vals[d*4+0], sf, 0);
pk = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk, vals[d*4+1], sf, 1);
pk = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk, vals[d*4+2], sf, 2);
pk = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(pk, vals[d*4+3], sf, 3);
ap[d] = pk;
}
} else { ap[0]=ap[1]=ap[2]=ap[3]=0; a_scale_out=127; }
}
// ============================================================
// Narrow 16x16 fused kernel (unchanged from v742)
// ============================================================
__global__ __launch_bounds__(64, 4)
void fused_quant_mfma_narrow(
const __hip_bfloat16* __restrict__ A, const uint8_t* __restrict__ B_shuffle,
const uint8_t* __restrict__ B_scale_sh, __hip_bfloat16* __restrict__ out,
int M, int N, int K, int N_tiles, int M_tiles)
{
const int tid = threadIdx.x;
const int mn_tile = blockIdx.x;
const int m_tile = mn_tile % M_tiles, n_tile = mn_tile / M_tiles;
const int n_base = n_tile * 16, m_base = m_tile * 16;
const int K_half = K / 2;
const int kg_pad = ((K / 32) + 7) & ~7;
const int bscale_d0_stride = (kg_pad / 8) * 256;
const int a_m_row = m_base + (tid & 15);
const int a_k_block = tid >> 4;
const bool valid_a = (a_m_row < M);
const i32x4 b_srd = make_buffer_srd(B_shuffle, (uint32_t)((int64_t)N * K_half));
const int b_col = tid & 15, b_k_block = tid >> 4;
const int b_voffset = b_k_block * 256 + b_col * 16;
const int b_soff_base = n_base * K_half;
const int b_global_n = n_base + b_col;
const int bs_sd0 = b_global_n / 32, bs_sd1 = (b_global_n & 31) >> 4, bs_sd2 = b_global_n & 15;
__shared__ uint8_t B_lds[4096];
v4f32 acc = {0,0,0,0};
bf16x2_t a0_bf16[16], a1_bf16[16];
if (0 < K) {
const bf16x2_t* a0_src = (const bf16x2_t*)(A + (size_t)(valid_a ? a_m_row : 0) * K + a_k_block * 32);
const bf16x2_t* a1_src = (const bf16x2_t*)(A + (size_t)(valid_a ? a_m_row : 0) * K + 128 + a_k_block * 32);
if (valid_a) {
#pragma unroll
for (int i = 0; i < 16; i++) a0_bf16[i] = a0_src[i];
#pragma unroll
for (int i = 0; i < 16; i++) a1_bf16[i] = a1_src[i];
}
llvm_amdgcn_raw_buffer_load_lds(b_srd, SPTR(&B_lds[0]), 16, b_voffset, b_soff_base, 0, 0);
llvm_amdgcn_raw_buffer_load_lds(b_srd, SPTR(&B_lds[1024]), 16, b_voffset, b_soff_base + 1024, 0, 0);
asm volatile("s_waitcnt vmcnt(0)" ::: "memory");
}
for (int k_step = 0; k_step < K; k_step += 256) {
const int cur_base = ((k_step >> 8) & 1) ? 2048 : 0;
const int nxt_base = 2048 - cur_base;
const bool has_next = (k_step + 256 < K);
uint32_t ap0[4]; int as0;
quant_32bf16_to_fp4(a0_bf16, valid_a, ap0, as0);
v8i32 a_reg = {}; a_reg[0]=ap0[0]; a_reg[1]=ap0[1]; a_reg[2]=ap0[2]; a_reg[3]=ap0[3];
const uint32_t* bl0 = (const uint32_t*)(&B_lds[cur_base + tid * 16]);
v8i32 b_reg = {}; b_reg[0]=bl0[0]; b_reg[1]=bl0[1]; b_reg[2]=bl0[2]; b_reg[3]=bl0[3];
int bs0;
{ int bkg = k_step/32 + b_k_block; int s3=bkg>>3; int s4=(bkg&7)>>2; int s5=bkg&3;
bs0 = (b_global_n < N) ? (int)B_scale_sh[bs_sd0*bscale_d0_stride + s3*256 + s5*64 + bs_sd2*4 + s4*2 + bs_sd1] : 127; }
if (has_next) llvm_amdgcn_raw_buffer_load_lds(b_srd, SPTR(&B_lds[nxt_base]), 16, b_voffset, b_soff_base + (k_step+256)*8, 0, 0);
if (has_next && valid_a) {
const bf16x2_t* an = (const bf16x2_t*)(A + (size_t)a_m_row * K + k_step + 256 + a_k_block * 32);
#pragma unroll
for (int i = 0; i < 16; i++) a0_bf16[i] = an[i];
}
acc = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(a_reg, b_reg, acc, 4, 4, 0, as0, 0, bs0);
uint32_t ap1[4]; int as1;
quant_32bf16_to_fp4(a1_bf16, valid_a, ap1, as1);
a_reg[0]=ap1[0]; a_reg[1]=ap1[1]; a_reg[2]=ap1[2]; a_reg[3]=ap1[3];
const uint32_t* bl1 = (const uint32_t*)(&B_lds[cur_base + 1024 + tid * 16]);
b_reg[0]=bl1[0]; b_reg[1]=bl1[1]; b_reg[2]=bl1[2]; b_reg[3]=bl1[3];
int bs1;
{ int bkg = (k_step+128)/32 + b_k_block; int s3=bkg>>3; int s4=(bkg&7)>>2; int s5=bkg&3;
bs1 = (b_global_n < N) ? (int)B_scale_sh[bs_sd0*bscale_d0_stride + s3*256 + s5*64 + bs_sd2*4 + s4*2 + bs_sd1] : 127; }
if (has_next) llvm_amdgcn_raw_buffer_load_lds(b_srd, SPTR(&B_lds[nxt_base + 1024]), 16, b_voffset, b_soff_base + (k_step+256)*8 + 1024, 0, 0);
if (has_next && valid_a) {
const bf16x2_t* an = (const bf16x2_t*)(A + (size_t)a_m_row * K + k_step + 256 + 128 + a_k_block * 32);
#pragma unroll
for (int i = 0; i < 16; i++) a1_bf16[i] = an[i];
}
acc = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(a_reg, b_reg, acc, 4, 4, 0, as1, 0, bs1);
if (has_next) { asm volatile("s_waitcnt vmcnt(0)" ::: "memory"); }
}
const int out_row_base = m_base + (tid >> 4) * 4;
const int out_col = n_base + (tid & 15);
if (out_col < N) {
#pragma unroll
for (int v = 0; v < 4; v++) { int row = out_row_base + v; if (row < M) out[(int64_t)row * N + out_col] = __float2bfloat16(acc[v]); }
}
}
// ============================================================
// Wide 16x32 fused kernel (unchanged from v742 — for S3/S4/S9/S10)
// ============================================================
__global__ __launch_bounds__(64, 4)
void fused_quant_mfma_wide(
const __hip_bfloat16* __restrict__ A, const uint8_t* __restrict__ B_shuffle,
const uint8_t* __restrict__ B_scale_sh, __hip_bfloat16* __restrict__ out,
int M, int N, int K, int N_tiles, int M_tiles)
{
const int tid = threadIdx.x;
const int mn_tile = blockIdx.x;
const int m_tile = mn_tile % M_tiles, n_tile = mn_tile / M_tiles;
const int n_base = n_tile * 32, m_base = m_tile * 16;
const int K_half = K / 2;
const int kg_pad = ((K / 32) + 7) & ~7;
const int bscale_d0_stride = (kg_pad / 8) * 256;
const int a_m_row = m_base + (tid & 15);
const int a_k_block = tid >> 4;
const bool valid_a = (a_m_row < M);
const i32x4 b_srd = make_buffer_srd(B_shuffle, (uint32_t)((int64_t)N * K_half));
const int b_col = tid & 15, b_k_block = tid >> 4;
const int b_voffset = b_k_block * 256 + b_col * 16;
const int b_soff_base_0 = n_base * K_half, b_soff_base_1 = (n_base + 16) * K_half;
const int b_global_n_0 = n_base + b_col, b_global_n_1 = n_base + 16 + b_col;
const int bs_sd0_0 = b_global_n_0/32, bs_sd1_0 = (b_global_n_0&31)>>4, bs_sd2_0 = b_global_n_0&15;
const int bs_sd0_1 = b_global_n_1/32, bs_sd1_1 = (b_global_n_1&31)>>4, bs_sd2_1 = b_global_n_1&15;
__shared__ uint8_t B_lds[8192];
v4f32 acc0={0,0,0,0}, acc1={0,0,0,0};
bf16x2_t a0_bf16[16], a1_bf16[16];
if (0 < K) {
if (valid_a) {
const bf16x2_t* a0_src = (const bf16x2_t*)(A + (size_t)a_m_row * K + a_k_block * 32);
const bf16x2_t* a1_src = (const bf16x2_t*)(A + (size_t)a_m_row * K + 128 + a_k_block * 32);
#pragma unroll
for (int i = 0; i < 16; i++) a0_bf16[i] = a0_src[i];
#pragma unroll
for (int i = 0; i < 16; i++) a1_bf16[i] = a1_src[i];
}
llvm_amdgcn_raw_buffer_load_lds(b_srd, SPTR(&B_lds[0]), 16, b_voffset, b_soff_base_0, 0, 0);
llvm_amdgcn_raw_buffer_load_lds(b_srd, SPTR(&B_lds[1024]), 16, b_voffset, b_soff_base_1, 0, 0);
llvm_amdgcn_raw_buffer_load_lds(b_srd, SPTR(&B_lds[2048]), 16, b_voffset, b_soff_base_0 + 1024, 0, 0);
llvm_amdgcn_raw_buffer_load_lds(b_srd, SPTR(&B_lds[3072]), 16, b_voffset, b_soff_base_1 + 1024, 0, 0);
asm volatile("s_waitcnt vmcnt(0)" ::: "memory");
}
for (int k_step = 0; k_step < K; k_step += 256) {
const int cur_base = ((k_step >> 8) & 1) ? 4096 : 0;
const int nxt_base = 4096 - cur_base;
const bool has_next = (k_step + 256 < K);
if (has_next) {
llvm_amdgcn_raw_buffer_load_lds(b_srd, SPTR(&B_lds[nxt_base]), 16, b_voffset, b_soff_base_0 + (k_step+256)*8, 0, 0);
llvm_amdgcn_raw_buffer_load_lds(b_srd, SPTR(&B_lds[nxt_base + 1024]), 16, b_voffset, b_soff_base_1 + (k_step+256)*8, 0, 0);
llvm_amdgcn_raw_buffer_load_lds(b_srd, SPTR(&B_lds[nxt_base + 2048]), 16, b_voffset, b_soff_base_0 + (k_step+256)*8 + 1024, 0, 0);
llvm_amdgcn_raw_buffer_load_lds(b_srd, SPTR(&B_lds[nxt_base + 3072]), 16, b_voffset, b_soff_base_1 + (k_step+256)*8 + 1024, 0, 0);
}
__builtin_amdgcn_sched_barrier(0);
uint32_t ap[4]; int as0;
quant_32bf16_to_fp4(a0_bf16, valid_a, ap, as0);
v8i32 a_reg = {}; a_reg[0]=ap[0]; a_reg[1]=ap[1]; a_reg[2]=ap[2]; a_reg[3]=ap[3];
int bs0s0, bs0s1;
{ int bkg=k_step/32+b_k_block; int s3=bkg>>3; int s4=(bkg&7)>>2; int s5=bkg&3;
bs0s0 = (b_global_n_0<N) ? (int)B_scale_sh[bs_sd0_0*bscale_d0_stride+s3*256+s5*64+bs_sd2_0*4+s4*2+bs_sd1_0] : 127;
bs0s1 = (b_global_n_1<N) ? (int)B_scale_sh[bs_sd0_1*bscale_d0_stride+s3*256+s5*64+bs_sd2_1*4+s4*2+bs_sd1_1] : 127; }
{ const uint32_t* bl=(const uint32_t*)(&B_lds[cur_base+tid*16]); v8i32 br={};
br[0]=bl[0]; br[1]=bl[1]; br[2]=bl[2]; br[3]=bl[3];
acc0 = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(a_reg, br, acc0, 4, 4, 0, as0, 0, bs0s0); }
{ const uint32_t* bl=(const uint32_t*)(&B_lds[cur_base+1024+tid*16]); v8i32 br={};
br[0]=bl[0]; br[1]=bl[1]; br[2]=bl[2]; br[3]=bl[3];
acc1 = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(a_reg, br, acc1, 4, 4, 0, as0, 0, bs0s1); }
int as1;
quant_32bf16_to_fp4(a1_bf16, valid_a, ap, as1);
a_reg[0]=ap[0]; a_reg[1]=ap[1]; a_reg[2]=ap[2]; a_reg[3]=ap[3];
int bs1s0, bs1s1;
{ int bkg=(k_step+128)/32+b_k_block; int s3=bkg>>3; int s4=(bkg&7)>>2; int s5=bkg&3;
bs1s0 = (b_global_n_0<N) ? (int)B_scale_sh[bs_sd0_0*bscale_d0_stride+s3*256+s5*64+bs_sd2_0*4+s4*2+bs_sd1_0] : 127;
bs1s1 = (b_global_n_1<N) ? (int)B_scale_sh[bs_sd0_1*bscale_d0_stride+s3*256+s5*64+bs_sd2_1*4+s4*2+bs_sd1_1] : 127; }
{ const uint32_t* bl=(const uint32_t*)(&B_lds[cur_base+2048+tid*16]); v8i32 br={};
br[0]=bl[0]; br[1]=bl[1]; br[2]=bl[2]; br[3]=bl[3];
acc0 = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(a_reg, br, acc0, 4, 4, 0, as1, 0, bs1s0); }
{ const uint32_t* bl=(const uint32_t*)(&B_lds[cur_base+3072+tid*16]); v8i32 br={};
br[0]=bl[0]; br[1]=bl[1]; br[2]=bl[2]; br[3]=bl[3];
acc1 = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(a_reg, br, acc1, 4, 4, 0, as1, 0, bs1s1); }
if (has_next && valid_a) {
const bf16x2_t* an0 = (const bf16x2_t*)(A + (size_t)a_m_row * K + k_step + 256 + a_k_block * 32);
const bf16x2_t* an1 = (const bf16x2_t*)(A + (size_t)a_m_row * K + k_step + 256 + 128 + a_k_block * 32);
#pragma unroll
for (int i = 0; i < 16; i++) a0_bf16[i] = an0[i];
#pragma unroll
for (int i = 0; i < 16; i++) a1_bf16[i] = an1[i];
}
if (has_next) { asm volatile("s_waitcnt vmcnt(0)" ::: "memory"); }
}
const int out_row_base = m_base + (tid >> 4) * 4;
const int out_col_0 = n_base + (tid & 15), out_col_1 = n_base + 16 + (tid & 15);
if (out_col_0 < N) {
#pragma unroll
for (int v = 0; v < 4; v++) { int r=out_row_base+v; if(r<M) out[(int64_t)r*N+out_col_0]=__float2bfloat16(acc0[v]); }
}
if (out_col_1 < N) {
#pragma unroll
for (int v = 0; v < 4; v++) { int r=out_row_base+v; if(r<M) out[(int64_t)r*N+out_col_1]=__float2bfloat16(acc1[v]); }
}
}
// ============================================================
// v788: Template-specialized K-loop + paired B_scale loads
// template<K_STEPS> gives compiler full visibility for cross-iteration
// scheduling. #pragma unroll with compile-time bound → branch-free code.
// B_scale: exploit n_base 64-alignment so sub-tile pairs (0,1) and (2,3)
// have contiguous 4-byte scale blocks. 2 uint32_t loads per K-step
// replace 8 scattered byte loads + all s3/s4/s5 address math.
// ============================================================
template<int K_STEPS>
__global__ __launch_bounds__(64, 2)
void fused_quant_mfma_asm_sched(
const __hip_bfloat16* __restrict__ A, const uint8_t* __restrict__ B_shuffle,
const uint8_t* __restrict__ B_scale_sh, __hip_bfloat16* __restrict__ out,
int M, int N, int N_tiles, int M_tiles)
{
constexpr int K = K_STEPS * 256;
constexpr int K_half = K / 2;
constexpr int kg_pad = ((K / 32) + 7) & ~7;
constexpr int bscale_d0_stride = (kg_pad / 8) * 256;
const int tid = threadIdx.x;
// ========== XCD-aware tile remapping ==========
int wgid = blockIdx.x;
const int NUM_WGS = gridDim.x;
if constexpr (K_STEPS <= 2) {
const int NUM_XCDS = 8;
int pids_per_xcd = (NUM_WGS + NUM_XCDS - 1) / NUM_XCDS;
int tall_xcds = NUM_WGS % NUM_XCDS;
if (tall_xcds == 0) tall_xcds = NUM_XCDS;
int xcd = wgid % NUM_XCDS;
int local_pid = wgid / NUM_XCDS;
if (xcd < tall_xcds) {
wgid = xcd * pids_per_xcd + local_pid;
} else {
wgid = tall_xcds * pids_per_xcd + (xcd - tall_xcds) * (pids_per_xcd - 1) + local_pid;
}
if (wgid >= NUM_WGS) return;
}
// Large-K: no XCD remap — preserve natural CU→XCD L2 locality
// L2-aware 2D tile grouping
// Group tiles into GROUP_W × GROUP_H super-tiles for B-data L2 reuse
// Within each super-tile, iterate N-first so consecutive blocks share B columns
int n_tile, m_tile;
if constexpr (K_STEPS <= 2) {
// N-inner: M-adjacent blocks share B in L2 — good for small-K
n_tile = wgid % N_tiles;
m_tile = wgid / N_tiles;
} else {
// 2D swizzled grouping for large-K shapes
constexpr int GROUP_W = 8; // N-tiles per super-tile column (wider for more L2 reuse)
const int tiles_per_group = GROUP_W * M_tiles;
const int group_id = wgid / tiles_per_group;
const int local_id = wgid % tiles_per_group;
// Within group: N varies fastest (local_id % GROUP_W), then M
const int local_n = local_id % GROUP_W;
const int local_m = local_id / GROUP_W;
n_tile = group_id * GROUP_W + local_n;
m_tile = local_m;
// Clamp n_tile for edge groups
if (n_tile >= N_tiles) {
n_tile = N_tiles - 1;
}
}
const int n_base = n_tile * 64; // 64-column tile
const int m_base = m_tile * 16;
const int a_m_row = m_base + (tid & 15);
const int a_k_block = tid >> 4;
const bool valid_a = (a_m_row < M);
const i32x4 b_srd = make_buffer_srd(B_shuffle, (uint32_t)((int64_t)N * K_half));
const int b_col = tid & 15, b_k_block = tid >> 4;
const int b_voffset = b_k_block * 256 + b_col * 16;
// 4 N-sub-tiles: n_base+0, n_base+16, n_base+32, n_base+48
const int b_soff_base_0 = n_base * K_half;
const int b_soff_base_1 = (n_base + 16) * K_half;
const int b_soff_base_2 = (n_base + 32) * K_half;
const int b_soff_base_3 = (n_base + 48) * K_half;
// Paired B_scale precomputation (exploiting n_base 64-alignment)
// Sub-tiles 0,1 share sd0=D; sub-tiles 2,3 share sd0=D+1
// Within each D-group, scale bytes for both K-halves × both sd1 values
// are contiguous: offset = step*256 + b_k_block*64 + b_col*4 + {0,1,2,3}
const int bs_D = n_base / 32;
const int bs_base_01 = bs_D * bscale_d0_stride + b_col * 4;
const int bs_base_23 = (bs_D + 1) * bscale_d0_stride + b_col * 4;
// ========== MAIN LOOP: Fused quant + MFMA ==========
__shared__ uint8_t B_lds[16384]; // Double-buffered: 2 × 8192
v4f32 acc0={0,0,0,0}, acc1={0,0,0,0}, acc2={0,0,0,0}, acc3={0,0,0,0}; __builtin_amdgcn_sched_barrier(0);
// Use float4 union for vectorized A loads (4×128-bit instead of 16×32-bit)
union { bf16x2_t bf[16]; float4 f4[4]; } a0_u, a1_u;
#define a0_bf16 a0_u.bf
#define a1_bf16 a1_u.bf
// Initial loads: A data via float4 + B to LDS for k_step=0
{
const float4* a0_src = (const float4*)(A + (size_t)a_m_row * K + a_k_block * 32);
const float4* a1_src = (const float4*)(A + (size_t)a_m_row * K + 128 + a_k_block * 32);
#pragma unroll
for (int i = 0; i < 4; i++) a0_u.f4[i] = a0_src[i];
#pragma unroll
for (int i = 0; i < 4; i++) a1_u.f4[i] = a1_src[i];
}
llvm_amdgcn_raw_buffer_load_lds(b_srd, SPTR(&B_lds[0]), 16, b_voffset, b_soff_base_0, 0, 0);
llvm_amdgcn_raw_buffer_load_lds(b_srd, SPTR(&B_lds[1024]), 16, b_voffset, b_soff_base_1, 0, 0);
llvm_amdgcn_raw_buffer_load_lds(b_srd, SPTR(&B_lds[2048]), 16, b_voffset, b_soff_base_2, 0, 0);
llvm_amdgcn_raw_buffer_load_lds(b_srd, SPTR(&B_lds[3072]), 16, b_voffset, b_soff_base_3, 0, 0);
llvm_amdgcn_raw_buffer_load_lds(b_srd, SPTR(&B_lds[4096]), 16, b_voffset, b_soff_base_0 + 1024, 0, 0);
llvm_amdgcn_raw_buffer_load_lds(b_srd, SPTR(&B_lds[5120]), 16, b_voffset, b_soff_base_1 + 1024, 0, 0);
llvm_amdgcn_raw_buffer_load_lds(b_srd, SPTR(&B_lds[6144]), 16, b_voffset, b_soff_base_2 + 1024, 0, 0);
llvm_amdgcn_raw_buffer_load_lds(b_srd, SPTR(&B_lds[7168]), 16, b_voffset, b_soff_base_3 + 1024, 0, 0);
// Prefetch B_scale for step 0 (issued before vmcnt so it starts early)
uint32_t packed_01 = *reinterpret_cast<const uint32_t*>(B_scale_sh + bs_base_01 + b_k_block * 64);
uint32_t packed_23 = *reinterpret_cast<const uint32_t*>(B_scale_sh + bs_base_23 + b_k_block * 64);
asm volatile("s_waitcnt vmcnt(0)" ::: "memory");
#pragma unroll
for (int step = 0; step < K_STEPS; step++) {
const int k_step = step * 256;
const int cur_base = (step & 1) ? 8192 : 0;
const int nxt_base = 8192 - cur_base;
// ---- PHASE 1: Issue 8 buffer_load_lds for next K-step (dead on last iter) ----
if (step + 1 < K_STEPS) {
const int nk = (k_step + 256) * 8;
llvm_amdgcn_raw_buffer_load_lds(b_srd, SPTR(&B_lds[nxt_base]), 16, b_voffset, b_soff_base_0 + nk, 0, 0);
llvm_amdgcn_raw_buffer_load_lds(b_srd, SPTR(&B_lds[nxt_base+1024]), 16, b_voffset, b_soff_base_1 + nk, 0, 0);
llvm_amdgcn_raw_buffer_load_lds(b_srd, SPTR(&B_lds[nxt_base+2048]), 16, b_voffset, b_soff_base_2 + nk, 0, 0);
llvm_amdgcn_raw_buffer_load_lds(b_srd, SPTR(&B_lds[nxt_base+3072]), 16, b_voffset, b_soff_base_3 + nk, 0, 0);
llvm_amdgcn_raw_buffer_load_lds(b_srd, SPTR(&B_lds[nxt_base+4096]), 16, b_voffset, b_soff_base_0 + nk + 1024, 0, 0);
llvm_amdgcn_raw_buffer_load_lds(b_srd, SPTR(&B_lds[nxt_base+5120]), 16, b_voffset, b_soff_base_1 + nk + 1024, 0, 0);
llvm_amdgcn_raw_buffer_load_lds(b_srd, SPTR(&B_lds[nxt_base+6144]), 16, b_voffset, b_soff_base_2 + nk + 1024, 0, 0);
llvm_amdgcn_raw_buffer_load_lds(b_srd, SPTR(&B_lds[nxt_base+7168]), 16, b_voffset, b_soff_base_3 + nk + 1024, 0, 0);
}
// ---- PHASE 2: Quant A half-0 (B_scale already in packed_01/packed_23 from prolog or previous step) ----
uint32_t ap[4]; int as0;
quant_32bf16_to_fp4(a0_bf16, true, ap, as0);
v8i32 a_reg = {}; a_reg[0]=ap[0]; a_reg[1]=ap[1]; a_reg[2]=ap[2]; a_reg[3]=ap[3];
// Extract B_scale for first K-half: byte layout [s4=0,sd1=0 | s4=0,sd1=1 | s4=1,sd1=0 | s4=1,sd1=1]
int bs0_h0 = packed_01 & 0xFF; // sub-tile 0
int bs1_h0 = (packed_01 >> 8) & 0xFF; // sub-tile 1
int bs2_h0 = packed_23 & 0xFF; // sub-tile 2
int bs3_h0 = (packed_23 >> 8) & 0xFF; // sub-tile 3
// ---- PHASE 3: 4 MFMAs for first K-half ----
{ const uint32_t* bl=(const uint32_t*)(&B_lds[cur_base+tid*16]); v8i32 br={};
br[0]=bl[0]; br[1]=bl[1]; br[2]=bl[2]; br[3]=bl[3];
acc0 = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(a_reg, br, acc0, 4, 4, 0, as0, 0, bs0_h0); }
{ const uint32_t* bl=(const uint32_t*)(&B_lds[cur_base+1024+tid*16]); v8i32 br={};
br[0]=bl[0]; br[1]=bl[1]; br[2]=bl[2]; br[3]=bl[3];
acc1 = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(a_reg, br, acc1, 4, 4, 0, as0, 0, bs1_h0); }
{ const uint32_t* bl=(const uint32_t*)(&B_lds[cur_base+2048+tid*16]); v8i32 br={};
br[0]=bl[0]; br[1]=bl[1]; br[2]=bl[2]; br[3]=bl[3];
acc2 = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(a_reg, br, acc2, 4, 4, 0, as0, 0, bs2_h0); }
{ const uint32_t* bl=(const uint32_t*)(&B_lds[cur_base+3072+tid*16]); v8i32 br={};
br[0]=bl[0]; br[1]=bl[1]; br[2]=bl[2]; br[3]=bl[3];
acc3 = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(a_reg, br, acc3, 4, 4, 0, as0, 0, bs3_h0); }
// ---- PHASE 3.5: Load A half-0 for NEXT k_step via float4 (dead on last iter) ----
if (step + 1 < K_STEPS) {
const float4* an0 = (const float4*)(A + (size_t)a_m_row * K + k_step + 256 + a_k_block * 32);
#pragma unroll
for (int i = 0; i < 4; i++) a0_u.f4[i] = an0[i];
}
// ---- PHASE 4: Quant A half-1 + 4 MFMAs for second K-half ----
int as1;
quant_32bf16_to_fp4(a1_bf16, true, ap, as1);
a_reg[0]=ap[0]; a_reg[1]=ap[1]; a_reg[2]=ap[2]; a_reg[3]=ap[3];
// Extract B_scale for second K-half
int bs0_h1 = (packed_01 >> 16) & 0xFF; // sub-tile 0
int bs1_h1 = (packed_01 >> 24) & 0xFF; // sub-tile 1
int bs2_h1 = (packed_23 >> 16) & 0xFF; // sub-tile 2
int bs3_h1 = (packed_23 >> 24) & 0xFF; // sub-tile 3
{ const uint32_t* bl=(const uint32_t*)(&B_lds[cur_base+4096+tid*16]); v8i32 br={};
br[0]=bl[0]; br[1]=bl[1]; br[2]=bl[2]; br[3]=bl[3];
acc0 = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(a_reg, br, acc0, 4, 4, 0, as1, 0, bs0_h1); }
{ const uint32_t* bl=(const uint32_t*)(&B_lds[cur_base+5120+tid*16]); v8i32 br={};
br[0]=bl[0]; br[1]=bl[1]; br[2]=bl[2]; br[3]=bl[3];
acc1 = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(a_reg, br, acc1, 4, 4, 0, as1, 0, bs1_h1); }
{ const uint32_t* bl=(const uint32_t*)(&B_lds[cur_base+6144+tid*16]); v8i32 br={};
br[0]=bl[0]; br[1]=bl[1]; br[2]=bl[2]; br[3]=bl[3];
acc2 = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(a_reg, br, acc2, 4, 4, 0, as1, 0, bs2_h1); }
{ const uint32_t* bl=(const uint32_t*)(&B_lds[cur_base+7168+tid*16]); v8i32 br={};
br[0]=bl[0]; br[1]=bl[1]; br[2]=bl[2]; br[3]=bl[3];
acc3 = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(a_reg, br, acc3, 4, 4, 0, as1, 0, bs3_h1); }
// ---- PHASE 4.5: Load A half-1 for NEXT k_step via float4 (dead on last iter) ----
if (step + 1 < K_STEPS) {
const float4* an1 = (const float4*)(A + (size_t)a_m_row * K + k_step + 256 + 128 + a_k_block * 32);
#pragma unroll
for (int i = 0; i < 4; i++) a1_u.f4[i] = an1[i];
}
// ---- PHASE 5: Wait for B loads + prefetch B_scale for next step ----
if (step + 1 < K_STEPS) {
// Prefetch B_scale for next step (overlaps with vmcnt wait)
const int bs_off_next = (step + 1) * 256 + b_k_block * 64;
packed_01 = *reinterpret_cast<const uint32_t*>(B_scale_sh + bs_base_01 + bs_off_next);
packed_23 = *reinterpret_cast<const uint32_t*>(B_scale_sh + bs_base_23 + bs_off_next);
asm volatile("s_waitcnt vmcnt(8)" ::: "memory");
}
__builtin_amdgcn_sched_barrier(0);
}
#undef a0_bf16
#undef a1_bf16
// ========== EPILOG: Write outputs for 4 N-sub-tiles ==========
const int out_row_base = m_base + (tid >> 4) * 4;
#pragma unroll
for (int sub = 0; sub < 4; sub++) {
const int out_col = n_base + sub * 16 + (tid & 15);
v4f32& acc = (sub == 0) ? acc0 : (sub == 1) ? acc1 : (sub == 2) ? acc2 : acc3;
if (out_col < N) {
#pragma unroll
for (int v = 0; v < 4; v++) {
int r = out_row_base + v;
if (r < M) out[(int64_t)r * N + out_col] = __float2bfloat16(acc[v]);
}
}
}
}
// Explicit template instantiations
template __global__ void fused_quant_mfma_asm_sched<8>(
const __hip_bfloat16*, const uint8_t*, const uint8_t*, __hip_bfloat16*,
int, int, int, int);
template __global__ void fused_quant_mfma_asm_sched<6>(
const __hip_bfloat16*, const uint8_t*, const uint8_t*, __hip_bfloat16*,
int, int, int, int);
template __global__ void fused_quant_mfma_asm_sched<2>(
const __hip_bfloat16*, const uint8_t*, const uint8_t*, __hip_bfloat16*,
int, int, int, int);
// ============================================================
// 2-Wave Intra-WG Split-K: 128 threads = 2 waves/SIMD guaranteed
// Each wave handles K_STEPS/2 steps. Reduces via 4KB LDS at end.
// ============================================================
template<int K_STEPS_TOTAL>
__global__ __launch_bounds__(128, 1)
void fused_quant_mfma_2wave_sk(
const __hip_bfloat16* __restrict__ A, const uint8_t* __restrict__ B_shuffle,
const uint8_t* __restrict__ B_scale_sh, __hip_bfloat16* __restrict__ out,
int M, int N, int N_tiles, int M_tiles)
{
constexpr int K_STEPS = K_STEPS_TOTAL / 2; // Per-wave steps
constexpr int K = K_STEPS_TOTAL * 256;
constexpr int K_half = K / 2;
constexpr int kg_pad = ((K / 32) + 7) & ~7;
constexpr int bscale_d0_stride = (kg_pad / 8) * 256;
// readfirstlane: wave_id is uniform within a wavefront but compiler
// treats threadIdx.x as VGPR. Without this, all wave_id-derived LDS
// pointers and soffsets stay in VGPRs → waterfall scatter loops.
const int wave_id = __builtin_amdgcn_readfirstlane(threadIdx.x >> 6);
const int tid = threadIdx.x & 63;
// Tile assignment with XCD remap ALL
int wgid = blockIdx.x;
const int NUM_WGS = gridDim.x;
// XCD remap ALL for better L2 locality
{
const int NUM_XCDS = 8;
int pids_per_xcd = (NUM_WGS + NUM_XCDS - 1) / NUM_XCDS;
int tall_xcds = NUM_WGS % NUM_XCDS;
if (tall_xcds == 0) tall_xcds = NUM_XCDS;
int xcd = wgid % NUM_XCDS;
int local_pid = wgid / NUM_XCDS;
if (xcd < tall_xcds) {
wgid = xcd * pids_per_xcd + local_pid;
} else {
wgid = tall_xcds * pids_per_xcd + (xcd - tall_xcds) * (pids_per_xcd - 1) + local_pid;
}
if (wgid >= NUM_WGS) return;
}
// Tiling: flat for short-K, GROUP for large-K
int n_tile, m_tile;
if constexpr (K_STEPS <= 1) {
// Flat N-inner tiling for small K (e.g. K=512)
n_tile = wgid % N_tiles;
m_tile = wgid / N_tiles;
} else {
constexpr int GROUP_W = 8;
const int tiles_per_group = GROUP_W * M_tiles;
const int group_id = wgid / tiles_per_group;
const int local_id = wgid % tiles_per_group;
const int local_n = local_id % GROUP_W;
const int local_m = local_id / GROUP_W;
n_tile = group_id * GROUP_W + local_n;
m_tile = local_m;
if (n_tile >= N_tiles) return;
}
const int n_base = n_tile * 64;
const int m_base = m_tile * 16;
const int a_m_row = m_base + (tid & 15);
const int a_k_block = tid >> 4;
const bool valid_a = (a_m_row < M);
const i32x4 b_srd = make_buffer_srd(B_shuffle, (uint32_t)((int64_t)N * K_half));
const int b_col = tid & 15, b_k_block = tid >> 4;
const int b_voffset = b_k_block * 256 + b_col * 16;
const int b_soff_base_0 = n_base * K_half;
const int b_soff_base_1 = (n_base + 16) * K_half;
const int b_soff_base_2 = (n_base + 32) * K_half;
const int b_soff_base_3 = (n_base + 48) * K_half;
const int bs_D = n_base / 32;
const int bs_base_01 = bs_D * bscale_d0_stride + b_col * 4;
const int bs_base_23 = (bs_D + 1) * bscale_d0_stride + b_col * 4;
// LDS layout: wave0 B[0..16383], wave1 B[16384..32767]
// Reduction reuses wave0's B buffer after main loop (no separate allocation)
__shared__ uint8_t B_lds[32768];
const int lds_b = wave_id * 16384; // Per-wave B double-buffer base
v4f32 acc0={0,0,0,0}, acc1={0,0,0,0}, acc2={0,0,0,0}, acc3={0,0,0,0};
union { bf16x2_t bf[16]; float4 f4[4]; } a0_u, a1_u;
#define a0_bf16 a0_u.bf
#define a1_bf16 a1_u.bf
// K offset for this wave
const int k_base = wave_id * K_STEPS * 256;
const int k_b_off = k_base * 8; // Packed B offset (k_base/2 * 16-byte unit factor)
// Initial A loads
{
const float4* a0_src = (const float4*)(A + (size_t)a_m_row * K + k_base + a_k_block * 32);
const float4* a1_src = (const float4*)(A + (size_t)a_m_row * K + k_base + 128 + a_k_block * 32);
#pragma unroll
for (int i = 0; i < 4; i++) a0_u.f4[i] = a0_src[i];
#pragma unroll
for (int i = 0; i < 4; i++) a1_u.f4[i] = a1_src[i];
}
// Initial B loads to wave-local LDS
llvm_amdgcn_raw_buffer_load_lds(b_srd, SPTR(&B_lds[lds_b]), 16, b_voffset, b_soff_base_0 + k_b_off, 0, 0);
llvm_amdgcn_raw_buffer_load_lds(b_srd, SPTR(&B_lds[lds_b+1024]), 16, b_voffset, b_soff_base_1 + k_b_off, 0, 0);
llvm_amdgcn_raw_buffer_load_lds(b_srd, SPTR(&B_lds[lds_b+2048]), 16, b_voffset, b_soff_base_2 + k_b_off, 0, 0);
llvm_amdgcn_raw_buffer_load_lds(b_srd, SPTR(&B_lds[lds_b+3072]), 16, b_voffset, b_soff_base_3 + k_b_off, 0, 0);
llvm_amdgcn_raw_buffer_load_lds(b_srd, SPTR(&B_lds[lds_b+4096]), 16, b_voffset, b_soff_base_0 + k_b_off + 1024, 0, 0);
llvm_amdgcn_raw_buffer_load_lds(b_srd, SPTR(&B_lds[lds_b+5120]), 16, b_voffset, b_soff_base_1 + k_b_off + 1024, 0, 0);
llvm_amdgcn_raw_buffer_load_lds(b_srd, SPTR(&B_lds[lds_b+6144]), 16, b_voffset, b_soff_base_2 + k_b_off + 1024, 0, 0);
llvm_amdgcn_raw_buffer_load_lds(b_srd, SPTR(&B_lds[lds_b+7168]), 16, b_voffset, b_soff_base_3 + k_b_off + 1024, 0, 0);
// Prefetch B_scale for first step of this wave's K range
// B_scale offset = absolute_step * 256 + b_k_block * 64
// absolute_step = wave_id * K_STEPS
const int bs_wave_base = wave_id * K_STEPS * 256; // B_scale K offset for this wave
uint32_t packed_01 = *reinterpret_cast<const uint32_t*>(B_scale_sh + bs_base_01 + bs_wave_base + b_k_block * 64);
uint32_t packed_23 = *reinterpret_cast<const uint32_t*>(B_scale_sh + bs_base_23 + bs_wave_base + b_k_block * 64);
asm volatile("s_waitcnt vmcnt(0)" ::: "memory");
// Main loop: K_STEPS iterations (half the full K)
#pragma unroll
for (int step = 0; step < K_STEPS; step++) {
const int k_step = k_base + step * 256;
const int cur_base = lds_b + ((step & 1) ? 8192 : 0);
const int nxt_base = lds_b + 8192 - ((step & 1) ? 8192 : 0);
// Issue B loads for next step
if (step + 1 < K_STEPS) {
const int nk = (k_step + 256) * 8;
llvm_amdgcn_raw_buffer_load_lds(b_srd, SPTR(&B_lds[nxt_base]), 16, b_voffset, b_soff_base_0 + nk, 0, 0);
llvm_amdgcn_raw_buffer_load_lds(b_srd, SPTR(&B_lds[nxt_base+1024]), 16, b_voffset, b_soff_base_1 + nk, 0, 0);
llvm_amdgcn_raw_buffer_load_lds(b_srd, SPTR(&B_lds[nxt_base+2048]), 16, b_voffset, b_soff_base_2 + nk, 0, 0);
llvm_amdgcn_raw_buffer_load_lds(b_srd, SPTR(&B_lds[nxt_base+3072]), 16, b_voffset, b_soff_base_3 + nk, 0, 0);
llvm_amdgcn_raw_buffer_load_lds(b_srd, SPTR(&B_lds[nxt_base+4096]), 16, b_voffset, b_soff_base_0 + nk + 1024, 0, 0);
llvm_amdgcn_raw_buffer_load_lds(b_srd, SPTR(&B_lds[nxt_base+5120]), 16, b_voffset, b_soff_base_1 + nk + 1024, 0, 0);
llvm_amdgcn_raw_buffer_load_lds(b_srd, SPTR(&B_lds[nxt_base+6144]), 16, b_voffset, b_soff_base_2 + nk + 1024, 0, 0);
llvm_amdgcn_raw_buffer_load_lds(b_srd, SPTR(&B_lds[nxt_base+7168]), 16, b_voffset, b_soff_base_3 + nk + 1024, 0, 0);
}
// Quant A half-0
uint32_t ap[4]; int as0;
quant_32bf16_to_fp4(a0_bf16, true, ap, as0);
v8i32 a_reg = {}; a_reg[0]=ap[0]; a_reg[1]=ap[1]; a_reg[2]=ap[2]; a_reg[3]=ap[3];
int bs0_h0 = packed_01 & 0xFF;
int bs1_h0 = (packed_01 >> 8) & 0xFF;
int bs2_h0 = packed_23 & 0xFF;
int bs3_h0 = (packed_23 >> 8) & 0xFF;
// 4 MFMAs for K-half 0 — no priority boost, let compiler schedule freely
{ const uint32_t* bl=(const uint32_t*)(&B_lds[cur_base+tid*16]); v8i32 br={};
br[0]=bl[0]; br[1]=bl[1]; br[2]=bl[2]; br[3]=bl[3];
acc0 = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(a_reg, br, acc0, 4, 4, 0, as0, 0, bs0_h0); }
{ const uint32_t* bl=(const uint32_t*)(&B_lds[cur_base+1024+tid*16]); v8i32 br={};
br[0]=bl[0]; br[1]=bl[1]; br[2]=bl[2]; br[3]=bl[3];
acc1 = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(a_reg, br, acc1, 4, 4, 0, as0, 0, bs1_h0); }
{ const uint32_t* bl=(const uint32_t*)(&B_lds[cur_base+2048+tid*16]); v8i32 br={};
br[0]=bl[0]; br[1]=bl[1]; br[2]=bl[2]; br[3]=bl[3];
acc2 = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(a_reg, br, acc2, 4, 4, 0, as0, 0, bs2_h0); }
{ const uint32_t* bl=(const uint32_t*)(&B_lds[cur_base+3072+tid*16]); v8i32 br={};
br[0]=bl[0]; br[1]=bl[1]; br[2]=bl[2]; br[3]=bl[3];
acc3 = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(a_reg, br, acc3, 4, 4, 0, as0, 0, bs3_h0); }
// Load A half-0 for next step
if (step + 1 < K_STEPS) {
const float4* an0 = (const float4*)(A + (size_t)a_m_row * K + k_step + 256 + a_k_block * 32);
#pragma unroll
for (int i = 0; i < 4; i++) a0_u.f4[i] = an0[i];
}
// Quant A half-1 + 4 MFMAs
int as1;
quant_32bf16_to_fp4(a1_bf16, true, ap, as1);
a_reg[0]=ap[0]; a_reg[1]=ap[1]; a_reg[2]=ap[2]; a_reg[3]=ap[3];
int bs0_h1 = (packed_01 >> 16) & 0xFF;
int bs1_h1 = (packed_01 >> 24) & 0xFF;
int bs2_h1 = (packed_23 >> 16) & 0xFF;
int bs3_h1 = (packed_23 >> 24) & 0xFF;
{ const uint32_t* bl=(const uint32_t*)(&B_lds[cur_base+4096+tid*16]); v8i32 br={};
br[0]=bl[0]; br[1]=bl[1]; br[2]=bl[2]; br[3]=bl[3];
acc0 = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(a_reg, br, acc0, 4, 4, 0, as1, 0, bs0_h1); }
{ const uint32_t* bl=(const uint32_t*)(&B_lds[cur_base+5120+tid*16]); v8i32 br={};
br[0]=bl[0]; br[1]=bl[1]; br[2]=bl[2]; br[3]=bl[3];
acc1 = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(a_reg, br, acc1, 4, 4, 0, as1, 0, bs1_h1); }
{ const uint32_t* bl=(const uint32_t*)(&B_lds[cur_base+6144+tid*16]); v8i32 br={};
br[0]=bl[0]; br[1]=bl[1]; br[2]=bl[2]; br[3]=bl[3];
acc2 = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(a_reg, br, acc2, 4, 4, 0, as1, 0, bs2_h1); }
{ const uint32_t* bl=(const uint32_t*)(&B_lds[cur_base+7168+tid*16]); v8i32 br={};
br[0]=bl[0]; br[1]=bl[1]; br[2]=bl[2]; br[3]=bl[3];
acc3 = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(a_reg, br, acc3, 4, 4, 0, as1, 0, bs3_h1); }
// Load A half-1 for next step
if (step + 1 < K_STEPS) {
const float4* an1 = (const float4*)(A + (size_t)a_m_row * K + k_step + 256 + 128 + a_k_block * 32);
#pragma unroll
for (int i = 0; i < 4; i++) a1_u.f4[i] = an1[i];
}
// Wait for B loads + prefetch B_scale
if (step + 1 < K_STEPS) {
const int bs_off_next = bs_wave_base + (step + 1) * 256 + b_k_block * 64;
packed_01 = *reinterpret_cast<const uint32_t*>(B_scale_sh + bs_base_01 + bs_off_next);
packed_23 = *reinterpret_cast<const uint32_t*>(B_scale_sh + bs_base_23 + bs_off_next);
asm volatile("s_waitcnt vmcnt(8)" ::: "memory");
}
}
#undef a0_bf16
#undef a1_bf16
// ===== Reduction: wave 1 writes partials to LDS, wave 0 adds =====
__syncthreads();
if (wave_id == 1) {
float* lds_red = (float*)&B_lds[0]; // Reuse wave0's B buffer space
int base = tid * 16;
lds_red[base+ 0]=acc0[0]; lds_red[base+ 1]=acc0[1]; lds_red[base+ 2]=acc0[2]; lds_red[base+ 3]=acc0[3];
lds_red[base+ 4]=acc1[0]; lds_red[base+ 5]=acc1[1]; lds_red[base+ 6]=acc1[2]; lds_red[base+ 7]=acc1[3];
lds_red[base+ 8]=acc2[0]; lds_red[base+ 9]=acc2[1]; lds_red[base+10]=acc2[2]; lds_red[base+11]=acc2[3];
lds_red[base+12]=acc3[0]; lds_red[base+13]=acc3[1]; lds_red[base+14]=acc3[2]; lds_red[base+15]=acc3[3];
}
__syncthreads();
if (wave_id == 0) {
const float* lds_red = (const float*)&B_lds[0]; // Reuse wave0's B buffer space
int base = tid * 16;
acc0[0]+=lds_red[base+ 0]; acc0[1]+=lds_red[base+ 1]; acc0[2]+=lds_red[base+ 2]; acc0[3]+=lds_red[base+ 3];
acc1[0]+=lds_red[base+ 4]; acc1[1]+=lds_red[base+ 5]; acc1[2]+=lds_red[base+ 6]; acc1[3]+=lds_red[base+ 7];
acc2[0]+=lds_red[base+ 8]; acc2[1]+=lds_red[base+ 9]; acc2[2]+=lds_red[base+10]; acc2[3]+=lds_red[base+11];
acc3[0]+=lds_red[base+12]; acc3[1]+=lds_red[base+13]; acc3[2]+=lds_red[base+14]; acc3[3]+=lds_red[base+15];
// Write output
const int out_row_base = m_base + (tid >> 4) * 4;
#pragma unroll
for (int sub = 0; sub < 4; sub++) {
const int out_col = n_base + sub * 16 + (tid & 15);
v4f32& acc = (sub == 0) ? acc0 : (sub == 1) ? acc1 : (sub == 2) ? acc2 : acc3;
if (out_col < N) {
#pragma unroll
for (int v = 0; v < 4; v++) {
int r = out_row_base + v;
if (r < M) out[(int64_t)r * N + out_col] = __float2bfloat16(acc[v]);
}
}
}
}
}
template __global__ void fused_quant_mfma_2wave_sk<8>(
const __hip_bfloat16*, const uint8_t*, const uint8_t*, __hip_bfloat16*,
int, int, int, int);
template __global__ void fused_quant_mfma_2wave_sk<6>(
const __hip_bfloat16*, const uint8_t*, const uint8_t*, __hip_bfloat16*,
int, int, int, int);
template __global__ void fused_quant_mfma_2wave_sk<2>(
const __hip_bfloat16*, const uint8_t*, const uint8_t*, __hip_bfloat16*,
int, int, int, int);
// ============================================================
// Dispatch infrastructure
// ============================================================
struct DoubleCfg {
torch::Tensor out;
int M, N, K, N_tiles, M_tiles, grid_size;
bool wide;
};
static std::unordered_map<uint64_t, DoubleCfg> g_double;
struct FusedAsmCfg {
torch::Tensor out;
int M, N, K, N_tiles, M_tiles, grid_size;
int K_STEPS;
};
static std::unordered_map<uint64_t, FusedAsmCfg> g_fused_asm;
struct TriCfg {
torch::Tensor workspace, out;
int M, N, K;
};
static std::unordered_map<uint64_t, TriCfg> g_tri;
static uint64_t shape_key(int M, int N, int K) {
return ((uint64_t)M << 40) | ((uint64_t)N << 20) | (uint64_t)K;
}
void register_double_shape(int64_t M, int64_t N, int64_t K, bool wide) {
uint64_t key = shape_key(M, N, K);
if (g_double.count(key)) return;
DoubleCfg c;
c.M=M; c.N=N; c.K=K; c.wide=wide;
c.N_tiles = wide ? (N+31)/32 : (N+15)/16;
c.M_tiles = (M+15)/16;
c.grid_size = c.N_tiles * c.M_tiles;
c.out = torch::empty({M, N}, torch::TensorOptions().dtype(torch::kBFloat16).device(torch::kCUDA));
g_double[key] = std::move(c);
}
void register_fused_asm_shape(int64_t M, int64_t N, int64_t K) {
uint64_t key = shape_key(M, N, K);
if (g_fused_asm.count(key)) return;
FusedAsmCfg c;
c.M=M; c.N=N; c.K=K;
c.N_tiles = (N+63)/64;
c.M_tiles = (M+15)/16;
c.grid_size = c.N_tiles * c.M_tiles;
c.K_STEPS = K / 256;
c.out = torch::empty({M, N}, torch::TensorOptions().dtype(torch::kBFloat16).device(torch::kCUDA));
g_fused_asm[key] = std::move(c);
fprintf(stderr, "[v936] Fused 16x64 registered: M=%lld N=%lld K=%lld K_STEPS=%d grid=%d\n",
(long long)M, (long long)N, (long long)K, c.K_STEPS, c.grid_size);
}
void register_tri_shape(int64_t M, int64_t N, int64_t K, int64_t split_k) {
uint64_t key = shape_key(M, N, K);
if (g_tri.count(key)) return;
TriCfg c; c.M=M; c.N=N; c.K=K;
c.workspace = torch::zeros({split_k, M, N}, torch::TensorOptions().dtype(torch::kFloat32).device(torch::kCUDA));
c.out = torch::empty({M, N}, torch::TensorOptions().dtype(torch::kBFloat16).device(torch::kCUDA));
g_tri[key] = std::move(c);
}
torch::Tensor dispatch_double(torch::Tensor A, torch::Tensor B_shuffle, torch::Tensor B_scale_sh) {
int M=A.size(0), K=A.size(1), N=B_shuffle.size(0);
auto it = g_double.find(shape_key(M, N, K));
if (it == g_double.end()) return A;
DoubleCfg& c = it->second;
if (c.wide) {
fused_quant_mfma_wide<<<c.grid_size, 64, 0, 0>>>(
(const __hip_bfloat16*)A.data_ptr(), (const uint8_t*)B_shuffle.data_ptr(),
(const uint8_t*)B_scale_sh.data_ptr(), (__hip_bfloat16*)c.out.data_ptr(),
c.M, c.N, c.K, c.N_tiles, c.M_tiles);
} else {
fused_quant_mfma_narrow<<<c.grid_size, 64, 0, 0>>>(
(const __hip_bfloat16*)A.data_ptr(), (const uint8_t*)B_shuffle.data_ptr(),
(const uint8_t*)B_scale_sh.data_ptr(), (__hip_bfloat16*)c.out.data_ptr(),
c.M, c.N, c.K, c.N_tiles, c.M_tiles);
}
return c.out;
}
torch::Tensor dispatch_fused_asm(torch::Tensor A, torch::Tensor B_shuffle, torch::Tensor B_scale_sh) {
int M=A.size(0), K=A.size(1), N=B_shuffle.size(0);
auto it = g_fused_asm.find(shape_key(M, N, K));
if (it == g_fused_asm.end()) return A;
FusedAsmCfg& c = it->second;
const auto* a_ptr = (const __hip_bfloat16*)A.data_ptr();
const auto* b_ptr = (const uint8_t*)B_shuffle.data_ptr();
const auto* bs_ptr = (const uint8_t*)B_scale_sh.data_ptr();
auto* o_ptr = (__hip_bfloat16*)c.out.data_ptr();
if (c.K_STEPS == 8) {
fused_quant_mfma_2wave_sk<8><<<c.grid_size, 128, 0, 0>>>(
a_ptr, b_ptr, bs_ptr, o_ptr, c.M, c.N, c.N_tiles, c.M_tiles);
} else if (c.K_STEPS == 6) {
fused_quant_mfma_2wave_sk<6><<<c.grid_size, 128, 0, 0>>>(
a_ptr, b_ptr, bs_ptr, o_ptr, c.M, c.N, c.N_tiles, c.M_tiles);
} else if (c.K_STEPS == 2) {
fused_quant_mfma_2wave_sk<2><<<c.grid_size, 128, 0, 0>>>(
a_ptr, b_ptr, bs_ptr, o_ptr, c.M, c.N, c.N_tiles, c.M_tiles);
}
return c.out;
}
torch::Tensor dispatch_all(torch::Tensor A, torch::Tensor B_shuffle, torch::Tensor B_scale_sh) {
int M = A.size(0), K = A.size(1), N = B_shuffle.size(0);
uint64_t key = shape_key(M, N, K);
// Try fused ASM-scheduled first (S5/S6)
auto it_fa = g_fused_asm.find(key);
if (it_fa != g_fused_asm.end()) {
return dispatch_fused_asm(A, B_shuffle, B_scale_sh);
}
// Try MFMA (narrow/wide)
auto it_dbl = g_double.find(key);
if (it_dbl != g_double.end()) {
return dispatch_double(A, B_shuffle, B_scale_sh);
}
// Triton needed — return empty tensor as sentinel
return torch::Tensor();
}
"""
CPP_SRC = """
void register_double_shape(int64_t M, int64_t N, int64_t K, bool wide);
void register_fused_asm_shape(int64_t M, int64_t N, int64_t K);
void register_tri_shape(int64_t M, int64_t N, int64_t K, int64_t split_k);
torch::Tensor dispatch_double(torch::Tensor A, torch::Tensor B_q, torch::Tensor B_scale_sh);
torch::Tensor dispatch_fused_asm(torch::Tensor A, torch::Tensor B_shuffle, torch::Tensor B_scale_sh);
torch::Tensor dispatch_all(torch::Tensor A, torch::Tensor B_shuffle, torch::Tensor B_scale_sh);
"""
print("[v936] Compiling wide-group-tight kernels...")
module = load_inline(
name='fused_v205', cpp_sources=[CPP_SRC], cuda_sources=[HIP_SRC],
functions=['register_double_shape', 'register_fused_asm_shape', 'register_tri_shape',
'dispatch_double', 'dispatch_fused_asm', 'dispatch_all'],
extra_cuda_cflags=["--offload-arch=gfx950", "-O3", "-std=c++20",
"-ffast-math",
"-mllvm", "-amdgpu-early-inline-all",
"-mllvm", "-amdgpu-function-calls=0",
],
)
for m, n, k in SHAPES:
if (m, n, k) in FUSED_ASM_SHAPES:
module.register_fused_asm_shape(m, n, k)
print(f" [{m}x{n}x{k}] -> FUSED_ASM_SCHED (16x64)")
elif (m, n, k) in TWOWAVE_SHAPES:
module.register_fused_asm_shape(m, n, k)
print(f" [{m}x{n}x{k}] -> 2WAVE_SK (16x64, K_STEPS=2)")
elif (m, n, k) in NARROW_MFMA_SHAPES:
module.register_double_shape(m, n, k, False)
print(f" [{m}x{n}x{k}] -> NARROW (16x16)")
elif (m, n, k) in WIDE_MFMA_SHAPES:
module.register_double_shape(m, n, k, True)
print(f" [{m}x{n}x{k}] -> WIDE (16x32)")
elif (m, n, k) in TRITON_SHAPES:
module.register_tri_shape(m, n, k, NUM_KSPLIT)
print(f" [{m}x{n}x{k}] -> TRITON (split_k={NUM_KSPLIT})")
else:
module.register_double_shape(m, n, k, True)
print(f" [{m}x{n}x{k}] -> WIDE (default)")
_dispatch_all = module.dispatch_all
# Triton warmup
_tri_cfg = {}
print("[v936] Warming Triton...")
for m, n, k in TRITON_SHAPES:
_dummy_A = torch.randn(m, k, dtype=torch.bfloat16, device='cuda')
_dummy_Bq = torch.randint(0, 255, (n, k//2), dtype=torch.uint8, device='cuda')
_kg = k // 32; _kg_pad = (_kg + 7) & ~7
_dummy_Bs = torch.randint(0, 255, (((n+31)//32)*32, _kg_pad), dtype=torch.uint8, device='cuda')
workspace = torch.empty((NUM_KSPLIT, m, n), dtype=torch.float32, device='cuda')
out = torch.empty((m, n), dtype=torch.bfloat16, device='cuda')
bsd0_stride = (_kg_pad // 8) * 256
_m, _n = m, n
grid_fn = lambda META, _m=_m, _n=_n: (
((_m+BLOCK_M-1)//BLOCK_M) * ((_n+META['BLOCK_N']-1)//META['BLOCK_N']) * NUM_KSPLIT,)
_triton_fused_quant_gemm_splitk[grid_fn](
_dummy_A, _dummy_Bq, workspace, _dummy_Bs, m, n, k,
_dummy_A.stride(0), _dummy_A.stride(1), _dummy_Bq.stride(1), _dummy_Bq.stride(0),
workspace.stride(0), workspace.stride(1), workspace.stride(2),
bsd0_stride, BLOCK_M=BLOCK_M, BLOCK_K=BLOCK_K, NUM_KSPLIT=NUM_KSPLIT, SPLITK_SIZE=SPLITK_BLOCK)
reduce_grid = (m, (n+REDUCE_BN-1)//REDUCE_BN)
_splitk_reduce[reduce_grid](workspace, out, m, n,
workspace.stride(0), workspace.stride(1), workspace.stride(2),
out.stride(0), out.stride(1), NUM_KSPLIT=NUM_KSPLIT, BLOCK_RN=REDUCE_BN)
_tri_cfg[(m,n,k)] = {
'workspace': workspace, 'out': out, 'grid_fn': grid_fn, 'reduce_grid': reduce_grid,
'ws_s0': workspace.stride(0), 'ws_s1': workspace.stride(1), 'ws_s2': workspace.stride(2),
'out_s0': out.stride(0), 'out_s1': out.stride(1),
'bq_s0': k//2, 'bq_s1': 1, 'bsd0_stride': bsd0_stride}
print(f" [{m}x{n}x{k}] warmup OK")
del _dummy_A, _dummy_Bq, _dummy_Bs
torch.cuda.empty_cache()
print("[v936] Setup complete")
def custom_kernel(data: input_t) -> output_t:
A = data[0]
# Fast path: single C++ call handles fused ASM + MFMA shapes
result = _dispatch_all(A, data[3], data[4])
if result is not None and result.numel() > 0:
return result
# Triton fallback for split-K shapes (S2/S7)
B_q = data[2].view(torch.uint8); B_scale_sh = data[4].view(torch.uint8)
m = A.shape[0]; k = A.shape[1]; n = B_q.shape[0]
c = _tri_cfg.get((m, n, k))
if c is not None:
_triton_fused_quant_gemm_splitk[c['grid_fn']](
A, B_q, c['workspace'], B_scale_sh,
m, n, k, A.stride(0), A.stride(1), c['bq_s1'], c['bq_s0'],
c['ws_s0'], c['ws_s1'], c['ws_s2'], c['bsd0_stride'],
BLOCK_M=BLOCK_M, BLOCK_K=BLOCK_K, NUM_KSPLIT=NUM_KSPLIT, SPLITK_SIZE=SPLITK_BLOCK)
_splitk_reduce[c['reduce_grid']](
c['workspace'], c['out'], m, n,
c['ws_s0'], c['ws_s1'], c['ws_s2'], c['out_s0'], c['out_s1'],
NUM_KSPLIT=NUM_KSPLIT, BLOCK_RN=REDUCE_BN)
return c['out']
# Fallback
return module.dispatch_double(A, data[3], data[4])scrolls · 1127 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