COURSE / SOURCE

cutlass_gemm.cu

All lessons
Source filecode/day84-cutlass/cutlass_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 84 part 2: one CUTLASS device GEMM, instantiated twice from one source.
//
// FP16 in, FP32 accumulate, FP32 out. A is row major, B is column major, D is
// row major, alpha = 1 and beta = 0, so the only thing that moves between the
// two builds is the hardware the kernel is compiled for:
//
//                    default build            -DDAY84_SM80
//   arch tag         cutlass::arch::Sm75      cutlass::arch::Sm80
//   InstructionShape 16 x 8 x 8               16 x 8 x 16
//   Stages           2                        3
//
// Those two rows are CUTLASS's own defaults for the two architectures, from
// DefaultGemmConfiguration<OpClassTensorOp, Sm75|Sm80, ...> in
// include/cutlass/gemm/device/default_gemm_configuration.h at v4.7.1. They
// are written out here rather than defaulted so that the page, the SASS dump
// and the binary cannot disagree about what was built.
//
// The threadblock and warp tiles are held at 128 x 128 x 32 and 64 x 64 x 32
// in both builds, which is not the CUTLASS default for either. The default
// deepens the K slice to 64 on Sm80, and 128 x 256 x 64 at three stages asks
// for 147,456 bytes of shared memory per block: an A100 has that and an
// RTX 30 or 40, capped at 99 KB per block, does not. Holding the tile fixed
// keeps this comparison to two moving parts and keeps both builds inside the
// 48 KiB every card in the course's matrix has.
//
// The default build runs on the course's Tesla T4. The -DDAY84_SM80 build
// does not: three-stage mainloops fetch through cp.async, which "Requires
// sm_80 or higher" (PTX ISA, cp.async Target ISA Notes), so the capability
// gate below refuses the T4 and says which build to use. Both builds compile
// on the T4 node, and the SASS this day reads comes out of the binary with
// cuobjdump, which needs no GPU at all. Day 46 is where that habit started.
//
// Needs CUTLASS's headers, which are header only and build nothing:
//   git clone --depth 1 --branch v4.7.1 https://github.com/NVIDIA/cutlass.git
//
// Build (Turing and newer, runs on a T4):
//   nvcc -std=c++17 -O3 -arch=sm_75 --expt-relaxed-constexpr \
//        -I cutlass/include -o cutlass_gemm_sm75 cutlass_gemm.cu
// Build (Ampere and newer, needs CC 8.0 to run):
//   nvcc -std=c++17 -O3 -arch=sm_80 --expt-relaxed-constexpr -DDAY84_SM80 \
//        -I cutlass/include -o cutlass_gemm_sm80 cutlass_gemm.cu
//
// VERIFIED 2026-09-02: the sm_75 build passed its numerical gate on a Tesla
// T4. The sm_80 build passed compilation and SASS inspection, then produced
// the expected CC 8.0 capability refusal on the T4. See evidence/.

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

#include <cuda_runtime.h>

#include <cutlass/cutlass.h>
#include <cutlass/epilogue/thread/linear_combination.h>
#include <cutlass/gemm/device/gemm.h>
#include <cutlass/gemm/gemm.h>
#include <cutlass/gemm/threadblock/threadblock_swizzle.h>
#include <cutlass/layout/matrix.h>
#include <cutlass/numeric_types.h>
#include <cutlass/tensor_ref.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)

// CUTLASS returns cutlass::Status, not cudaError_t, so CUDA_CHECK cannot wrap
// it and CUDA-CODE-STYLE.md allows no second error macro. Every CUTLASS call
// in this file hands its status to this function and the caller branches on
// the answer. Day 44 shipped a bug by dropping a library status inside a
// lambda, which is the reason none of them is discarded here.
static bool cutlassOk(cutlass::Status st, const char* what) {
    if (st == cutlass::Status::kSuccess) {
        return true;
    }
    std::fprintf(stderr, "CUTLASS error: %s: %s\n", what,
                 cutlassGetStatusString(st));
    return false;
}

using ElementInput = cutlass::half_t;
using ElementOutput = float;
using ElementAccumulator = float;
using ElementCompute = float;

using LayoutInputA = cutlass::layout::RowMajor;
using LayoutInputB = cutlass::layout::ColumnMajor;
using LayoutOutput = cutlass::layout::RowMajor;

// OpClassTensorOp is the choice that decides whether HMMA appears in the SASS
// at all. OpClassSimt would compile on the same card and issue FFMA instead,
// which is day 44's kernel with more template arguments.
using MmaOpClass = cutlass::arch::OpClassTensorOp;

#if defined(DAY84_SM80)
using SmArch = cutlass::arch::Sm80;
using InstructionShape = cutlass::gemm::GemmShape<16, 8, 16>;
constexpr int kStages = 3;
constexpr int kMinMajor = 8;
constexpr int kMinMinor = 0;
constexpr const char* kBuildName = "sm_80, mma.m16n8k16, 3 stages";
#else
using SmArch = cutlass::arch::Sm75;
using InstructionShape = cutlass::gemm::GemmShape<16, 8, 8>;
constexpr int kStages = 2;
constexpr int kMinMajor = 7;
constexpr int kMinMinor = 5;
constexpr const char* kBuildName = "sm_75, mma.m16n8k8, 2 stages";
#endif

using ThreadblockShape = cutlass::gemm::GemmShape<128, 128, 32>;
using WarpShape = cutlass::gemm::GemmShape<64, 64, 32>;

// 2 x 2 warps per block tile, so 128 threads. The block tile divided by the
// warp tile is the warp grid, and the warp tile divided by the instruction
// shape is how many mma instructions each warp issues per K slice.
constexpr int kWarpsPerBlock = (128 / 64) * (128 / 64);
constexpr int kThreadsPerBlock = kWarpsPerBlock * 32;

// The mainloop's shared memory per block: both operand tiles, in half, times
// the stage count. 32 KiB at two stages and 48 KiB at three. CUTLASS unions
// the epilogue's shared storage with the mainloop's, so this is the figure
// that decides whether the kernel fits; the run prints it beside the card so
// a reader can check it against Nsight Compute.
constexpr int kSharedBytes =
    (128 * 32 + 128 * 32) * static_cast<int>(sizeof(ElementInput)) * kStages;

static_assert(kThreadsPerBlock % 32 == 0,
              "block size must be a whole number of warps");
static_assert(kSharedBytes <= 48 * 1024,
              "over 48 KiB per block needs the dynamic shared memory opt-in "
              "and stops fitting some cards; shrink the tile or the stages");

using EpilogueOp = cutlass::epilogue::thread::LinearCombination<
    ElementOutput, 128 / cutlass::sizeof_bits<ElementOutput>::value,
    ElementAccumulator, ElementCompute>;

using SwizzleThreadBlock =
    cutlass::gemm::threadblock::GemmIdentityThreadblockSwizzle<>;

using Gemm = cutlass::gemm::device::Gemm<
    ElementInput, LayoutInputA, ElementInput, LayoutInputB, ElementOutput,
    LayoutOutput, ElementAccumulator, MmaOpClass, SmArch, ThreadblockShape,
    WarpShape, InstructionShape, EpilogueOp, SwizzleThreadBlock, kStages>;

// 512 is a whole number of 128 x 128 tiles in M and N and of 32-element K
// slices, so no ragged edge runs and nothing here measures the epilogue path.
constexpr int kM = 512;
constexpr int kN = 512;
constexpr int kK = 512;

// FP32 accumulation, so eps is 2^-23. Day 66's rule: rounding across a K-term
// dot product grows like sqrt(K) because the errors are uncorrelated, and the
// factor 4 is slack for a summation order that is valid and different. The
// inputs are read as the same half values on both sides, so this tolerance
// covers accumulation order and nothing else.
constexpr double kBaseRtol = 1e-5;
constexpr int kF32MantissaBits = 23;

static_assert(kM % 128 == 0 && kN % 128 == 0 && kK % 32 == 0,
              "the problem must be a whole number of threadblock tiles");
static_assert(kK % 8 == 0,
              "a 128-bit aligned half operand needs K divisible by 8");

static double scaledRtol(double tableRtol, int mantissaBits, int kDim) {
    const double eps = std::ldexp(1.0, -mantissaBits);
    const double grown = 4.0 * eps * std::sqrt(static_cast<double>(kDim));
    return grown > tableRtol ? grown : tableRtol;
}

// A deterministic host filler. No RNG library, so the same values appear on
// every machine and a mismatch report is reproducible.
static float nextValue(unsigned int* state) {
    *state = *state * 1664525u + 1013904223u;
    const float unit = static_cast<float>((*state >> 8) & 0xFFFFu) / 65535.0f;
    return 2.0f * unit - 1.0f;
}

// The reference. Plain loops, accumulating in double, reading the same half
// values the GPU read: h_a is row major (M, K) and h_b is column major
// (K, N), so B's element (k, n) sits at n * K + k.
static void gemmCpu(const std::vector<ElementInput>& h_a,
                    const std::vector<ElementInput>& h_b,
                    std::vector<double>& h_want) {
    for (int m = 0; m < kM; ++m) {
        for (int n = 0; n < kN; ++n) {
            double acc = 0.0;
            for (int k = 0; k < kK; ++k) {
                const double a = static_cast<double>(
                    static_cast<float>(h_a[static_cast<size_t>(m) * kK + k]));
                const double b = static_cast<double>(
                    static_cast<float>(h_b[static_cast<size_t>(n) * kK + k]));
                acc += a * b;
            }
            h_want[static_cast<size_t>(m) * kN + n] = acc;
        }
    }
}

int main() {
    cudaDeviceProp prop;
    CUDA_CHECK(cudaGetDeviceProperties(&prop, 0));
    std::printf("GPU: %s (compute capability %d.%d), %d SMs\n", prop.name,
                prop.major, prop.minor, prop.multiProcessorCount);
    std::printf("build: %s\n", kBuildName);
    std::printf(
        "threadblock 128x128x32, warp 64x64x32, %d threads, "
        "%d B shared per block\n",
        kThreadsPerBlock, kSharedBytes);
    std::printf("problem: %d x %d x %d, half in, float accumulate\n\n", kM, kN,
                kK);

    if (prop.major * 10 + prop.minor < kMinMajor * 10 + kMinMinor) {
        std::fprintf(stderr,
                     "this build needs compute capability %d.%d and this "
                     "device is %d.%d.\nBuild without -DDAY84_SM80 at "
                     "-arch=sm_75 to run the Turing path here, or run this "
                     "binary on a CC 8.0 card.\n",
                     kMinMajor, kMinMinor, prop.major, prop.minor);
        return EXIT_FAILURE;
    }

    const size_t elemsA = static_cast<size_t>(kM) * kK;
    const size_t elemsB = static_cast<size_t>(kK) * kN;
    const size_t elemsD = static_cast<size_t>(kM) * kN;

    std::vector<ElementInput> h_a(elemsA);
    std::vector<ElementInput> h_b(elemsB);
    std::vector<float> h_got(elemsD);
    std::vector<double> h_want(elemsD);

    unsigned int state = 20260901u;
    for (size_t e = 0; e < elemsA; ++e) {
        h_a[e] = ElementInput(nextValue(&state));
    }
    for (size_t e = 0; e < elemsB; ++e) {
        h_b[e] = ElementInput(nextValue(&state));
    }

    ElementInput* d_a = nullptr;
    ElementInput* d_b = nullptr;
    ElementOutput* d_c = nullptr;
    ElementOutput* d_d = nullptr;
    uint8_t* d_workspace = nullptr;
    CUDA_CHECK(cudaMalloc(&d_a, elemsA * sizeof(ElementInput)));
    CUDA_CHECK(cudaMalloc(&d_b, elemsB * sizeof(ElementInput)));
    CUDA_CHECK(cudaMalloc(&d_c, elemsD * sizeof(ElementOutput)));
    CUDA_CHECK(cudaMalloc(&d_d, elemsD * sizeof(ElementOutput)));
    CUDA_CHECK(cudaMemcpy(d_a, h_a.data(), elemsA * sizeof(ElementInput),
                          cudaMemcpyHostToDevice));
    CUDA_CHECK(cudaMemcpy(d_b, h_b.data(), elemsB * sizeof(ElementInput),
                          cudaMemcpyHostToDevice));
    CUDA_CHECK(cudaMemset(d_c, 0, elemsD * sizeof(ElementOutput)));
    CUDA_CHECK(cudaMemset(d_d, 0, elemsD * sizeof(ElementOutput)));

    int status = EXIT_SUCCESS;

    // The leading dimensions are the strides the layouts carry: K for a row
    // major A, K for a column major B, N for a row major D.
    // snippet: instantiate
    cutlass::TensorRef<ElementInput const, LayoutInputA> refA(d_a,
                                                              LayoutInputA(kK));
    cutlass::TensorRef<ElementInput const, LayoutInputB> refB(d_b,
                                                              LayoutInputB(kK));
    cutlass::TensorRef<ElementOutput const, LayoutOutput> refC(
        d_c, LayoutOutput(kN));
    cutlass::TensorRef<ElementOutput, LayoutOutput> refD(d_d, LayoutOutput(kN));

    typename Gemm::Arguments args(cutlass::gemm::GemmCoord(kM, kN, kK), refA,
                                  refB, refC, refD,
                                  {ElementCompute(1.0f), ElementCompute(0.0f)},
                                  1);  // split-k slices
    Gemm gemmOp;
    // end snippet

    if (!cutlassOk(Gemm::can_implement(args), "can_implement")) {
        status = EXIT_FAILURE;
    }

    size_t workspaceBytes = 0;
    if (status == EXIT_SUCCESS) {
        workspaceBytes = Gemm::get_workspace_size(args);
        std::printf("CUTLASS workspace: %zu B\n", workspaceBytes);
        if (workspaceBytes > 0) {
            CUDA_CHECK(cudaMalloc(&d_workspace, workspaceBytes));
        }
        if (!cutlassOk(gemmOp.initialize(args, d_workspace), "initialize")) {
            status = EXIT_FAILURE;
        }
    }

    if (status == EXIT_SUCCESS) {
        if (!cutlassOk(gemmOp(), "gemm launch")) {
            status = EXIT_FAILURE;
        }
    }

    if (status == EXIT_SUCCESS) {
        CUDA_CHECK(cudaGetLastError());
        CUDA_CHECK(cudaDeviceSynchronize());
        CUDA_CHECK(cudaMemcpy(h_got.data(), d_d, elemsD * sizeof(ElementOutput),
                              cudaMemcpyDeviceToHost));

        gemmCpu(h_a, h_b, h_want);

        const double rtol = scaledRtol(kBaseRtol, kF32MantissaBits, kK);
        const double atol = 1e-4;
        std::printf("tolerance: atol %.1e, rtol %.6e (4 * 2^-%d * sqrt(%d))\n",
                    atol, rtol, kF32MantissaBits, kK);

        double worst = 0.0;
        size_t worstAt = 0;
        for (size_t e = 0; e < elemsD; ++e) {
            const double got = static_cast<double>(h_got[e]);
            const double want = h_want[e];
            if (!std::isfinite(got)) {
                std::fprintf(stderr, "element %zu is not finite\n", e);
                status = EXIT_FAILURE;
                break;
            }
            const double budget = atol + rtol * std::fabs(want);
            const double share = std::fabs(got - want) / budget;
            if (share > worst) {
                worst = share;
                worstAt = e;
            }
        }
        if (status == EXIT_SUCCESS) {
            std::printf(
                "worst error is %.3f of the budget, at element %zu "
                "(row %zu, col %zu)\n",
                worst, worstAt, worstAt / kN, worstAt % kN);
            if (worst > 1.0) {
                std::fprintf(stderr,
                             "CUTLASS output is outside the tolerance at "
                             "element %zu\n",
                             worstAt);
                status = EXIT_FAILURE;
            } else {
                std::printf(
                    "all %zu output elements are inside the "
                    "tolerance\n",
                    elemsD);
            }
        }
    }

    if (d_workspace != nullptr) {
        CUDA_CHECK(cudaFree(d_workspace));
    }
    CUDA_CHECK(cudaFree(d_a));
    CUDA_CHECK(cudaFree(d_b));
    CUDA_CHECK(cudaFree(d_c));
    CUDA_CHECK(cudaFree(d_d));
    return status;
}