COURSE / SOURCE

fusion.cu

All lessons
Source filecode/day48-fusion/fusion.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 48: kernel fusion and launch overhead, priced apart.
//
// One four-stage elementwise chain, written twice:
//
//   staged   four kernels. Each reads its input from global memory and
//            writes its output back, so the chain moves 36 bytes per
//            element and pays four launches.
//   fused    one kernel. The three intermediates never leave a register, so
//            the chain moves 12 bytes per element and pays one launch.
//
// Fusing deletes two different things, and they are not the same size. This
// program prices them separately.
//
//   Part 1  what a launch costs on its own, from a kernel that does nothing.
//   Part 2  the chain at 16 Mi elements, where the traffic is milliseconds
//           and the launches are microseconds.
//   Part 3  the same two chains swept down in size until that reverses.
//   Part 4  what fusion costs. The fused kernel is widened until each thread
//           holds enough live values to change how many warps fit on an SM.
//
// The byte counts are the prediction, so they are computed from the
// constants and printed above the times. Stage 3 reads the residual as well
// as its own input, so the staged chain moves 9 floats per element and not
// 8: a model that takes one stage and multiplies by four gets 32 bytes and
// is wrong before it ever meets a GPU.
//
// Build: nvcc -std=c++17 -O3 -arch=sm_75 -o fusion fusion.cu
// Run:   ./fusion
//

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

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

// 16 Mi elements is 64 MiB per buffer and six buffers on the device. It is
// chosen large enough that the chain's traffic is milliseconds while its
// launches stay microseconds, which is the gap part 2 is about.
constexpr size_t kElems = 16777216;

// The correctness size. 611 is 13 x 47, so no block size and no elements-
// per-thread width divides it and every bounds check runs on every launch.
// The timed size above is a power of two on purpose, so no guard fires
// inside a timed kernel and the timings are not measuring tail handling.
constexpr size_t kCheckElems = 1048576 + 611;

constexpr int kThreadsPerBlock = 256;  // 8 warps
constexpr int kWarmupRuns = 3;
constexpr int kTimedRuns = 10;

// Part 1 launches an empty kernel this many times inside one timed call, so
// the event pair brackets something far larger than its own half-microsecond
// resolution. timeKernel then repeats that batch, which is why the total is
// (kWarmupRuns + kTimedRuns) x kLaunchesPerBatch launches.
constexpr int kLaunchesPerBatch = 1000;

// Part 3. Every value must fit the buffers allocated for kElems.
constexpr size_t kSweep[] = {1024, 16384, 262144, 4194304, 16777216};
constexpr int kSweepCount =
    static_cast<int>(sizeof(kSweep) / sizeof(kSweep[0]));

// Part 4. A template argument cannot come from a loop, so these six widths
// are instantiated by hand in main and this array exists only so that a
// static_assert can hold the two lists together.
constexpr int kWideK[] = {1, 2, 4, 8, 16, 32};
constexpr int kWideCount = static_cast<int>(sizeof(kWideK) / sizeof(kWideK[0]));

// The chain's parameters. kHi is 3 rather than something safely above the
// range so that the clamp actually fires on about a quarter of the elements;
// a stage that never does anything is a stage the compiler may delete, and
// then the byte model is describing a program that no longer exists.
constexpr float kGamma = 1.5f;
constexpr float kBeta = -0.25f;
constexpr float kLo = -6.0f;
constexpr float kHi = 3.0f;

// The tanh form of GELU, which is what GPT-2 used and what PyTorch calls
// approximate='tanh'. kGeluA is sqrt(2 / pi) rounded to float; the CPU
// reference promotes these same rounded floats to double rather than using
// the exact constants, so the reference is checking the kernel's arithmetic
// and not the author's choice of literal.
constexpr float kGeluA = 0.7978845608f;
constexpr float kGeluB = 0.044715f;

// Floats moved per element. Staged: 2 for stage 1, 2 for stage 2, 3 for
// stage 3 because it also reads the residual, 2 for stage 4. Fused: x, the
// residual, and the output. These two numbers are the whole prediction.
constexpr double kStagedFloatsPerElem = 9.0;
constexpr double kFusedFloatsPerElem = 3.0;
constexpr double kCopyFloatsPerElem = 2.0;

constexpr double kRelTolerance = 1e-5;
constexpr double kAbsTolerance = 1e-6;

static_assert(kThreadsPerBlock % 32 == 0,
              "the block is a whole number of warps");
static_assert(kThreadsPerBlock <= 1024,
              "1024 threads is the per-block maximum on every compute "
              "capability this course targets");
static_assert(kElems % 32 == 0,
              "the timed size must be a whole number of blocks at every "
              "block size this file invites, down to one warp, so that no "
              "bounds check runs inside a timed kernel");
static_assert(kCheckElems % 2 == 1,
              "the correctness size must be odd, so that no power-of-two "
              "launch shape covers it exactly and every bounds check fires "
              "on every launch");
static_assert(kSweep[kSweepCount - 1] <= kElems,
              "the sweep reuses the buffers allocated for kElems, so its "
              "largest size cannot exceed them");
static_assert(kWideK[0] == 1 && kWideK[1] == 2 && kWideK[2] == 4 &&
                  kWideK[3] == 8 && kWideK[4] == 16 && kWideK[5] == 32,
              "the six explicit template instantiations in main must match "
              "this list, because a template argument cannot come from a "
              "loop and nothing else keeps the two in step");
static_assert(kHi > kLo, "the clamp range must be a range");

// The four stages, as ordinary device functions, so that the staged kernels
// and every fused kernel below run byte-identical arithmetic and the only
// thing that differs between them is where the intermediates live.
__device__ __forceinline__ float scaleBias(float v) {
    return v * kGamma + kBeta;
}

__device__ __forceinline__ float gelu(float v) {
    const float inner = kGeluA * (v + kGeluB * v * v * v);
    return 0.5f * v * (1.0f + tanhf(inner));
}

__device__ __forceinline__ float addResidual(float v, float r) {
    return v + r;
}

__device__ __forceinline__ float clampRange(float v) {
    return fminf(fmaxf(v, kLo), kHi);
}

// Does nothing, and is the point of part 1.
//
// Memory: none. It issues no loads and no stores, so whatever the clock says
// about a run of these is the cost of getting a kernel started and finished.
//
// Launch assumption: none. Any grid, any block.
__global__ void doNothing() {}

// One read and one write per element, no arithmetic.
//
// Memory: consecutive lanes take consecutive elements, so one warp's 32
// addresses cover 128 contiguous bytes and cost four 32-byte sectors.
//
// The floor row. No elementwise chain over this buffer can beat a kernel
// that only reads it and writes it, and measuring the floor here rather than
// borrowing day 11's number keeps this program's block shape, buffer size
// and clock state out of the comparison.
//
// Launch assumption: any 1D grid that covers n.
__global__ void copyFloor(const float* __restrict__ in, float* __restrict__ out,
                          size_t n) {
    const size_t i = blockIdx.x * static_cast<size_t>(blockDim.x) + threadIdx.x;
    if (i < n) {
        out[i] = in[i];
    }
}

// Stage 1 of 4. One thread scales and shifts one element.
//
// Memory: 32 consecutive lanes read 32 consecutive floats and write 32
// consecutive floats, so both sides coalesce into four sectors per warp. The
// access pattern is not what this program is measuring; the number of times
// the pattern happens is.
//
// Launch assumption: gridDim.x * blockDim.x >= n.
__global__ void applyScaleBias(const float* __restrict__ in,
                               float* __restrict__ out, size_t n) {
    const size_t i = blockIdx.x * static_cast<size_t>(blockDim.x) + threadIdx.x;
    if (i < n) {
        out[i] = scaleBias(in[i]);
    }
}

// Stage 2 of 4. One thread takes the GELU of one element.
//
// Memory: as stage 1, one coalesced read and one coalesced write. This is
// the only stage with a transcendental in it, and it is here so that the
// chain has some real arithmetic to be memory-bound in spite of.
//
// Launch assumption: gridDim.x * blockDim.x >= n.
__global__ void applyGelu(const float* __restrict__ in, float* __restrict__ out,
                          size_t n) {
    const size_t i = blockIdx.x * static_cast<size_t>(blockDim.x) + threadIdx.x;
    if (i < n) {
        out[i] = gelu(in[i]);
    }
}

// Stage 3 of 4. One thread adds one residual element to one input element.
//
// Memory: two coalesced reads and one coalesced write, so this stage moves
// three floats per element where the others move two. That asymmetry is
// deliberate. A staged chain of four two-float stages would move 32 bytes
// per element and a reader could get the right ratio from the wrong model.
//
// Launch assumption: gridDim.x * blockDim.x >= n.
__global__ void applyAddResidual(const float* __restrict__ in,
                                 const float* __restrict__ r,
                                 float* __restrict__ out, size_t n) {
    const size_t i = blockIdx.x * static_cast<size_t>(blockDim.x) + threadIdx.x;
    if (i < n) {
        out[i] = addResidual(in[i], r[i]);
    }
}

// Stage 4 of 4. One thread clamps one element into [kLo, kHi].
//
// Memory: one coalesced read and one coalesced write.
//
// Launch assumption: gridDim.x * blockDim.x >= n.
__global__ void applyClampRange(const float* __restrict__ in,
                                float* __restrict__ out, size_t n) {
    const size_t i = blockIdx.x * static_cast<size_t>(blockDim.x) + threadIdx.x;
    if (i < n) {
        out[i] = clampRange(in[i]);
    }
}

// All four stages in one thread, with the three intermediates living in
// registers instead of in global memory.
//
// Memory: two coalesced reads and one coalesced write per element, 12 bytes
// against the staged chain's 36. The arithmetic is the same four device
// functions in the same order, so anything the clock finds between this and
// the staged chain is traffic and launches, not maths.
//
// Launch assumption: gridDim.x * blockDim.x >= n.
// snippet: fused-chain
__global__ void chainFused(const float* __restrict__ x,
                           const float* __restrict__ r, float* __restrict__ out,
                           size_t n) {
    const size_t i = blockIdx.x * static_cast<size_t>(blockDim.x) + threadIdx.x;
    if (i < n) {
        float v = scaleBias(x[i]);
        v = gelu(v);
        v = addResidual(v, r[i]);
        out[i] = clampRange(v);
    }
}
// end snippet: fused-chain

// The same fused chain, with each thread carrying kWide elements instead of
// one. Same total work, same total traffic, same arithmetic, same order.
// The only thing that changes is how many values are live at once.
//
// Memory: for a fixed k the 32 lanes of a warp read 32 consecutive floats,
// because the elements a thread owns are one whole grid apart rather than
// adjacent. That is the grid-stride layout from day 8, unrolled: a thread
// takes elements base, base + stride, base + 2 * stride, and so on, so every
// one of the kWide loads is a coalesced 128-byte span per warp.
//
// The three loops are separate on purpose. Loading all kWide values, then
// transforming all kWide, then storing all kWide, is what puts kWide
// requests in flight per thread, and it is also what makes the compiler hold
// kWide floats live across the middle loop. Both halves of that trade are
// the point: more work in flight per thread, fewer threads that fit.
//
// Launch assumption: gridDim.x * blockDim.x * kWide >= n. With that, the
// union of the indices below covers [0, n) exactly once.
// snippet: fused-wide
template <int kWide>
__global__ void chainFusedWide(const float* __restrict__ x,
                               const float* __restrict__ r,
                               float* __restrict__ out, size_t n) {
    const size_t base =
        blockIdx.x * static_cast<size_t>(blockDim.x) + threadIdx.x;
    const size_t stride = gridDim.x * static_cast<size_t>(blockDim.x);

    float v[kWide];
#pragma unroll
    for (int k = 0; k < kWide; ++k) {
        const size_t i = base + static_cast<size_t>(k) * stride;
        v[k] = (i < n) ? scaleBias(x[i]) : 0.0f;
    }
#pragma unroll
    for (int k = 0; k < kWide; ++k) {
        const size_t i = base + static_cast<size_t>(k) * stride;
        v[k] = addResidual(gelu(v[k]), (i < n) ? r[i] : 0.0f);
    }
#pragma unroll
    for (int k = 0; k < kWide; ++k) {
        const size_t i = base + static_cast<size_t>(k) * stride;
        if (i < n) {
            out[i] = clampRange(v[k]);
        }
    }
}
// end snippet: fused-wide

// The chain in double on the host. Written for obvious correctness, not
// speed: plain loop, no OpenMP, no intrinsics. It never allocates; the
// caller owns every buffer.
//
// The float constants promote to double rather than being respelled as
// double literals, so the reference checks the kernel's arithmetic and not
// the author's choice of how many digits to type.
static void chainCpu(const float* x, const float* r, double* out, size_t n) {
    for (size_t i = 0; i < n; ++i) {
        double v = static_cast<double>(x[i]) * kGamma + kBeta;
        const double inner = kGeluA * (v + kGeluB * v * v * v);
        v = 0.5 * v * (1.0 + std::tanh(inner));
        v = v + static_cast<double>(r[i]);
        if (v < kLo) {
            v = kLo;
        }
        if (v > kHi) {
            v = kHi;
        }
        out[i] = v;
    }
}

// Returns the first index where got and want differ by more than the
// tolerance, or n if they agree everywhere. Returning the index rather than
// a bool is the whole point: "wrong at 512" names the block, "wrong" does
// not.
static size_t firstMismatch(const float* got, const double* want, size_t n) {
    for (size_t i = 0; i < n; ++i) {
        const double scale = std::fabs(want[i]);
        const double allowed = kAbsTolerance + kRelTolerance * scale;
        if (std::fabs(static_cast<double>(got[i]) - want[i]) > allowed) {
            return i;
        }
    }
    return n;
}

// Times a launch with CUDA events and returns the mean milliseconds per run.
// A host-side clock around a launch measures the launch, not the kernel,
// because launches are asynchronous. Day 9 takes that apart.
//
// This is the one template and the one lambda allowed in module 1 to 3 code.
// Copy it verbatim.
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.
    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;
}

static int blocksFor(size_t n) {
    return static_cast<int>((n + kThreadsPerBlock - 1) / kThreadsPerBlock);
}

// Effective bandwidth from the compulsory byte count, not from a profiler.
// This is what the program can know: the bytes the algorithm has to move.
// What actually crossed the bus is dram__bytes in the shipped Nsight Compute
// report, and the gap between the two is the caches.
static double bandwidthGBs(double floatsPerElem, size_t n, float ms) {
    const double bytes = floatsPerElem * static_cast<double>(n) * sizeof(float);
    return bytes / (static_cast<double>(ms) * 1.0e-3) / 1.0e9;
}

// The four staged launches, in order, on the default stream. Enqueued and
// not synchronized: the caller decides where the measurement ends.
static void launchStaged(const float* d_x, const float* d_r, float* d_t1,
                         float* d_t2, float* d_t3, float* d_y, size_t n) {
    const int blocks = blocksFor(n);
    applyScaleBias<<<blocks, kThreadsPerBlock>>>(d_x, d_t1, n);
    applyGelu<<<blocks, kThreadsPerBlock>>>(d_t1, d_t2, n);
    applyAddResidual<<<blocks, kThreadsPerBlock>>>(d_t2, d_r, d_t3, n);
    applyClampRange<<<blocks, kThreadsPerBlock>>>(d_t3, d_y, n);
}

struct WideRow {
    int wide;
    int blocks;
    int regs;
    size_t spillBytes;
    int blocksPerSm;
    double occupancy;
    float ms;
    size_t bad;
};

// Runs one width twice: once at the correctness size, where the result is
// compared against the double reference, and once at the timed size, where
// it is measured. Registers and theoretical occupancy come from the runtime
// API rather than from a profiler, so this table needs no performance
// counters and no root.
template <int kWide>
static void probeWide(WideRow* row, const float* d_x, const float* d_r,
                      float* d_out, float* h_scratch, const double* h_want,
                      int maxThreadsPerSm) {
    const size_t perBlock =
        static_cast<size_t>(kWide) * static_cast<size_t>(kThreadsPerBlock);
    const int checkBlocks =
        static_cast<int>((kCheckElems + perBlock - 1) / perBlock);
    const int blocks = static_cast<int>((kElems + perBlock - 1) / perBlock);

    chainFusedWide<kWide>
        <<<checkBlocks, kThreadsPerBlock>>>(d_x, d_r, d_out, kCheckElems);
    CUDA_CHECK(cudaGetLastError());
    CUDA_CHECK(cudaDeviceSynchronize());
    CUDA_CHECK(cudaMemcpy(h_scratch, d_out, kCheckElems * sizeof(float),
                          cudaMemcpyDeviceToHost));
    row->bad = firstMismatch(h_scratch, h_want, kCheckElems);

    cudaFuncAttributes attr;
    CUDA_CHECK(cudaFuncGetAttributes(&attr, chainFusedWide<kWide>));
    int blocksPerSm = 0;
    CUDA_CHECK(cudaOccupancyMaxActiveBlocksPerMultiprocessor(
        &blocksPerSm, chainFusedWide<kWide>, kThreadsPerBlock, 0));

    row->wide = kWide;
    row->blocks = blocks;
    row->regs = attr.numRegs;
    row->spillBytes = attr.localSizeBytes;
    row->blocksPerSm = blocksPerSm;
    row->occupancy = 100.0 * static_cast<double>(blocksPerSm) *
                     kThreadsPerBlock / static_cast<double>(maxThreadsPerSm);
    row->ms = timeKernel([&] {
        chainFusedWide<kWide>
            <<<blocks, kThreadsPerBlock>>>(d_x, d_r, d_out, kElems);
    });
}

// Runs one kernel or chain at the correctness size and compares the result
// against the double reference. Returns the first bad index, or kCheckElems
// when everything matched.
static size_t checkAtCheckSize(bool fused, const float* d_x, const float* d_r,
                               float* d_t1, float* d_t2, float* d_t3,
                               float* d_y, float* h_scratch,
                               const double* h_want) {
    if (fused) {
        chainFused<<<blocksFor(kCheckElems), kThreadsPerBlock>>>(d_x, d_r, d_y,
                                                                 kCheckElems);
    } else {
        launchStaged(d_x, d_r, d_t1, d_t2, d_t3, d_y, kCheckElems);
    }
    CUDA_CHECK(cudaGetLastError());
    CUDA_CHECK(cudaDeviceSynchronize());
    CUDA_CHECK(cudaMemcpy(h_scratch, d_y, kCheckElems * sizeof(float),
                          cudaMemcpyDeviceToHost));
    return firstMismatch(h_scratch, h_want, kCheckElems);
}

int main() {
    const int device = 0;
    CUDA_CHECK(cudaSetDevice(device));
    cudaDeviceProp prop;
    CUDA_CHECK(cudaGetDeviceProperties(&prop, device));

    const size_t bytes = kElems * sizeof(float);
    const double mib = static_cast<double>(bytes) / (1024.0 * 1024.0);

    // Registers per thread at which this card still fits its full complement
    // of resident threads on one SM. Derived from the card, not asserted:
    // the number is 64 on a T4 and it is not 64 everywhere.
    const int regsAtFullOccupancy =
        prop.regsPerMultiprocessor / prop.maxThreadsPerMultiProcessor;

    std::vector<float> h_x(kElems);
    std::vector<float> h_r(kElems);
    std::vector<float> h_staged(kElems);
    std::vector<float> h_fused(kElems);
    std::vector<double> h_want(kCheckElems);

    // x sweeps [-4, 4) so that stage 1 produces [-6.25, 5.75) and the clamp
    // at kHi = 3 catches a real share of the elements rather than none.
    for (size_t i = 0; i < kElems; ++i) {
        h_x[i] = (static_cast<float>(i % 2048) / 1024.0f - 1.0f) * 4.0f;
        h_r[i] = static_cast<float>(i % 97) / 97.0f - 0.5f;
    }
    chainCpu(h_x.data(), h_r.data(), h_want.data(), kCheckElems);

    // How much of the reference sits on a clamp. A stage that never fires is
    // a stage that might not be there.
    size_t clamped = 0;
    for (size_t i = 0; i < kCheckElems; ++i) {
        if (h_want[i] == static_cast<double>(kHi) ||
            h_want[i] == static_cast<double>(kLo)) {
            ++clamped;
        }
    }

    float* d_x = nullptr;
    float* d_r = nullptr;
    float* d_t1 = nullptr;
    float* d_t2 = nullptr;
    float* d_t3 = nullptr;
    float* d_y = nullptr;
    CUDA_CHECK(cudaMalloc(&d_x, bytes));
    CUDA_CHECK(cudaMalloc(&d_r, bytes));
    CUDA_CHECK(cudaMalloc(&d_t1, bytes));
    CUDA_CHECK(cudaMalloc(&d_t2, bytes));
    CUDA_CHECK(cudaMalloc(&d_t3, bytes));
    CUDA_CHECK(cudaMalloc(&d_y, bytes));
    CUDA_CHECK(cudaMemcpy(d_x, h_x.data(), bytes, cudaMemcpyHostToDevice));
    CUDA_CHECK(cudaMemcpy(d_r, h_r.data(), bytes, cudaMemcpyHostToDevice));

    int status = EXIT_SUCCESS;

    const size_t badStaged = checkAtCheckSize(
        false, d_x, d_r, d_t1, d_t2, d_t3, d_y, h_staged.data(), h_want.data());
    const size_t badFused = checkAtCheckSize(
        true, d_x, d_r, d_t1, d_t2, d_t3, d_y, h_fused.data(), h_want.data());

    // Real branches returning EXIT_FAILURE, not assert(). CI builds Release,
    // Release defines NDEBUG, and NDEBUG deletes assert(), so a check written
    // that way would vanish in exactly the build that matters.
    if (badStaged != kCheckElems) {
        std::fprintf(stderr, "staged chain wrong at %zu: got %.9g, want %.9g\n",
                     badStaged, static_cast<double>(h_staged[badStaged]),
                     h_want[badStaged]);
        status = EXIT_FAILURE;
    }
    if (badFused != kCheckElems) {
        std::fprintf(stderr, "chainFused wrong at %zu: got %.9g, want %.9g\n",
                     badFused, static_cast<double>(h_fused[badFused]),
                     h_want[badFused]);
        status = EXIT_FAILURE;
    }

    float copyMs = 0.0f;
    float stageMs[4] = {0.0f, 0.0f, 0.0f, 0.0f};
    float stagedChainMs = 0.0f;
    float fusedChainMs = 0.0f;
    double launchUsSmall = 0.0;
    double launchUsFull = 0.0;
    float sweepStagedMs[kSweepCount] = {0.0f};
    float sweepFusedMs[kSweepCount] = {0.0f};
    WideRow wide[kWideCount] = {};
    size_t differing = 0;
    double maxAbsDiff = 0.0;

    if (status == EXIT_SUCCESS) {
        // Part 1. A batch of empty launches, back to back on one stream.
        // With the host free to run ahead, what this measures is the gap the
        // device leaves between one kernel ending and the next beginning,
        // which is the cost a chain of dependent kernels actually pays. A
        // single cold launch costs more; day 9 measured that one.
        const float smallBatchMs = timeKernel([&] {
            for (int k = 0; k < kLaunchesPerBatch; ++k) {
                doNothing<<<1, 1>>>();
            }
        });
        const float fullBatchMs = timeKernel([&] {
            for (int k = 0; k < kLaunchesPerBatch; ++k) {
                doNothing<<<blocksFor(kElems), kThreadsPerBlock>>>();
            }
        });
        launchUsSmall =
            static_cast<double>(smallBatchMs) * 1000.0 / kLaunchesPerBatch;
        launchUsFull =
            static_cast<double>(fullBatchMs) * 1000.0 / kLaunchesPerBatch;

        // Part 2. Every timed launch reads one buffer and writes another and
        // the two are never swapped, so each kernel sees the same input on
        // every run and no reset has to sit inside the measurement.
        const int blocks = blocksFor(kElems);
        copyMs = timeKernel(
            [&] { copyFloor<<<blocks, kThreadsPerBlock>>>(d_x, d_y, kElems); });
        stageMs[0] = timeKernel([&] {
            applyScaleBias<<<blocks, kThreadsPerBlock>>>(d_x, d_t1, kElems);
        });
        stageMs[1] = timeKernel([&] {
            applyGelu<<<blocks, kThreadsPerBlock>>>(d_t1, d_t2, kElems);
        });
        stageMs[2] = timeKernel([&] {
            applyAddResidual<<<blocks, kThreadsPerBlock>>>(d_t2, d_r, d_t3,
                                                           kElems);
        });
        stageMs[3] = timeKernel([&] {
            applyClampRange<<<blocks, kThreadsPerBlock>>>(d_t3, d_y, kElems);
        });
        stagedChainMs = timeKernel(
            [&] { launchStaged(d_x, d_r, d_t1, d_t2, d_t3, d_y, kElems); });
        CUDA_CHECK(
            cudaMemcpy(h_staged.data(), d_y, bytes, cudaMemcpyDeviceToHost));
        fusedChainMs = timeKernel([&] {
            chainFused<<<blocks, kThreadsPerBlock>>>(d_x, d_r, d_y, kElems);
        });
        CUDA_CHECK(
            cudaMemcpy(h_fused.data(), d_y, bytes, cudaMemcpyDeviceToHost));

        // Does fusion change the answer? Both versions round to float at the
        // same four points, so they should agree exactly, but the compiler
        // is free to contract a multiply and an add across a stage boundary
        // in the fused kernel and not in the staged one. Counted rather than
        // assumed, and reported rather than gated: a difference here is a
        // fact about fusion, not a failure.
        for (size_t i = 0; i < kElems; ++i) {
            if (h_staged[i] != h_fused[i]) {
                ++differing;
                const double d = std::fabs(static_cast<double>(h_staged[i]) -
                                           static_cast<double>(h_fused[i]));
                if (d > maxAbsDiff) {
                    maxAbsDiff = d;
                }
            }
        }

        // Part 3. The same two chains, smaller and smaller, until four
        // launches cost more than the bytes they save.
        for (int s = 0; s < kSweepCount; ++s) {
            const size_t n = kSweep[s];
            const int sweepBlocks = blocksFor(n);
            sweepStagedMs[s] = timeKernel(
                [&] { launchStaged(d_x, d_r, d_t1, d_t2, d_t3, d_y, n); });
            sweepFusedMs[s] = timeKernel([&] {
                chainFused<<<sweepBlocks, kThreadsPerBlock>>>(d_x, d_r, d_y, n);
            });
        }

        // Part 4. Six widths, instantiated by hand because a template
        // argument cannot come from a loop. The static_assert above keeps
        // this list and kWideK in step.
        //
        // h_staged is reused as the correctness scratch buffer here. The
        // staged-against-fused comparison above has already run, so nothing
        // still needs what it held.
        probeWide<1>(&wide[0], d_x, d_r, d_y, h_staged.data(), h_want.data(),
                     prop.maxThreadsPerMultiProcessor);
        probeWide<2>(&wide[1], d_x, d_r, d_y, h_staged.data(), h_want.data(),
                     prop.maxThreadsPerMultiProcessor);
        probeWide<4>(&wide[2], d_x, d_r, d_y, h_staged.data(), h_want.data(),
                     prop.maxThreadsPerMultiProcessor);
        probeWide<8>(&wide[3], d_x, d_r, d_y, h_staged.data(), h_want.data(),
                     prop.maxThreadsPerMultiProcessor);
        probeWide<16>(&wide[4], d_x, d_r, d_y, h_staged.data(), h_want.data(),
                      prop.maxThreadsPerMultiProcessor);
        probeWide<32>(&wide[5], d_x, d_r, d_y, h_staged.data(), h_want.data(),
                      prop.maxThreadsPerMultiProcessor);

        for (int w = 0; w < kWideCount; ++w) {
            if (wide[w].bad != kCheckElems) {
                std::fprintf(stderr,
                             "chainFusedWide<%d> wrong at %zu, want %.9g\n",
                             wide[w].wide, wide[w].bad, h_want[wide[w].bad]);
                status = EXIT_FAILURE;
            }
        }
    }

    // Arithmetic, not a claim about the card: a zero or negative elapsed time
    // means the event pair never separated, and every figure derived from it
    // would be infinite or nonsense.
    if (status == EXIT_SUCCESS) {
        if (copyMs <= 0.0f || stagedChainMs <= 0.0f || fusedChainMs <= 0.0f ||
            launchUsSmall <= 0.0 || launchUsFull <= 0.0) {
            std::fprintf(stderr,
                         "timing broken: copy %.6f, staged %.6f, fused %.6f, "
                         "launch %.6f us\n",
                         static_cast<double>(copyMs),
                         static_cast<double>(stagedChainMs),
                         static_cast<double>(fusedChainMs), launchUsSmall);
            status = EXIT_FAILURE;
        }
        for (int s = 0; s < kSweepCount; ++s) {
            if (sweepStagedMs[s] <= 0.0f || sweepFusedMs[s] <= 0.0f) {
                std::fprintf(stderr, "timing broken in the sweep at n = %zu\n",
                             kSweep[s]);
                status = EXIT_FAILURE;
            }
        }
        for (int w = 0; w < kWideCount; ++w) {
            if (wide[w].ms <= 0.0f) {
                std::fprintf(stderr, "timing broken at width %d\n",
                             wide[w].wide);
                status = EXIT_FAILURE;
            }
        }
    }

    // The staged chain has to be slower than the fused one on this hardware
    // and at this size, because it does the same arithmetic and moves three
    // times the bytes. If it is not, the two are not running the same chain
    // and no ratio on the page means anything.
    if (status == EXIT_SUCCESS && stagedChainMs <= fusedChainMs) {
        std::fprintf(stderr,
                     "the staged chain (%.6f ms) is not slower than the fused "
                     "one (%.6f ms), which cannot be right at 3x the traffic\n",
                     static_cast<double>(stagedChainMs),
                     static_cast<double>(fusedChainMs));
        status = EXIT_FAILURE;
    }

    if (status == EXIT_SUCCESS) {
        std::printf("GPU: %s (compute capability %d.%d, %d SMs)\n", prop.name,
                    prop.major, prop.minor, prop.multiProcessorCount);
        std::printf(
            "%d resident threads/SM, %d registers/SM, %d KiB L2, %d "
            "threads/block\n",
            prop.maxThreadsPerMultiProcessor, prop.regsPerMultiprocessor,
            prop.l2CacheSize / 1024, kThreadsPerBlock);
        std::printf(
            "Timed size %zu elements, %.1f MiB per buffer, six buffers\n",
            kElems, mib);
        std::printf("Correctness size %zu elements, %.1f%% of them clamped\n",
                    kCheckElems,
                    100.0 * static_cast<double>(clamped) /
                        static_cast<double>(kCheckElems));
        std::printf(
            "  staged chain, chainFused and all %d widths match the CPU\n"
            "  reference at that size, which no block size divides\n\n",
            kWideCount);

        std::printf("Bytes per element, from the constants\n");
        std::printf("  staged chain    %5.0f   (2 + 2 + 3 + 2 floats)\n",
                    kStagedFloatsPerElem * sizeof(float));
        std::printf("  chainFused      %5.0f   (x, residual, out)\n",
                    kFusedFloatsPerElem * sizeof(float));
        std::printf("  ratio           %5.2f\n\n",
                    kStagedFloatsPerElem / kFusedFloatsPerElem);

        std::printf(
            "Part 1: what one launch costs, %d back-to-back launches per "
            "timed run\n",
            kLaunchesPerBatch);
        std::printf("%-38s %16s\n", "launch shape", "per launch (us)");
        std::printf("%-38s %16s\n", "-------------------------------------",
                    "---------------");
        std::printf("%-38s %16.3f\n", "doNothing, 1 block of 1 thread",
                    launchUsSmall);
        std::printf("%-38s %16.3f\n\n", "doNothing, the chain's own grid",
                    launchUsFull);

        std::printf(
            "Part 2: the chain at %zu elements, mean of %d runs after %d "
            "warm-ups\n",
            kElems, kTimedRuns, kWarmupRuns);
        std::printf("%-28s %8s %11s %9s %8s\n", "kernel or chain", "ms",
                    "bytes/elem", "GB/s", "x copy");
        std::printf("%-28s %8s %11s %9s %8s\n", "---------------------------",
                    "-------", "----------", "--------", "-------");

        // Every row's effective bandwidth is its own compulsory bytes over
        // its own time, and the last column divides that by the floor row's,
        // so two rows moving different byte counts are still comparable.
        const double copyGBs = bandwidthGBs(kCopyFloatsPerElem, kElems, copyMs);
        std::printf("%-28s %8.3f %11.0f %9.1f %8.2f\n", "copyFloor (the floor)",
                    static_cast<double>(copyMs),
                    kCopyFloatsPerElem * sizeof(float), copyGBs, 1.00);

        const char* stageNames[4] = {"applyScaleBias", "applyGelu",
                                     "applyAddResidual", "applyClampRange"};
        const double stageFloats[4] = {2.0, 2.0, 3.0, 2.0};
        double summed = 0.0;
        for (int s = 0; s < 4; ++s) {
            summed += static_cast<double>(stageMs[s]);
            const double gbs = bandwidthGBs(stageFloats[s], kElems, stageMs[s]);
            std::printf("%-28s %8.3f %11.0f %9.1f %8.2f\n", stageNames[s],
                        static_cast<double>(stageMs[s]),
                        stageFloats[s] * sizeof(float), gbs, gbs / copyGBs);
        }
        const double summedGBs = bandwidthGBs(kStagedFloatsPerElem, kElems,
                                              static_cast<float>(summed));
        std::printf("%-28s %8.3f %11.0f %9.1f %8.2f\n", "  the four, summed",
                    summed, kStagedFloatsPerElem * sizeof(float), summedGBs,
                    summedGBs / copyGBs);
        const double stagedGBs =
            bandwidthGBs(kStagedFloatsPerElem, kElems, stagedChainMs);
        std::printf("%-28s %8.3f %11.0f %9.1f %8.2f\n",
                    "staged chain, timed as one",
                    static_cast<double>(stagedChainMs),
                    kStagedFloatsPerElem * sizeof(float), stagedGBs,
                    stagedGBs / copyGBs);
        const double fusedGBs =
            bandwidthGBs(kFusedFloatsPerElem, kElems, fusedChainMs);
        std::printf("%-28s %8.3f %11.0f %9.1f %8.2f\n\n", "chainFused",
                    static_cast<double>(fusedChainMs),
                    kFusedFloatsPerElem * sizeof(float), fusedGBs,
                    fusedGBs / copyGBs);
        std::printf("staged / fused = %.2f, and the byte ratio is %.2f\n",
                    static_cast<double>(stagedChainMs / fusedChainMs),
                    kStagedFloatsPerElem / kFusedFloatsPerElem);
        if (differing == 0) {
            std::printf(
                "The two chains agree bit for bit on all %zu elements.\n\n",
                kElems);
        } else {
            std::printf(
                "The two chains differ on %zu of %zu elements, at most by "
                "%.3g.\n\n",
                differing, kElems, maxAbsDiff);
        }

        std::printf("Part 3: the same two chains, swept over n\n");
        std::printf("%11s %12s %12s %13s %17s\n", "n", "staged (ms)",
                    "fused (ms)", "staged/fused", "staged us/launch");
        std::printf("%11s %12s %12s %13s %17s\n", "----------", "-----------",
                    "-----------", "------------", "----------------");
        for (int s = 0; s < kSweepCount; ++s) {
            std::printf("%11zu %12.4f %12.4f %13.2f %17.3f\n", kSweep[s],
                        static_cast<double>(sweepStagedMs[s]),
                        static_cast<double>(sweepFusedMs[s]),
                        static_cast<double>(sweepStagedMs[s] / sweepFusedMs[s]),
                        static_cast<double>(sweepStagedMs[s]) * 1000.0 / 4.0);
        }
        std::printf(
            "\nThe last column is the staged chain's time divided by its four\n"
            "launches. As n falls it has to approach part 1's per-launch "
            "figure,\nbecause that is all the chain is still doing.\n\n");

        std::printf(
            "Part 4: chainFused widened, same work and same bytes every "
            "row\n");
        std::printf("%11s %8s %6s %8s %10s %10s %8s %9s\n", "per thread",
                    "blocks", "regs", "spill B", "blocks/SM", "occupancy", "ms",
                    "GB/s");
        std::printf("%11s %8s %6s %8s %10s %10s %8s %9s\n", "----------",
                    "-------", "-----", "-------", "---------", "---------",
                    "-------", "--------");
        for (int w = 0; w < kWideCount; ++w) {
            std::printf("%11d %8d %6d %8zu %10d %9.1f%% %8.3f %9.1f\n",
                        wide[w].wide, wide[w].blocks, wide[w].regs,
                        wide[w].spillBytes, wide[w].blocksPerSm,
                        wide[w].occupancy, static_cast<double>(wide[w].ms),
                        bandwidthGBs(kFusedFloatsPerElem, kElems, wide[w].ms));
        }
        std::printf(
            "\nThis card keeps all %d of its resident threads on an SM only "
            "while a\nthread uses %d registers or fewer, which is %d "
            "registers per SM\ndivided by %d threads. Above that, warps come "
            "off the SM.\n",
            prop.maxThreadsPerMultiProcessor, regsAtFullOccupancy,
            prop.regsPerMultiprocessor, prop.maxThreadsPerMultiProcessor);
    }

    CUDA_CHECK(cudaFree(d_x));
    CUDA_CHECK(cudaFree(d_r));
    CUDA_CHECK(cudaFree(d_t1));
    CUDA_CHECK(cudaFree(d_t2));
    CUDA_CHECK(cudaFree(d_t3));
    CUDA_CHECK(cudaFree(d_y));
    return status;
}