COURSE / SOURCE

matmul_registers.cu

All lessons
Source filecode/day44-matmul-2/matmul_registers.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 44: register tiling, vectorised loads, and how close a hand-written
// SGEMM gets to cuBLAS on a Tesla T4.
//
// Five ways to compute the same square C = A * B, FP32 in and FP32 out:
//
//   matmulSharedTiled     where day 43 stops. A 16 x 16 shared tile, one
//                         output element per thread, two shared-memory load
//                         instructions per multiply-add.
//   matmulRegTile1d       step 4. A 64 x 64 block tile, 512 threads, and one
//                         thread owns a column of 8 outputs held in
//                         registers. Nine shared loads per eight FMAs.
//   matmulRegTile2d       step 5. A 128 x 128 block tile, 256 threads, and
//                         one thread owns an 8 x 8 micro-tile: 64
//                         accumulators in registers, 16 shared loads per 64
//                         FMAs.
//   matmulRegTile2dVec4   step 6. The same shape, with float4 loads out of
//                         global memory, the A tile stored transposed so the
//                         inner loop reads it four floats at a time, and
//                         float4 stores to C. Four shared load instructions
//                         per 64 FMAs.
//   cublasGemmEx          the target. CUDA_R_32F in and out with compute
//                         type CUBLAS_COMPUTE_32F, so it is the same
//                         arithmetic and no TF32 tensor-core path is in
//                         play. A comparison against a different compute
//                         type is a different question.
//
// The four kernels carry no ragged-edge guard. Day 16 has that guard and it
// belongs in the inner loop, which is the one place this file cannot pay
// for it, so every case size here is a whole number of every block tile and
// a static_assert enforces it. Real libraries handle the edge with a
// separate epilogue kernel or by padding the allocation. That is the honest
// cost of these three steps and the page says so.
//
// The program also prints registers per thread, local memory per thread and
// blocks per SM for each kernel, read from cudaFuncGetAttributes and
// cudaOccupancyMaxActiveBlocksPerMultiprocessor. Those two calls need no
// performance counters and no root, so the occupancy half of this lesson
// works on a tier where Nsight Compute does not.
//
// Build: nvcc -std=c++17 -O3 -arch=sm_75 -o matmul_registers \
//            matmul_registers.cu -lcublas
// Run:   ./matmul_registers
//

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

#include <cuda_runtime.h>

#include <cublas_v2.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)

// 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, which is also what
// the correctness gates below do.

// Where day 43 stops: a 16 x 16 shared tile, one output element per thread.
constexpr int kTileDim = 16;

// Step 4. A 64 x 64 block tile, 512 threads, one thread owns a column of 8.
constexpr int kBlock1dM = 64;
constexpr int kBlock1dN = 64;
constexpr int kBlock1dK = 8;
constexpr int kThread1dM = 8;
constexpr int kThreads1d = (kBlock1dM * kBlock1dN) / kThread1dM;

// Steps 5 and 6. A 128 x 128 block tile and an 8 x 8 micro-tile per thread.
// kThreadM and kThreadN are the exercise's two knobs: the thread count falls
// out of them, so 4 x 4 gives 1024 threads and 16 x 16 gives 64.
// snippet: knobs
constexpr int kBlockM = 128;
constexpr int kBlockN = 128;
constexpr int kBlockK = 8;
constexpr int kThreadM = 8;
constexpr int kThreadN = 8;
constexpr int kRegThreads = (kBlockM * kBlockN) / (kThreadM * kThreadN);
// end snippet

// Three square sizes. 512 is deliberate: a 128 x 128 block tile gives it 16
// blocks, and this card has 40 SMs, so more than half the machine has no
// work. The 16 x 16 baseline launches 1024 blocks for the same problem.
constexpr int kNumCases = 3;
constexpr int kSize0 = 512;
constexpr int kSize1 = 1024;
constexpr int kSize2 = 2048;
constexpr size_t kMaxElems =
    static_cast<size_t>(kSize2) * static_cast<size_t>(kSize2);

constexpr int kNumKernels = 5;
constexpr int kCublasIndex = 4;

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

// Every product here is an integer in [-6, 6] and the largest dot product
// any case can reach is 6 * 2048 = 12,288, well under the 2^24 where a float
// stops representing consecutive integers. So both processors are exact and
// any difference at all is a bug, not rounding. The tolerance is small
// rather than zero only so that a future case with larger inputs does not
// silently become a bitwise test.
constexpr double kRelTolerance = 1e-5;

// Every case size must be a whole number of every block tile in this file,
// because none of these kernels guards its edges.
constexpr bool tilesDivide(int size) {
    return size % kTileDim == 0 && size % kBlock1dM == 0 &&
           size % kBlock1dN == 0 && size % kBlock1dK == 0 &&
           size % kBlockM == 0 && size % kBlockN == 0 && size % kBlockK == 0;
}

static_assert(tilesDivide(kSize0), "case 0 must divide every block tile");
static_assert(tilesDivide(kSize1), "case 1 must divide every block tile");
static_assert(tilesDivide(kSize2), "case 2 must divide every block tile");

static_assert(kTileDim * kTileDim % 32 == 0 && kThreads1d % 32 == 0 &&
                  kRegThreads % 32 == 0,
              "every block in this file is a whole number of warps");
static_assert(kTileDim * kTileDim <= 1024 && kThreads1d <= 1024 &&
                  kRegThreads <= 1024,
              "1024 threads is the per-block maximum on every compute "
              "capability this course targets");
static_assert(kBlock1dM % kThread1dM == 0,
              "the 64 by 64 tile must split into whole columns of "
              "kThread1dM outputs");
static_assert(kBlockM % kThreadM == 0 && kBlockN % kThreadN == 0,
              "the 128 by 128 tile must split into whole micro-tiles");
static_assert(kThreadM % 4 == 0 && kThreadN % 4 == 0,
              "the vectorised kernel moves four floats at a time, so both "
              "micro-tile edges must be a multiple of four");
static_assert(kBlockK % 4 == 0,
              "the transposed A load reads four consecutive k values per "
              "thread");
static_assert(kBlockN % 4 == 0, "the B tile is filled with float4 stores");
static_assert((kBlock1dM * kBlock1dK + kBlock1dK * kBlock1dN) * sizeof(float) <=
                  48u * 1024u,
              "the 64 by 64 tiles must fit the 48 KiB of shared memory a "
              "Turing block gets without an opt-in");
static_assert((kBlockM * kBlockK + kBlockK * kBlockN) * sizeof(float) <=
                  48u * 1024u,
              "the 128 by 128 tiles must fit the same 48 KiB");

static const char* const kKernelNames[kNumKernels] = {
    "matmulSharedTiled", "matmulRegTile1d", "matmulRegTile2d",
    "matmulRegTile2dVec4", "cublasGemmEx"};

// The block tile each kernel uses, for the traffic arithmetic below. cuBLAS
// picks its own and does not tell us, so its row is left blank.
constexpr int kKernelTileM[kNumKernels] = {kTileDim, kBlock1dM, kBlockM,
                                           kBlockM, 0};
constexpr int kKernelTileN[kNumKernels] = {kTileDim, kBlock1dN, kBlockN,
                                           kBlockN, 0};

// Shared-memory load instructions per multiply-add, counted from the inner
// loops in this file. Arithmetic, not a measurement. The vectorised kernel
// issues a quarter of the loads of the kernel above it and moves four floats
// with each one, so the values read are identical and the instructions are
// not.
constexpr double kSmemPerFma[kNumKernels] = {
    2.0, static_cast<double>(1 + kThread1dM) / kThread1dM,
    static_cast<double>(kThreadM + kThreadN) / (kThreadM * kThreadN),
    static_cast<double>(kThreadM / 4 + kThreadN / 4) / (kThreadM * kThreadN),
    0.0};

// Global floats moved per output element, from the tile shape alone: a block
// re-reads a strip of A once per block column and a strip of B once per
// block row, and writes C once. Day 30 uses the same formula, so the two
// pages agree on what a FLOP per byte figure means.
//
// It is arithmetic, not a measurement. The L1 and L2 caches serve part of
// this traffic, which is why a measured GB/s can read above the ceiling this
// number implies.
static double floatsPerOutput(int tileM, int tileN, int size) {
    const double depth = static_cast<double>(size);
    return depth / static_cast<double>(tileN) +
           depth / static_cast<double>(tileM) + 1.0;
}

static double flopPerByte(int tileM, int tileN, int size) {
    return 2.0 * static_cast<double>(size) /
           (4.0 * floatsPerOutput(tileM, tileN, size));
}

// Small integers, and the value depends on both the row and the column, so a
// kernel that swaps them or shifts either one produces a different answer
// rather than the same one.
static void fillInputs(float* a, float* b, 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);
        }
    }
}

// CPU reference. Written for obvious correctness, not speed, and it never
// allocates. The loop order is row, p, col rather than row, col, p: it walks
// b and out along their rows instead of down a column, and it adds the terms
// of each dot product in the same p order, so the result is identical.
//
// It accumulates in double even though every kernel accumulates in float,
// because the reference's job is to be right rather than to match bit for
// bit.
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;
}

// Where day 43 stops. One thread owns one element of C, and a 16 x 16 tile
// of A and of B live in shared memory, so a thread issues 2 * K / 16 global
// loads instead of 2 * K.
//
// Memory: consecutive threads take consecutive columns, so the tileB load
// and the C store are each 32 consecutive floats per warp, four 32-byte
// sectors. The inner product reads two shared values for every multiply-add
// it performs, and that ratio is what the rest of this file attacks.
//
// There is no `m` parameter in any kernel here: the grid covers exactly M
// rows because M is a whole number of block tiles, which tilesDivide checks
// at compile time.
__global__ void matmulSharedTiled(const float* __restrict__ a,
                                  const float* __restrict__ b,
                                  float* __restrict__ c, size_t n, size_t k) {
    __shared__ float tileA[kTileDim][kTileDim];
    __shared__ float tileB[kTileDim][kTileDim];

    const size_t col = blockIdx.x * static_cast<size_t>(kTileDim) + threadIdx.x;
    const size_t row = blockIdx.y * static_cast<size_t>(kTileDim) + threadIdx.y;

    float acc = 0.0f;
    for (size_t kIdx = 0; kIdx < k; kIdx += kTileDim) {
        tileA[threadIdx.y][threadIdx.x] = a[row * k + kIdx + threadIdx.x];
        tileB[threadIdx.y][threadIdx.x] = b[(kIdx + threadIdx.y) * n + col];
        __syncthreads();

        for (int p = 0; p < kTileDim; ++p) {
            acc += tileA[threadIdx.y][p] * tileB[p][threadIdx.x];
        }
        __syncthreads();
    }

    c[row * n + col] = acc;
}

// Step 4, one-dimensional register tiling. One thread owns a column of
// kThread1dM = 8 elements of C inside a 64 x 64 block tile, so the block
// needs 512 threads instead of 4096.
//
// Memory: threadCol is tid % 64, so a warp's 32 lanes take 32 consecutive
// columns and both the tileB load and the C store stay contiguous. The tile
// fill loops walk the tile linearly, one element per thread per pass, which
// keeps the B fill fully coalesced and the A fill coalesced in runs of
// kBlock1dK.
//
// The eight accumulators live in registers, and that is the whole change:
// the value read out of tileB is used eight times instead of once, so the
// nine shared loads per pass buy eight multiply-adds rather than four.
__global__ void matmulRegTile1d(const float* __restrict__ a,
                                const float* __restrict__ b,
                                float* __restrict__ c, size_t n, size_t k) {
    __shared__ float tileA[kBlock1dM * kBlock1dK];
    __shared__ float tileB[kBlock1dK * kBlock1dN];

    const int tid = static_cast<int>(threadIdx.x);
    const int threadCol = tid % kBlock1dN;
    const int threadRow = tid / kBlock1dN;

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

    float acc[kThread1dM];
    for (int i = 0; i < kThread1dM; ++i) {
        acc[i] = 0.0f;
    }

    for (size_t kIdx = 0; kIdx < k; kIdx += kBlock1dK) {
        for (int elem = tid; elem < kBlock1dM * kBlock1dK; elem += kThreads1d) {
            const int rowA = elem / kBlock1dK;
            const int colA = elem % kBlock1dK;
            tileA[elem] = a[(blockRow + rowA) * k + kIdx + colA];
        }
        for (int elem = tid; elem < kBlock1dK * kBlock1dN; elem += kThreads1d) {
            const int rowB = elem / kBlock1dN;
            const int colB = elem % kBlock1dN;
            tileB[elem] = b[(kIdx + rowB) * n + blockCol + colB];
        }
        __syncthreads();

        for (int p = 0; p < kBlock1dK; ++p) {
            const float bVal = tileB[p * kBlock1dN + threadCol];
            for (int i = 0; i < kThread1dM; ++i) {
                acc[i] +=
                    tileA[(threadRow * kThread1dM + i) * kBlock1dK + p] * bVal;
            }
        }
        __syncthreads();
    }

    for (int i = 0; i < kThread1dM; ++i) {
        const size_t row = blockRow + threadRow * kThread1dM + i;
        c[row * n + blockCol + threadCol] = acc[i];
    }
}

// Step 5, two-dimensional register tiling. One thread owns an 8 x 8
// micro-tile of C inside a 128 x 128 block tile: 64 accumulators, 8 values
// of A and 8 values of B in registers, and 64 multiply-adds for the 16
// shared loads that filled them.
//
// Memory: threadCol is tid % 16 and each thread writes 8 consecutive
// columns, so 16 threads cover the 128 columns of the tile and a warp's
// store is two runs of 128 consecutive floats.
//
// Launch requirement: kRegThreads threads, which is 128 * 128 divided by the
// micro-tile area. Change kThreadM or kThreadN and the block size changes
// with it.
__global__ void matmulRegTile2d(const float* __restrict__ a,
                                const float* __restrict__ b,
                                float* __restrict__ c, size_t n, size_t k) {
    __shared__ float tileA[kBlockM * kBlockK];
    __shared__ float tileB[kBlockK * kBlockN];

    const int tid = static_cast<int>(threadIdx.x);
    const int threadCol = tid % (kBlockN / kThreadN);
    const int threadRow = tid / (kBlockN / kThreadN);

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

    float acc[kThreadM][kThreadN];
    for (int i = 0; i < kThreadM; ++i) {
        for (int j = 0; j < kThreadN; ++j) {
            acc[i][j] = 0.0f;
        }
    }
    float regM[kThreadM];
    float regN[kThreadN];

    for (size_t kIdx = 0; kIdx < k; kIdx += kBlockK) {
        for (int elem = tid; elem < kBlockM * kBlockK; elem += kRegThreads) {
            const int rowA = elem / kBlockK;
            const int colA = elem % kBlockK;
            tileA[elem] = a[(blockRow + rowA) * k + kIdx + colA];
        }
        for (int elem = tid; elem < kBlockK * kBlockN; elem += kRegThreads) {
            const int rowB = elem / kBlockN;
            const int colB = elem % kBlockN;
            tileB[elem] = b[(kIdx + rowB) * n + blockCol + colB];
        }
        __syncthreads();

        // snippet: reg-tile-inner
        for (int p = 0; p < kBlockK; ++p) {
            for (int i = 0; i < kThreadM; ++i) {
                regM[i] = tileA[(threadRow * kThreadM + i) * kBlockK + p];
            }
            for (int j = 0; j < kThreadN; ++j) {
                regN[j] = tileB[p * kBlockN + threadCol * kThreadN + j];
            }
            for (int i = 0; i < kThreadM; ++i) {
                for (int j = 0; j < kThreadN; ++j) {
                    acc[i][j] += regM[i] * regN[j];
                }
            }
        }
        // end snippet
        __syncthreads();
    }

    for (int i = 0; i < kThreadM; ++i) {
        const size_t row = blockRow + threadRow * kThreadM + i;
        for (int j = 0; j < kThreadN; ++j) {
            c[row * n + blockCol + threadCol * kThreadN + j] = acc[i][j];
        }
    }
}

// Step 6, the same shape with the loads widened.
//
// Three changes from the kernel above and no fourth:
//
// 1. Both tile fills read 16 bytes per lane with float4 instead of 4.
// 2. The A tile is stored transposed, kBlockK rows of kBlockM, so the inner
//    loop's eight A values sit next to each other and can be read as two
//    float4 instead of eight scalars strided by kBlockK.
// 3. C is written with float4.
//
// Alignment: a float4 access must sit on a 16-byte boundary. The shared
// arrays are declared as float4 because float4 carries an alignment of 16
// and a float array carries 4. On the global side, cudaMalloc returns at
// least 256-byte aligned memory and every offset used here is a multiple of
// four floats, which holds because K and N are whole numbers of block tiles.
// Break either and the run fails with `misaligned address`.
__global__ void matmulRegTile2dVec4(const float* __restrict__ a,
                                    const float* __restrict__ b,
                                    float* __restrict__ c, size_t n, size_t k) {
    __shared__ float4 tileA4[(kBlockK * kBlockM) / 4];
    __shared__ float4 tileB4[(kBlockK * kBlockN) / 4];
    float* tileA = reinterpret_cast<float*>(tileA4);
    float* tileB = reinterpret_cast<float*>(tileB4);

    const int tid = static_cast<int>(threadIdx.x);
    const int threadCol = tid % (kBlockN / kThreadN);
    const int threadRow = tid / (kBlockN / kThreadN);

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

    float acc[kThreadM][kThreadN];
    for (int i = 0; i < kThreadM; ++i) {
        for (int j = 0; j < kThreadN; ++j) {
            acc[i][j] = 0.0f;
        }
    }
    float regM[kThreadM];
    float regN[kThreadN];

    for (size_t kIdx = 0; kIdx < k; kIdx += kBlockK) {
        // snippet: vec4-load-a
        // One 16-byte load pulls four consecutive k values of one row of A,
        // and the four scalar stores scatter them down one column of the
        // transposed tile. The scatter is the price of the transpose, and
        // the inner loop is what buys it back.
        for (int elem = tid; elem < (kBlockM * kBlockK) / 4;
             elem += kRegThreads) {
            const int rowA = elem / (kBlockK / 4);
            const int colA = (elem % (kBlockK / 4)) * 4;
            const float4 v = *reinterpret_cast<const float4*>(
                &a[(blockRow + rowA) * k + kIdx + colA]);
            tileA[(colA + 0) * kBlockM + rowA] = v.x;
            tileA[(colA + 1) * kBlockM + rowA] = v.y;
            tileA[(colA + 2) * kBlockM + rowA] = v.z;
            tileA[(colA + 3) * kBlockM + rowA] = v.w;
        }
        // end snippet
        // B needs no transpose: four consecutive columns of one row of B are
        // four consecutive floats of the tile, so the float4 goes straight
        // in.
        for (int elem = tid; elem < (kBlockK * kBlockN) / 4;
             elem += kRegThreads) {
            const int rowB = elem / (kBlockN / 4);
            const int colB = (elem % (kBlockN / 4)) * 4;
            tileB4[elem] = *reinterpret_cast<const float4*>(
                &b[(kIdx + rowB) * n + blockCol + colB]);
        }
        __syncthreads();

        for (int p = 0; p < kBlockK; ++p) {
            // The float4 is unpacked into the scalar arrays rather than the
            // arrays being aliased, because a local float array carries an
            // alignment of 4 and a 128-bit access through it would be
            // undefined.
            for (int i = 0; i < kThreadM; i += 4) {
                const float4 v = *reinterpret_cast<const float4*>(
                    &tileA[p * kBlockM + threadRow * kThreadM + i]);
                regM[i + 0] = v.x;
                regM[i + 1] = v.y;
                regM[i + 2] = v.z;
                regM[i + 3] = v.w;
            }
            for (int j = 0; j < kThreadN; j += 4) {
                const float4 v = *reinterpret_cast<const float4*>(
                    &tileB[p * kBlockN + threadCol * kThreadN + j]);
                regN[j + 0] = v.x;
                regN[j + 1] = v.y;
                regN[j + 2] = v.z;
                regN[j + 3] = v.w;
            }
            for (int i = 0; i < kThreadM; ++i) {
                for (int j = 0; j < kThreadN; ++j) {
                    acc[i][j] += regM[i] * regN[j];
                }
            }
        }
        __syncthreads();
    }

    for (int i = 0; i < kThreadM; ++i) {
        const size_t row = blockRow + threadRow * kThreadM + i;
        for (int j = 0; j < kThreadN; j += 4) {
            const float4 v = make_float4(acc[i][j + 0], acc[i][j + 1],
                                         acc[i][j + 2], acc[i][j + 3]);
            *reinterpret_cast<float4*>(
                &c[row * n + blockCol + threadCol * kThreadN + j]) = v;
        }
    }
}

// cuBLAS is column major and this program is row major. Rather than
// transpose anything, ask for C^T = B^T * A^T with both operands
// untransposed: the column-major result of that call occupies exactly the
// bytes of the row-major C = A * B. Swapping the two operands and swapping
// m and n is the whole trick, and it costs nothing.
//
// CUDA_R_32F in and out with CUBLAS_COMPUTE_32F is the comparison this page
// promises. CUBLAS_COMPUTE_32F is documented as using "compute and
// intermediate storage precisions of at least 32-bits", so it cannot be the
// TF32 path, which lives behind CUBLAS_COMPUTE_32F_FAST_TF32. Change that
// one argument on an Ampere or newer card and the percentage this program
// prints changes more than any kernel edit in this file.
// snippet: cublas-call
static cublasStatus_t callCublas(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);
}
// end snippet

// Launches one of the five paths. Returns the cuBLAS status for the cuBLAS
// path and CUBLAS_STATUS_SUCCESS for the four kernels, so one call site can
// check both kinds of failure.
static cublasStatus_t launchMatmul(int which, cublasHandle_t handle,
                                   const float* d_a, const float* d_b,
                                   float* d_c, int size) {
    const size_t n = static_cast<size_t>(size);
    const size_t k = static_cast<size_t>(size);
    const unsigned int tiles16 = static_cast<unsigned int>(size / kTileDim);
    const unsigned int tiles1dX = static_cast<unsigned int>(size / kBlock1dN);
    const unsigned int tiles1dY = static_cast<unsigned int>(size / kBlock1dM);
    const unsigned int tilesX = static_cast<unsigned int>(size / kBlockN);
    const unsigned int tilesY = static_cast<unsigned int>(size / kBlockM);

    if (which == 0) {
        const dim3 block(kTileDim, kTileDim);
        const dim3 grid(tiles16, tiles16);
        matmulSharedTiled<<<grid, block>>>(d_a, d_b, d_c, n, k);
    } else if (which == 1) {
        const dim3 grid(tiles1dX, tiles1dY);
        matmulRegTile1d<<<grid, kThreads1d>>>(d_a, d_b, d_c, n, k);
    } else if (which == 2) {
        const dim3 grid(tilesX, tilesY);
        matmulRegTile2d<<<grid, kRegThreads>>>(d_a, d_b, d_c, n, k);
    } else if (which == 3) {
        const dim3 grid(tilesX, tilesY);
        matmulRegTile2dVec4<<<grid, kRegThreads>>>(d_a, d_b, d_c, n, k);
    } else {
        return callCublas(handle, d_a, d_b, d_c, size);
    }
    return CUBLAS_STATUS_SUCCESS;
}

// Times a launch with CUDA events and returns the mean milliseconds per run.
// Copy it verbatim; the alternative is five 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 per thread, static shared memory per
// block, and how many blocks of that shape the hardware can keep resident.
// None of it needs performance counters, so this table is available on a
// tier where Nsight Compute is not.
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_a, float* d_b,
                       float* d_c, float* h_a, float* h_b, float* h_c,
                       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 size_t bytes = elems * sizeof(float);
        const double flop = 2.0 * static_cast<double>(size) *
                            static_cast<double>(size) *
                            static_cast<double>(size);

        fillInputs(h_a, h_b, size);
        matmulCpu(h_a, h_b, h_want, size);
        CUDA_CHECK(cudaMemcpy(d_a, h_a, bytes, cudaMemcpyHostToDevice));
        CUDA_CHECK(cudaMemcpy(d_b, h_b, bytes, cudaMemcpyHostToDevice));

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

        // Correctness first, for every path, before anything is timed. That
        // is not only hygiene: it puts the case's five correctness launches
        // back to back at the start of the case, which is what lets one
        // `ncu --launch-skip S --launch-count 4` collect exactly the four
        // hand-written kernels. The program prints S per case below.
        for (int which = 0; which < kNumKernels; ++which) {
            // Zero C first, so a kernel that writes only part of it is
            // caught by the reference check rather than reading as correct
            // because the previous kernel left the right answer behind.
            CUDA_CHECK(cudaMemset(d_c, 0, bytes));

            const cublasStatus_t st =
                launchMatmul(which, handle, d_a, d_b, d_c, size);
            if (st != CUBLAS_STATUS_SUCCESS) {
                std::fprintf(stderr, "%s failed at size %d: %s\n",
                             kKernelNames[which], size,
                             cublasGetStatusString(st));
                return EXIT_FAILURE;
            }
            CUDA_CHECK(cudaGetLastError());
            CUDA_CHECK(cudaDeviceSynchronize());
            CUDA_CHECK(cudaMemcpy(h_c, d_c, bytes, 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",
                             kKernelNames[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;
            }
        }

        // Launches of the four hand-written kernels per case, counted the
        // way an `ncu --kernel-name regex:'matmul'` filter counts them: one
        // correctness launch each, then kWarmupRuns + kTimedRuns each in
        // the timing pass. cuBLAS launches kernels under its own names, so
        // the filter never sees them.
        const int perCase = 4 * (1 + kWarmupRuns + kTimedRuns);
        std::printf(
            "Profiler: case %d correctness launches are matmul* "
            "launches %d to %d\n(--launch-skip %d --launch-count 4 "
            "under --kernel-name regex:'matmul')\n\n",
            size, caseIdx * perCase, caseIdx * perCase + 3, caseIdx * perCase);

        for (int which = 0; which < kNumKernels; ++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_a, d_b, d_c, 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",
                             kKernelNames[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 %9s %9s %8s %8s %7s %9s\n", "kernel", "ms",
                    "GFLOP/s", "%cuBLAS", "gf/out", "F/byte", "smem/FMA");
        std::printf("%-21s %9s %9s %8s %8s %7s %9s\n", "---------------------",
                    "---------", "---------", "--------", "--------", "-------",
                    "---------");
        for (int which = 0; which < kNumKernels; ++which) {
            const double gflops = flop / (ms[which] * 1.0e-3) / 1.0e9;
            const double pct = 100.0 * ms[kCublasIndex] / ms[which];
            if (which == kCublasIndex) {
                std::printf("%-21s %9.3f %9.1f %8.1f %8s %7s %9s\n",
                            kKernelNames[which], ms[which], gflops, pct, "-",
                            "-", "-");
            } else {
                std::printf(
                    "%-21s %9.3f %9.1f %8.1f %8.1f %7.2f %9.4f\n",
                    kKernelNames[which], ms[which], gflops, pct,
                    floatsPerOutput(kKernelTileM[which], kKernelTileN[which],
                                    size),
                    flopPerByte(kKernelTileM[which], kKernelTileN[which], size),
                    kSmemPerFma[which]);
            }
        }
        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);
    std::printf(
        "%zu bytes of shared memory per block, %d resident threads per SM\n\n",
        prop.sharedMemPerBlock, prop.maxThreadsPerMultiProcessor);

    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("matmulSharedTiled", matmulSharedTiled, kTileDim * kTileDim,
                    prop.maxThreadsPerMultiProcessor);
    printAttributes("matmulRegTile1d", matmulRegTile1d, kThreads1d,
                    prop.maxThreadsPerMultiProcessor);
    printAttributes("matmulRegTile2d", matmulRegTile2d, kRegThreads,
                    prop.maxThreadsPerMultiProcessor);
    printAttributes("matmulRegTile2dVec4", matmulRegTile2dVec4, kRegThreads,
                    prop.maxThreadsPerMultiProcessor);
    std::printf(
        "\nlocal B above is spill: anything but zero means the "
        "micro-tile did not fit\nthe register file and the "
        "accumulators went to local memory.\n\n");

    std::vector<float> h_a(kMaxElems);
    std::vector<float> h_b(kMaxElems);
    std::vector<float> h_c(kMaxElems);
    std::vector<double> h_want(kMaxElems);

    float* d_a = nullptr;
    float* d_b = nullptr;
    float* d_c = nullptr;
    const size_t maxBytes = kMaxElems * sizeof(float);
    CUDA_CHECK(cudaMalloc(&d_a, maxBytes));
    CUDA_CHECK(cudaMalloc(&d_b, maxBytes));
    CUDA_CHECK(cudaMalloc(&d_c, maxBytes));

    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, even though it is the default, so the comparison
        // this program makes is stated in the program and not only in the
        // prose. CUBLAS_DEFAULT_MATH plus CUBLAS_COMPUTE_32F is FP32 all the
        // way through.
        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_a, d_b, d_c, h_a.data(), h_b.data(),
                             h_c.data(), h_want.data());
    }

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

    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_a));
    CUDA_CHECK(cudaFree(d_b));
    CUDA_CHECK(cudaFree(d_c));
    return status;
}