COURSE / SOURCE

attention.cu

All lessons
Source filecode/day97-attention/attention.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 97: attention two ways, one head, FP32, no tensor cores.
//
// Path 1 materialises the whole N x N score matrix: a tiled Q K^T (day 16's
// tile shape), a three-pass safe softmax over each row, then a tiled P V.
// Path 2 is one kernel that never allocates that matrix at all. It walks the
// keys in tiles of 32, carries a running maximum, a running denominator and
// a running output per query row, and rescales all three whenever a tile
// brings a larger maximum. That is the online normalizer calculation, and
// using it tile-wise is what FlashAttention is.
//
// Both paths are judged against the same double-precision CPU reference,
// which computes the softmax in double with the maximum subtracted. The
// tolerance is day 66's K-scaled rule with K = N, because every output
// element is a sum over N keys.
//
// The comparison this file exists to support is memory, not milliseconds.
// The program prints both byte counts before it prints a single time, so a
// reader can see that the traffic result is arithmetic and the time result
// is a measurement that may not follow it. Day 34 is the precedent: a change
// that cut reads made the kernel slower.
//
// Build: nvcc -std=c++17 -O3 -arch=sm_75 -o attention attention.cu
// Run:   ./attention

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

// One head of width 64, which is what a 4096-wide model with 64 heads uses.
// Everything else in this file is expressed against it.
constexpr int kHeadDim = 64;
constexpr int kWarpSize = 32;
constexpr int kThreadsPerBlock = 256;
constexpr int kWarpsPerBlock = kThreadsPerBlock / kWarpSize;

// The tile shape of the online kernel. One lane scores one key, so the key
// tile is exactly a warp. One warp owns 8 query rows and one block owns 64,
// which is what decides how often K and V are re-read: once per query tile.
constexpr int kKeyTile = kWarpSize;
constexpr int kRowsPerWarp = 8;
constexpr int kQueryTile = kWarpsPerBlock * kRowsPerWarp;
constexpr int kDimsPerLane = kHeadDim / kWarpSize;

// Day 16's tile, used unchanged by the two matmuls of the naive path, so the
// baseline is a tiled matmul and not a straw man. The only thing the naive
// path is being blamed for on this page is the N x N buffer.
constexpr int kTileDim = 16;

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

// Three sequence lengths. 256 keeps the whole score matrix (256 KiB) inside
// the T4's 4 MiB L2, 1024 fills that L2 exactly, and 4096 needs 64 MiB
// of DRAM per pass. The double-precision reference is O(N^2 d) on one core,
// which is what stops this ladder at 4096 rather than at a size where the
// naive path stops fitting at all.
constexpr int kNumCases = 3;
constexpr size_t kCaseN[kNumCases] = {256, 1024, 4096};
constexpr size_t kMaxN = 4096;

// Sequence lengths the page quotes as out of reach for a materialised score
// matrix. Printed rather than asserted, so the arithmetic on the page comes
// from the program.
constexpr int kNumImpossible = 4;
constexpr size_t kImpossibleN[kNumImpossible] = {8192, 32768, 131072, 1048576};

// f32 from the course tolerance table, scaled by the reduction depth below.
constexpr double kF32Rtol = 1e-5;
constexpr int kF32MantissaBits = 23;

static_assert(kThreadsPerBlock % kWarpSize == 0,
              "block size must be a whole number of warps");
static_assert(kKeyTile == kWarpSize,
              "one lane scores one key, so the key tile is one warp wide");
static_assert(kHeadDim % kWarpSize == 0,
              "each lane owns kDimsPerLane output dimensions of the head");
static_assert(kCaseN[0] % kQueryTile == 0 && kCaseN[1] % kQueryTile == 0 &&
                  kCaseN[2] % kQueryTile == 0,
              "every case divides by the query tile, so the online kernel "
              "needs no ragged-edge guard and the table compares algorithms");
static_assert(kCaseN[0] % kKeyTile == 0 && kCaseN[1] % kKeyTile == 0 &&
                  kCaseN[2] % kKeyTile == 0,
              "every case divides by the key tile for the same reason");
static_assert(kCaseN[0] % kTileDim == 0 && kCaseN[1] % kTileDim == 0 &&
                  kCaseN[2] % kTileDim == 0,
              "every case divides by day 16's tile, so the naive path's two "
              "matmuls run the same clean geometry at every size");
static_assert(kCaseN[2] == kMaxN,
              "the largest case sizes every allocation in main()");
// 16,384 bytes of queries, 8,320 of keys with the pad, 8,192 of values.
static_assert((kQueryTile * kHeadDim + kKeyTile * (kHeadDim + 1) +
               kKeyTile * kHeadDim) *
                      sizeof(float) <=
                  48 * 1024,
              "the online kernel's tiles must fit the 48 KiB a block gets on "
              "a T4 without the cudaFuncSetAttribute opt-in");

static constexpr size_t ceilDiv(size_t a, size_t b) {
    return (a + b - 1) / b;
}

// S = (Q K^T) * scale, written to global memory in full. Day 16's 16 x 16
// tile, with the K tile stored transposed in shared memory so the global
// read stays coalesced: consecutive lanes read consecutive elements of one
// key vector, not one element of 16 different key vectors.
//
// Memory: one warp's 32 addresses cover two rows of a tile, 64 contiguous
// bytes each. The pad to 17 keeps the transposed shared read off a single
// bank, which is day 15's arithmetic: stride 17 and 32 banks share no
// factor.
//
// Launch assumption: 16 x 16 threads, a grid that covers n x n, and n a
// multiple of 16 so the guards below are never taken in this program.
__global__ void scoresQKt(const float* __restrict__ q,
                          const float* __restrict__ k, float* __restrict__ s,
                          size_t n, float scale) {
    __shared__ float tileQ[kTileDim][kTileDim];
    __shared__ float tileK[kTileDim][kTileDim + 1];

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

    float acc = 0.0f;
    for (int t = 0; t < kHeadDim / kTileDim; ++t) {
        const size_t dim = static_cast<size_t>(t) * kTileDim + threadIdx.x;
        tileQ[threadIdx.y][threadIdx.x] =
            (row < n) ? q[row * kHeadDim + dim] : 0.0f;
        tileK[threadIdx.y][threadIdx.x] =
            (keyRow < n) ? k[keyRow * kHeadDim + dim] : 0.0f;
        __syncthreads();

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

    if (row < n && col < n) {
        s[row * n + col] = acc * scale;
    }
}

// The safe softmax, one block per row: read the row to find its maximum,
// then read it again to write exp(x - max) into a second buffer and total
// it. Two reads and one write of the row, which is the cost the online
// version removes. The normalisation is left to the kernel below, so this
// one does not spend a further pass rewriting the row.
//
// Writing to a second buffer rather than in place is not free (it is the
// second N x N allocation this path needs) and it is not decoration: a
// kernel that updates its input in place cannot be timed in a loop, because
// the second run would exponentiate the first run's output.
//
// Memory: threads sweep the row with a grid-stride loop, so one warp's 32
// addresses are 128 contiguous bytes at every step.
//
// Launch assumption: one block per row, kThreadsPerBlock threads, and the
// shared scratch sized for the block's warps.
__global__ void softmaxRows(const float* __restrict__ s, float* __restrict__ p,
                            float* __restrict__ rowSum, size_t n) {
    __shared__ float scratch[kWarpsPerBlock];

    const size_t row = blockIdx.x;
    const unsigned int tid = threadIdx.x;
    const unsigned int lane = tid % kWarpSize;
    const unsigned int warp = tid / kWarpSize;
    const float* line = s + row * n;
    float* out = p + row * n;

    float best = -INFINITY;
    for (size_t j = tid; j < n; j += kThreadsPerBlock) {
        best = fmaxf(best, line[j]);
    }
    for (int off = kWarpSize / 2; off > 0; off >>= 1) {
        best = fmaxf(best, __shfl_xor_sync(0xffffffffu, best, off));
    }
    if (lane == 0) {
        scratch[warp] = best;
    }
    __syncthreads();
    best = (tid < kWarpsPerBlock) ? scratch[tid] : -INFINITY;
    for (int off = kWarpsPerBlock / 2; off > 0; off >>= 1) {
        best = fmaxf(best, __shfl_xor_sync(0xffffffffu, best, off));
    }
    if (tid == 0) {
        scratch[0] = best;
    }
    __syncthreads();
    const float rowMax = scratch[0];
    __syncthreads();

    float total = 0.0f;
    for (size_t j = tid; j < n; j += kThreadsPerBlock) {
        const float e = expf(line[j] - rowMax);
        out[j] = e;
        total += e;
    }
    for (int off = kWarpSize / 2; off > 0; off >>= 1) {
        total += __shfl_xor_sync(0xffffffffu, total, off);
    }
    if (lane == 0) {
        scratch[warp] = total;
    }
    __syncthreads();
    total = (tid < kWarpsPerBlock) ? scratch[tid] : 0.0f;
    for (int off = kWarpsPerBlock / 2; off > 0; off >>= 1) {
        total += __shfl_xor_sync(0xffffffffu, total, off);
    }
    if (tid == 0) {
        rowSum[row] = total;
    }
}

// O = (P V) / rowSum, day 16's tile again. P is the N x N buffer the softmax
// left behind, so this kernel reads every one of its bytes back.
//
// Memory: both tile fills are coalesced, and the division by the row total
// costs one broadcast read of an N-element vector.
//
// Launch assumption: 16 x 16 threads, grid covering n rows by kHeadDim
// columns, n a multiple of 16.
__global__ void weightedV(const float* __restrict__ p,
                          const float* __restrict__ v,
                          const float* __restrict__ rowSum,
                          float* __restrict__ o, size_t n) {
    __shared__ float tileP[kTileDim][kTileDim];
    __shared__ float tileV[kTileDim][kTileDim];

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

    float acc = 0.0f;
    const size_t tiles = (n + kTileDim - 1) / kTileDim;
    for (size_t t = 0; t < tiles; ++t) {
        const size_t key = t * kTileDim + threadIdx.x;
        const size_t keyRow = t * kTileDim + threadIdx.y;
        tileP[threadIdx.y][threadIdx.x] =
            (row < n && key < n) ? p[row * n + key] : 0.0f;
        tileV[threadIdx.y][threadIdx.x] =
            (keyRow < n) ? v[keyRow * kHeadDim + col] : 0.0f;
        __syncthreads();

        for (int e = 0; e < kTileDim; ++e) {
            acc += tileP[threadIdx.y][e] * tileV[e][threadIdx.x];
        }
        __syncthreads();
    }

    if (row < n) {
        o[row * kHeadDim + col] = acc / rowSum[row];
    }
}

// The whole of attention for 64 query rows, in one kernel, with no N x N
// buffer anywhere. Each warp owns 8 query rows; each lane owns one key of
// the current 32-key tile and two of the 64 output dimensions.
//
// Per tile, per query row: lane j scores key j, the warp reduces that to the
// tile maximum, the running maximum grows to cover it, and everything
// carried from earlier tiles is multiplied by exp(oldMax - newMax) before
// this tile's contribution is added. The running denominator and the running
// output are rescaled by the same factor, which is why the answer at the end
// equals the answer a three-pass softmax would have produced.
//
// Memory: the query tile is read once for the whole sweep. The key and value
// tiles are read once per query tile, so K and V are re-read n / 64 times
// across the grid; the program prints that ask next to the floor. Global
// reads are coalesced (consecutive threads take consecutive dimensions) and
// the key tile is padded by one float so lane j's walk down key j does not
// collide with lane j+1's, day 15's rule with stride 65 against 32 banks.
//
// Launch assumption: kThreadsPerBlock threads, one block per 64 query rows,
// n a multiple of both 64 and 32. No thread returns early, because every
// thread has to reach every barrier below.
__global__ void attentionOnline(const float* __restrict__ q,
                                const float* __restrict__ k,
                                const float* __restrict__ v,
                                float* __restrict__ o, size_t n, float scale) {
    __shared__ float tileQ[kQueryTile][kHeadDim];
    __shared__ float tileK[kKeyTile][kHeadDim + 1];
    __shared__ float tileV[kKeyTile][kHeadDim];

    const unsigned int tid = threadIdx.x;
    const unsigned int lane = tid % kWarpSize;
    const unsigned int warp = tid / kWarpSize;
    const size_t queryBase = blockIdx.x * static_cast<size_t>(kQueryTile);

    for (int e = tid; e < kQueryTile * kHeadDim; e += kThreadsPerBlock) {
        tileQ[e / kHeadDim][e % kHeadDim] = q[queryBase * kHeadDim + e];
    }

    // snippet: online-state
    float acc[kRowsPerWarp][kDimsPerLane];
    float runMax[kRowsPerWarp];
    float runSum[kRowsPerWarp];
    for (int r = 0; r < kRowsPerWarp; ++r) {
        runMax[r] = -INFINITY;
        runSum[r] = 0.0f;
        for (int c = 0; c < kDimsPerLane; ++c) {
            acc[r][c] = 0.0f;
        }
    }
    // end snippet

    for (size_t base = 0; base < n; base += kKeyTile) {
        __syncthreads();
        for (int e = tid; e < kKeyTile * kHeadDim; e += kThreadsPerBlock) {
            const size_t key = base + static_cast<size_t>(e / kHeadDim);
            const int dim = e % kHeadDim;
            tileK[e / kHeadDim][dim] = k[key * kHeadDim + dim];
            tileV[e / kHeadDim][dim] = v[key * kHeadDim + dim];
        }
        __syncthreads();

        for (int r = 0; r < kRowsPerWarp; ++r) {
            const int row = warp * kRowsPerWarp + r;
            float score = 0.0f;
            for (int e = 0; e < kHeadDim; ++e) {
                score += tileQ[row][e] * tileK[lane][e];
            }
            score *= scale;

            // snippet: online-update
            float tileMax = score;
            for (int off = kWarpSize / 2; off > 0; off >>= 1) {
                tileMax =
                    fmaxf(tileMax, __shfl_xor_sync(0xffffffffu, tileMax, off));
            }
            const float newMax = fmaxf(runMax[r], tileMax);
            // exp(-inf - newMax) is 0, so the first tile's rescale wipes the
            // initialised state instead of special-casing it.
            const float rescale = expf(runMax[r] - newMax);
            const float weight = expf(score - newMax);

            float tileSum = weight;
            for (int off = kWarpSize / 2; off > 0; off >>= 1) {
                tileSum += __shfl_xor_sync(0xffffffffu, tileSum, off);
            }
            runSum[r] = runSum[r] * rescale + tileSum;
            runMax[r] = newMax;

            for (int c = 0; c < kDimsPerLane; ++c) {
                acc[r][c] *= rescale;
            }
            for (int j = 0; j < kKeyTile; ++j) {
                const float w = __shfl_sync(0xffffffffu, weight, j);
                for (int c = 0; c < kDimsPerLane; ++c) {
                    acc[r][c] += w * tileV[j][c * kWarpSize + lane];
                }
            }
            // end snippet
        }
    }

    for (int r = 0; r < kRowsPerWarp; ++r) {
        const size_t row =
            queryBase + static_cast<size_t>(warp * kRowsPerWarp + r);
        for (int c = 0; c < kDimsPerLane; ++c) {
            o[row * kHeadDim + c * kWarpSize + lane] = acc[r][c] / runSum[r];
        }
    }
}

// CPU reference in double, one query row at a time, with the maximum
// subtracted before any exp. It holds one row of scores, not the matrix, so
// the reference is not itself limited by the thing this page is about.
// Written for obvious correctness: plain loops, no blocking.
static void attentionCpu(const float* q, const float* k, const float* v,
                         double* o, double* scratch, size_t n, double scale) {
    for (size_t i = 0; i < n; ++i) {
        double best = -HUGE_VAL;
        for (size_t j = 0; j < n; ++j) {
            double dot = 0.0;
            for (int e = 0; e < kHeadDim; ++e) {
                dot += static_cast<double>(q[i * kHeadDim + e]) *
                       static_cast<double>(k[j * kHeadDim + e]);
            }
            scratch[j] = dot * scale;
            if (scratch[j] > best) {
                best = scratch[j];
            }
        }
        double total = 0.0;
        for (size_t j = 0; j < n; ++j) {
            scratch[j] = std::exp(scratch[j] - best);
            total += scratch[j];
        }
        for (int e = 0; e < kHeadDim; ++e) {
            double sum = 0.0;
            for (size_t j = 0; j < n; ++j) {
                sum += scratch[j] * static_cast<double>(v[j * kHeadDim + e]);
            }
            o[i * kHeadDim + e] = sum / total;
        }
    }
}

// Deterministic inputs with no exact binary values in them, so no path can
// pass by accident. The periods 97, 89 and 83 divide none of the case sizes,
// so the pattern does not line up with any tile boundary.
//
// One key is deliberately amplified: key n - 3 sits in the last tile of
// every sweep and scores far above the rest, so every query row's running
// maximum is still moving on the final tile. Without it the online kernel
// could rescale once at the start and never again, and the branch this day
// is about would go untested. That key is the `spike-late` case of the page's
// widget, written into the data.
static void makeInputs(std::vector<float>* h_q, std::vector<float>* h_k,
                       std::vector<float>* h_v, size_t n) {
    for (size_t i = 0; i < n; ++i) {
        for (int e = 0; e < kHeadDim; ++e) {
            const size_t at = i * kHeadDim + static_cast<size_t>(e);
            (*h_q)[at] = static_cast<float>(
                             static_cast<int>((i * 31 + e * 7) % 97) - 48) *
                         0.021f;
            const float kv =
                static_cast<float>(static_cast<int>((i * 17 + e * 11) % 89) -
                                   44) *
                0.019f;
            (*h_k)[at] = (i + 3 == n) ? kv * 3.0f : kv;
            (*h_v)[at] = static_cast<float>(
                             static_cast<int>((i * 13 + e * 5) % 83) - 41) *
                         0.011f;
        }
    }
}

// The K-scaled tolerance from day 66. Every output element here is a sum
// over n keys, so the reduction depth is n, not the head width: rounding
// errors across the sum are uncorrelated and grow like sqrt(n), and the
// factor 4 is slack for a different-but-valid summation order. The two GPU
// paths sum in different orders on purpose, so this bound is what makes them
// comparable at all.
static double kScaledRtol(double tableRtol, int mantissaBits, size_t depth) {
    const double eps = std::ldexp(1.0, -mantissaBits);
    const double grown = 4.0 * eps * std::sqrt(static_cast<double>(depth));
    return grown > tableRtol ? grown : tableRtol;
}

// Worst error as a fraction of the gate |got - ref| <= atol + rtol * |ref|.
// A fraction above 1 anywhere fails, and the caller gets the first offending
// index so the report can name a row and a dimension. A non-finite output
// fails outright: inf here means an overflowed exp, which is the bug the
// maximum subtraction exists to prevent.
static double worstGateFraction(const float* got, const double* want, size_t n,
                                double rtol, double atol, size_t* firstBad) {
    double worst = 0.0;
    *firstBad = n;
    for (size_t i = 0; i < n; ++i) {
        if (!std::isfinite(got[i])) {
            *firstBad = i;
            return HUGE_VAL;
        }
        const double err = std::fabs(static_cast<double>(got[i]) - want[i]);
        const double frac = err / (atol + rtol * std::fabs(want[i]));
        if (frac > worst) {
            worst = frac;
        }
        if (frac > 1.0 && *firstBad == n) {
            *firstBad = i;
        }
    }
    return worst;
}

// The two GPU paths against each other, which is the check the exercise
// grades on: the naive path is only available as a reference while the score
// matrix still fits.
static double worstPairFraction(const float* got, const float* other, size_t n,
                                double rtol, double atol, size_t* firstBad) {
    double worst = 0.0;
    *firstBad = n;
    for (size_t i = 0; i < n; ++i) {
        const double a = static_cast<double>(got[i]);
        const double b = static_cast<double>(other[i]);
        if (!std::isfinite(a) || !std::isfinite(b)) {
            *firstBad = i;
            return HUGE_VAL;
        }
        const double frac = std::fabs(a - b) / (atol + rtol * std::fabs(b));
        if (frac > worst) {
            worst = frac;
        }
        if (frac > 1.0 && *firstBad == n) {
            *firstBad = i;
        }
    }
    return worst;
}

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

static double maxAbs(const double* x, size_t n) {
    double worst = 0.0;
    for (size_t i = 0; i < n; ++i) {
        if (std::fabs(x[i]) > worst) {
            worst = std::fabs(x[i]);
        }
    }
    return worst;
}

// Times a launch with CUDA events and returns the mean milliseconds per run.
//
// This is the one template and the one lambda allowed in module 1 to 3 code.
// Copy it verbatim; the alternative is a second copy 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.
    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, static shared memory and resident blocks per SM, from the
// runtime API rather than from a profiler: no performance counters, no root,
// so a learner on a hosted tier gets this table too. It matters here because
// the online kernel asks for 32 KiB of shared memory per block, and on a card
// with 64 KiB per SM that ceiling, not the register file, is what decides how
// many blocks are resident.
template <typename Kernel>
static void printKernelStats(const char* name, Kernel fn, int threads,
                             int threadsPerSm) {
    cudaFuncAttributes attr;
    CUDA_CHECK(cudaFuncGetAttributes(&attr, reinterpret_cast<const void*>(fn)));
    int blocksPerSm = 0;
    CUDA_CHECK(cudaOccupancyMaxActiveBlocksPerMultiprocessor(&blocksPerSm, fn,
                                                             threads, 0));
    const double occupancy =
        100.0 * blocksPerSm * threads / static_cast<double>(threadsPerSm);
    std::printf("  %-16s %5d %10zu %10d %9.1f%%\n", name, attr.numRegs,
                attr.sharedSizeBytes, blocksPerSm, occupancy);
}

// What the score matrix costs at sequence lengths this program cannot run.
// Every figure is 4 * n * n bytes of FP32, printed in mebibytes so the point
// is legible: the term is quadratic and nothing about the hardware changes
// that.
static void printScoreMatrixCost() {
    std::printf("FP32 score matrix, one head, by sequence length:\n");
    for (int c = 0; c < kNumImpossible; ++c) {
        const size_t n = kImpossibleN[c];
        const size_t bytes = 4 * n * n;
        std::printf("  n = %7zu   N^2 floats = %14zu   %10zu MiB\n", n, n * n,
                    bytes / (1024 * 1024));
    }
    std::printf("\n");
}

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)\n", prop.name, prop.major,
                prop.minor);
    std::printf("head dim %d, query tile %d, key tile %d, %d threads/block\n\n",
                kHeadDim, kQueryTile, kKeyTile, kThreadsPerBlock);

    printScoreMatrixCost();

    std::printf("  %-16s %5s %10s %10s %10s\n", "kernel", "regs", "shared B",
                "blocks/SM", "occupancy");
    printKernelStats("scoresQKt", scoresQKt, kTileDim * kTileDim,
                     prop.maxThreadsPerMultiProcessor);
    printKernelStats("softmaxRows", softmaxRows, kThreadsPerBlock,
                     prop.maxThreadsPerMultiProcessor);
    printKernelStats("weightedV", weightedV, kTileDim * kTileDim,
                     prop.maxThreadsPerMultiProcessor);
    printKernelStats("attentionOnline", attentionOnline, kThreadsPerBlock,
                     prop.maxThreadsPerMultiProcessor);
    std::printf("\n");

    const size_t maxHead = kMaxN * kHeadDim;
    const double scale = 1.0 / std::sqrt(static_cast<double>(kHeadDim));

    std::vector<float> h_q(maxHead);
    std::vector<float> h_k(maxHead);
    std::vector<float> h_v(maxHead);
    std::vector<float> h_naive(maxHead);
    std::vector<float> h_online(maxHead);
    std::vector<double> h_want(maxHead);
    std::vector<double> h_scratch(kMaxN);

    float* d_q = nullptr;
    float* d_k = nullptr;
    float* d_v = nullptr;
    float* d_s = nullptr;
    float* d_p = nullptr;
    float* d_rowSum = nullptr;
    float* d_naive = nullptr;
    float* d_online = nullptr;
    CUDA_CHECK(cudaMalloc(&d_q, maxHead * sizeof(float)));
    CUDA_CHECK(cudaMalloc(&d_k, maxHead * sizeof(float)));
    CUDA_CHECK(cudaMalloc(&d_v, maxHead * sizeof(float)));
    CUDA_CHECK(cudaMalloc(&d_s, kMaxN * kMaxN * sizeof(float)));
    CUDA_CHECK(cudaMalloc(&d_p, kMaxN * kMaxN * sizeof(float)));
    CUDA_CHECK(cudaMalloc(&d_rowSum, kMaxN * sizeof(float)));
    CUDA_CHECK(cudaMalloc(&d_naive, maxHead * sizeof(float)));
    CUDA_CHECK(cudaMalloc(&d_online, maxHead * sizeof(float)));

    // Failure is recorded rather than returned, so every path falls through
    // to the one cleanup block below and no cudaMalloc escapes its cudaFree.
    int badCase = kNumCases;
    const char* badPath = nullptr;
    size_t badIndex = 0;

    for (int caseIdx = 0; caseIdx < kNumCases && badPath == nullptr;
         ++caseIdx) {
        const size_t n = kCaseN[caseIdx];
        const size_t outElems = n * kHeadDim;
        const size_t outBytes = outElems * sizeof(float);

        makeInputs(&h_q, &h_k, &h_v, n);
        CUDA_CHECK(
            cudaMemcpy(d_q, h_q.data(), outBytes, cudaMemcpyHostToDevice));
        CUDA_CHECK(
            cudaMemcpy(d_k, h_k.data(), outBytes, cudaMemcpyHostToDevice));
        CUDA_CHECK(
            cudaMemcpy(d_v, h_v.data(), outBytes, cudaMemcpyHostToDevice));
        attentionCpu(h_q.data(), h_k.data(), h_v.data(), h_want.data(),
                     h_scratch.data(), n, scale);

        const double refMax = maxAbs(h_want.data(), outElems);
        const double rtol = kScaledRtol(kF32Rtol, kF32MantissaBits, n);
        const double atol = rtol * refMax;

        // The bytes, before any time. The score matrix is written once by
        // scoresQKt, read twice and written once by softmaxRows, and read
        // again by weightedV: five passes over N^2 floats. The online kernel
        // moves Q, K, V and O once each and allocates nothing quadratic, but it
        // does ask for K and V once per query tile, and that ask is printed too
        // because whether it reaches DRAM is a cache question this program
        // cannot answer.
        const size_t scoreBytes = 5 * 4 * n * n;
        const size_t onlineFloor = 4 * 4 * n * kHeadDim;
        const size_t onlineAsk =
            2 * 4 * n * kHeadDim * (n / kQueryTile) + 2 * 4 * n * kHeadDim;
        std::printf("n = %zu, max |reference| = %.6f\n", n, refMax);
        std::printf("  gate: |got-ref| <= %.3e + %.3e*|ref|", atol, rtol);
        std::printf("  (rtol = max(1e-5, 4*2^-23*sqrt(n)))\n");
        std::printf("  naive score-matrix traffic   %12zu bytes\n", scoreBytes);
        std::printf("  online Q/K/V/O floor         %12zu bytes\n",
                    onlineFloor);
        std::printf(
            "  online K/V re-read ask       %12zu bytes"
            " (%zu query tiles)\n",
            onlineAsk, n / kQueryTile);

        const dim3 mmBlock(kTileDim, kTileDim);
        const dim3 scoreGrid(static_cast<unsigned int>(ceilDiv(n, kTileDim)),
                             static_cast<unsigned int>(ceilDiv(n, kTileDim)));
        const dim3 outGrid(
            static_cast<unsigned int>(ceilDiv(kHeadDim, kTileDim)),
            static_cast<unsigned int>(ceilDiv(n, kTileDim)));
        const unsigned int onlineBlocks =
            static_cast<unsigned int>(n / kQueryTile);
        const float fscale = static_cast<float>(scale);

        // Correctness first, timing second, so a wrong kernel is never
        // reported as a fast one.
        CUDA_CHECK(cudaMemset(d_naive, 0, outBytes));
        scoresQKt<<<scoreGrid, mmBlock>>>(d_q, d_k, d_s, n, fscale);
        softmaxRows<<<static_cast<unsigned int>(n), kThreadsPerBlock>>>(
            d_s, d_p, d_rowSum, n);
        weightedV<<<outGrid, mmBlock>>>(d_p, d_v, d_rowSum, d_naive, n);
        CUDA_CHECK(cudaGetLastError());
        CUDA_CHECK(cudaDeviceSynchronize());
        CUDA_CHECK(cudaMemcpy(h_naive.data(), d_naive, outBytes,
                              cudaMemcpyDeviceToHost));

        CUDA_CHECK(cudaMemset(d_online, 0, outBytes));
        attentionOnline<<<onlineBlocks, kThreadsPerBlock>>>(
            d_q, d_k, d_v, d_online, n, fscale);
        CUDA_CHECK(cudaGetLastError());
        CUDA_CHECK(cudaDeviceSynchronize());
        CUDA_CHECK(cudaMemcpy(h_online.data(), d_online, outBytes,
                              cudaMemcpyDeviceToHost));

        size_t firstBad = outElems;
        const double naiveFrac = worstGateFraction(
            h_naive.data(), h_want.data(), outElems, rtol, atol, &firstBad);
        if (firstBad != outElems) {
            badCase = caseIdx;
            badPath = "naive (materialised S)";
            badIndex = firstBad;
            break;
        }
        const double onlineFrac = worstGateFraction(
            h_online.data(), h_want.data(), outElems, rtol, atol, &firstBad);
        if (firstBad != outElems) {
            badCase = caseIdx;
            badPath = "online (tiled, no S)";
            badIndex = firstBad;
            break;
        }
        const double pairFrac = worstPairFraction(
            h_online.data(), h_naive.data(), outElems, rtol, atol, &firstBad);
        if (firstBad != outElems) {
            badCase = caseIdx;
            badPath = "online against naive";
            badIndex = firstBad;
            break;
        }

        const double naiveErr =
            maxAbsError(h_naive.data(), h_want.data(), outElems);
        const double onlineErr =
            maxAbsError(h_online.data(), h_want.data(), outElems);

        const float msScores = timeKernel([&] {
            scoresQKt<<<scoreGrid, mmBlock>>>(d_q, d_k, d_s, n, fscale);
        });
        const float msSoftmax = timeKernel([&] {
            softmaxRows<<<static_cast<unsigned int>(n), kThreadsPerBlock>>>(
                d_s, d_p, d_rowSum, n);
        });
        const float msWeighted = timeKernel([&] {
            weightedV<<<outGrid, mmBlock>>>(d_p, d_v, d_rowSum, d_naive, n);
        });
        // The three stages timed as one region, because three kernels timed
        // in three loops is not one path timed once (day 48).
        const float msNaive = timeKernel([&] {
            scoresQKt<<<scoreGrid, mmBlock>>>(d_q, d_k, d_s, n, fscale);
            softmaxRows<<<static_cast<unsigned int>(n), kThreadsPerBlock>>>(
                d_s, d_p, d_rowSum, n);
            weightedV<<<outGrid, mmBlock>>>(d_p, d_v, d_rowSum, d_naive, n);
        });
        const float msOnline = timeKernel([&] {
            attentionOnline<<<onlineBlocks, kThreadsPerBlock>>>(
                d_q, d_k, d_v, d_online, n, fscale);
        });

        std::printf("  %-24s %13s %10s %10s\n", "path", "max abs err",
                    "gate frac", "time (ms)");
        std::printf("  %-24s %13.3e %10.4f %10.4f\n", "naive, 3 kernels",
                    naiveErr, naiveFrac, msNaive);
        std::printf("  %-24s %13.3e %10.4f %10.4f\n", "online, 1 kernel",
                    onlineErr, onlineFrac, msOnline);
        std::printf("  naive breakdown: QK^T %.4f, softmax %.4f, PV %.4f ms\n",
                    msScores, msSoftmax, msWeighted);
        std::printf("  online vs naive worst gate fraction %.4f\n", pairFrac);
        std::printf("  naive/online time %.3fx, traffic %.1fx\n\n",
                    static_cast<double>(msNaive) / msOnline,
                    static_cast<double>(scoreBytes + onlineFloor) /
                        static_cast<double>(onlineFloor));
    }

    CUDA_CHECK(cudaFree(d_q));
    CUDA_CHECK(cudaFree(d_k));
    CUDA_CHECK(cudaFree(d_v));
    CUDA_CHECK(cudaFree(d_s));
    CUDA_CHECK(cudaFree(d_p));
    CUDA_CHECK(cudaFree(d_rowSum));
    CUDA_CHECK(cudaFree(d_naive));
    CUDA_CHECK(cudaFree(d_online));

    if (badPath != nullptr) {
        std::fprintf(stderr,
                     "%s failed its gate at %zu (query row %zu, dim %zu) "
                     "with n = %zu\n",
                     badPath, badIndex, badIndex / kHeadDim,
                     badIndex % kHeadDim, kCaseN[badCase]);
        return EXIT_FAILURE;
    }

    std::printf("all %d cases passed their gates\n", kNumCases);
    return EXIT_SUCCESS;
}