KernelBench hard · RTX PRO 6000
W4A16 GEMM GLM-5.3
36.5%geomean peak fraction across shapes
manually audited: clean
Isolated regrade 0.3653 (in-run 0.3558). W4A16 unpack+GEMM: CUDA gemv1 for M==1 and Triton tl.dot for M>4. CUDAGraph is keyed on x.data_ptr; same-buffer overwrite on both paths reads live bytes. Lint CLEAN. Numeric stress on.
harnesszai-claudeagent session1h 35mtotal wall1h 43mcheck4sbenchmark2soutput tokens202,809cost$16.00regimememory
Per-shape vs governing ceilingeach shape graded against whichever binds — bf16 compute or HBM bandwidth
1×12288×40960.024 ms62.6%1.13 TB/s · 63% of 1.8 TB/s HBM · also 4 TFLOPS (1% of compute)
32×12288×40960.035 ms43.5%0.78 TB/s · 44% of 1.8 TB/s HBM · also 91 TFLOPS (18% of compute)
256×12288×40960.146 ms13.3%176 TFLOPS · 35% of 500 TF bf16 peak · also 0.24 TB/s (13% of HBM)
1×4096×40960.015 ms34.2%0.62 TB/s · 34% of 1.8 TB/s HBM · also 2 TFLOPS (0% of compute)
16×14336×40960.034 ms52.3%0.94 TB/s · 52% of 1.8 TB/s HBM · also 56 TFLOPS (11% of compute)
compute-bound memory-bound · bar + right column = official fraction of the ceiling (the geomean input)
geomean(62.6% · 43.5% · 13.3% · 34.2% · 52.3%) = 36.5%
Kernel source (redacted)
"""W4A16 GEMM (AWQ-style int4 weights, bf16 activations), fused unpack+dequant+GEMM.
M == 1: custom CUDA GEMV kernel (persistent, warp-per-column half-K-split for
small N), replayed through a CUDA graph to strip launch overhead.
M > 1: Triton tensor-core GEMM with fused dequant epilogue on the B operand.
"""
from __future__ import annotations
import os
os.environ.setdefault("CC", "clang")
os.environ.setdefault("CXX", "clang++")
import torch
import torch.nn as nn
import triton
import triton.language as tl
from torch.utils.cpp_extension import load_inline
# ---------------------------------------------------------------------------
# CUDA GEMV for M == 1, K == 4096
# ---------------------------------------------------------------------------
_CUDA_SRC = r"""
#include <torch/extension.h>
#include <cuda_bf16.h>
#include <cuda_fp16.h>
#include <ATen/cuda/CUDAContext.h>
// Persistent M==1, K==4096 kernel. Repacked layouts (built once at load):
// w4: (N, 128) uint4 chunks, chunk c of column n holds k in [32c, 32c+32)
// sz: (N, 32) bf16 (scale, zero) pairs, pair g covers k in [128g, 128g+128)
// Block = 256 threads = 8 warps. WPC warps cooperate on one column (K split),
// so small-N shapes still fill the GPU. fp16x2 SIMD dequant: bits 0x6400|q
// form half(1024+q), subtract (1024+z)*2^0 bias, scale, FMA against x pairs.
template <int WPC, int CPL, int MINB>
__global__ void __launch_bounds__(256, MINB)
gemv1_p(const __nv_bfloat16* __restrict__ x,
const uint4* __restrict__ w4,
const __nv_bfloat162* __restrict__ sz,
__nv_bfloat16* __restrict__ out,
int N)
{
constexpr int CHUNKS = 128; // 16B chunks per column (K=4096)
constexpr int WCHUNKS = CHUNKS / WPC;
constexpr int PASSES = WCHUNKS / (32 * CPL);
constexpr int COLS_PER_BLOCK = 8 / WPC;
const int warp = threadIdx.x >> 5;
const int lane = threadIdx.x & 31;
const int grp = warp / WPC;
const int sub = warp % WPC;
__shared__ float red[8];
const int col0 = blockIdx.x * COLS_PER_BLOCK + grp;
const int stride = gridDim.x * COLS_PER_BLOCK;
uint4 wreg[PASSES * CPL];
for (int colc = col0; colc < N; colc += stride) {
{
const uint4* wc = w4 + (size_t)colc * CHUNKS + sub * WCHUNKS + lane * CPL;
#pragma unroll
for (int p = 0; p < PASSES; p++)
#pragma unroll
for (int c = 0; c < CPL; c++) wreg[p * CPL + c] = __ldg(wc + p * 32 * CPL + c);
}
__half2 acc[2 * CPL];
#pragma unroll
for (int i = 0; i < 2 * CPL; i++) acc[i] = __half2half2(__ushort_as_half(0));
const uint4* xc = reinterpret_cast<const uint4*>(x);
const __nv_bfloat162* szc = sz + (size_t)colc * 32;
#pragma unroll
for (int pass = 0; pass < PASSES; pass++) {
const int cbase = sub * WCHUNKS + pass * 32 * CPL + lane * CPL;
const uint4* xp = xc + cbase * 4;
uint4 raw[4 * CPL];
#pragma unroll
for (int i = 0; i < 4 * CPL; i++) raw[i] = __ldg(xp + i);
const __nv_bfloat162* rb = reinterpret_cast<const __nv_bfloat162*>(raw);
__nv_bfloat162 szv = __ldg(szc + (cbase >> 2));
__half2 c2 = __float2half2_rn(1024.0f + __bfloat162float(szv.y));
__half2 s2 = __float2half2_rn(__bfloat162float(szv.x));
#pragma unroll
for (int c = 0; c < CPL; c++) {
const unsigned* wv = reinterpret_cast<const unsigned*>(&wreg[pass * CPL + c]);
#pragma unroll
for (int j = 0; j < 4; j++) {
unsigned word = wv[j];
unsigned lo4 = word & 0x0F0F0F0Fu;
unsigned hi4 = (word >> 4) & 0x0F0F0F0Fu;
__half2 xw[4];
#pragma unroll
for (int p = 0; p < 4; p++)
xw[p] = __float22half2_rn(__bfloat1622float2(rb[(4 * c + j) * 4 + p]));
#pragma unroll
for (int p = 0; p < 4; p++) {
unsigned l = (lo4 >> (8 * p)) & 0xFu;
unsigned h = (hi4 >> (8 * p)) & 0xFu;
__half2 u2 = __halves2half2(__ushort_as_half((unsigned short)(0x6400u | l)),
__ushort_as_half((unsigned short)(0x6400u | h)));
__half2 d2 = __hmul2(__hsub2(u2, c2), s2);
acc[2 * c + (j / 2)] = __hfma2(d2, xw[p], acc[2 * c + (j / 2)]);
}
}
}
}
__half2 s2 = __half2half2(__ushort_as_half(0));
#pragma unroll
for (int i = 0; i < 2 * CPL; i++) s2 = __hadd2(s2, acc[i]);
float s = __low2float(s2) + __high2float(s2);
#pragma unroll
for (int off = 16; off > 0; off >>= 1)
s += __shfl_down_sync(0xffffffffu, s, off);
if (WPC > 1) {
if (lane == 0) red[warp] = s;
__syncthreads();
if (sub == 0 && lane == 0) {
float t = 0.f;
#pragma unroll
for (int i = 0; i < WPC; i++) t += red[grp * WPC + i];
out[colc] = __float2bfloat16(t);
}
__syncthreads();
} else {
if (lane == 0) out[colc] = __float2bfloat16(s);
}
}
}
// Generic small-M fallback (M 1..4, K % 32 == 0); not on any benchmark path.
template <int MT>
__global__ void __launch_bounds__(256)
gemv_gen(const __nv_bfloat16* __restrict__ x,
const uint8_t* __restrict__ w4r,
const __nv_bfloat162* __restrict__ sz,
__nv_bfloat16* __restrict__ out,
int N, int K)
{
const int warp = threadIdx.x >> 5;
const int lane = threadIdx.x & 31;
const int col = blockIdx.x * 8 + warp;
const int ng = K / 128;
const bool col_ok = col < N;
__half2 acc[MT];
#pragma unroll
for (int m = 0; m < MT; m++) acc[m] = __half2half2(__ushort_as_half(0));
if (col_ok) {
const uint4* wcol = reinterpret_cast<const uint4*>(w4r + (size_t)col * (K / 2));
const __nv_bfloat162* szc = sz + (size_t)col * ng;
const int chunks = K / 32;
for (int c = lane; c < chunks; c += 32) {
uint4 w = __ldg(wcol + c);
__nv_bfloat162 szv = __ldg(szc + c / 4);
__half2 c2 = __float2half2_rn(1024.0f + __bfloat162float(szv.y));
__half2 s2 = __float2half2_rn(__bfloat162float(szv.x));
const unsigned wv[4] = {w.x, w.y, w.z, w.w};
const int kbase = c * 32;
#pragma unroll
for (int j = 0; j < 4; j++) {
unsigned word = wv[j];
unsigned lo4 = word & 0x0F0F0F0Fu;
unsigned hi4 = (word >> 4) & 0x0F0F0F0Fu;
#pragma unroll
for (int p = 0; p < 4; p++) {
unsigned l = (lo4 >> (8 * p)) & 0xFu;
unsigned h = (hi4 >> (8 * p)) & 0xFu;
__half2 u2 = __halves2half2(__ushort_as_half((unsigned short)(0x6400u | l)),
__ushort_as_half((unsigned short)(0x6400u | h)));
__half2 d2 = __hmul2(__hsub2(u2, c2), s2);
const int k = kbase + 8 * j + 2 * p;
#pragma unroll
for (int m = 0; m < MT; m++) {
__nv_bfloat162 xb = *reinterpret_cast<const __nv_bfloat162*>(x + (size_t)m * K + k);
__half2 xq = __float22half2_rn(__bfloat1622float2(xb));
acc[m] = __hfma2(d2, xq, acc[m]);
}
}
}
}
}
#pragma unroll
for (int m = 0; m < MT; m++) {
float s = __low2float(acc[m]) + __high2float(acc[m]);
#pragma unroll
for (int off = 16; off > 0; off >>= 1)
s += __shfl_down_sync(0xffffffffu, s, off);
if (lane == 0 && col_ok) out[(size_t)m * N + col] = __float2bfloat16(s);
}
}
static int g_max_blocks = 0;
void gemv1(torch::Tensor x, torch::Tensor w4r, torch::Tensor sz, torch::Tensor out)
{
const int N = (int)w4r.size(0);
auto stream = at::cuda::getCurrentCUDAStream();
if (g_max_blocks == 0) {
int dev; cudaGetDevice(&dev);
cudaDeviceProp prop; cudaGetDeviceProperties(&prop, dev);
int b1 = 0, b2 = 0;
cudaOccupancyMaxActiveBlocksPerMultiprocessor(&b1, gemv1_p<1, 1, 6>, 256, 0);
cudaOccupancyMaxActiveBlocksPerMultiprocessor(&b2, gemv1_p<2, 1, 5>, 256, 0);
int bpv = std::min(b1, b2);
if (bpv <= 0) bpv = 1;
g_max_blocks = prop.multiProcessorCount * bpv;
}
{
const int nblocks = std::min((N + 7) / 8, g_max_blocks);
gemv1_p<1, 1, 6><<<nblocks, 256, 0, stream>>>(
reinterpret_cast<const __nv_bfloat16*>(x.data_ptr()),
reinterpret_cast<const uint4*>(w4r.data_ptr()),
reinterpret_cast<const __nv_bfloat162*>(sz.data_ptr()),
reinterpret_cast<__nv_bfloat16*>(out.data_ptr()), N);
}
}
void gemv_gen_launch(torch::Tensor x, torch::Tensor w4r, torch::Tensor sz, torch::Tensor out)
{
const int M = (int)x.size(0);
const int K = (int)x.size(1);
const int N = (int)w4r.size(0);
auto stream = at::cuda::getCurrentCUDAStream();
dim3 grid((N + 7) / 8);
#define CASE(MT) \
gemv_gen<MT><<<grid, 256, 0, stream>>>( \
reinterpret_cast<const __nv_bfloat16*>(x.data_ptr()), \
reinterpret_cast<const uint8_t*>(w4r.data_ptr()), \
reinterpret_cast<const __nv_bfloat162*>(sz.data_ptr()), \
reinterpret_cast<__nv_bfloat16*>(out.data_ptr()), N, K)
switch (M) {
case 1: CASE(1); break;
case 2: CASE(2); break;
case 3: CASE(3); break;
case 4: CASE(4); break;
default: TORCH_CHECK(false, "M must be 1..4");
}
#undef CASE
}
"""
_CPP_SRC = r"""
void gemv1(torch::Tensor x, torch::Tensor w4r, torch::Tensor sz, torch::Tensor out);
void gemv_gen_launch(torch::Tensor x, torch::Tensor w4r, torch::Tensor sz, torch::Tensor out);
"""
_ext = load_inline(
name="w4gemv_sol_v1",
cpp_sources=_CPP_SRC,
cuda_sources=_CUDA_SRC,
functions=["gemv1", "gemv_gen_launch"],
extra_cuda_cflags=["-O3", "--use_fast_math", "-gencode=arch=compute_120a,code=sm_120a", "-ccbin=clang"],
verbose=False,
)
# ---------------------------------------------------------------------------
# Triton tensor-core GEMM with fused dequant (M > 1)
# ---------------------------------------------------------------------------
@triton.jit
def _w4a16_kernel(
x_ptr, wq_ptr, s_ptr, z_ptr, part_ptr,
M, N, K,
stride_xm, stride_partm, stride_partk,
BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr,
BLOCK_K: tl.constexpr, SPLIT_K: tl.constexpr, GROUP_K: tl.constexpr,
):
pid_n = tl.program_id(0)
pid_m = tl.program_id(1)
pid_k = tl.program_id(2)
offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
offs_kh = tl.arange(0, BLOCK_K // 2)
acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
k_per_split = K // SPLIT_K
k_start = pid_k * k_per_split
n_w = K // 2
for k0 in range(k_start, k_start + k_per_split, BLOCK_K):
offs_k = k0 + tl.arange(0, BLOCK_K)
w = tl.load(wq_ptr + offs_n[:, None] * n_w + (k0 // 2) + offs_kh[None, :])
wlo = (w & 0xF).to(tl.float32)
whi = (w >> 4).to(tl.float32)
wq2 = tl.interleave(wlo, whi)
g = k0 // GROUP_K
s = tl.load(s_ptr + g * N + offs_n).to(tl.float32)
z = tl.load(z_ptr + g * N + offs_n).to(tl.float32)
wdeq = (wq2 - z[:, None]) * s[:, None]
xb = tl.load(
x_ptr + offs_m[:, None] * stride_xm + offs_k[None, :],
mask=offs_m[:, None] < M, other=0.0,
)
acc += tl.dot(xb, tl.trans(wdeq.to(tl.bfloat16)), out_dtype=tl.float32)
part_ptr += pid_k * stride_partk + offs_m[:, None] * stride_partm + offs_n[None, :]
tl.store(part_ptr, acc, mask=offs_m[:, None] < M)
@triton.jit
def _reduce_kernel(
part_ptr, out_ptr, N,
stride_partm, stride_partk, stride_on,
SPLIT_K: tl.constexpr, BLOCK_N: tl.constexpr,
):
pid_n = tl.program_id(0)
pid_m = tl.program_id(1)
offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
acc = tl.zeros((BLOCK_N,), dtype=tl.float32)
for sk in tl.static_range(SPLIT_K):
acc += tl.load(part_ptr + sk * stride_partk + pid_m * stride_partm + offs_n)
tl.store(out_ptr + pid_m * stride_on + offs_n, acc.to(tl.bfloat16))
def pick_cfg(M, N, K):
if M <= 16:
return (16, 64, 128, 2, 4, 3)
return (32, 64, 128, 2, 4, 3)
# ---------------------------------------------------------------------------
# Model
# ---------------------------------------------------------------------------
class Model(nn.Module):
def __init__(self, M, N, K, group_size=128):
super().__init__()
self.M, self.N, self.K = M, N, K
self.group_size = group_size
self.register_buffer("w_q", torch.zeros(K // 2, N, dtype=torch.uint8))
self.register_buffer("scales", torch.zeros(K // group_size, N, dtype=torch.bfloat16))
self.register_buffer("zeros", torch.zeros(K // group_size, N, dtype=torch.bfloat16))
self.w4r = None
self.szv = None
self.out = None
self.part = None
self._g1 = None
self._g1_key = -1
self._g1_replay = None
self._gt = None
self._gt_key = -1
self._gt_replay = None
def _prep(self):
dev = self.w_q.device
# (K//2, N) -> (N, K//2): each column's packed bytes contiguous
self.w4r = self.w_q.t().contiguous()
ng = self.scales.shape[0]
self.szv = (
torch.stack([self.scales, self.zeros], dim=-1)
.transpose(0, 1)
.contiguous()
.view(torch.uint8)
) # (N, ng, 2) bf16 pairs
self.out = torch.empty(max(self.M, 4), self.N, dtype=torch.bfloat16, device=dev)
self.out1 = self.out[:1]
M = self.M
if M > 1:
BM, BN, BK, SK, nw, ns = pick_cfg(M, self.N, self.K)
self.part = torch.empty(SK, M, self.N, dtype=torch.float32, device=dev)
def _gemv1_graph(self, x):
"""M==1 fast path with CUDA-graph replay keyed on the input pointer.
On a new pointer we recapture (the graph bakes the input address);
a reused address holding a new tensor still replays correctly
because the kernel reads device memory at replay time.
"""
key = x.data_ptr()
if self._g1_key == key and self._g1_replay is not None:
self._g1_replay()
return
out1 = self.out1
try:
s = torch.cuda.Stream()
s.wait_stream(torch.cuda.current_stream())
with torch.cuda.stream(s):
_ext.gemv1(x, self.w4r, self.szv, out1)
torch.cuda.current_stream().wait_stream(s)
g = torch.cuda.CUDAGraph()
with torch.cuda.graph(g):
_ext.gemv1(x, self.w4r, self.szv, out1)
self._g1 = g
self._g1_key = key
self._g1_replay = g.replay
self._g1_replay()
return
except Exception:
self._g1 = None
self._g1_key = -1
self._g1_replay = None
_ext.gemv1(x, self.w4r, self.szv, out1)
def _triton_graph(self, x, M):
key = x.data_ptr()
cfg = pick_cfg(M, self.N, self.K)
if M != self.M:
# off-nominal batch: plain uncached path
BM, BN, BK, SK, nw, ns = cfg
part = torch.empty(SK, M, self.N, dtype=torch.float32, device=x.device)
out = torch.empty(M, self.N, dtype=torch.bfloat16, device=x.device)
self._triton_run(x, part, out, M, cfg)
return out
if self._gt_key == key and self._gt_replay is not None:
self._gt_replay()
return self.out[:M]
part, out = self.part, self.out
try:
s = torch.cuda.Stream()
s.wait_stream(torch.cuda.current_stream())
with torch.cuda.stream(s):
self._triton_run(x, part, out, M, cfg)
torch.cuda.current_stream().wait_stream(s)
g = torch.cuda.CUDAGraph()
with torch.cuda.graph(g):
self._triton_run(x, part, out, M, cfg)
self._gt = g
self._gt_key = key
self._gt_replay = g.replay
self._gt_replay()
return out[:M]
except Exception:
self._gt = None
self._gt_key = -1
self._gt_replay = None
self._triton_run(x, part, out, M, cfg)
return out[:M]
def _triton_run(self, x, part, out, M, cfg):
BM, BN, BK, SK, nw, ns = cfg
N, K = self.N, self.K
grid = (triton.cdiv(N, BN), triton.cdiv(M, BM), SK)
_w4a16_kernel[grid](
x, self.w4r, self.scales, self.zeros, part,
M, N, K, x.stride(0), part.stride(1), part.stride(0),
BLOCK_M=BM, BLOCK_N=BN, BLOCK_K=BK, SPLIT_K=SK, GROUP_K=self.group_size,
num_warps=nw, num_stages=ns,
)
grid2 = (triton.cdiv(N, 256), M)
_reduce_kernel[grid2](
part, out, N, part.stride(1), part.stride(0), out.stride(0),
SPLIT_K=SK, BLOCK_N=256, num_warps=1,
)
def forward(self, x):
if self.w4r is None:
self._prep()
M = x.shape[0]
if M == 1:
self._gemv1_graph(x)
return self.out[:1]
if M <= 4:
_ext.gemv_gen_launch(x, self.w4r, self.szv, self.out[:M].contiguous())
return self.out[:M]
return self._triton_graph(x, M)
# skip nn.Module.__call__ hook machinery on the hot path
__call__ = forward
M = 1
N = 12288
K = 4096
def get_inputs():
x = torch.randn(M, K, dtype=torch.bfloat16)
return [x]
def get_init_inputs():
return [M, N, K]
20260822_075321_zai-claude_glm-5.3_07_w4a16_gemm