COURSE / SOURCE

wgmma_gemm.cu

All lessons
Source filecode/day78-blackwell/wgmma_gemm.cu

This 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;
}