code/day78-blackwell/wgmma_gemm.cuThis is the source used by the lesson and its recorded evidence. Compile commands and expected output live in the directory README.
// SPDX-License-Identifier: MIT
//
// Day 78: a Hopper GEMM to annotate, built on wgmma.mma_async.
//
// One CTA of 256 threads is two warpgroups with different jobs:
//
// warpgroup 0 (threads 0..127) consumer. Issues the asynchronous
// wgmma.mma_async instructions and owns
// the 64x64 accumulator in registers.
// warpgroup 1 (threads 128..255) producer. Copies the next K-tile of A
// and B from global into shared memory
// while the consumer computes on the
// current one.
//
// The exercise on the page is to annotate this file: which warpgroup
// issues, which instruction waits, and what each barrier separates. The
// quiz in content/quizzes/day78.toml is answerable from this listing and
// the PTX ISA, without running anything.
//
// This program runs on exactly one architecture. wgmma requires the
// architecture-specific target sm_90a ("Requires sm_90a", PTX ISA 9.7.16,
// https://docs.nvidia.com/cuda/parallel-thread-execution/index.html
// #asynchronous-warpgroup-level-matrix-instructions-wgmma-mma), sm_90a
// binaries carry no forward-portable PTX, and no Blackwell card, consumer
// or datacenter, implements wgmma. The capability gate below says so
// instead of dying inside the driver.
//
// Build: nvcc -std=c++17 -O3 -arch=sm_90a -o wgmma_gemm wgmma_gemm.cu
// Run: ./wgmma_gemm
//
// UNVERIFIED: not yet run on real hardware; the project's Tesla T4
// verification node is CC 7.5 and cannot load this binary. Do not publish
// any output as this program's output until a Hopper run exists. See
// research/REVIEW-PROCESS.md.
#include <cmath>
#include <cstdint>
#include <cstdio>
#include <cstdlib>
#include <vector>
#include <cuda_fp16.h>
#include <cuda_runtime.h>
// The one error macro. This file is standalone, the way a Compiler Explorer
// embed is, so it carries its own verbatim copy. `err_` carries a trailing
// underscore so it cannot collide with a variable at the call site, and the
// do/while makes the macro one statement so it survives a braceless `if`.
#define CUDA_CHECK(call) \
do { \
cudaError_t err_ = (call); \
if (err_ != cudaSuccess) { \
std::fprintf(stderr, "CUDA error %s:%d: %s: %s\n", __FILE__, \
__LINE__, #call, cudaGetErrorString(err_)); \
std::exit(EXIT_FAILURE); \
} \
} while (0)
// Square problem, one size. 256 is 4 K-tiles of 64: enough iterations for
// the producer/consumer pipeline to be real, small enough that the CPU
// reference in double is instant.
constexpr int kDim = 256;
constexpr int kTileM = 64; // rows of D one CTA owns
constexpr int kTileN = 64; // columns of D one CTA owns
constexpr int kTileK = 64; // K depth staged in shared memory per tile
constexpr int kWgmmaK = 16; // K depth of one wgmma.mma_async.m64n64k16
constexpr int kThreadsPerBlock = 256; // two warpgroups of 128
constexpr int kAccRegs = 32; // 64x64 floats / 128 threads
// Inputs are drawn from {-8..8}/8, all exactly representable in FP16, so
// the only rounding in the whole computation is the FP32 accumulation
// order inside the tensor core, which the PTX ISA leaves unspecified.
// K = 256 accumulation steps at FP32 epsilon leaves relative error orders
// of magnitude below 1e-3 for these inputs; day 66 is the tolerance
// discipline this follows.
constexpr float kRelTolerance = 1e-3f;
static_assert(kThreadsPerBlock == 256,
"two warpgroups of 128; the role split below assumes it");
static_assert(kDim % kTileM == 0 && kDim % kTileN == 0 && kDim % kTileK == 0,
"one size, chosen to tile evenly; no ragged edge on this day");
// Shared memory tiles live in the canonical K-major no-swizzle layout the
// matrix descriptor describes (PTX ISA "Shared Memory Matrix Layouts"):
// 8x8 FP16 core matrices, each 128 contiguous bytes, cores consecutive
// along K then along M (or N for the B tile). Element (mn, k) of a 64x64
// tile lands at this element offset.
__device__ __forceinline__ int coreOffset(int mn, int k) {
return ((mn / 8) * (kTileK / 8) + k / 8) * 64 + (mn % 8) * 8 + (k % 8);
}
// snippet: descriptor
// A wgmma matrix descriptor is a 64-bit value, not a pointer: bits 13..0
// encode the shared memory start address, 29..16 the leading dimension
// byte offset (core to core along K: 128 bytes here), 45..32 the stride
// dimension byte offset (8 rows to the next 8 rows: 1024 bytes here),
// 63..62 the swizzle mode (0, none). Encode is (x & 0x3FFFF) >> 4, so one
// descriptor + 16 addresses the tile 256 bytes further on. PTX ISA
// "Matrix Descriptor Format"; the Hopper run's harness pass is what
// certifies these two constants.
__device__ __forceinline__ uint64_t tileDesc(const __half* p) {
const uint32_t addr = static_cast<uint32_t>(__cvta_generic_to_shared(p));
return ((addr & 0x3FFFFu) >> 4) | (uint64_t(128 >> 4) << 16) |
(uint64_t(1024 >> 4) << 32);
}
// end snippet
// One m64n64k16 wgmma, both operands from shared memory. The .sync.aligned
// qualifiers are mandatory: all 128 threads of the warpgroup issue this
// instruction together, and it computes with their combined registers.
// scale-d = 1 keeps D += A*B; trans-a = trans-b = 0 because A is stored
// row-major and B column-major, exactly what the instruction defines.
__device__ __forceinline__ void wgmmaM64n64k16(float acc[kAccRegs],
uint64_t descA, uint64_t descB) {
asm volatile(
"wgmma.mma_async.sync.aligned.m64n64k16.f32.f16.f16\n\t"
"{%0,%1,%2,%3,%4,%5,%6,%7,"
"%8,%9,%10,%11,%12,%13,%14,%15,"
"%16,%17,%18,%19,%20,%21,%22,%23,"
"%24,%25,%26,%27,%28,%29,%30,%31},\n\t"
"%32, %33, 1, 1, 1, 0, 0;"
: "+f"(acc[0]), "+f"(acc[1]), "+f"(acc[2]), "+f"(acc[3]), "+f"(acc[4]),
"+f"(acc[5]), "+f"(acc[6]), "+f"(acc[7]), "+f"(acc[8]), "+f"(acc[9]),
"+f"(acc[10]), "+f"(acc[11]), "+f"(acc[12]), "+f"(acc[13]),
"+f"(acc[14]), "+f"(acc[15]), "+f"(acc[16]), "+f"(acc[17]),
"+f"(acc[18]), "+f"(acc[19]), "+f"(acc[20]), "+f"(acc[21]),
"+f"(acc[22]), "+f"(acc[23]), "+f"(acc[24]), "+f"(acc[25]),
"+f"(acc[26]), "+f"(acc[27]), "+f"(acc[28]), "+f"(acc[29]),
"+f"(acc[30]), "+f"(acc[31])
: "l"(descA), "l"(descB));
}
// D = A * B in FP16 with FP32 accumulate. One CTA owns a 64x64 tile of D.
//
// One thread of the consumer warpgroup ends the kernel holding 32 floats
// of D; one thread of the producer warpgroup copies 32 elements of A and
// 32 of B per K-tile. Producer reads of global memory are coalesced along
// K (consecutive threads read consecutive addresses in each 64-element
// row); its shared stores scatter into the core-matrix layout, which
// shared memory tolerates and the descriptor requires.
//
// Launch assumption: gridDim = (kDim/kTileN, kDim/kTileM), blockDim = 256.
__global__ void wgmmaGemm(const __half* __restrict__ a,
const __half* __restrict__ b, float* __restrict__ d) {
// Double-buffered stages so the producer can fill one while the
// consumer drains the other. 2 x (8 KiB + 8 KiB) = 32 KiB, well under
// Hopper's 227 KiB per-block ceiling; the page says why the same
// budget question kills bigger pipelines on a 99 KiB consumer card.
__shared__ __half smA[2][kTileM * kTileK];
__shared__ __half smB[2][kTileN * kTileK];
const int wg = static_cast<int>(threadIdx.x) / 128; // 0 eats, 1 feeds
const int t = static_cast<int>(threadIdx.x) % 128;
const int rowBase = static_cast<int>(blockIdx.y) * kTileM;
const int colBase = static_cast<int>(blockIdx.x) * kTileN;
constexpr int kTiles = kDim / kTileK;
float acc[kAccRegs] = {}; // consumer's 64x64 quarter, all zeros
// The producer's copy of one K-tile into stage s. 64x64 elements over
// 128 threads is 32 per thread per matrix. A is row-major MxK, B is
// row-major KxN in global memory; the B store transposes it into the
// column-major (K-contiguous) layout wgmma defines for B.
auto loadStage = [&](int s, int kt) {
for (int e = 0; e < (kTileM * kTileK) / 128; ++e) {
const int idx = t + e * 128;
const int mn = idx / kTileK;
const int k = idx % kTileK;
smA[s][coreOffset(mn, k)] =
a[(rowBase + mn) * kDim + kt * kTileK + k];
smB[s][coreOffset(mn, k)] =
b[(kt * kTileK + k) * kDim + colBase + mn];
}
// The stores above went through the generic proxy; wgmma reads
// shared memory through the async proxy. This fence orders the
// two, and the __syncthreads() after it carries the ordering to
// the consumer warpgroup.
asm volatile("fence.proxy.async.shared::cta;");
};
// snippet: pipeline
if (wg == 1) {
loadStage(0, 0); // prologue: stage 0 filled before anyone computes
}
__syncthreads();
for (int kt = 0; kt < kTiles; ++kt) {
const int s = kt & 1;
if (wg == 1) {
// Producer: fill the other stage while the consumer eats
// this one. Nothing here waits on the tensor cores.
if (kt + 1 < kTiles) {
loadStage(s ^ 1, kt + 1);
}
} else {
// Consumer: 4 wgmma issues walk the 64-deep tile in k16
// steps. wgmma.fence first (register accesses, mandatory),
// then the issues, then one commit_group, then wait_group 0,
// which blocks until every committed wgmma has finished
// reading shared memory and writing acc.
asm volatile("wgmma.fence.sync.aligned;");
for (int ks = 0; ks < kTileK / kWgmmaK; ++ks) {
const uint64_t off = uint64_t(ks) * (2 * 128 >> 4);
wgmmaM64n64k16(acc, tileDesc(smA[s]) + off,
tileDesc(smB[s]) + off);
}
asm volatile("wgmma.commit_group.sync.aligned;");
asm volatile("wgmma.wait_group.sync.aligned 0;");
}
// Both warpgroups: the consumer is done reading stage s and the
// producer is done writing stage s^1, so the next iteration may
// swap them. Move the wait_group after this barrier and the
// producer can overwrite a tile the tensor cores are still
// reading: a race with no error message.
__syncthreads();
}
// end snippet
if (wg == 0) {
// Accumulator layout per PTX ISA Figure 149 ("WGMMA .m64nNk16
// register fragment layout for accumulator matrix D"): warp w
// owns rows 16w..16w+15; regs come in groups of four per
// 8-column block, the same m16n8 pattern day 73 stored. The
// harness pass on Hopper is what certifies this mapping.
const int warp = t / 32;
const int lane = t % 32;
for (int r = 0; r < kAccRegs; ++r) {
const int row = rowBase + 16 * warp + lane / 4 + (r % 4 / 2) * 8;
const int col = colBase + 8 * (r / 4) + 2 * (lane % 4) + r % 2;
d[row * kDim + col] = acc[r];
}
}
}
// Plain triple loop in double. The reference's job is to be right, not to
// match bit for bit; kRelTolerance covers the FP32 accumulation order.
static void gemmCpu(const std::vector<__half>& a, const std::vector<__half>& b,
std::vector<double>& d) {
for (int m = 0; m < kDim; ++m) {
for (int n = 0; n < kDim; ++n) {
double sum = 0.0;
for (int k = 0; k < kDim; ++k) {
sum += static_cast<double>(__half2float(a[m * kDim + k])) *
static_cast<double>(__half2float(b[k * kDim + n]));
}
d[m * kDim + n] = sum;
}
}
}
int main() {
// snippet: gate
// wgmma exists on sm_90a and nowhere else, and an sm_90a binary
// carries no PTX a newer card could JIT. Every non-Hopper card,
// including all of Blackwell, fails here by design; the page's table
// says what each card lacks.
cudaDeviceProp prop;
CUDA_CHECK(cudaGetDeviceProperties(&prop, 0));
if (prop.major != 9 || prop.minor != 0) {
std::fprintf(stderr,
"%s is CC %d.%d. This program needs CC 9.0 (Hopper): "
"wgmma is sm_90a-only and no Blackwell card has it.\n",
prop.name, prop.major, prop.minor);
return EXIT_FAILURE;
}
// end snippet
// Inputs from a fixed LCG, mapped to {-8..8}/8: exactly representable
// in FP16, so the device/reference difference is accumulation order
// alone.
const size_t elems = static_cast<size_t>(kDim) * kDim;
std::vector<__half> h_a(elems);
std::vector<__half> h_b(elems);
uint32_t state = 20260901u;
for (size_t i = 0; i < 2 * elems; ++i) {
state = state * 1664525u + 1013904223u;
const float v =
static_cast<float>(static_cast<int>(state >> 24) % 17 - 8) / 8.0f;
if (i < elems) {
h_a[i] = __float2half(v);
} else {
h_b[i - elems] = __float2half(v);
}
}
__half* d_a = nullptr;
__half* d_b = nullptr;
float* d_d = nullptr;
CUDA_CHECK(cudaMalloc(&d_a, elems * sizeof(__half)));
CUDA_CHECK(cudaMalloc(&d_b, elems * sizeof(__half)));
CUDA_CHECK(cudaMalloc(&d_d, elems * sizeof(float)));
CUDA_CHECK(cudaMemcpy(d_a, h_a.data(), elems * sizeof(__half),
cudaMemcpyHostToDevice));
CUDA_CHECK(cudaMemcpy(d_b, h_b.data(), elems * sizeof(__half),
cudaMemcpyHostToDevice));
const dim3 grid(kDim / kTileN, kDim / kTileM);
wgmmaGemm<<<grid, kThreadsPerBlock>>>(d_a, d_b, d_d);
CUDA_CHECK(cudaGetLastError());
CUDA_CHECK(cudaDeviceSynchronize());
std::vector<float> h_d(elems);
CUDA_CHECK(cudaMemcpy(h_d.data(), d_d, elems * sizeof(float),
cudaMemcpyDeviceToHost));
std::vector<double> h_ref(elems);
gemmCpu(h_a, h_b, h_ref);
// Correctness gate: a real branch, never assert(). First mismatch is
// named so a wrong descriptor constant or fragment index is
// debuggable from the transcript alone.
bool ok = true;
float worst = 0.0f;
for (size_t i = 0; i < elems; ++i) {
const float ref = static_cast<float>(h_ref[i]);
const float rel =
std::fabs(h_d[i] - ref) / std::fmax(std::fabs(ref), 1.0f);
worst = std::fmax(worst, rel);
if (rel > kRelTolerance && ok) {
ok = false;
std::fprintf(stderr, "FAIL at [%zu]: gpu %f, ref %f, rel %g > %g\n",
i, static_cast<double>(h_d[i]),
static_cast<double>(ref), static_cast<double>(rel),
static_cast<double>(kRelTolerance));
}
}
if (ok) {
std::printf(
"wgmma GEMM %dx%dx%d: PASS, max rel error %g "
"(tolerance %g)\n",
kDim, kDim, kDim, static_cast<double>(worst),
static_cast<double>(kRelTolerance));
}
CUDA_CHECK(cudaFree(d_a));
CUDA_CHECK(cudaFree(d_b));
CUDA_CHECK(cudaFree(d_d));
return ok ? EXIT_SUCCESS : EXIT_FAILURE;
}