COURSE / SOURCE

wmma_matmul.cu

All lessons
Source filecode/day72-wmma/wmma_matmul.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 72: the same square matmul as day 44, moved onto the tensor cores
// with nvcuda::wmma, and scored against cublasGemmEx at both dtypes.
//
// Four ways to compute the same C = A * B:
//
//   matmulWmmaGlobal   one warp owns one 16 x 16 tile of C and walks K in
//                      16-wide steps, loading both f16 fragments straight
//                      from global memory every step. The MMA is a tensor
//                      core; the feed is the naive memory path day 43
//                      started from.
//   matmulWmmaStaged   the same fragments, fed from a 32 x 16 A tile and a
//                      16 x 32 B tile staged in shared memory by a 2 x 2
//                      arrangement of warps, which is day 16's move
//                      replayed at warp granularity.
//   cublasGemmEx, FP32 in and out with CUBLAS_COMPUTE_32F: day 44's exact
//                      call, the SIMT baseline, no tensor cores on a T4.
//   cublasGemmEx, f16 in, f32 out, CUBLAS_COMPUTE_32F: the library on the
//                      same units and the same dtypes as the two WMMA
//                      kernels. This row is the denominator of the percent
//                      column, because it is the only fair one.
//
// The inputs are day 44's integer trick carried into f16: every value is a
// small integer, exactly representable as __half, every product is an
// integer with magnitude at most 6, and the largest dot product any case
// reaches is 6 * 2048 = 12,288, far under the 2^24 where a float stops
// representing consecutive integers. Every partial sum is therefore exact
// in an f32 accumulator no matter what order the tensor core adds it, so
// all four paths must agree with the CPU reference and any mismatch is a
// real bug, not rounding. Storage in f16 changed nothing about these
// numbers; on real-valued data it would, and that story is day 71's.
//
// Build: nvcc -std=c++17 -O3 -arch=sm_75 -o wmma_matmul \
//            wmma_matmul.cu -lcublas
// Run:   ./wmma_matmul

#include <cmath>
#include <cstdio>
#include <cstdlib>
#include <vector>

#include <cuda_fp16.h>
#include <cuda_runtime.h>
#include <mma.h>

#include <cublas_v2.h>

namespace wmma = nvcuda::wmma;

// 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)

// cuBLAS returns cublasStatus_t, not cudaError_t, so CUDA_CHECK cannot wrap
// it and CUDA-CODE-STYLE.md allows no second error macro. Every cuBLAS call
// in this file returns its status to a real `if` instead.

// The one WMMA shape this file uses: 16 x 16 x 16, __half in, AccT out.
constexpr int kWmmaM = 16;
constexpr int kWmmaN = 16;
constexpr int kWmmaK = 16;

// The exercise's knob. float is the shipped value; the exercise rebuilds
// with __half and reads what the correctness gate says about it. Both are
// rows of the 16x16x16 supported-combinations table in the programming
// guide's warp matrix functions section.
// snippet: acc-knob
using AccT = float;
// end snippet

// Four warps per block in both kernels. 256 is the course default block
// size, but WMMA hands work out per warp, not per thread, and these tile
// shapes give a block of four warps exactly four 16 x 16 tiles of C.
constexpr int kWarpsPerBlock = 4;
constexpr int kBlockThreads = 32 * kWarpsPerBlock;

// The staged kernel's block tile: 2 x 2 warps covering 32 x 32 of C, fed
// from one 32 x 16 A tile and one 16 x 32 B tile in shared memory.
constexpr int kStageM = 32;
constexpr int kStageN = 32;

// Three square sizes. 256 keeps the CPU reference quick while still giving
// the staged kernel an 8 x 8 grid of blocks; 2048 is day 44's headline
// size, so the wall-time comparison against its evidence file is like for
// like.
constexpr int kNumCases = 3;
constexpr int kSize0 = 256;
constexpr int kSize1 = 1024;
constexpr int kSize2 = 2048;
constexpr size_t kMaxElems =
    static_cast<size_t>(kSize2) * static_cast<size_t>(kSize2);

constexpr int kNumPaths = 4;
constexpr int kGemmExF16Index = 3;

constexpr int kWarmupRuns = 3;
constexpr int kTimedRuns = 10;

// Small rather than zero only so a future case with non-integer inputs
// does not silently become a bitwise test. With the shipped inputs every
// path is exact and any miss at all is a bug.
constexpr double kRelTolerance = 1e-5;

// Every case size must be a whole number of every tile in this file: the
// 16 x 16 WMMA tile, the global kernel's four stacked warps (64 rows), and
// the staged kernel's 32 x 32 block tile. No tile shape depends on AccT,
// so the asserts hold for the shipped float accumulator and for the
// exercise's __half one alike.
constexpr bool tilesDivide(int size) {
    return size % kWmmaM == 0 && size % kWmmaN == 0 && size % kWmmaK == 0 &&
           size % (kWmmaM * kWarpsPerBlock) == 0 && size % kStageM == 0 &&
           size % kStageN == 0;
}

static_assert(tilesDivide(kSize0), "case 0 must divide every tile");
static_assert(tilesDivide(kSize1), "case 1 must divide every tile");
static_assert(tilesDivide(kSize2), "case 2 must divide every tile");
static_assert(kStageM * kStageN == kWarpsPerBlock * kWmmaM * kWmmaN,
              "the 2 x 2 warp arrangement must cover the staged tile");
static_assert((kStageM * kWmmaK + kWmmaK * kStageN) * sizeof(__half) <=
                  48u * 1024u,
              "the staged tiles must fit the 48 KiB a Turing block gets "
              "without an opt-in");
// load_matrix_sync requires ldm to be a multiple of 8 __half elements, and
// both shared tiles are read with their own row length as ldm.
static_assert(kWmmaK % 8 == 0 && kStageN % 8 == 0,
              "shared-tile leading dimensions must be multiples of 8 halves");

static const char* const kPathNames[kNumPaths] = {
    "matmulWmmaGlobal", "matmulWmmaStaged", "cublasGemmEx fp32",
    "cublasGemmEx f16"};

// Day 44's input pattern, unchanged, so the two evidence files describe the
// same problem. Every value is an integer in [-3, 3] or [-2, 2], and both
// depend on row and column so a transposed load produces a wrong answer,
// not the same one. Integers this small are exact in __half, so the f16
// copy of each matrix is the same matrix, not an approximation of it.
// __float2half is the documented host-and-device conversion from
// cuda_fp16.h.
static void fillInputs(float* a, float* b, __half* a16, __half* b16, int size) {
    const size_t s = static_cast<size_t>(size);
    for (size_t row = 0; row < s; ++row) {
        for (size_t col = 0; col < s; ++col) {
            const size_t i = row * s + col;
            a[i] =
                static_cast<float>(static_cast<int>((row + 2 * col) % 7) - 3);
            b[i] =
                static_cast<float>(static_cast<int>((3 * row + col) % 5) - 2);
            a16[i] = __float2half(a[i]);
            b16[i] = __float2half(b[i]);
        }
    }
}

// CPU reference, identical to day 44's: double accumulation, row-p-col
// loop order so it walks b and out along their rows. One reference serves
// all four paths, because with these inputs there is exactly one right
// answer.
static void matmulCpu(const float* a, const float* b, double* out, int size) {
    const size_t s = static_cast<size_t>(size);
    for (size_t i = 0; i < s * s; ++i) {
        out[i] = 0.0;
    }
    for (size_t row = 0; row < s; ++row) {
        for (size_t p = 0; p < s; ++p) {
            const double av = static_cast<double>(a[row * s + p]);
            for (size_t col = 0; col < s; ++col) {
                out[row * s + col] += av * static_cast<double>(b[p * s + col]);
            }
        }
    }
}

static size_t firstMismatch(const float* got, const double* want, size_t n,
                            double relTolerance) {
    for (size_t i = 0; i < n; ++i) {
        const double scale = (want[i] == 0.0) ? 1.0 : std::fabs(want[i]);
        if (std::fabs(static_cast<double>(got[i]) - want[i]) >
            relTolerance * scale) {
            return i;
        }
    }
    return n;
}

// The WMMA kernels write AccT; the comparison runs in float. __half2float
// is the documented host-side conversion, so the same two lines serve both
// builds of the exercise. Exactly one of the two is called in any build, so
// both carry [[maybe_unused]]; without it nvcc emits warning #177-D for the
// other one and a warning in a shipped file trains people to ignore
// warnings.
[[maybe_unused]] static float accToFloat(float v) {
    return v;
}
[[maybe_unused]] static float accToFloat(__half v) {
    return __half2float(v);
}

// One warp computes one 16 x 16 tile of C = A * B. No thread in this
// kernel owns an element: the warp owns the tile, and which lane holds
// which piece of a fragment is unspecified by the API on purpose.
//
// Memory: each load_matrix_sync pulls a 16 x 16 f16 tile, sixteen rows of
// 32 contiguous bytes, one 32-byte sector per row. Both operands come from
// global memory on every K step, so a warp reads 1 KiB per 16 x 16 x 16
// MMA and the only reuse is whatever L1 and L2 catch across warps. That is
// the deliberate flaw this kernel exists to price.
//
// Launch requirement: kBlockThreads threads, grid
// (n / 16, n / (16 * kWarpsPerBlock)); n and k whole multiples of 16,
// which tilesDivide checks at compile time. Alignment holds because every
// fragment pointer offset is a multiple of 16 halves (32 bytes) and
// load_matrix_sync wants 256-bit.
// snippet: wmma-inner
__global__ void matmulWmmaGlobal(const __half* __restrict__ a,
                                 const __half* __restrict__ b,
                                 AccT* __restrict__ c, size_t n, size_t k) {
    const unsigned int warp = threadIdx.x / 32;
    const size_t row0 =
        (static_cast<size_t>(blockIdx.y) * kWarpsPerBlock + warp) * kWmmaM;
    const size_t col0 = blockIdx.x * static_cast<size_t>(kWmmaN);

    wmma::fragment<wmma::matrix_a, kWmmaM, kWmmaN, kWmmaK, __half,
                   wmma::row_major>
        aFrag;
    wmma::fragment<wmma::matrix_b, kWmmaM, kWmmaN, kWmmaK, __half,
                   wmma::row_major>
        bFrag;
    wmma::fragment<wmma::accumulator, kWmmaM, kWmmaN, kWmmaK, AccT> accFrag;
    wmma::fill_fragment(accFrag, static_cast<AccT>(0.0f));

    for (size_t p = 0; p < k; p += kWmmaK) {
        wmma::load_matrix_sync(aFrag, a + row0 * k + p,
                               static_cast<unsigned int>(k));
        wmma::load_matrix_sync(bFrag, b + p * n + col0,
                               static_cast<unsigned int>(n));
        wmma::mma_sync(accFrag, aFrag, bFrag, accFrag);
    }

    wmma::store_matrix_sync(c + row0 * n + col0, accFrag,
                            static_cast<unsigned int>(n), wmma::mem_row_major);
}
// end snippet

// Four warps in a 2 x 2 arrangement compute a 32 x 32 tile of C, with the
// operand tiles staged in shared memory: each A fragment is loaded from
// global once and multiplied by two warps, likewise each B fragment.
//
// Memory: the cooperative fill gives consecutive threads consecutive
// halves, 256 contiguous bytes per warp per pass. The fragment loads then
// hit shared memory, ldm 16 for A and 32 for B, both multiples of 8 halves
// as load_matrix_sync requires. alignas(32) satisfies its 256-bit pointer
// alignment, which a bare __half array would not.
//
// Launch requirement: kBlockThreads threads, grid
// (n / kStageN, n / kStageM); every size a whole multiple of 32.
__global__ void matmulWmmaStaged(const __half* __restrict__ a,
                                 const __half* __restrict__ b,
                                 AccT* __restrict__ c, size_t n, size_t k) {
    __shared__ alignas(32) __half tileA[kStageM * kWmmaK];
    __shared__ alignas(32) __half tileB[kWmmaK * kStageN];

    const unsigned int tid = threadIdx.x;
    const unsigned int warp = tid / 32;
    const unsigned int warpRow = warp / 2u;
    const unsigned int warpCol = warp % 2u;

    const size_t blockRow = blockIdx.y * static_cast<size_t>(kStageM);
    const size_t blockCol = blockIdx.x * static_cast<size_t>(kStageN);

    wmma::fragment<wmma::matrix_a, kWmmaM, kWmmaN, kWmmaK, __half,
                   wmma::row_major>
        aFrag;
    wmma::fragment<wmma::matrix_b, kWmmaM, kWmmaN, kWmmaK, __half,
                   wmma::row_major>
        bFrag;
    wmma::fragment<wmma::accumulator, kWmmaM, kWmmaN, kWmmaK, AccT> accFrag;
    wmma::fill_fragment(accFrag, static_cast<AccT>(0.0f));

    for (size_t p = 0; p < k; p += kWmmaK) {
        for (unsigned int elem = tid; elem < kStageM * kWmmaK;
             elem += kBlockThreads) {
            const unsigned int rowA = elem / kWmmaK;
            const unsigned int colA = elem % kWmmaK;
            tileA[elem] = a[(blockRow + rowA) * k + p + colA];
        }
        for (unsigned int elem = tid; elem < kWmmaK * kStageN;
             elem += kBlockThreads) {
            const unsigned int rowB = elem / kStageN;
            const unsigned int colB = elem % kStageN;
            tileB[elem] = b[(p + rowB) * n + blockCol + colB];
        }
        __syncthreads();

        wmma::load_matrix_sync(aFrag, tileA + warpRow * kWmmaM * kWmmaK,
                               kWmmaK);
        wmma::load_matrix_sync(bFrag, tileB + warpCol * kWmmaN, kStageN);
        wmma::mma_sync(accFrag, aFrag, bFrag, accFrag);
        __syncthreads();
    }

    const size_t row0 = blockRow + warpRow * kWmmaM;
    const size_t col0 = blockCol + warpCol * kWmmaN;
    wmma::store_matrix_sync(c + row0 * n + col0, accFrag,
                            static_cast<unsigned int>(n), wmma::mem_row_major);
}

// cuBLAS is column major and this program is row major, so both calls ask
// for C^T = B^T * A^T by swapping the operands, day 44's trick unchanged.
//
// This call is day 44's baseline byte for byte: FP32 in and out,
// CUBLAS_COMPUTE_32F, no tensor cores on a T4.
static cublasStatus_t callCublasFp32(cublasHandle_t handle, const float* d_a,
                                     const float* d_b, float* d_c, int size) {
    const float alpha = 1.0f;
    const float beta = 0.0f;
    return cublasGemmEx(handle, CUBLAS_OP_N, CUBLAS_OP_N, size, size, size,
                        &alpha, d_b, CUDA_R_32F, size, d_a, CUDA_R_32F, size,
                        &beta, d_c, CUDA_R_32F, size, CUBLAS_COMPUTE_32F,
                        CUBLAS_GEMM_DEFAULT);
}

// The same call with the inputs in f16 and the output and compute type
// still 32-bit: the same dtypes the WMMA kernels use, so this is the row
// the percent column divides by. alpha and beta stay float because the
// compute type, not the storage type, decides their type.
// snippet: gemmex-f16
static cublasStatus_t callCublasF16(cublasHandle_t handle, const __half* d_a,
                                    const __half* d_b, float* d_c, int size) {
    const float alpha = 1.0f;
    const float beta = 0.0f;
    return cublasGemmEx(handle, CUBLAS_OP_N, CUBLAS_OP_N, size, size, size,
                        &alpha, d_b, CUDA_R_16F, size, d_a, CUDA_R_16F, size,
                        &beta, d_c, CUDA_R_32F, size, CUBLAS_COMPUTE_32F,
                        CUBLAS_GEMM_DEFAULT);
}
// end snippet

// Launches one of the four paths. The WMMA kernels write d_cAcc (AccT);
// the two cuBLAS paths write d_c32 (float). Returns the cuBLAS status for
// the library paths and CUBLAS_STATUS_SUCCESS for the kernels, so one call
// site can check both kinds of failure.
static cublasStatus_t launchMatmul(int which, cublasHandle_t handle,
                                   const float* d_a32, const float* d_b32,
                                   const __half* d_a16, const __half* d_b16,
                                   float* d_c32, AccT* d_cAcc, int size) {
    const size_t n = static_cast<size_t>(size);
    const size_t k = static_cast<size_t>(size);

    if (which == 0) {
        const dim3 grid(
            static_cast<unsigned int>(size / kWmmaN),
            static_cast<unsigned int>(size / (kWmmaM * kWarpsPerBlock)));
        matmulWmmaGlobal<<<grid, kBlockThreads>>>(d_a16, d_b16, d_cAcc, n, k);
    } else if (which == 1) {
        const dim3 grid(static_cast<unsigned int>(size / kStageN),
                        static_cast<unsigned int>(size / kStageM));
        matmulWmmaStaged<<<grid, kBlockThreads>>>(d_a16, d_b16, d_cAcc, n, k);
    } else if (which == 2) {
        return callCublasFp32(handle, d_a32, d_b32, d_c32, size);
    } else {
        return callCublasF16(handle, d_a16, d_b16, d_c32, size);
    }
    return CUBLAS_STATUS_SUCCESS;
}

// Times a launch with CUDA events and returns the mean milliseconds per
// run. Copy it verbatim; the alternative is four copies of the event
// boilerplate, which is how a warm-up goes missing from one of them.
template <typename LaunchFn>
static float timeKernel(LaunchFn launch) {
    cudaEvent_t start, stop;
    CUDA_CHECK(cudaEventCreate(&start));
    CUDA_CHECK(cudaEventCreate(&stop));

    // Warm up this kernel, not just the first kernel in the program. Lazy
    // module loading has been the default since CUDA 12.2 on Linux, so the
    // first launch of each kernel pays its own load, and cuBLAS picks and
    // caches its algorithm on the first call as well.
    for (int i = 0; i < kWarmupRuns; ++i) {
        launch();
    }
    CUDA_CHECK(cudaDeviceSynchronize());
    CUDA_CHECK(cudaGetLastError());

    CUDA_CHECK(cudaEventRecord(start));
    for (int i = 0; i < kTimedRuns; ++i) {
        launch();
    }
    CUDA_CHECK(cudaEventRecord(stop));
    CUDA_CHECK(cudaEventSynchronize(stop));
    CUDA_CHECK(cudaGetLastError());

    float ms = 0.0f;
    CUDA_CHECK(cudaEventElapsedTime(&ms, start, stop));
    CUDA_CHECK(cudaEventDestroy(start));
    CUDA_CHECK(cudaEventDestroy(stop));
    return ms / kTimedRuns;
}

// Registers per thread, spill bytes, shared memory and resident blocks per
// SM, from the same two profiler-free calls day 44 used, so the launch
// shape of a WMMA kernel can be read on a tier where Nsight Compute
// cannot run.
template <typename KernelFn>
static void printAttributes(const char* name, KernelFn kernel, int threads,
                            int maxThreadsPerSm) {
    cudaFuncAttributes attr;
    CUDA_CHECK(cudaFuncGetAttributes(&attr, kernel));
    int blocksPerSm = 0;
    CUDA_CHECK(cudaOccupancyMaxActiveBlocksPerMultiprocessor(
        &blocksPerSm, kernel, threads, 0));
    const double occupancy =
        100.0 * blocksPerSm * threads / static_cast<double>(maxThreadsPerSm);
    std::printf("%-21s %7d %5d %8zu %8zu %7d %6.1f\n", name, threads,
                attr.numRegs, attr.localSizeBytes, attr.sharedSizeBytes,
                blocksPerSm, occupancy);
}

static int runAllCases(cublasHandle_t handle, float* d_a32, float* d_b32,
                       __half* d_a16, __half* d_b16, float* d_c32, AccT* d_cAcc,
                       float* h_a32, float* h_b32, __half* h_a16, __half* h_b16,
                       float* h_c, AccT* h_cAcc, double* h_want) {
    const int sizes[kNumCases] = {kSize0, kSize1, kSize2};

    for (int caseIdx = 0; caseIdx < kNumCases; ++caseIdx) {
        const int size = sizes[caseIdx];
        const size_t elems =
            static_cast<size_t>(size) * static_cast<size_t>(size);
        const double flop = 2.0 * static_cast<double>(size) *
                            static_cast<double>(size) *
                            static_cast<double>(size);

        fillInputs(h_a32, h_b32, h_a16, h_b16, size);
        matmulCpu(h_a32, h_b32, h_want, size);
        CUDA_CHECK(cudaMemcpy(d_a32, h_a32, elems * sizeof(float),
                              cudaMemcpyHostToDevice));
        CUDA_CHECK(cudaMemcpy(d_b32, h_b32, elems * sizeof(float),
                              cudaMemcpyHostToDevice));
        CUDA_CHECK(cudaMemcpy(d_a16, h_a16, elems * sizeof(__half),
                              cudaMemcpyHostToDevice));
        CUDA_CHECK(cudaMemcpy(d_b16, h_b16, elems * sizeof(__half),
                              cudaMemcpyHostToDevice));

        double ms[kNumPaths] = {0.0, 0.0, 0.0, 0.0};

        // Correctness first, for every path, before anything is timed.
        // Each path's output buffer is zeroed before its launch, so a
        // kernel that writes only part of C is caught by the reference
        // check rather than reading as correct because a previous path
        // left the right answer behind.
        for (int which = 0; which < kNumPaths; ++which) {
            const bool wmmaPath = which < 2;
            CUDA_CHECK(cudaMemset(d_c32, 0, elems * sizeof(float)));
            CUDA_CHECK(cudaMemset(d_cAcc, 0, elems * sizeof(AccT)));

            const cublasStatus_t st = launchMatmul(
                which, handle, d_a32, d_b32, d_a16, d_b16, d_c32, d_cAcc, size);
            if (st != CUBLAS_STATUS_SUCCESS) {
                std::fprintf(stderr, "%s failed at size %d: %s\n",
                             kPathNames[which], size,
                             cublasGetStatusString(st));
                return EXIT_FAILURE;
            }
            CUDA_CHECK(cudaGetLastError());
            CUDA_CHECK(cudaDeviceSynchronize());

            if (wmmaPath) {
                CUDA_CHECK(cudaMemcpy(h_cAcc, d_cAcc, elems * sizeof(AccT),
                                      cudaMemcpyDeviceToHost));
                for (size_t i = 0; i < elems; ++i) {
                    h_c[i] = accToFloat(h_cAcc[i]);
                }
            } else {
                CUDA_CHECK(cudaMemcpy(h_c, d_c32, elems * sizeof(float),
                                      cudaMemcpyDeviceToHost));
            }

            const size_t bad = firstMismatch(h_c, h_want, elems, kRelTolerance);
            if (bad != elems) {
                std::fprintf(stderr,
                             "%s wrong at size %d, element %zu "
                             "(row %zu, col %zu): got %.9g, want %.9g\n",
                             kPathNames[which], size, bad,
                             bad / static_cast<size_t>(size),
                             bad % static_cast<size_t>(size),
                             static_cast<double>(h_c[bad]), h_want[bad]);
                return EXIT_FAILURE;
            }
        }

        for (int which = 0; which < kNumPaths; ++which) {
            // The lambda runs many launches; record the first failure and
            // check it after timeKernel returns, so a cuBLAS error during
            // timing cannot silently time an error path.
            cublasStatus_t timedStatus = CUBLAS_STATUS_SUCCESS;
            ms[which] = static_cast<double>(timeKernel([&] {
                const cublasStatus_t st =
                    launchMatmul(which, handle, d_a32, d_b32, d_a16, d_b16,
                                 d_c32, d_cAcc, size);
                if (st != CUBLAS_STATUS_SUCCESS &&
                    timedStatus == CUBLAS_STATUS_SUCCESS) {
                    timedStatus = st;
                }
            }));
            if (timedStatus != CUBLAS_STATUS_SUCCESS) {
                std::fprintf(stderr, "%s failed while timed at size %d: %s\n",
                             kPathNames[which], size,
                             cublasGetStatusString(timedStatus));
                return EXIT_FAILURE;
            }
        }

        std::printf(
            "size %d, %.3f GFLOP per call, mean of %d timed runs "
            "after %d warm-ups\n",
            size, flop / 1.0e9, kTimedRuns, kWarmupRuns);
        std::printf("%-21s %10s %10s %9s %9s\n", "path", "in/acc", "ms",
                    "GFLOP/s", "%f16row");
        std::printf("%-21s %10s %10s %9s %9s\n", "---------------------",
                    "----------", "----------", "---------", "---------");
        const char* const wmmaDtype = sizeof(AccT) == 4 ? "f16/f32" : "f16/f16";
        const char* const dtypes[kNumPaths] = {wmmaDtype, wmmaDtype, "f32/f32",
                                               "f16/f32"};
        for (int which = 0; which < kNumPaths; ++which) {
            const double gflops = flop / (ms[which] * 1.0e-3) / 1.0e9;
            const double pct = 100.0 * ms[kGemmExF16Index] / ms[which];
            std::printf("%-21s %10s %10.3f %9.1f %9.1f\n", kPathNames[which],
                        dtypes[which], ms[which], gflops, pct);
        }
        std::printf("\n");
    }
    return EXIT_SUCCESS;
}

int main() {
    const int device = 0;
    CUDA_CHECK(cudaSetDevice(device));
    cudaDeviceProp prop;
    CUDA_CHECK(cudaGetDeviceProperties(&prop, device));
    std::printf("GPU: %s (compute capability %d.%d), %d SMs\n", prop.name,
                prop.major, prop.minor, prop.multiProcessorCount);
    if (prop.major < 7) {
        // WMMA with __half needs the tensor cores that arrived with
        // compute capability 7.0.
        std::fprintf(stderr,
                     "this program needs compute capability 7.0 or higher\n");
        return EXIT_FAILURE;
    }
    std::printf("accumulator type: %s\n\n",
                sizeof(AccT) == 4 ? "float" : "__half");

    std::printf("Launch shape, and what the compiler did with it\n");
    std::printf("%-21s %7s %5s %8s %8s %7s %6s\n", "kernel", "thr/blk", "regs",
                "local B", "smem B", "blk/SM", "occ%");
    std::printf("%-21s %7s %5s %8s %8s %7s %6s\n", "---------------------",
                "-------", "-----", "--------", "--------", "-------",
                "------");
    printAttributes("matmulWmmaGlobal", matmulWmmaGlobal, kBlockThreads,
                    prop.maxThreadsPerMultiProcessor);
    printAttributes("matmulWmmaStaged", matmulWmmaStaged, kBlockThreads,
                    prop.maxThreadsPerMultiProcessor);
    std::printf("\n");

    std::vector<float> h_a32(kMaxElems);
    std::vector<float> h_b32(kMaxElems);
    std::vector<__half> h_a16(kMaxElems);
    std::vector<__half> h_b16(kMaxElems);
    std::vector<float> h_c(kMaxElems);
    std::vector<AccT> h_cAcc(kMaxElems);
    std::vector<double> h_want(kMaxElems);

    float* d_a32 = nullptr;
    float* d_b32 = nullptr;
    __half* d_a16 = nullptr;
    __half* d_b16 = nullptr;
    float* d_c32 = nullptr;
    AccT* d_cAcc = nullptr;
    CUDA_CHECK(cudaMalloc(&d_a32, kMaxElems * sizeof(float)));
    CUDA_CHECK(cudaMalloc(&d_b32, kMaxElems * sizeof(float)));
    CUDA_CHECK(cudaMalloc(&d_a16, kMaxElems * sizeof(__half)));
    CUDA_CHECK(cudaMalloc(&d_b16, kMaxElems * sizeof(__half)));
    CUDA_CHECK(cudaMalloc(&d_c32, kMaxElems * sizeof(float)));
    CUDA_CHECK(cudaMalloc(&d_cAcc, kMaxElems * sizeof(AccT)));

    int status = EXIT_SUCCESS;
    cublasHandle_t handle = nullptr;
    const cublasStatus_t created = cublasCreate(&handle);
    if (created != CUBLAS_STATUS_SUCCESS) {
        std::fprintf(stderr, "cublasCreate failed: %s\n",
                     cublasGetStatusString(created));
        status = EXIT_FAILURE;
    }

    if (status == EXIT_SUCCESS) {
        // Set explicitly, day 44's habit: CUBLAS_DEFAULT_MATH keeps the
        // FP32 row FP32 all the way through. The f16-input row uses tensor
        // cores under the same mode, because for f16 inputs that is the
        // default path, not an opt-in.
        const cublasStatus_t mode =
            cublasSetMathMode(handle, CUBLAS_DEFAULT_MATH);
        if (mode != CUBLAS_STATUS_SUCCESS) {
            std::fprintf(stderr, "cublasSetMathMode failed: %s\n",
                         cublasGetStatusString(mode));
            status = EXIT_FAILURE;
        }
    }

    if (status == EXIT_SUCCESS) {
        status =
            runAllCases(handle, d_a32, d_b32, d_a16, d_b16, d_c32, d_cAcc,
                        h_a32.data(), h_b32.data(), h_a16.data(), h_b16.data(),
                        h_c.data(), h_cAcc.data(), h_want.data());
    }

    if (status == EXIT_SUCCESS) {
        std::printf(
            "all %d paths match the CPU reference at every element "
            "and every size\n",
            kNumPaths);
    }

    if (handle != nullptr) {
        const cublasStatus_t destroyed = cublasDestroy(handle);
        if (destroyed != CUBLAS_STATUS_SUCCESS) {
            std::fprintf(stderr, "cublasDestroy failed: %s\n",
                         cublasGetStatusString(destroyed));
            status = EXIT_FAILURE;
        }
    }
    CUDA_CHECK(cudaFree(d_a32));
    CUDA_CHECK(cudaFree(d_b32));
    CUDA_CHECK(cudaFree(d_a16));
    CUDA_CHECK(cudaFree(d_b16));
    CUDA_CHECK(cudaFree(d_c32));
    CUDA_CHECK(cudaFree(d_cAcc));
    return status;
}