KernelBench hard · RTX PRO 6000
W4A16 GEMM Grok 4.6
23.4%geomean peak fraction across shapes
manually audited: clean
Fused W4A16 Triton GEMM: unpack + dequant + matmul, AWQ/GPTQ group-128, weights stay packed on the HBM path. No graph cache. template_mutated=false.
harnessgrokagent session1h 15mtotal wall1h 17mcheck3sbenchmark3soutput tokens—regimememory
Per-shape vs governing ceilingeach shape graded against whichever binds — bf16 compute or HBM bandwidth
1×12288×40960.045 ms33.4%0.60 TB/s · 33% of 1.8 TB/s HBM · also 2 TFLOPS (0% of compute)
32×12288×40960.052 ms29.8%0.54 TB/s · 30% of 1.8 TB/s HBM · also 62 TFLOPS (12% of compute)
256×12288×40960.182 ms10.8%142 TFLOPS · 28% of 500 TF bf16 peak · also 0.19 TB/s (11% of HBM)
1×4096×40960.027 ms18.4%0.33 TB/s · 18% of 1.8 TB/s HBM · also 1 TFLOPS (0% of compute)
16×14336×40960.049 ms36.0%0.65 TB/s · 36% of 1.8 TB/s HBM · also 38 TFLOPS (8% of compute)
compute-bound memory-bound · bar + right column = official fraction of the ceiling (the geomean input)
geomean(33.4% · 29.8% · 10.8% · 18.4% · 36.0%) = 23.4%
Kernel source (redacted)
"""Fused W4A16 GEMM: unpack + dequant + matmul in one Triton kernel.
AWQ/GPTQ-style asymmetric int4, group size 128, bf16 activations.
Weights stay packed on the HBM path; nibbles expand in registers.
"""
from __future__ import annotations
import torch
import torch.nn as nn
import triton
import triton.language as tl
GROUP_SIZE = 128
@triton.jit
def _gemm_tma_kernel(
A, B, Scales, Zeros, C,
M, N, K,
stride_am, stride_ak, stride_bk, stride_bn,
stride_sk, stride_sn, stride_zk, stride_zn,
stride_cm, stride_cn,
GROUP_SIZE: tl.constexpr,
BLOCK_M: tl.constexpr,
BLOCK_N: tl.constexpr,
BLOCK_K: tl.constexpr,
):
pid_m = tl.program_id(0)
pid_n = tl.program_id(1)
a_ptr = tl.make_block_ptr(
A, shape=(M, K), strides=(stride_am, stride_ak),
offsets=(pid_m * BLOCK_M, 0), block_shape=(BLOCK_M, BLOCK_K), order=(1, 0),
)
b_ptr = tl.make_block_ptr(
B, shape=(K // 2, N), strides=(stride_bk, stride_bn),
offsets=(0, pid_n * BLOCK_N), block_shape=(BLOCK_K // 2, BLOCK_N), order=(1, 0),
)
offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
shifts = tl.arange(0, 2) * 4
acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
for k in range(0, K, BLOCK_K):
a = tl.load(a_ptr, boundary_check=(0, 1), padding_option="zero")
pk = tl.load(b_ptr, boundary_check=(0, 1), padding_option="zero")
nib = (pk[:, None, :].to(tl.int32) >> shifts[None, :, None]) & 0xF
b = tl.reshape(nib, (BLOCK_K, BLOCK_N)).to(tl.bfloat16)
gid = k // GROUP_SIZE
s = tl.load(Scales + gid * stride_sk + offs_n * stride_sn, mask=offs_n < N, other=1.0)
z = tl.load(Zeros + gid * stride_zk + offs_n * stride_zn, mask=offs_n < N, other=0.0)
b = (b - z[None, :].to(tl.bfloat16)) * s[None, :].to(tl.bfloat16)
acc = tl.dot(a.to(tl.bfloat16), b, acc, out_dtype=tl.float32)
a_ptr = tl.advance(a_ptr, (0, BLOCK_K))
b_ptr = tl.advance(b_ptr, (BLOCK_K // 2, 0))
offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
tl.store(
C + offs_m[:, None] * stride_cm + offs_n[None, :] * stride_cn,
acc.to(tl.bfloat16),
mask=(offs_m[:, None] < M) & (offs_n[None, :] < N),
)
@triton.jit
def _gemm_kernel(
A, B, Scales, Zeros, C,
M, N, K,
stride_am, stride_ak, stride_bk, stride_bn,
stride_sk, stride_sn, stride_zk, stride_zn,
stride_cm, stride_cn,
GROUP_SIZE: tl.constexpr,
BLOCK_M: tl.constexpr,
BLOCK_N: tl.constexpr,
BLOCK_K: tl.constexpr,
):
pid_m = tl.program_id(0)
pid_n = tl.program_id(1)
offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
offs_kp = tl.arange(0, BLOCK_K // 2)
shifts = tl.arange(0, 2) * 4
a_ptrs = A + offs_m[:, None] * stride_am + tl.arange(0, BLOCK_K)[None, :] * stride_ak
b_ptrs = B + offs_kp[:, None] * stride_bk + offs_n[None, :] * stride_bn
acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
for k in range(0, K, BLOCK_K):
kk = k + tl.arange(0, BLOCK_K)
a = tl.load(
a_ptrs,
mask=(offs_m[:, None] < M) & (kk[None, :] < K),
other=0.0,
eviction_policy="evict_last",
)
pk_row = k // 2 + offs_kp
pk = tl.load(
b_ptrs,
mask=(pk_row[:, None] < (K // 2)) & (offs_n[None, :] < N),
other=0,
eviction_policy="evict_first",
)
nib = (pk[:, None, :].to(tl.int32) >> shifts[None, :, None]) & 0xF
b = tl.reshape(nib, (BLOCK_K, BLOCK_N)).to(tl.bfloat16)
gid = k // GROUP_SIZE
s = tl.load(Scales + gid * stride_sk + offs_n * stride_sn, mask=offs_n < N, other=1.0)
z = tl.load(Zeros + gid * stride_zk + offs_n * stride_zn, mask=offs_n < N, other=0.0)
b = (b - z[None, :].to(tl.bfloat16)) * s[None, :].to(tl.bfloat16)
acc = tl.dot(a.to(tl.bfloat16), b, acc, out_dtype=tl.float32)
a_ptrs += BLOCK_K * stride_ak
b_ptrs += (BLOCK_K // 2) * stride_bk
tl.store(
C + offs_m[:, None] * stride_cm + offs_n[None, :] * stride_cn,
acc.to(tl.bfloat16),
mask=(offs_m[:, None] < M) & (offs_n[None, :] < N),
)
@triton.jit
def _splitk_kernel(
A, B, Scales, Zeros, Partials,
M, N, K,
stride_am, stride_ak, stride_bk, stride_bn,
stride_sk, stride_sn, stride_zk, stride_zn,
stride_pk, stride_pm, stride_pn,
GROUP_SIZE: tl.constexpr,
BLOCK_M: tl.constexpr,
BLOCK_N: tl.constexpr,
BLOCK_K: tl.constexpr,
SPLIT_K: tl.constexpr,
):
pid_m = tl.program_id(0)
pid_n = tl.program_id(1)
pid_k = tl.program_id(2)
offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
offs_kp = tl.arange(0, BLOCK_K // 2)
shifts = tl.arange(0, 2) * 4
a_ptrs = A + offs_m[:, None] * stride_am + (pid_k * BLOCK_K + tl.arange(0, BLOCK_K))[None, :] * stride_ak
b_ptrs = B + (pid_k * (BLOCK_K // 2) + offs_kp)[:, None] * stride_bk + offs_n[None, :] * stride_bn
acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
for k in range(pid_k * BLOCK_K, K, SPLIT_K * BLOCK_K):
kk = k + tl.arange(0, BLOCK_K)
a = tl.load(a_ptrs, mask=(offs_m[:, None] < M) & (kk[None, :] < K), other=0.0)
pk = tl.load(b_ptrs, mask=offs_n[None, :] < N, other=0)
nib = (pk[:, None, :].to(tl.int32) >> shifts[None, :, None]) & 0xF
b = tl.reshape(nib, (BLOCK_K, BLOCK_N)).to(tl.bfloat16)
gid = k // GROUP_SIZE
s = tl.load(Scales + gid * stride_sk + offs_n * stride_sn, mask=offs_n < N, other=1.0)
z = tl.load(Zeros + gid * stride_zk + offs_n * stride_zn, mask=offs_n < N, other=0.0)
b = (b - z[None, :].to(tl.bfloat16)) * s[None, :].to(tl.bfloat16)
acc = tl.dot(a.to(tl.bfloat16), b, acc, out_dtype=tl.float32)
a_ptrs += SPLIT_K * BLOCK_K * stride_ak
b_ptrs += SPLIT_K * (BLOCK_K // 2) * stride_bk
# Unique slot per split — no atomics, no pre-zero.
tl.store(
Partials + pid_k * stride_pk + offs_m[:, None] * stride_pm + offs_n[None, :] * stride_pn,
acc,
mask=(offs_m[:, None] < M) & (offs_n[None, :] < N),
)
@triton.jit
def _reduce_kernel(Partials, Out, M, N, SPLIT_K: tl.constexpr, BLOCK: tl.constexpr):
offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK)
mask = offs < (M * N)
acc = tl.zeros([BLOCK], dtype=tl.float32)
for sk in range(SPLIT_K):
acc += tl.load(Partials + sk * M * N + offs, mask=mask, other=0.0)
tl.store(Out + offs, acc.to(tl.bfloat16), mask=mask)
_WS: dict[tuple, torch.Tensor] = {}
def _workspace(sk: int, M: int, N: int, device) -> torch.Tensor:
key = (sk, M, N, device)
ws = _WS.get(key)
if ws is None:
ws = torch.empty((sk, M, N), device=device, dtype=torch.float32)
_WS[key] = ws
return ws
def _splitk_launch(x, w_q, scales, zeros, M, N, K, BM, BN, BK, SK, nw, ns):
ws = _workspace(SK, M, N, x.device)
_splitk_kernel[(triton.cdiv(M, BM), triton.cdiv(N, BN), SK)](
x, w_q, scales, zeros, ws, M, N, K,
x.stride(0), x.stride(1), w_q.stride(0), w_q.stride(1),
scales.stride(0), scales.stride(1), zeros.stride(0), zeros.stride(1),
ws.stride(0), ws.stride(1), ws.stride(2),
GROUP_SIZE=GROUP_SIZE, BLOCK_M=BM, BLOCK_N=BN, BLOCK_K=BK, SPLIT_K=SK,
num_warps=nw, num_stages=ns,
)
y = torch.empty((M, N), device=x.device, dtype=torch.bfloat16)
BLOCK = 256
_reduce_kernel[(triton.cdiv(M * N, BLOCK),)](
ws, y, M, N, SPLIT_K=SK, BLOCK=BLOCK, num_warps=4,
)
return y
def _w4a16(x: torch.Tensor, w_q: torch.Tensor, scales: torch.Tensor, zeros: torch.Tensor) -> torch.Tensor:
M, K = x.shape
N = w_q.shape[1]
x = x.contiguous()
if M <= 4:
sk = 16
return _splitk_launch(x, w_q, scales, zeros, M, N, K, 16, 128, 64, sk, 4, 3)
if M <= 16:
return _splitk_launch(x, w_q, scales, zeros, M, N, K, 16, 128, 64, 4, 4, 3)
if M <= 32:
return _splitk_launch(x, w_q, scales, zeros, M, N, K, 32, 128, 64, 4, 4, 3)
if M >= 128:
BM, BN, BK = 256, 128, 32
y = torch.empty((M, N), device=x.device, dtype=torch.bfloat16)
_gemm_tma_kernel[(triton.cdiv(M, BM), triton.cdiv(N, BN))](
x, w_q, scales, zeros, y, M, N, K,
x.stride(0), x.stride(1), w_q.stride(0), w_q.stride(1),
scales.stride(0), scales.stride(1), zeros.stride(0), zeros.stride(1),
y.stride(0), y.stride(1),
GROUP_SIZE=GROUP_SIZE, BLOCK_M=BM, BLOCK_N=BN, BLOCK_K=BK,
num_warps=8, num_stages=3,
)
return y
BM, BN, BK, nw, ns = 32, 128, 64, 4, 4
y = torch.empty((M, N), device=x.device, dtype=torch.bfloat16)
_gemm_kernel[(triton.cdiv(M, BM), triton.cdiv(N, BN))](
x, w_q, scales, zeros, y, M, N, K,
x.stride(0), x.stride(1), w_q.stride(0), w_q.stride(1),
scales.stride(0), scales.stride(1), zeros.stride(0), zeros.stride(1),
y.stride(0), y.stride(1),
GROUP_SIZE=GROUP_SIZE, BLOCK_M=BM, BLOCK_N=BN, BLOCK_K=BK,
num_warps=nw, num_stages=ns,
)
return y
class Model(nn.Module):
def __init__(self, M: int, N: int, K: int, group_size: int = GROUP_SIZE):
super().__init__()
assert K % group_size == 0
assert K % 2 == 0
self.M, self.N, self.K = M, N, K
self.group_size = group_size
n_groups = K // group_size
self.register_buffer("w_q", torch.empty(K // 2, N, dtype=torch.uint8))
self.register_buffer("scales", torch.empty(n_groups, N, dtype=torch.bfloat16))
self.register_buffer("zeros", torch.empty(n_groups, N, dtype=torch.bfloat16))
def forward(self, x: torch.Tensor) -> torch.Tensor:
return _w4a16(x.to(torch.bfloat16), self.w_q, self.scales, self.zeros)
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]
20260813_073002_grok_grok-4.6_07_w4a16_gemm