submission 754951
dc1312 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 683 lines, June 9 Researcher Reciprocity License v1.0.
submission_v160.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-754951?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:dc02f98a78e809d952f31eb83bd58739c23ff3f9aabfbcf6fd7665036f4ec254
license declaredunknown
license concludedunknown
authorsdc1312
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
autotune
Strategy: HIP assembly kernel for small-K shapes, autotuned Triton for large-K.fp4
MXFP4 GEMM: hybrid HIP + Triton approach for mixed-precision FP4 matmul.num-warps = 4
num_warps=4, num_stages=2,shared-memory
__shared__ __align__(16) uint8_t Alds[2 * BLOCK_M * LDS_ROW];split-k
Split-K decomposition for M=16 with high K to maximize CU utilization.stages = 2
num_warps=4, num_stages=2,tile-m = 16
constexpr int MFMA_K=128, DOUBLE_K=256, LDS_ROW=128, HALF_K=64, BLOCK_M=16;tile-n = 64
constexpr int BLOCK_N = 64, K_FIXED = 512;vector-width = int4
int4 ra = reinterpret_cast<const int4*>(s)[0];Kernel source
submission_v160.py683 lines
"""
MXFP4 GEMM: hybrid HIP + Triton approach for mixed-precision FP4 matmul.
Optimized for MI355X (gfx950/CDNA4) with block-scaled FP4 MFMA.
Strategy: HIP assembly kernel for small-K shapes, autotuned Triton for large-K.
Split-K decomposition for M=16 with high K to maximize CU utilization.
"""
import torch
import triton
import triton.language as tl
from task import input_t, output_t
import os
os.environ["PYTORCH_ROCM_ARCH"] = "gfx950"
# ===================== HIP kernel for M<=32 K=512 =====================
_HIP_CPP = r"""
#include <torch/extension.h>
void launch_sk2_k512(int64_t a, int64_t bq, int64_t bsc, int64_t c, int M, int N);
"""
_HIP_SRC = r"""
#include <hip/hip_runtime.h>
#include <hip/hip_bf16.h>
#include <torch/extension.h>
constexpr int MFMA_K=128, DOUBLE_K=256, LDS_ROW=128, HALF_K=64, BLOCK_M=16;
typedef int __attribute__((ext_vector_type(4))) int4_vec;
typedef float __attribute__((ext_vector_type(4))) float4_vec;
typedef int32_t __attribute__((ext_vector_type(4))) i32x4;
typedef uint32_t __attribute__((address_space(3)))* as3_uint32_ptr;
extern "C" __device__ void llvm_amdgcn_raw_buffer_load_lds(
i32x4 rsrc, as3_uint32_ptr lds_ptr, int size, int voffset,
int soffset, int offset, int aux) __asm("llvm.amdgcn.raw.buffer.load.lds");
struct buffer_resource { uint64_t ptr; uint32_t range; uint32_t config; };
__device__ __forceinline__ i32x4 make_srsrc(const void* p, uint32_t r) {
buffer_resource s = {reinterpret_cast<uint64_t>(p), r, 0x110000};
return *reinterpret_cast<const i32x4*>(&s);
}
__device__ __forceinline__ float4_vec mfma_fp4(int4_vec A, int4_vec B, float4_vec C, int sA, int sB) {
float4_vec D;
asm("v_mfma_scale_f32_16x16x128_f8f6f4 %0,%1,%2,%3,%4,%5 cbsz:4 blgp:4"
: "=v"(D) : "v"(A), "v"(B), "v"(C), "v"(sA), "v"(sB));
return D;
}
__device__ __forceinline__ int lds_swz(int o) {
return o ^ (((o & 2047) >> 8) << 4);
}
__device__ __forceinline__ void compute_scale(float mx, uint8_t& sc, float& sf) {
if (mx > 0.f) {
uint32_t b = __float_as_uint(mx);
b = (b + 0x200000u) & 0xFF800000u;
int su = ((b >> 23) & 0xFF) - 129;
su = su < -127 ? -127 : (su > 127 ? 127 : su);
sc = (uint8_t)(su + 127);
sf = __uint_as_float((uint32_t)(su + 127) << 23);
} else {
sc = 0; sf = 0.f;
}
}
__device__ __forceinline__ float tree_max8(const float* v) {
float a0 = fmaxf(fabsf(v[0]), fabsf(v[1]));
float a1 = fmaxf(fabsf(v[2]), fabsf(v[3]));
float a2 = fmaxf(fabsf(v[4]), fabsf(v[5]));
float a3 = fmaxf(fabsf(v[6]), fabsf(v[7]));
return fmaxf(fmaxf(a0, a1), fmaxf(a2, a3));
}
__global__ __launch_bounds__(512, 3)
void gemm_bn64_sk2_k512(
const __hip_bfloat16* __restrict__ A,
const uint8_t* __restrict__ Bq,
const uint8_t* __restrict__ Bsc,
__hip_bfloat16* __restrict__ C,
int M, int N)
{
constexpr int BLOCK_N = 64, K_FIXED = 512;
constexpr int BSTRIDE = K_FIXED >> 1, SCS = K_FIXED >> 5;
const int group = threadIdx.x >> 8;
const int ltid = threadIdx.x & 255;
const int wid_local = ltid >> 6;
const int lm = ltid & 15;
const int lk = (ltid >> 4) & 3;
const int bn = blockIdx.x * BLOCK_N, bm = blockIdx.y * BLOCK_M;
const int wn = bn + (wid_local << 4);
__shared__ __align__(16) uint8_t Alds[2 * BLOCK_M * LDS_ROW];
__shared__ uint8_t Asclds[2 * BLOCK_M * 8];
__shared__ __align__(16) uint8_t Blds[2 * BLOCK_N * LDS_ROW];
__shared__ uint8_t Bslds[2 * BLOCK_N * 8];
const int a_base = group * BLOCK_M * LDS_ROW;
const int as_base = group * BLOCK_M * 8;
const int b_base = group * BLOCK_N * LDS_ROW;
const int bs_base = group * BLOCK_N * 8;
const int ke = group * DOUBLE_K;
const int kb = group * LDS_ROW;
const i32x4 srsrc = make_srsrc(Bq, N * BSTRIDE);
float4_vec acc = {0, 0, 0, 0};
if (ltid < 128) {
const int qr = ltid >> 3, qg = ltid & 7;
const int gr = bm + qr, ko = ke + qg * 32;
uint32_t p0 = 0, p1 = 0, p2 = 0, p3 = 0;
uint8_t asc = 0x7f;
if (gr < M) {
const __hip_bfloat16* s = A + gr * K_FIXED + ko;
int4 ra = reinterpret_cast<const int4*>(s)[0];
int4 rb = reinterpret_cast<const int4*>(s)[1];
int4 rc = reinterpret_cast<const int4*>(s)[2];
int4 rd = reinterpret_cast<const int4*>(s)[3];
const __hip_bfloat16* bfa = reinterpret_cast<const __hip_bfloat16*>(&ra);
const __hip_bfloat16* bfb = reinterpret_cast<const __hip_bfloat16*>(&rb);
const __hip_bfloat16* bfc = reinterpret_cast<const __hip_bfloat16*>(&rc);
const __hip_bfloat16* bfd = reinterpret_cast<const __hip_bfloat16*>(&rd);
float v0[8], v1[8], v2[8], v3[8];
for (int i = 0; i < 8; i++) v0[i] = __bfloat162float(bfa[i]);
for (int i = 0; i < 8; i++) v1[i] = __bfloat162float(bfb[i]);
for (int i = 0; i < 8; i++) v2[i] = __bfloat162float(bfc[i]);
for (int i = 0; i < 8; i++) v3[i] = __bfloat162float(bfd[i]);
float l = fmaxf(fmaxf(tree_max8(v0), tree_max8(v1)),
fmaxf(tree_max8(v2), tree_max8(v3)));
float sf;
compute_scale(l, asc, sf);
p0 = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(p0, v0[0], v0[1], sf, 0);
p0 = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(p0, v0[2], v0[3], sf, 1);
p0 = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(p0, v0[4], v0[5], sf, 2);
p0 = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(p0, v0[6], v0[7], sf, 3);
p1 = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(p1, v1[0], v1[1], sf, 0);
p1 = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(p1, v1[2], v1[3], sf, 1);
p1 = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(p1, v1[4], v1[5], sf, 2);
p1 = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(p1, v1[6], v1[7], sf, 3);
p2 = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(p2, v2[0], v2[1], sf, 0);
p2 = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(p2, v2[2], v2[3], sf, 1);
p2 = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(p2, v2[4], v2[5], sf, 2);
p2 = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(p2, v2[6], v2[7], sf, 3);
p3 = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(p3, v3[0], v3[1], sf, 0);
p3 = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(p3, v3[2], v3[3], sf, 1);
p3 = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(p3, v3[4], v3[5], sf, 2);
p3 = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(p3, v3[6], v3[7], sf, 3);
}
int4 wr;
wr.x = (int)p0; wr.y = (int)p1; wr.z = (int)p2; wr.w = (int)p3;
*reinterpret_cast<int4*>(&Alds[a_base + lds_swz(qr * LDS_ROW + qg * 16)]) = wr;
Asclds[as_base + qr * 8 + qg] = asc;
}
if (ltid >= 128) {
const int btid = ltid - 128;
constexpr int BNT = 128;
constexpr int BL = (BLOCK_N * LDS_ROW / 16 + BNT - 1) / BNT;
for (int ld = 0; ld < BL; ld++) {
const int f = (ld * BNT + btid) << 4;
const int r = f >> 7;
const int gn = bn + r;
if (r < BLOCK_N && gn < N) {
const int sf = lds_swz(f), sc = sf & 127, ac = kb + sc;
llvm_amdgcn_raw_buffer_load_lds(srsrc,
(as3_uint32_ptr)(reinterpret_cast<uintptr_t>(Blds) + b_base + f),
16,
(gn >> 4) * (BSTRIDE << 4) + (ac >> 5) * 512 + ((ac >> 4) & 1) * 256 + (gn & 15) * 16,
0, 0, 0);
}
}
}
{ const int so = group << 3;
if (ltid >= 128) {
const int stid = ltid - 128;
const int row = stid & (BLOCK_N - 1), gb = (stid >> 6) << 2;
const int gn = bn + row;
if (gn < N) {
const int rb = (gn >> 5) * 32 * SCS + (gn & 15) * 4 + ((gn >> 4) & 1);
for (int g = 0; g < 4; g++) {
const int grp = gb + g, ac = so + grp;
Bslds[bs_base + row * 8 + grp] = Bsc[rb + (ac & 3) * 64 + ((ac & 7) >> 2) * 2 + (ac >> 3) * 256];
}
} else {
for (int g = 0; g < 4; g++) Bslds[bs_base + row * 8 + gb + g] = 0x7f;
}
}}
asm volatile("s_waitcnt vmcnt(0)");
__syncthreads();
const int br = (wid_local << 4) + lm;
{
int4_vec A0, B0; int as0, bs0;
{ const int o = lds_swz(lm * LDS_ROW + (lk << 4));
const int4 t = *reinterpret_cast<const int4*>(&Alds[a_base + o]);
A0 = {t.x, t.y, t.z, t.w}; }
as0 = (int)Asclds[as_base + (lm << 3) + lk];
{ const int o = lds_swz(br * LDS_ROW + (lk << 4));
const int4 t = *reinterpret_cast<const int4*>(&Blds[b_base + o]);
B0 = {t.x, t.y, t.z, t.w}; }
bs0 = (int)Bslds[bs_base + br * 8 + lk];
acc = mfma_fp4(A0, B0, acc, as0, bs0);
}
{
int4_vec A1, B1; int as1, bs1;
{ const int o = lds_swz(lm * LDS_ROW + HALF_K + (lk << 4));
const int4 t = *reinterpret_cast<const int4*>(&Alds[a_base + o]);
A1 = {t.x, t.y, t.z, t.w}; }
as1 = (int)Asclds[as_base + (lm << 3) + 4 + lk];
{ const int o = lds_swz(br * LDS_ROW + HALF_K + (lk << 4));
const int4 t = *reinterpret_cast<const int4*>(&Blds[b_base + o]);
B1 = {t.x, t.y, t.z, t.w}; }
bs1 = (int)Bslds[bs_base + br * 8 + 4 + lk];
acc = mfma_fp4(A1, B1, acc, as1, bs1);
}
__syncthreads();
float* reduce_buf = reinterpret_cast<float*>(Alds);
if (group == 1) {
const float* ap = reinterpret_cast<const float*>(&acc);
for (int r = 0; r < 4; r++)
reduce_buf[ltid * 4 + r] = ap[r];
}
__syncthreads();
if (group == 0) {
float* mp = reinterpret_cast<float*>(&acc);
for (int r = 0; r < 4; r++)
mp[r] += reduce_buf[ltid * 4 + r];
const int or_ = bm + (lk << 2), oc = wn + lm;
if (oc < N) {
for (int r = 0; r < 4; r++) {
int g = or_ + r;
if (g < M) C[g * N + oc] = __float2bfloat16(mp[r]);
}
}
}
}
void launch_sk2_k512(int64_t a, int64_t bq, int64_t bsc, int64_t c, int M, int N) {
dim3 b(512), g((N + 63) / 64, (M + 15) / 16);
hipLaunchKernelGGL(gemm_bn64_sk2_k512, g, b, 0, 0,
reinterpret_cast<const __hip_bfloat16*>(a),
reinterpret_cast<const uint8_t*>(bq),
reinterpret_cast<const uint8_t*>(bsc),
reinterpret_cast<__hip_bfloat16*>(c), M, N);
}
"""
from torch.utils.cpp_extension import load_inline
_hip_mod = load_inline(
name="hip_sk2_k512",
cpp_sources=_HIP_CPP,
cuda_sources=_HIP_SRC,
functions=["launch_sk2_k512"],
verbose=False,
extra_cuda_cflags=["-O3", "-std=c++17", "-fno-gpu-rdc", "-ffp-contract=fast",
"--offload-arch=gfx950", "-ffast-math",
"-funsafe-math-optimizations",
"-mllvm", "-amdgpu-max-memory-clause=64",
"-mllvm", "-amdgpu-load-store-vectorizer",
"-mllvm", "-amdgpu-early-ifcvt",
"-mllvm", "-amdgpu-early-inline-all",
"-mllvm", "-amdgpu-internalize-symbols",
"-mllvm", "-amdgpu-scalarize-global-loads",
"-mllvm", "-amdgpu-dpp-combine",
"-mllvm", "-amdgpu-enable-pre-ra-optimizations",
"-mllvm", "-amdgpu-promote-alloca-to-vector-limit=256"],
)
_hip_dispatch = _hip_mod.launch_sk2_k512
# ---- Triton MXFP4 GEMM kernels ----
@triton.jit
def _xcd_reorder(pid, GRID_SIZE, NUM_XCDS: tl.constexpr = 8):
pids_per_xcd = (GRID_SIZE + NUM_XCDS - 1) // NUM_XCDS
tall_xcds = GRID_SIZE % NUM_XCDS
tall_xcds = tl.where(tall_xcds == 0, NUM_XCDS, tall_xcds)
xcd = pid % NUM_XCDS
local_pid = pid // NUM_XCDS
new_pid = tl.where(
xcd < tall_xcds,
xcd * pids_per_xcd + local_pid,
tall_xcds * pids_per_xcd + (xcd - tall_xcds) * (pids_per_xcd - 1) + local_pid,
)
return new_pid
@triton.jit
def _quantize_to_mxfp4(x_f32, BLOCK_M: tl.constexpr, NK_SC: tl.constexpr):
SZ16: tl.constexpr = NK_SC * 16
x_g = x_f32.reshape(BLOCK_M, NK_SC, 32)
amax = tl.max(tl.abs(x_g), axis=2, keep_dims=True)
amax_u32 = amax.to(tl.int32, bitcast=True)
amax_u32 = ((amax_u32 + 0x200000).to(tl.uint32, bitcast=True)) & 0xFF800000
exp_bits = (amax_u32 >> 23)
raw_exp = tl.maximum(exp_bits.to(tl.int32) - 2, 0)
a_scale = raw_exp.to(tl.uint8).reshape(BLOCK_M, NK_SC)
sf_exp = tl.maximum(raw_exp, 1).to(tl.uint32)
sf = (sf_exp << 23).to(tl.float32, bitcast=True)
x_pairs = x_g.reshape(BLOCK_M, NK_SC, 16, 2)
a_elems, b_elems = tl.split(x_pairs)
sf_bc = tl.broadcast_to(sf, (BLOCK_M, NK_SC, 16))
raw = tl.inline_asm_elementwise(
"v_cvt_scalef32_pk_fp4_f32 $0, $1, $2, $3",
"=v,v,v,v",
[a_elems, b_elems, sf_bc],
dtype=tl.uint32, is_pure=True, pack=1,
)
a_fp4 = (raw & 0xFF).to(tl.uint8).reshape(BLOCK_M, SZ16)
return a_fp4, a_scale
# Split-K path: distributes K dimension across CTAs, then reduces
@triton.jit
def _splitk_fp4_gemm(
a_ptr, b_ptr, c_ptr, b_sc_ptr,
M, N, K, N16, N32,
stride_am, stride_ak, stride_bn, stride_bk,
stride_cm, stride_cn, stride_bsn, stride_bsk,
BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr,
BLOCK_K: tl.constexpr, NUM_KSPLIT: tl.constexpr,
):
SG: tl.constexpr = 32
NK_SC: tl.constexpr = BLOCK_K // SG
tl.assume(stride_am > 0)
tl.assume(stride_ak > 0)
tl.assume(stride_bn > 0)
tl.assume(stride_bk > 0)
tl.assume(stride_cm > 0)
tl.assume(stride_cn > 0)
tl.assume(stride_bsn > 0)
tl.assume(stride_bsk > 0)
num_pid_n = tl.cdiv(N, BLOCK_N)
num_pid_m = tl.cdiv(M, BLOCK_M)
GRID_MN = num_pid_m * num_pid_n
pid = tl.program_id(0)
pid = _xcd_reorder(pid, GRID_MN * NUM_KSPLIT)
pid_k = pid % NUM_KSPLIT
pid_mn = pid // NUM_KSPLIT
pid_m = pid_mn // num_pid_n
pid_n = pid_mn % num_pid_n
nk_total = K // BLOCK_K
nk_base = nk_total // NUM_KSPLIT
nk_rem = nk_total % NUM_KSPLIT
my_nk = nk_base + tl.where(pid_k < nk_rem, 1, 0)
k_start_iter = pid_k * nk_base + tl.where(pid_k < nk_rem, pid_k, nk_rem)
k_offset = k_start_iter * BLOCK_K
offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
m_mask = offs_m < M
offs_k_bf16 = tl.arange(0, BLOCK_K)
a_ptrs = a_ptr + offs_m[:, None] * stride_am + (k_offset + offs_k_bf16[None, :]) * stride_ak
offs_bn_sh = (pid_n * (BLOCK_N // 16) + tl.arange(0, BLOCK_N // 16)) % N16
offs_k_sh = (k_offset // 2) * 16 + tl.arange(0, (BLOCK_K // 2) * 16)
b_ptrs = b_ptr + offs_bn_sh[:, None].to(tl.int64) * stride_bn + offs_k_sh[None, :].to(tl.int64) * stride_bk
offs_bn_sc = (pid_n * (BLOCK_N // 32) + tl.arange(0, BLOCK_N // 32)) % N32
offs_ks_raw = k_offset + tl.arange(0, BLOCK_K)
b_sc_ptrs = b_sc_ptr + offs_bn_sc[:, None].to(tl.int64) * stride_bsn + offs_ks_raw[None, :].to(tl.int64) * stride_bsk
offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
n_mask = offs_n < N
acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
for _ in range(my_nk):
a_bf16 = tl.load(a_ptrs, mask=m_mask[:, None], other=0)
b_sc_raw = tl.load(b_sc_ptrs, cache_modifier=".cg")
b_raw = tl.load(b_ptrs, cache_modifier=".cg")
a_fp4, a_sc = _quantize_to_mxfp4(a_bf16.to(tl.float32), BLOCK_M, NK_SC)
b_sc = (b_sc_raw
.reshape(BLOCK_N // 32, BLOCK_K // 256, 4, 16, 2, 2, 1)
.permute(0, 5, 3, 1, 4, 2, 6)
.reshape(BLOCK_N, NK_SC))
b = (b_raw
.reshape(1, BLOCK_N // 16, BLOCK_K // 64, 2, 16, 16)
.permute(0, 1, 4, 2, 3, 5)
.reshape(BLOCK_N, BLOCK_K // 2)
.trans(1, 0))
acc = tl.dot_scaled(a_fp4, a_sc, "e2m1", b, b_sc, "e2m1", acc)
a_ptrs += BLOCK_K * stride_ak
b_ptrs += (BLOCK_K // 2) * 16 * stride_bk
b_sc_ptrs += BLOCK_K * stride_bsk
c_ptrs = c_ptr + pid_k.to(tl.int64) * (M * N) + offs_m[:, None].to(tl.int64) * stride_cm + offs_n[None, :].to(tl.int64) * stride_cn
c_mask = m_mask[:, None] & n_mask[None, :]
if my_nk > 0:
tl.store(c_ptrs, acc, mask=c_mask, cache_modifier=".wt")
else:
tl.store(c_ptrs, tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32), mask=c_mask, cache_modifier=".wt")
@triton.jit
def _reduce_partials(
partial_ptr, c_ptr, total_elems,
NUM_KSPLIT: tl.constexpr, BLOCK: tl.constexpr,
):
pid = tl.program_id(0)
offs = pid * BLOCK + tl.arange(0, BLOCK)
mask = offs < total_elems
acc = tl.zeros((BLOCK,), dtype=tl.float32)
for ks in range(NUM_KSPLIT):
p = tl.load(partial_ptr + ks * total_elems + offs, mask=mask, other=0.0)
acc += p
tl.store(c_ptr + offs, acc.to(tl.bfloat16), mask=mask)
# Autotuned FP4 matmul with XCD-aware tile scheduling
def _cfg(bm, bn, bk, gsm, xcds, w, s):
return triton.Config(
{'BLOCK_M': bm, 'BLOCK_N': bn, 'BLOCK_K': bk,
'GROUP_SIZE_M': gsm, 'NUM_XCDS': xcds},
num_warps=w, num_stages=s,
)
_MATMUL_CONFIGS = [
_cfg(16, 128, 256, 4, 8, 4, 2),
_cfg(16, 128, 512, 4, 8, 4, 2),
_cfg(16, 128, 256, 1, 8, 4, 2),
_cfg(16, 128, 512, 1, 8, 4, 2),
_cfg(16, 128, 1024, 4, 8, 4, 1),
_cfg(16, 128, 1024, 4, 8, 4, 2),
_cfg(16, 128, 1024, 1, 8, 4, 1),
_cfg(16, 128, 256, 1, 8, 4, 1),
_cfg(16, 128, 512, 1, 8, 4, 1),
_cfg(16, 256, 256, 1, 8, 4, 2),
_cfg(16, 256, 512, 1, 8, 4, 2),
_cfg(16, 256, 256, 4, 8, 4, 2),
_cfg(16, 128, 256, 8, 8, 4, 2),
_cfg(16, 256, 256, 8, 8, 4, 2),
_cfg(16, 256, 256, 16, 8, 4, 2),
_cfg(16, 128, 256, 16, 8, 4, 2),
_cfg(16, 256, 256, 8, 8, 8, 2),
_cfg(16, 256, 256, 16, 8, 8, 2),
_cfg(16, 128, 256, 8, 8, 8, 2),
_cfg(16, 128, 512, 8, 8, 8, 2),
_cfg(16, 128, 1024, 4, 8, 8, 1),
_cfg(16, 256, 512, 8, 8, 8, 2),
_cfg(32, 128, 256, 4, 8, 4, 2),
_cfg(32, 128, 512, 4, 8, 4, 2),
_cfg(32, 128, 256, 4, 8, 8, 2),
_cfg(32, 128, 512, 4, 8, 8, 2),
_cfg(32, 128, 1024, 4, 8, 8, 1),
_cfg(32, 256, 256, 4, 8, 8, 2),
_cfg(32, 256, 512, 4, 8, 8, 2),
_cfg(32, 128, 256, 8, 8, 4, 2),
_cfg(32, 128, 256, 8, 8, 8, 2),
_cfg(32, 256, 256, 8, 8, 8, 2),
_cfg(32, 256, 512, 8, 8, 8, 2),
_cfg(64, 128, 256, 4, 8, 8, 2),
_cfg(64, 128, 512, 4, 8, 8, 2),
_cfg(64, 256, 256, 4, 8, 8, 2),
_cfg(64, 128, 256, 8, 8, 8, 2),
_cfg(16, 128, 512, 8, 8, 4, 2),
_cfg(16, 256, 512, 4, 8, 4, 2),
_cfg(16, 256, 512, 8, 8, 4, 2),
_cfg(16, 128, 512, 8, 1, 4, 2),
_cfg(16, 128, 512, 4, 1, 4, 2),
_cfg(16, 256, 512, 8, 1, 8, 2),
_cfg(16, 128, 256, 8, 1, 4, 2),
_cfg(16, 128, 256, 4, 1, 4, 2),
_cfg(32, 128, 512, 4, 1, 4, 2),
_cfg(32, 128, 256, 8, 1, 8, 2),
_cfg(64, 128, 256, 4, 1, 8, 2),
_cfg(64, 128, 512, 4, 1, 8, 2),
_cfg(16, 128, 512, 8, 4, 4, 2),
_cfg(16, 256, 512, 8, 4, 8, 2),
_cfg(32, 128, 256, 8, 4, 8, 2),
_cfg(16, 128, 512, 8, 8, 4, 3),
_cfg(16, 128, 512, 4, 8, 4, 3),
_cfg(16, 256, 512, 8, 8, 8, 3),
_cfg(16, 128, 256, 8, 8, 4, 3),
_cfg(32, 128, 512, 4, 8, 4, 3),
_cfg(32, 128, 256, 8, 8, 8, 3),
_cfg(64, 128, 256, 4, 8, 8, 3),
_cfg(16, 128, 512, 8, 1, 4, 3),
_cfg(16, 256, 512, 8, 1, 8, 3),
# Wider grid variants for better CU utilization at M=64
_cfg(16, 64, 512, 8, 8, 4, 2),
_cfg(16, 64, 512, 4, 8, 4, 2),
_cfg(16, 64, 256, 8, 8, 4, 2),
_cfg(16, 64, 256, 4, 8, 4, 2),
_cfg(16, 64, 1024, 4, 8, 4, 1),
_cfg(16, 64, 1024, 4, 8, 4, 2),
_cfg(16, 64, 512, 8, 8, 2, 2),
_cfg(16, 64, 256, 8, 8, 2, 2),
_cfg(16, 64, 512, 8, 1, 4, 2),
_cfg(16, 64, 512, 8, 8, 4, 3),
# Larger BM tiles with narrow BN for high-M shapes
_cfg(32, 64, 512, 8, 8, 4, 2),
_cfg(32, 64, 512, 4, 8, 4, 2),
_cfg(32, 64, 256, 8, 8, 4, 2),
_cfg(32, 64, 512, 8, 8, 2, 2),
_cfg(64, 64, 512, 4, 8, 4, 2),
_cfg(64, 64, 256, 4, 8, 4, 2),
_cfg(64, 64, 512, 4, 8, 8, 2),
_cfg(64, 64, 256, 4, 8, 8, 2),
_cfg(64, 64, 512, 8, 8, 8, 2),
_cfg(32, 64, 256, 8, 8, 2, 2),
]
def _filter_valid_configs(configs, named_args, **kwargs):
K = named_args['K']
M = named_args['M']
return [c for c in configs
if K % c.kwargs['BLOCK_K'] == 0
and (c.kwargs['BLOCK_M'] == 16 or c.kwargs['BLOCK_M'] <= M)]
@triton.autotune(configs=_MATMUL_CONFIGS, key=['M', 'N', 'K'],
prune_configs_by={'early_config_prune': _filter_valid_configs})
@triton.heuristics({
'EVEN_M': lambda args: args['M'] % args['BLOCK_M'] == 0,
'EVEN_N': lambda args: args['N'] % args['BLOCK_N'] == 0,
'NUM_ITERS': lambda args: args['K'] // args['BLOCK_K'] if args['K'] >= 1024 else 0,
})
@triton.jit
def _fp4_matmul_kernel(
a_ptr, b_ptr, c_ptr, b_sc_ptr,
M, N, K, N16, N32,
stride_am, stride_ak, stride_bn, stride_bk,
stride_cm, stride_cn, stride_bsn, stride_bsk,
BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr,
BLOCK_K: tl.constexpr, GROUP_SIZE_M: tl.constexpr,
NUM_XCDS: tl.constexpr,
EVEN_M: tl.constexpr, EVEN_N: tl.constexpr,
NUM_ITERS: tl.constexpr,
):
SG: tl.constexpr = 32
NK_SC: tl.constexpr = BLOCK_K // SG
tl.assume(stride_am > 0)
tl.assume(stride_ak > 0)
tl.assume(stride_bn > 0)
tl.assume(stride_bk > 0)
tl.assume(stride_cm > 0)
tl.assume(stride_cn > 0)
tl.assume(stride_bsn > 0)
tl.assume(stride_bsk > 0)
pid = tl.program_id(0)
num_pid_m = tl.cdiv(M, BLOCK_M)
num_pid_n = tl.cdiv(N, BLOCK_N)
GRID_MN = num_pid_m * num_pid_n
pid = _xcd_reorder(pid, GRID_MN, NUM_XCDS)
num_pid_in_group = GROUP_SIZE_M * num_pid_n
group_id = pid // num_pid_in_group
first_pid_m = group_id * GROUP_SIZE_M
group_size_m = tl.minimum(num_pid_m - first_pid_m, GROUP_SIZE_M)
pid_m = first_pid_m + (pid % num_pid_in_group) % group_size_m
pid_n = (pid % num_pid_in_group) // group_size_m
offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
m_mask = offs_m < M
offs_k_bf16 = tl.arange(0, BLOCK_K)
a_ptrs = a_ptr + offs_m[:, None] * stride_am + offs_k_bf16[None, :] * stride_ak
offs_bn_sh = (pid_n * (BLOCK_N // 16) + tl.arange(0, BLOCK_N // 16)) % N16
offs_k_sh = tl.arange(0, (BLOCK_K // 2) * 16)
b_ptrs = b_ptr + offs_bn_sh[:, None].to(tl.int64) * stride_bn + offs_k_sh[None, :].to(tl.int64) * stride_bk
offs_bn_sc = (pid_n * (BLOCK_N // 32) + tl.arange(0, BLOCK_N // 32)) % N32
offs_ks_raw = tl.arange(0, BLOCK_K)
b_sc_ptrs = b_sc_ptr + offs_bn_sc[:, None].to(tl.int64) * stride_bsn + offs_ks_raw[None, :].to(tl.int64) * stride_bsk
offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
n_mask = offs_n < N
acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
niters = NUM_ITERS if NUM_ITERS > 0 else K // BLOCK_K
for _ in range(niters):
if EVEN_M:
a_bf16 = tl.load(a_ptrs)
else:
a_bf16 = tl.load(a_ptrs, mask=m_mask[:, None], other=0)
b_sc_raw = tl.load(b_sc_ptrs, cache_modifier=".cg")
b_raw = tl.load(b_ptrs, cache_modifier=".cg")
a_fp4, a_sc = _quantize_to_mxfp4(a_bf16.to(tl.float32), BLOCK_M, NK_SC)
b_sc = (b_sc_raw
.reshape(BLOCK_N // 32, BLOCK_K // 256, 4, 16, 2, 2, 1)
.permute(0, 5, 3, 1, 4, 2, 6)
.reshape(BLOCK_N, NK_SC))
b = (b_raw
.reshape(1, BLOCK_N // 16, BLOCK_K // 64, 2, 16, 16)
.permute(0, 1, 4, 2, 3, 5)
.reshape(BLOCK_N, BLOCK_K // 2)
.trans(1, 0))
acc = tl.dot_scaled(a_fp4, a_sc, "e2m1", b, b_sc, "e2m1", acc)
a_ptrs += BLOCK_K * stride_ak
b_ptrs += (BLOCK_K // 2) * 16 * stride_bk
b_sc_ptrs += BLOCK_K * stride_bsk
c = acc.to(tl.bfloat16)
offs_cn = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
c_ptrs = c_ptr + offs_m[:, None] * stride_cm + offs_cn[None, :] * stride_cn
if EVEN_M and EVEN_N:
tl.store(c_ptrs, c, cache_modifier=".wt")
elif EVEN_N:
tl.store(c_ptrs, c, mask=m_mask[:, None], cache_modifier=".wt")
else:
tl.store(c_ptrs, c, mask=m_mask[:, None] & n_mask[None, :], cache_modifier=".wt")
# ---- Entry point ----
_N_COMPUTE_UNITS = 304
def _run_splitk_path(A_in, B_sh, B_sc_raw, m, n, k, n_padded, dev):
"""Handle M<=16 large-K via split-K decomposition + partial reduction."""
tile_m, tile_n, tile_k = 16, 128, 512
num_k_tiles = k // tile_k
mn_grid = triton.cdiv(m, tile_m) * triton.cdiv(n, tile_n)
n_splits = min(num_k_tiles, max(1, _N_COMPUTE_UNITS // mn_grid))
while num_k_tiles % n_splits != 0 and n_splits > 1:
n_splits -= 1
if n_splits < 2:
return None
total = m * n
partials = torch.empty(n_splits * total, dtype=torch.float32, device=dev)
out = torch.empty(m, n, dtype=torch.bfloat16, device=dev)
_splitk_fp4_gemm[(n_splits * mn_grid,)](
A_in, B_sh, partials, B_sc_raw,
m, n, k, n // 16, n_padded // 32,
A_in.stride(0), A_in.stride(1),
B_sh.stride(0), B_sh.stride(1),
n, 1,
B_sc_raw.stride(0), B_sc_raw.stride(1),
BLOCK_M=tile_m, BLOCK_N=tile_n, BLOCK_K=tile_k, NUM_KSPLIT=n_splits,
num_warps=4, num_stages=2,
)
merge_blk = 256
_reduce_partials[(triton.cdiv(total, merge_blk),)](
partials, out, total,
NUM_KSPLIT=n_splits, BLOCK=merge_blk,
num_warps=4, num_stages=2,
)
return out
def custom_kernel(data: input_t) -> output_t:
A_in, B, B_q, B_shuffle, B_scale_sh = data
m, k = A_in.shape
n = B.shape[0]
dev = A_in.device
# Fast path: small M with K=512 uses hand-tuned HIP assembly kernel
if m <= 32 and k == 512:
out = torch.empty(m, n, dtype=torch.bfloat16, device=dev)
_hip_dispatch(A_in.data_ptr(), B_shuffle.data_ptr(), B_scale_sh.data_ptr(),
out.data_ptr(), m, n)
return out
# Prepare B operands for Triton path
B_sh = B_shuffle.contiguous().view(torch.uint8).view(n // 16, (k // 2) * 16)
n_pad = (n + 255) // 256 * 256
B_sc = B_scale_sh.contiguous().view(torch.uint8).reshape(n_pad // 32, k)
# Try split-K for skinny M with large K
if m <= 16 and k >= 1024:
result = _run_splitk_path(A_in, B_sh, B_sc, m, n, k, n_pad, dev)
if result is not None:
return result
# General autotuned GEMM path
out = torch.empty(m, n, dtype=torch.bfloat16, device=dev)
grid_fn = lambda META: (triton.cdiv(m, META['BLOCK_M']) * triton.cdiv(n, META['BLOCK_N']),)
_fp4_matmul_kernel[grid_fn](
A_in, B_sh, out, B_sc,
m, n, k, n // 16, n_pad // 32,
A_in.stride(0), A_in.stride(1),
B_sh.stride(0), B_sh.stride(1),
out.stride(0), out.stride(1),
B_sc.stride(0), B_sc.stride(1),
)
return out
scrolls · 683 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