KernelBench hard · H100
W4A16 GEMM GPT-5.6 Sol
18.9%geomean peak fraction across shapes
manually audited: clean
harnesscodexagent session52mtotal wall55mcheck3mbenchmark14soutput tokens47,622gpu-lock wait16mgpu-lock held23mregimememory
Per-shape vs governing ceilingeach shape graded against whichever binds — bf16 compute or HBM bandwidth
1×12288×40960.054 ms24.4%0.50 TB/s · 24% of 2.0 TB/s HBM · also 2 TFLOPS (0% of compute)
32×12288×40960.063 ms21.6%0.44 TB/s · 22% of 2.0 TB/s HBM · also 51 TFLOPS (7% of compute)
256×12288×40960.138 ms12.4%186 TFLOPS · 25% of 756 TF bf16 peak · also 0.25 TB/s (12% of HBM)
1×4096×40960.030 ms14.8%0.30 TB/s · 15% of 2.0 TB/s HBM · also 1 TFLOPS (0% of compute)
16×14336×40960.063 ms24.8%0.50 TB/s · 25% of 2.0 TB/s HBM · also 30 TFLOPS (4% of compute)
compute-bound memory-bound · bar + right column = official fraction of the ceiling (the geomean input)
geomean(24.4% · 21.6% · 12.4% · 14.8% · 24.8%) = 18.9%
Kernel source (redacted)
"""Fused AWQ-style W4A16 GEMM for Hopper.
The packed matrix is consumed directly by the kernels below. In particular,
there is deliberately no persistent (or temporary) dequantized K x N matrix.
"""
from __future__ import annotations
import os
import sys
from pathlib import Path
import torch
import torch.nn as nn
import triton
import triton.language as tl
GROUP_SIZE = 128
_CUTLASS_EXT = {}
def _cutlass_ext(small_m: bool = False):
"""Build/load the SM90a mixed-input mainloop on first prefill call."""
key = 32 if small_m else 128
if key in _CUTLASS_EXT:
return _CUTLASS_EXT[key]
root = Path(__file__).resolve().parent
# The archive's tool wrappers are intentionally ahead of the toolkit in
# PATH, so cpp_extension needs the real toolkit root explicitly.
os.environ["CUDA_HOME"] = "/usr/local/cuda"
os.environ["TORCH_CUDA_ARCH_LIST"] = "9.0a"
os.environ.setdefault("MAX_JOBS", "4")
deps = root / ".deps"
if str(deps) not in sys.path:
sys.path.insert(0, str(deps))
os.environ["PATH"] = str(deps / "bin") + os.pathsep + os.environ.get("PATH", "")
from torch.utils.cpp_extension import load
cuda_flags = [
"-O3", "--use_fast_math", "-DCUTLASS_ENABLE_TENSOR_CORE_MMA=1",
]
if small_m:
cuda_flags.append("-DW4_TILE_N=32")
module = load(
name=("w4a16_cutlass_hopper_n32_v1" if small_m
else "w4a16_cutlass_hopper_v1"),
sources=[str(root / "cutlass_kernel.cu")],
extra_include_paths=[
str(root / "cutlass" / "include"),
str(root / "cutlass" / "tools" / "util" / "include"),
],
extra_cflags=["-O3"],
extra_cuda_cflags=cuda_flags,
verbose=False,
)
_CUTLASS_EXT[key] = module
return module
@triton.jit
def _w4a16_gemv(
x_ptr, q_ptr, scale_ptr, zero_ptr, out_ptr,
N: tl.constexpr, K: tl.constexpr,
BLOCK_N: tl.constexpr,
):
"""M=1 path. A program owns BLOCK_N output columns."""
cols = tl.program_id(0) * BLOCK_N + tl.arange(0, BLOCK_N)
valid_n = cols < N
packed_k = tl.arange(0, 64)
acc = tl.zeros((BLOCK_N,), tl.float32)
# One iteration is exactly one quantization group (128 K values / 64 B).
for group in range(0, K // 128):
k0 = group * 128
q = tl.load(
q_ptr + (k0 // 2 + packed_k[:, None]) * N + cols[None, :],
mask=valid_n[None, :], other=0,
)
scale = tl.load(scale_ptr + group * N + cols, mask=valid_n, other=0.0)
zero = tl.load(zero_ptr + group * N + cols, mask=valid_n, other=0.0)
lo = q & 0x0F
hi = q >> 4
# Match the reference's bf16 dequantized matrix before accumulating in
# fp32, just as a bf16 tensor-core dot does.
w_lo = ((lo.to(tl.bfloat16) - zero[None, :]) * scale[None, :]).to(tl.bfloat16)
w_hi = ((hi.to(tl.bfloat16) - zero[None, :]) * scale[None, :]).to(tl.bfloat16)
x_lo = tl.load(x_ptr + k0 + 2 * packed_k)
x_hi = tl.load(x_ptr + k0 + 2 * packed_k + 1)
products = (
x_lo[:, None].to(tl.float32) * w_lo.to(tl.float32)
+ x_hi[:, None].to(tl.float32) * w_hi.to(tl.float32)
)
acc += tl.sum(products, axis=0)
tl.store(out_ptr + cols, acc.to(tl.bfloat16), mask=valid_n)
@triton.jit
def _w4a16_gemm(
x_ptr, q_ptr, scale_ptr, zero_ptr, out_ptr,
M: tl.constexpr, N: tl.constexpr, K: tl.constexpr,
BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr,
):
"""Tensor-core path; each loop consumes one complete quant group."""
pid_m = tl.program_id(0)
pid_n = tl.program_id(1)
rows = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
cols = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
packed_k = tl.arange(0, 64)
valid_m = rows < M
valid_n = cols < N
acc = tl.zeros((BLOCK_M, BLOCK_N), tl.float32)
for group in range(0, K // 128):
k0 = group * 128
q = tl.load(
q_ptr + (k0 // 2 + packed_k[:, None]) * N + cols[None, :],
mask=valid_n[None, :], other=0,
)
scale = tl.load(scale_ptr + group * N + cols, mask=valid_n, other=0.0)
zero = tl.load(zero_ptr + group * N + cols, mask=valid_n, other=0.0)
w_lo = ((q & 0x0F).to(tl.bfloat16) - zero[None, :]) * scale[None, :]
w_hi = (q >> 4).to(tl.bfloat16) - zero[None, :]
w_lo = w_lo.to(tl.bfloat16)
w_hi = (w_hi * scale[None, :]).to(tl.bfloat16)
a_lo = tl.load(
x_ptr + rows[:, None] * K + k0 + 2 * packed_k[None, :],
mask=valid_m[:, None], other=0.0,
)
a_hi = tl.load(
x_ptr + rows[:, None] * K + k0 + 2 * packed_k[None, :] + 1,
mask=valid_m[:, None], other=0.0,
)
acc = tl.dot(a_lo, w_lo, acc)
acc = tl.dot(a_hi, w_hi, acc)
out_offsets = rows[:, None] * N + cols[None, :]
tl.store(out_ptr + out_offsets, acc.to(tl.bfloat16),
mask=valid_m[:, None] & valid_n[None, :])
def _launch(x: torch.Tensor, q: torch.Tensor, scales: torch.Tensor,
zeros: torch.Tensor, M: int, N: int, K: int) -> torch.Tensor:
out = torch.empty((M, N), dtype=torch.bfloat16, device=x.device)
if M == 1:
block_n = 32
_w4a16_gemv[(triton.cdiv(N, block_n),)](
x, q, scales, zeros, out,
N=N, K=K, BLOCK_N=block_n,
num_warps=4, num_stages=2,
)
else:
if M <= 32:
block_m, block_n, warps = 16 if M <= 16 else 32, 64, 8
else:
block_m, block_n, warps = 64, 64, 8
_w4a16_gemm[(triton.cdiv(M, block_m), triton.cdiv(N, block_n))](
x, q, scales, zeros, out,
M=M, N=N, K=K, BLOCK_M=block_m, BLOCK_N=block_n,
num_warps=warps, num_stages=2,
)
return out
class Model(nn.Module):
def __init__(self, M: int, N: int, K: int, group_size: int = GROUP_SIZE):
super().__init__()
assert group_size == GROUP_SIZE
assert K % GROUP_SIZE == 0 and K % 2 == 0
self.M, self.N, self.K = M, N, K
self.group_size = group_size
self.register_buffer("w_q", torch.empty((K // 2, N), dtype=torch.uint8))
self.register_buffer("scales", torch.empty((K // GROUP_SIZE, N), dtype=torch.bfloat16))
self.register_buffer("zeros", torch.empty((K // GROUP_SIZE, N), dtype=torch.bfloat16))
# CUTLASS's narrow operand is K-major. These are lazily prepared once
# after state_dict loading and deliberately remain non-persistent.
self._q_kmajor = None
self._zero_bias = None
def forward(self, x: torch.Tensor) -> torch.Tensor:
if self.M == 1:
out = torch.empty((1, self.N), dtype=torch.bfloat16, device=x.device)
_cutlass_ext(small_m=True).gemv(
x, self.w_q, self.scales, self.zeros, out, self.N, self.K,
)
return out
if self._q_kmajor is None or self._q_kmajor.device != self.w_q.device:
# Offline weight-only packing: K becomes contiguous, but the data
# remains packed uint4 (no dequantized matrix is cached).
self._q_kmajor = self.w_q.t().contiguous()
# The customized CUTLASS register transform consumes -zero and
# evaluates (q + (-zero))*scale in the reference's bf16 order.
self._zero_bias = -self.zeros
out = torch.empty((self.M, self.N), dtype=torch.bfloat16, device=x.device)
_cutlass_ext(small_m=self.M <= 32).w4a16(
x, self._q_kmajor, self.scales, self._zero_bias, out,
self.M, self.N, self.K,
)
return out
M = 1
N = 12288
K = 4096
def get_inputs():
return [torch.randn(M, K, dtype=torch.bfloat16)]
def get_init_inputs():
return [M, N, K]
# ==================================================================
# ===== sidecar: cutlass_kernel.cu (7848 bytes, loaded by solution.py) =====
# ==================================================================
#include <torch/extension.h>
#include <ATen/cuda/CUDAContext.h>
#include "cutlass/cutlass.h"
#include "cute/tensor.hpp"
#include "cutlass/epilogue/collective/collective_builder.hpp"
#include "cutlass/gemm/collective/collective_builder.hpp"
#include "cutlass/gemm/device/gemm_universal_adapter.h"
#include "cutlass/gemm/kernel/gemm_universal.hpp"
#include "cutlass/util/packed_stride.hpp"
using namespace cute;
using Wide = cutlass::bfloat16_t;
using Narrow = cutlass::uint4b_t;
using ElementScale = Wide;
using ElementZero = Wide;
using LayoutA = cutlass::layout::RowMajor;
using LayoutB = cutlass::layout::ColumnMajor;
using LayoutC = cutlass::layout::RowMajor;
using LayoutD = cutlass::layout::RowMajor;
using LayoutAT = typename cutlass::layout::LayoutTranspose<LayoutA>::type;
using LayoutBT = typename cutlass::layout::LayoutTranspose<LayoutB>::type;
constexpr int AlignA = 8;
constexpr int AlignB = 32;
constexpr int AlignC = 8;
constexpr int AlignD = 8;
using Acc = float;
using Arch = cutlass::arch::Sm90;
using OpClass = cutlass::arch::OpClassTensorOp;
#ifndef W4_TILE_M
#define W4_TILE_M 128
#endif
#ifndef W4_TILE_N
#define W4_TILE_N 128
#endif
using TileShapeMNK = Shape<Int<W4_TILE_M>, Int<W4_TILE_N>, _64>;
using Cluster = Shape<_1, _1, _1>;
using MainSchedule = cutlass::gemm::KernelTmaWarpSpecializedCooperative;
using EpiSchedule = cutlass::epilogue::TmaWarpSpecializedCooperative;
using Epi = typename cutlass::epilogue::collective::CollectiveBuilder<
Arch, OpClass, TileShapeMNK, Cluster,
cutlass::epilogue::collective::EpilogueTileAuto,
Acc, Acc,
void, typename cutlass::layout::LayoutTranspose<LayoutC>::type, AlignC,
Wide, typename cutlass::layout::LayoutTranspose<LayoutD>::type, AlignD,
EpiSchedule>::CollectiveOp;
using Mainloop = typename cutlass::gemm::collective::CollectiveBuilder<
Arch, OpClass,
cute::tuple<Narrow, ElementScale, ElementZero>, LayoutBT, AlignB,
Wide, LayoutAT, AlignA,
Acc, TileShapeMNK, Cluster,
cutlass::gemm::collective::StageCountAutoCarveout<
static_cast<int>(sizeof(typename Epi::SharedStorage))>,
MainSchedule>::CollectiveOp;
using Kernel = cutlass::gemm::kernel::GemmUniversal<
Shape<int, int, int, int>, Mainloop, Epi>;
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<Kernel>;
using StrideA = cutlass::detail::TagToStrideA_t<LayoutA>;
using StrideB = cutlass::detail::TagToStrideB_t<LayoutB>;
using StrideS = typename Mainloop::StrideScale;
using StrideC = typename Kernel::StrideC;
using StrideD = typename Kernel::StrideD;
// Decode kernel: a CTA owns 32 adjacent output columns. Its 32 warps each
// consume one complete 128-K quantization group, preserving fully coalesced
// byte loads from the reference's native (K/2, N) packing. Warp 0 performs
// the final cross-group reduction.
__global__ __launch_bounds__(1024, 1)
void w4a16_gemv_kernel(
__nv_bfloat16 const* __restrict__ x,
uint8_t const* __restrict__ q,
__nv_bfloat16 const* __restrict__ scales,
__nv_bfloat16 const* __restrict__ zeros,
__nv_bfloat16* __restrict__ out,
int n_extent) {
int lane = int(threadIdx.x) & 31;
int warp = int(threadIdx.x) >> 5;
int n = int(blockIdx.x) * 32 + lane;
int group = warp;
float accum = 0.0f;
if (n < n_extent) {
__nv_bfloat16 scale = scales[group * n_extent + n];
__nv_bfloat16 zero = zeros[group * n_extent + n];
__nv_bfloat162 scale2 = __bfloat162bfloat162(scale);
__nv_bfloat162 zero2 = __bfloat162bfloat162(zero);
uint32_t magic_bits = 0x43004300u; // bf16x2 {128, 128}
__nv_bfloat162 magic = *reinterpret_cast<__nv_bfloat162*>(&magic_bits);
int packed_base = group * 64;
int k_base = group * 128;
#pragma unroll
for (int p = 0; p < 64; ++p) {
uint8_t packed = q[(packed_base + p) * n_extent + n];
// Form bf16x2 {128+lo, 128+hi} directly in integer registers,
// subtract the magic bias, then apply the exact bf16 affine order.
uint32_t q_bits = 0x43004300u |
uint32_t(packed & 15) | (uint32_t(packed & 0xf0) << 12);
__nv_bfloat162 q_magic = *reinterpret_cast<__nv_bfloat162*>(&q_bits);
__nv_bfloat162 q_pair = __hsub2(q_magic, magic);
__nv_bfloat162 w_pair = __hmul2(__hsub2(q_pair, zero2), scale2);
__nv_bfloat162 x_pair = reinterpret_cast<__nv_bfloat162 const*>(x)[
k_base / 2 + p];
float2 wf = __bfloat1622float2(w_pair);
float2 xf = __bfloat1622float2(x_pair);
accum = fmaf(xf.x, wf.x, accum);
accum = fmaf(xf.y, wf.y, accum);
}
}
__shared__ float partial[32][32];
partial[warp][lane] = accum;
__syncthreads();
if (warp == 0 && n < n_extent) {
float total = 0.0f;
#pragma unroll
for (int g = 0; g < 32; ++g) {
total += partial[g][lane];
}
out[n] = __float2bfloat16_rn(total);
}
}
void w4a16_gemv(torch::Tensor x, torch::Tensor q,
torch::Tensor scales, torch::Tensor zeros,
torch::Tensor out, int64_t n, int64_t k) {
TORCH_CHECK(k == 4096, "decode kernel requires K=4096");
cudaStream_t stream = at::cuda::getCurrentCUDAStream();
w4a16_gemv_kernel<<<(int(n) + 31) / 32, 1024, 0, stream>>>(
reinterpret_cast<__nv_bfloat16 const*>(x.data_ptr()),
reinterpret_cast<uint8_t const*>(q.data_ptr()),
reinterpret_cast<__nv_bfloat16 const*>(scales.data_ptr()),
reinterpret_cast<__nv_bfloat16 const*>(zeros.data_ptr()),
reinterpret_cast<__nv_bfloat16*>(out.data_ptr()), int(n));
}
void w4a16_cutlass(torch::Tensor x, torch::Tensor q_transposed,
torch::Tensor scales, torch::Tensor zero_bias,
torch::Tensor out, int64_t m, int64_t n, int64_t k) {
TORCH_CHECK(x.is_cuda() && q_transposed.is_cuda() && scales.is_cuda() &&
zero_bias.is_cuda() && out.is_cuda(), "all tensors must be CUDA");
TORCH_CHECK(q_transposed.is_contiguous(), "packed transposed weight must be contiguous");
auto stride_a = cutlass::make_cute_packed_stride(
StrideA{}, cute::make_shape(int(m), int(k), 1));
auto stride_b = cutlass::make_cute_packed_stride(
StrideB{}, cute::make_shape(int(n), int(k), 1));
auto stride_s = cutlass::make_cute_packed_stride(
StrideS{}, cute::make_shape(int(n), int(k / 128), 1));
auto stride_c = cutlass::make_cute_packed_stride(
StrideC{}, cute::make_shape(int(n), int(m), 1));
auto stride_d = cutlass::make_cute_packed_stride(
StrideD{}, cute::make_shape(int(n), int(m), 1));
typename Gemm::Arguments args{
cutlass::gemm::GemmUniversalMode::kGemm,
{int(n), int(m), int(k), 1},
{
reinterpret_cast<Narrow const*>(q_transposed.data_ptr()), stride_b,
reinterpret_cast<Wide const*>(x.data_ptr()), stride_a,
reinterpret_cast<ElementScale const*>(scales.data_ptr()), stride_s,
128,
reinterpret_cast<ElementZero const*>(zero_bias.data_ptr())
},
{
{1.0f, 0.0f},
nullptr, stride_c,
reinterpret_cast<Wide*>(out.data_ptr()), stride_d
}
};
Gemm gemm;
auto status = gemm.can_implement(args);
TORCH_CHECK(status == cutlass::Status::kSuccess,
"CUTLASS cannot implement problem, status=", int(status));
size_t workspace_size = Gemm::get_workspace_size(args);
TORCH_CHECK(workspace_size == 0, "unexpected CUTLASS workspace requirement");
cudaStream_t stream = at::cuda::getCurrentCUDAStream();
status = gemm.initialize(args, nullptr, stream);
TORCH_CHECK(status == cutlass::Status::kSuccess,
"CUTLASS initialize failed, status=", int(status));
status = gemm.run(stream);
TORCH_CHECK(status == cutlass::Status::kSuccess,
"CUTLASS launch failed, status=", int(status));
}
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.def("w4a16", &w4a16_cutlass);
m.def("gemv", &w4a16_gemv);
}
20260721_150450_codex_gpt-5.6-sol_07_w4a16_gemm