submission 737364
divc13 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 676 lines, June 9 Researcher Reciprocity License v1.0.
submission_v78.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-737364?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:86df4e4ffab700bd425d04881481b1d5a6f4130f427c7de32feff462d983850e
license declaredunknown
license concludedunknown
authorsdivc13
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
autotune
return triton.Config(num-warps = 4
num_warps=4, num_stages=2,shared-memory
__shared__ __align__(16) uint8_t Alds[2 * BLOCK_M * LDS_ROW];split-k
No split-K changes (MY_NK constexpr causes regression).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
v78: v76 + BN=64 configs for M=64/M=256 (better CU utilization).vector-width = int4
int4 ra = reinterpret_cast<const int4*>(s)[0];Kernel source
submission_v78.py676 lines
"""
v78: v76 + BN=64 configs for M=64/M=256 (better CU utilization).
M=64: grid 224→448 (CU 74%→147%). M=256: grid 192→768 (CU 63%→253%).
No split-K changes (MY_NK constexpr causes regression).
"""
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 kernels (v75) =====================
@triton.jit
def _remap_xcd(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 _quant_fp4_hw(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 GEMM (M=16) ---------------
@triton.jit
def _fused_splitk_gemm_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, 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 = _remap_xcd(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 = _quant_fp4_hw(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 _merge_splitk_kernel(
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 GEMM (v75 configs) ---------------
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,
)
_FUSED_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),
# --- v78: BN=64 for M=64 (grid 224→448) ---
_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),
# --- v78: BN=64 w/ BM=32/64 for M=256 ---
_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 _prune_fused_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=_FUSED_CONFIGS, key=['M', 'N', 'K'],
prune_configs_by={'early_config_prune': _prune_fused_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 _fused_gemm_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 = _remap_xcd(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 = _quant_fp4_hw(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")
# ===================== Host dispatch =====================
NUM_CUS = 304
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
# HIP path for M<=32 K=512 (M=4, M=32 benchmark shapes)
if m <= 32 and k == 512:
C = torch.empty(m, n, dtype=torch.bfloat16, device=dev)
_hip_dispatch(A_in.data_ptr(), B_shuffle.data_ptr(), B_scale_sh.data_ptr(),
C.data_ptr(), m, n)
return C
# Triton path for everything else (M=16 K=7168, M=64 K=2048, M=256 K=1536)
B_sh = B_shuffle.contiguous().view(torch.uint8).view(n // 16, (k // 2) * 16)
n_padded = (n + 255) // 256 * 256
B_sc_raw = B_scale_sh.contiguous().view(torch.uint8).reshape(n_padded // 32, k)
use_splitk = False
if m <= 16 and k >= 1024:
BM_sk, BN_sk, BK_sk = 16, 128, 512
nk = k // BK_sk
grid_mn = triton.cdiv(m, BM_sk) * triton.cdiv(n, BN_sk)
if nk > 1:
NUM_KSPLIT = min(nk, max(1, NUM_CUS // grid_mn))
while nk % NUM_KSPLIT != 0 and NUM_KSPLIT > 1:
NUM_KSPLIT -= 1
if NUM_KSPLIT >= 2:
use_splitk = True
if use_splitk:
total_elems = m * n
C_parts = torch.empty(NUM_KSPLIT * total_elems, dtype=torch.float32, device=dev)
C = torch.empty(m, n, dtype=torch.bfloat16, device=dev)
_fused_splitk_gemm_kernel[(NUM_KSPLIT * grid_mn,)](
A_in, B_sh, C_parts, 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=BM_sk, BLOCK_N=BN_sk, BLOCK_K=BK_sk, NUM_KSPLIT=NUM_KSPLIT,
num_warps=4, num_stages=2,
)
MERGE_BLOCK = 256
_merge_splitk_kernel[(triton.cdiv(total_elems, MERGE_BLOCK),)](
C_parts, C, total_elems,
NUM_KSPLIT=NUM_KSPLIT, BLOCK=MERGE_BLOCK,
num_warps=4, num_stages=2,
)
else:
C = torch.empty(m, n, dtype=torch.bfloat16, device=dev)
grid = lambda META: (triton.cdiv(m, META['BLOCK_M']) * triton.cdiv(n, META['BLOCK_N']),)
_fused_gemm_kernel[grid](
A_in, B_sh, C, 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),
C.stride(0), C.stride(1),
B_sc_raw.stride(0), B_sc_raw.stride(1),
)
return C
scrolls · 676 lines total
Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0
Changes from previous submission
Against this author's previous submission submission 572802.
"""- Phase 60: Hybrid dispatch — sub59 B_scale + sub57 BLOCK_N=128 for large-M.- Benchmarks 5/6 (M>=64) use BLOCK_N=128; benchmarks 1-4 (M<=32) use BLOCK_N=64.- Both paths use sub59's 2*BLOCK_N-thread B_scale loading (4 loads each).+ v78: v76 + BN=64 configs for M=64/M=256 (better CU utilization).+ M=64: grid 224→448 (CU 74%→147%). M=256: grid 192→768 (CU 63%→253%).+ No split-K changes (MY_NK constexpr causes regression)."""- from task import input_t, output_timport torch+ import triton+ import triton.language as tl+ from task import input_t, output_timport osos.environ["PYTORCH_ROCM_ARCH"] = "gfx950"- CPP_SOURCE = r"""+ # ===================== HIP kernel for M<=32 K=512 =====================++ _HIP_CPP = r"""#include <torch/extension.h>- // BLOCK_N=64 (sub53) path- void launch_n64_nosplit(- torch::Tensor A_bf16, torch::Tensor B_q, torch::Tensor B_scale,- torch::Tensor C, int M, int N, int K);- void launch_n64_splitk(- torch::Tensor A_bf16, torch::Tensor B_q, torch::Tensor B_scale,- torch::Tensor workspace, int M, int N, int K, int split_k);- // BLOCK_N=128 (sub57) path- void launch_n128_nosplit(- torch::Tensor A_bf16, torch::Tensor B_q, torch::Tensor B_scale,- torch::Tensor C, int M, int N, int K);- void launch_n128_splitk(- torch::Tensor A_bf16, torch::Tensor B_q, torch::Tensor B_scale,- torch::Tensor workspace, int M, int N, int K, int split_k);- // Reduce- void launch_reduce(torch::Tensor workspace, torch::Tensor C, int M, int N, int split_k);+ void launch_sk2_k512(int64_t a, int64_t bq, int64_t bsc, int64_t c, int M, int N);"""- HIP_SOURCE = r"""+ _HIP_SRC = r"""#include <hip/hip_runtime.h>#include <hip/hip_bf16.h>#include <torch/extension.h>- constexpr int WARP_SIZE = 64;- constexpr int MFMA_K = 128;- constexpr int DOUBLE_K = MFMA_K * 2;- constexpr int LDS_ROW = DOUBLE_K >> 1;- constexpr int HALF_K = MFMA_K >> 1;- constexpr int SCALE_GROUP = 32;- constexpr int BLOCK_M = 16;-+ 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");+ 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* ptr, uint32_t range_bytes) {- buffer_resource rsrc = {reinterpret_cast<uint64_t>(ptr), range_bytes, 0x110000};- return *reinterpret_cast<const i32x4*>(&rsrc);+ __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_scaled(- int4_vec A, int4_vec B, float4_vec C, int sA, int sB- ) {+ __device__ __forceinline__ float4_vec mfma_fp4(int4_vec A, int4_vec B, float4_vec C, int sA, int sB) {float4_vec D;- asm volatile(- "v_mfma_scale_f32_16x16x128_f8f6f4 %0, %1, %2, %3, %4, %5 cbsz:4 blgp:4"+ 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 offset) {- return offset ^ (((offset & 2047) >> 8) << 4);+ __device__ __forceinline__ int lds_swz(int o) {+ return o ^ (((o & 2047) >> 8) << 4);}- __device__ __forceinline__ void compute_scale(float max_abs, uint8_t& sc, float& scale_f) {- if (max_abs > 0.0f) {- uint32_t b = __float_as_uint(max_abs);+ __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);- scale_f = __uint_as_float((uint32_t)(su + 127) << 23);- } else { sc = 0; scale_f = 0.0f; }+ sf = __uint_as_float((uint32_t)(su + 127) << 23);+ } else {+ sc = 0; sf = 0.f;+ }}- // ============================================================- // Templated kernel — BLOCK_N and NUM_WARPS as template params- // ============================================================- template <int BLOCK_N, int NUM_WARPS>- __global__ __launch_bounds__(NUM_WARPS * 64, (NUM_WARPS == 4 ? 3 : 2))- void gemm_kernel(- const __hip_bfloat16* __restrict__ A_bf16,- const uint8_t* __restrict__ B_q,- const uint8_t* __restrict__ B_scale,- float* __restrict__ workspace,- __hip_bfloat16* __restrict__ C_out,- const int M, const int N, const int K,- const int k_steps_per_split- ) {- constexpr int NUM_THREADS = NUM_WARPS * 64;- const int warp_id = threadIdx.x >> 6;- const int lane_id = threadIdx.x & 63;- const int lane_m = lane_id & 15;- const int lane_k = lane_id >> 4;- const int tid = threadIdx.x;+ __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));+ }- const int block_m = blockIdx.y * BLOCK_M;- const int block_n = blockIdx.x * BLOCK_N;- const int warp_n = block_n + (warp_id << 4);- const int split_id = blockIdx.z;+ __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 b_stride = K >> 1;- const int sc_stride = K >> 5;+ const int bn = blockIdx.x * BLOCK_N, bm = blockIdx.y * BLOCK_M;+ const int wn = bn + (wid_local << 4);- __shared__ __align__(16) uint8_t A_lds[BLOCK_M * LDS_ROW];- __shared__ __align__(16) uint8_t B_lds[BLOCK_N * LDS_ROW];- __shared__ uint8_t A_scale_lds[BLOCK_M * 8];- __shared__ uint8_t B_scale_lds[BLOCK_N * 8];+ __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 i32x4 b_srsrc = make_srsrc(B_q, N * b_stride);- float4_vec acc = {0.0f, 0.0f, 0.0f, 0.0f};+ 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 ks_start = split_id * k_steps_per_split;- const int ks_end = ks_start + k_steps_per_split;+ const int ke = group * DOUBLE_K;+ const int kb = group * LDS_ROW;- for (int ks = ks_start; ks < ks_end; ks++) {- const int k_elem = ks * DOUBLE_K;- const int k_byte = ks * LDS_ROW;+ const i32x4 srsrc = make_srsrc(Bq, N * BSTRIDE);+ float4_vec acc = {0, 0, 0, 0};- // A quant: HW FP4 conversion (only 256 threads needed)- if (tid < 256) {- const int group_id = tid >> 1;- const int half = tid & 1;- const int q_row = group_id >> 3;- const int q_grp = group_id & 7;- const int g_row = block_m + q_row;- const int k_off = k_elem + q_grp * SCALE_GROUP + half * 16;-- uint32_t pk_lo = 0, pk_hi = 0;- uint8_t a_scale_val = 0x7f;-- if (g_row < M) {- const __hip_bfloat16* src = A_bf16 + g_row * K + k_off;- int4 raw[2];- #pragma unroll- for (int j = 0; j < 2; j++)- raw[j] = reinterpret_cast<const int4*>(src)[j];- const __hip_bfloat16* bf = reinterpret_cast<const __hip_bfloat16*>(raw);- float vals[16];- float local_max = 0.0f;- #pragma unroll- for (int i = 0; i < 16; i++) {- vals[i] = __bfloat162float(bf[i]);- local_max = fmaxf(local_max, fabsf(vals[i]));- }- float global_max = fmaxf(local_max, __shfl_xor(local_max, 1));- float scale_f;- compute_scale(global_max, a_scale_val, scale_f);- pk_lo = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk_lo, vals[0], vals[1], scale_f, 0);- pk_lo = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk_lo, vals[2], vals[3], scale_f, 1);- pk_lo = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk_lo, vals[4], vals[5], scale_f, 2);- pk_lo = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk_lo, vals[6], vals[7], scale_f, 3);- pk_hi = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk_hi, vals[8], vals[9], scale_f, 0);- pk_hi = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk_hi, vals[10], vals[11], scale_f, 1);- pk_hi = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk_hi, vals[12], vals[13], scale_f, 2);- pk_hi = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk_hi, vals[14], vals[15], scale_f, 3);- }- const int a_lds_off = lds_swz(q_row * LDS_ROW + q_grp * 16 + half * 8);- int2 packed; packed.x = (int)pk_lo; packed.y = (int)pk_hi;- *reinterpret_cast<int2*>(&A_lds[a_lds_off]) = packed;- if (half == 0) A_scale_lds[q_row * 8 + q_grp] = a_scale_val;+ 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;+ }- // B: buffer_load_lds from B_shuffle with swizzled global source- {- constexpr int B_LOADS = (BLOCK_N * LDS_ROW / 16 + NUM_THREADS - 1) / NUM_THREADS;- #pragma unroll- for (int ld = 0; ld < B_LOADS; ld++) {- const int flat = (ld * NUM_THREADS + tid) << 4;- const int row = flat >> 7;- const int g_row = block_n + row;- if (row < BLOCK_N && g_row < N) {- const int swz_flat = lds_swz(flat);- const int swz_col = swz_flat & 127;- const int abs_col = k_byte + swz_col;- const int tile_n = g_row >> 4;- const int inner_n = g_row & 15;- const int tile_k = abs_col >> 5;- const int inner_k_hi = (abs_col >> 4) & 1;- const int src_off = tile_n * (b_stride << 4) + tile_k * 512 + inner_k_hi * 256 + inner_n * 16;-- llvm_amdgcn_raw_buffer_load_lds(b_srsrc,- (as3_uint32_ptr)(reinterpret_cast<uintptr_t>(B_lds) + flat),- 16, src_off, 0, 0, 2); // aux=2: SLC (non-temporal)- }+ 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);}}+ }- // B_scale from B_scale_sh — sub59 style: 2*BLOCK_N threads, 4 loads each- // For BLOCK_N=64: 128 threads active, for BLOCK_N=128: 256 threads active- {- const int sc_off = ks << 3;- if (tid < BLOCK_N * 2) {- const int row = tid & (BLOCK_N - 1); // 0..BLOCK_N-1- constexpr int BN_SHIFT = (BLOCK_N == 64) ? 6 : 7;- const int grp_base = (tid >> BN_SHIFT) << 2; // 0 or 4- const int g_row = block_n + row;-- if (g_row < N) {- const int row_base = (g_row >> 5) * 32 * sc_stride- + (g_row & 15) * 4- + ((g_row >> 4) & 1);- #pragma unroll- for (int g = 0; g < 4; g++) {- const int grp = grp_base + g;- const int abs_col = sc_off + grp;- const int col_off = (abs_col & 3) * 64- + ((abs_col & 7) >> 2) * 2- + (abs_col >> 3) * 256;- B_scale_lds[row * 8 + grp] = B_scale[row_base + col_off];- }- } else {- #pragma unroll- for (int g = 0; g < 4; g++)- B_scale_lds[row * 8 + grp_base + g] = 0x7f;- }+ { 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();+ asm volatile("s_waitcnt vmcnt(0)");+ __syncthreads();- // MFMA- #pragma unroll- for (int half = 0; half < 2; half++) {- const int kh = half * HALF_K;- int4_vec A_reg;- {- const int a_off = lds_swz(lane_m * LDS_ROW + kh + (lane_k << 4));- const int4 tmp = *reinterpret_cast<const int4*>(&A_lds[a_off]);- A_reg.s0 = tmp.x; A_reg.s1 = tmp.y; A_reg.s2 = tmp.z; A_reg.s3 = tmp.w;- }- int4_vec B_reg;- {- const int b_row = (warp_id << 4) + lane_m;- const int b_off = lds_swz(b_row * LDS_ROW + kh + (lane_k << 4));- const int4 tmp = *reinterpret_cast<const int4*>(&B_lds[b_off]);- B_reg.s0 = tmp.x; B_reg.s1 = tmp.y; B_reg.s2 = tmp.z; B_reg.s3 = tmp.w;- }- const int a_sc = (int)A_scale_lds[(lane_m << 3) + (half << 2) + lane_k];- const int b_sc = (int)B_scale_lds[((warp_id << 4) + lane_m) * 8 + (half << 2) + lane_k];- acc = mfma_fp4_scaled(A_reg, B_reg, acc, a_sc, b_sc);- }- __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);+ }- // Store- const int out_row = block_m + (lane_k << 2);- const int out_col = warp_n + lane_m;- if (out_col < N) {+ __syncthreads();+ float* reduce_buf = reinterpret_cast<float*>(Alds);++ if (group == 1) {const float* ap = reinterpret_cast<const float*>(&acc);- if (C_out) {- #pragma unroll+ 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++) {- const int gm = out_row + r;- if (gm < M) C_out[gm * N + out_col] = __float2bfloat16(ap[r]);+ int g = or_ + r;+ if (g < M) C[g * N + oc] = __float2bfloat16(mp[r]);}- } else {- float* ws = workspace + split_id * M * N;- #pragma unroll- for (int r = 0; r < 4; r++) {- const int gm = out_row + r;- if (gm < M) ws[gm * N + out_col] = ap[r];- }}}}- // Reduction kernel- __global__ void reduce_kernel(- const float* __restrict__ workspace,- __hip_bfloat16* __restrict__ C,- const int M, const int N, const int split_k- ) {- const int idx = blockIdx.x * blockDim.x + threadIdx.x;- if (idx >= M * N) return;- float sum = 0.0f;- for (int s = 0; s < split_k; s++)- sum += workspace[s * M * N + idx];- C[idx] = __float2bfloat16(sum);+ 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);}+ """- // ---- Launch functions for BLOCK_N=64 (4 warps, 256 threads) ----- void launch_n64_nosplit(- torch::Tensor A_bf16, torch::Tensor B_q, torch::Tensor B_scale,- torch::Tensor C, int M, int N, int K- ) {- const int k_steps = K / (MFMA_K * 2);- dim3 block(256);- dim3 grid((N + 63) / 64, (M + BLOCK_M - 1) / BLOCK_M, 1);- hipLaunchKernelGGL((gemm_kernel<64, 4>), grid, block, 0, 0,- reinterpret_cast<const __hip_bfloat16*>(A_bf16.data_ptr()),- reinterpret_cast<const uint8_t*>(B_q.data_ptr()),- reinterpret_cast<const uint8_t*>(B_scale.data_ptr()),- (float*)nullptr,- reinterpret_cast<__hip_bfloat16*>(C.data_ptr()), M, N, K, k_steps);- }+ 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- void launch_n64_splitk(- torch::Tensor A_bf16, torch::Tensor B_q, torch::Tensor B_scale,- torch::Tensor workspace, int M, int N, int K, int split_k- ) {- const int k_steps = K / (MFMA_K * 2);- dim3 block(256);- dim3 grid((N + 63) / 64, (M + BLOCK_M - 1) / BLOCK_M, split_k);- hipLaunchKernelGGL((gemm_kernel<64, 4>), grid, block, 0, 0,- reinterpret_cast<const __hip_bfloat16*>(A_bf16.data_ptr()),- reinterpret_cast<const uint8_t*>(B_q.data_ptr()),- reinterpret_cast<const uint8_t*>(B_scale.data_ptr()),- reinterpret_cast<float*>(workspace.data_ptr()),- (__hip_bfloat16*)nullptr, M, N, K, k_steps / split_k);- }- // ---- Launch functions for BLOCK_N=128 (8 warps, 512 threads) ----- void launch_n128_nosplit(- torch::Tensor A_bf16, torch::Tensor B_q, torch::Tensor B_scale,- torch::Tensor C, int M, int N, int K- ) {- const int k_steps = K / (MFMA_K * 2);- dim3 block(512);- dim3 grid((N + 127) / 128, (M + BLOCK_M - 1) / BLOCK_M, 1);- hipLaunchKernelGGL((gemm_kernel<128, 8>), grid, block, 0, 0,- reinterpret_cast<const __hip_bfloat16*>(A_bf16.data_ptr()),- reinterpret_cast<const uint8_t*>(B_q.data_ptr()),- reinterpret_cast<const uint8_t*>(B_scale.data_ptr()),- (float*)nullptr,- reinterpret_cast<__hip_bfloat16*>(C.data_ptr()), M, N, K, k_steps);- }+ # ===================== Triton kernels (v75) =====================- void launch_n128_splitk(- torch::Tensor A_bf16, torch::Tensor B_q, torch::Tensor B_scale,- torch::Tensor workspace, int M, int N, int K, int split_k- ) {- const int k_steps = K / (MFMA_K * 2);- dim3 block(512);- dim3 grid((N + 127) / 128, (M + BLOCK_M - 1) / BLOCK_M, split_k);- hipLaunchKernelGGL((gemm_kernel<128, 8>), grid, block, 0, 0,- reinterpret_cast<const __hip_bfloat16*>(A_bf16.data_ptr()),- reinterpret_cast<const uint8_t*>(B_q.data_ptr()),- reinterpret_cast<const uint8_t*>(B_scale.data_ptr()),- reinterpret_cast<float*>(workspace.data_ptr()),- (__hip_bfloat16*)nullptr, M, N, K, k_steps / split_k);- }+ @triton.jit+ def _remap_xcd(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- void launch_reduce(torch::Tensor workspace, torch::Tensor C, int M, int N, int split_k) {- const int num = M * N;- hipLaunchKernelGGL(reduce_kernel, dim3((num+255)/256), dim3(256), 0, 0,- reinterpret_cast<const float*>(workspace.data_ptr()),- reinterpret_cast<__hip_bfloat16*>(C.data_ptr()), M, N, split_k);- }- """- from torch.utils.cpp_extension import load_inline+ @triton.jit+ def _quant_fp4_hw(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- _module = None- def _get_module():- global _module- if _module is None:- _module = load_inline(- name="hybrid_v74",- cpp_sources=CPP_SOURCE,- cuda_sources=HIP_SOURCE,- functions=[- "launch_n64_nosplit", "launch_n64_splitk",- "launch_n128_nosplit", "launch_n128_splitk",- "launch_reduce",- ],- verbose=False,- extra_cuda_cflags=["-O3", "-fno-gpu-rdc", "-ffp-contract=fast"],- )- return _module+ # --------------- Split-K GEMM (M=16) ---------------- def _pick_split_k(m, n, k, block_n):- k_steps = k // 256- blocks_mn = ((n + block_n - 1) // block_n) * ((m + 15) // 16)- if blocks_mn >= 304:- return 1- # Target ~912 total blocks (304 CUs × 3 blocks/CU)- target_split = max(1, (912 + blocks_mn - 1) // blocks_mn)- best = 1- for s in range(1, k_steps + 1):- if k_steps % s == 0 and s <= target_split:- best = s- while best > 1 and k_steps // best < 2:- best //= 2- return max(1, best)+ @triton.jit+ def _fused_splitk_gemm_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, 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 = _remap_xcd(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 = _quant_fp4_hw(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 _merge_splitk_kernel(+ 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 GEMM (v75 configs) ---------------++ 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,+ )++ _FUSED_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),+ # --- v78: BN=64 for M=64 (grid 224→448) ---+ _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),+ # --- v78: BN=64 w/ BM=32/64 for M=256 ---+ _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 _prune_fused_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=_FUSED_CONFIGS, key=['M', 'N', 'K'],+ prune_configs_by={'early_config_prune': _prune_fused_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 _fused_gemm_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 = _remap_xcd(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 = _quant_fp4_hw(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")+++ # ===================== Host dispatch =====================++ NUM_CUS = 304+def custom_kernel(data: input_t) -> output_t:- A, B, B_q, B_shuffle, B_scale_sh = data- A = A.contiguous()- m, k = A.shape+ A_in, B, B_q, B_shuffle, B_scale_sh = data+ m, k = A_in.shapen = B.shape[0]+ dev = A_in.device- mod = _get_module()+ # HIP path for M<=32 K=512 (M=4, M=32 benchmark shapes)+ if m <= 32 and k == 512:+ C = torch.empty(m, n, dtype=torch.bfloat16, device=dev)+ _hip_dispatch(A_in.data_ptr(), B_shuffle.data_ptr(), B_scale_sh.data_ptr(),+ C.data_ptr(), m, n)+ return C- B_sh_u8 = B_shuffle.contiguous().view(torch.uint8)- B_sc = B_scale_sh.contiguous().view(torch.uint8)+ # Triton path for everything else (M=16 K=7168, M=64 K=2048, M=256 K=1536)+ B_sh = B_shuffle.contiguous().view(torch.uint8).view(n // 16, (k // 2) * 16)+ n_padded = (n + 255) // 256 * 256+ B_sc_raw = B_scale_sh.contiguous().view(torch.uint8).reshape(n_padded // 32, k)- # Dispatch: use BLOCK_N=128 for large-M benchmarks (M>=64)- use_n128 = (m >= 64)+ use_splitk = False+ if m <= 16 and k >= 1024:+ BM_sk, BN_sk, BK_sk = 16, 128, 512+ nk = k // BK_sk+ grid_mn = triton.cdiv(m, BM_sk) * triton.cdiv(n, BN_sk)+ if nk > 1:+ NUM_KSPLIT = min(nk, max(1, NUM_CUS // grid_mn))+ while nk % NUM_KSPLIT != 0 and NUM_KSPLIT > 1:+ NUM_KSPLIT -= 1+ if NUM_KSPLIT >= 2:+ use_splitk = True- if use_n128:- block_n = 128- split_k = _pick_split_k(m, n, k, block_n)- if split_k == 1:- C = torch.empty((m, n), dtype=torch.bfloat16, device="cuda")- mod.launch_n128_nosplit(A, B_sh_u8, B_sc, C, m, n, k)- else:- workspace = torch.empty((split_k, m, n), dtype=torch.float32, device="cuda")- mod.launch_n128_splitk(A, B_sh_u8, B_sc, workspace, m, n, k, split_k)- C = torch.empty((m, n), dtype=torch.bfloat16, device="cuda")- mod.launch_reduce(workspace, C, m, n, split_k)+ if use_splitk:+ total_elems = m * n+ C_parts = torch.empty(NUM_KSPLIT * total_elems, dtype=torch.float32, device=dev)+ C = torch.empty(m, n, dtype=torch.bfloat16, device=dev)+ _fused_splitk_gemm_kernel[(NUM_KSPLIT * grid_mn,)](+ A_in, B_sh, C_parts, 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=BM_sk, BLOCK_N=BN_sk, BLOCK_K=BK_sk, NUM_KSPLIT=NUM_KSPLIT,+ num_warps=4, num_stages=2,+ )+ MERGE_BLOCK = 256+ _merge_splitk_kernel[(triton.cdiv(total_elems, MERGE_BLOCK),)](+ C_parts, C, total_elems,+ NUM_KSPLIT=NUM_KSPLIT, BLOCK=MERGE_BLOCK,+ num_warps=4, num_stages=2,+ )else:- block_n = 64- split_k = _pick_split_k(m, n, k, block_n)- if split_k == 1:- C = torch.empty((m, n), dtype=torch.bfloat16, device="cuda")- mod.launch_n64_nosplit(A, B_sh_u8, B_sc, C, m, n, k)- else:- workspace = torch.empty((split_k, m, n), dtype=torch.float32, device="cuda")- mod.launch_n64_splitk(A, B_sh_u8, B_sc, workspace, m, n, k, split_k)- C = torch.empty((m, n), dtype=torch.bfloat16, device="cuda")- mod.launch_reduce(workspace, C, m, n, split_k)+ C = torch.empty(m, n, dtype=torch.bfloat16, device=dev)+ grid = lambda META: (triton.cdiv(m, META['BLOCK_M']) * triton.cdiv(n, META['BLOCK_N']),)+ _fused_gemm_kernel[grid](+ A_in, B_sh, C, 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),+ C.stride(0), C.stride(1),+ B_sc_raw.stride(0), B_sc_raw.stride(1),+ )return C
scrolls · 1030 diff lines total
Best evidence level for this revision: reported
JSON