COURSE / SOURCE

mlp_train.cu

All lessons
Source filecode/day100-capstone-5/mlp_train.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 100: capstone 5, a 784-256-10 network trained end to end with kernels
// written in this course. Forward, loss, backward and the optimiser step are
// all in this file; nothing calls cuBLAS, cuDNN or Thrust.
//
// Eight cases, in this order, each a real branch that returns EXIT_FAILURE:
//
//   1 onebatch-forward  one image against a float64 CPU forward.
//   2 batch-forward one batch against a float64 CPU forward, tolerance
//                   scaled by the reduction depth.
//   3 loss-extreme  logits at plus and minus 60. A softmax without the max
//                   subtraction returns inf or nan here and nowhere else.
//   4 grad-check    every parameter gradient three ways: the GPU kernels, a float64
//                   analytic backward, and a central finite difference of
//                   the float64 loss. The finite difference is the
//                   independent one; it knows nothing about either backward.
//   5 step          one momentum SGD update against float64 arithmetic.
//   6 overfit-32    32 images, 200 steps, loss under 0.01. A sign error in
//                   the backward pass survives cases 1 to 3 and dies here.
//   7 train         the full recipe. Test accuracy is gated against a
//                   float64 CPU training run with the same seed, batch
//                   order and hyperparameters, not against a number.
//   8 perf          fused and unfused epoch timing. PyTorch is measured by
//                   pytorch_baseline.py on the same data and recipe.
//
// The dataset is the vendored Fashion-MNIST subset under data/day100: 6,000
// training images and the complete 10,000-image test split. The files are
// gzip-compressed IDX, so grading needs no network.
//
// Ceiling worth knowing: the batch order is a single seeded shuffle done
// once on the host before upload, not a reshuffle per epoch. On a real
// dataset that costs generalisation. The upgrade is an index buffer per
// epoch and a gather kernel in front of the first GEMM.
//
// Build: nvcc -std=c++17 -O3 -arch=sm_75 -lineinfo \
//            -I$CONDA_PREFIX/include -L$CONDA_PREFIX/lib \
//            -o mlp_train mlp_train.cu -lz
// Run:   ./mlp_train [--profile] [--data ../../data/day100]
//
// All eight cases, the PyTorch baseline and the registered-string Nsight
// Systems capture passed on a Tesla T4 on 2026-09-02.

#include <cfloat>
#include <cmath>
#include <cstdint>
#include <cstdio>
#include <cstdlib>
#include <cstring>
#include <string>
#include <vector>

#include <zlib.h>

#include <cuda_runtime.h>
#include <nvtx3/nvToolsExt.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)

// The architecture is fixed by the capstone specification: 784 inputs so a
// Fashion-MNIST image drops in unchanged, 256 hidden units, 10 classes.
constexpr int kInput = 784;
constexpr int kHidden = 256;
constexpr int kClasses = 10;

// The recipe. Fixing it is what makes one learner's accuracy comparable to
// another's and to the CPU reference: same seed, same batch size, same
// learning rate, same number of passes.
constexpr int kTrain = 6000;
constexpr int kTest = 10000;
constexpr int kBatch = 128;
constexpr int kEpochs = 20;
constexpr float kLearningRate = 0.1f;
constexpr float kMomentum = 0.9f;
constexpr uint64_t kSeed = 20260829ull;

constexpr int kTileDim = 16;
constexpr int kThreadsPerBlock = 256;  // 8 warps
constexpr int kWarpSize = 32;

// Gradient-check and overfit-case sizes.
constexpr int kGradBatch = 32;
constexpr int kGradCoordsPerTensor = 50;
constexpr int kOverfitSamples = 32;
constexpr int kOverfitSteps = 200;
constexpr double kOverfitLossGate = 0.01;

// The accuracy gate, in percentage points, against the float64 CPU run.
constexpr double kAccuracyGatePoints = 1.0;

// Epoch 0 is the warm-up: it launches every kernel the timed epochs launch,
// which is the only warm-up that covers lazy module loading for all of them.
constexpr int kWarmupEpochs = 1;

constexpr int kParamCount =
    kInput * kHidden + kHidden + kHidden * kClasses + kClasses;
constexpr int kOffW1 = 0;
constexpr int kOffB1 = kInput * kHidden;
constexpr int kOffW2 = kOffB1 + kHidden;
constexpr int kOffB2 = kOffW2 + kHidden * kClasses;

// The largest row count any activation buffer has to hold: the test set goes
// through the forward pass in one launch.
constexpr int kActRows = (kTest > kBatch) ? kTest : kBatch;

constexpr int kStepsPerEpoch = (kTrain + kBatch - 1) / kBatch;

static_assert(kThreadsPerBlock % kWarpSize == 0,
              "the block is a whole number of warps");
static_assert(kEpochs > kWarmupEpochs,
              "at least one epoch has to be timed after the warm-up epoch");
static_assert(kGradBatch <= kBatch,
              "the gradient check reuses the training activation buffers");
static_assert(kOverfitSamples <= kTrain, "the overfit case is a data subset");
static_assert(kInput % kTileDim == 0 && kHidden % kTileDim == 0,
              "the two large GEMM dimensions divide the tile, so the guards "
              "in gemmTiled only ever fire on the 10-wide class dimension");

// Kernels

// One tiled GEMM covering the four shapes this network needs, chosen at
// compile time so no branch survives into the inner loop:
//
//   C = A B      forward, both layers          kTransA=false kTransB=false
//   C = A^T B    dW1 = X^T dH, dW2 = A^T dY    kTransA=true
//   C = A B^T    dA = dY W2^T                  kTransB=true
//
// One thread owns one element of C and walks k in tiles of 16.
// One warp's 32 addresses: threadIdx.x is the column of C, so the store and
// the tileB load are contiguous in n. Under kTransA the tileA load reads
// a[k * m + row] with row contiguous across threadIdx.x, which is also
// contiguous; that is why the transpose lives in the index and not in a
// separate transpose kernel (day 12's lesson, arriving as a layout choice).
// Launch assumption: block is exactly kTileDim x kTileDim. The m, n and k
// guards handle the 10-wide class dimension, which divides no tile.
template <bool kTransA, bool kTransB, bool kBias, bool kRelu>
__global__ void gemmTiled(const float* __restrict__ a,
                          const float* __restrict__ b,
                          const float* __restrict__ bias, float* __restrict__ c,
                          int m, int n, int k) {
    // No padding on either tile: the inner loop reads tileA[ty][i], which is
    // a broadcast across the warp, and tileB[i][tx], which is contiguous.
    // Neither is the column-strided access day 15 pads against.
    __shared__ float tileA[kTileDim][kTileDim];
    __shared__ float tileB[kTileDim][kTileDim];

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

    float acc = 0.0f;
    for (int t = 0; t < k; t += kTileDim) {
        const int ka = t + tx;
        const int kb = t + ty;
        // Guard the loads, never the barrier: every thread in the block
        // reaches both __syncthreads() on every iteration.
        float va = 0.0f;
        if (row < m && ka < k) {
            va = kTransA ? a[ka * m + row] : a[row * k + ka];
        }
        float vb = 0.0f;
        if (col < n && kb < k) {
            vb = kTransB ? b[col * k + kb] : b[kb * n + col];
        }
        tileA[ty][tx] = va;
        tileB[ty][tx] = vb;
        __syncthreads();
        for (int i = 0; i < kTileDim; ++i) {
            acc += tileA[ty][i] * tileB[i][tx];
        }
        __syncthreads();
    }

    // snippet: epilogue
    if (row < m && col < n) {
        if (kBias) {
            acc += bias[col];
        }
        if (kRelu) {
            acc = fmaxf(acc, 0.0f);
        }
        c[row * n + col] = acc;
    }
    // end snippet
}

// Bias and ReLU as their own kernels. Nothing in the fused path calls these:
// they exist so case 6 can pay for the two extra launches and the two extra
// round trips through global memory that the epilogue above deletes, which
// is day 48's measurement arriving on a workload with a loss curve.
__global__ void addBias(float* __restrict__ y, const float* __restrict__ bias,
                        int m, int n) {
    const int i = blockIdx.x * blockDim.x + threadIdx.x;
    const int total = m * n;
    if (i < total) {
        y[i] += bias[i % n];
    }
}

__global__ void reluForward(float* __restrict__ y, int total) {
    const int i = blockIdx.x * blockDim.x + threadIdx.x;
    if (i < total) {
        y[i] = fmaxf(y[i], 0.0f);
    }
}

// Softmax and cross entropy over one row, one warp per row.
//
// One thread does 10/32 of a row here, which wastes most of the warp at this
// class count. The shape is the one that survives: at a 50,000-token vocab
// the loop below is the only version that fits, and it is day 24's reduction
// with day 23's shuffle doing the combining.
//
// The max subtraction is the whole numerical point. exp(60) overflows a
// float; exp(60 - 60) does not.
// snippet: softmax-ce
__global__ void softmaxCrossEntropy(const float* __restrict__ logits,
                                    const int* __restrict__ labels,
                                    float* __restrict__ dlogits,
                                    float* __restrict__ rowLoss, int m, int n,
                                    float scale) {
    const int lane = static_cast<int>(threadIdx.x) % kWarpSize;
    const int warp = (blockIdx.x * blockDim.x + threadIdx.x) / kWarpSize;
    // The whole warp shares `warp`, so this branch never splits a warp and
    // the shuffles below always have all 32 lanes.
    if (warp >= m) {
        return;
    }

    const float* row = logits + static_cast<size_t>(warp) * n;
    float best = -FLT_MAX;
    for (int c = lane; c < n; c += kWarpSize) {
        best = fmaxf(best, row[c]);
    }
    for (int off = kWarpSize / 2; off > 0; off >>= 1) {
        best = fmaxf(best, __shfl_down_sync(0xffffffffu, best, off));
    }
    best = __shfl_sync(0xffffffffu, best, 0);

    float sum = 0.0f;
    for (int c = lane; c < n; c += kWarpSize) {
        sum += __expf(row[c] - best);
    }
    for (int off = kWarpSize / 2; off > 0; off >>= 1) {
        sum += __shfl_down_sync(0xffffffffu, sum, off);
    }
    sum = __shfl_sync(0xffffffffu, sum, 0);

    const int label = labels[warp];
    const float logSum = __logf(sum);
    float* dst = dlogits + static_cast<size_t>(warp) * n;
    for (int c = lane; c < n; c += kWarpSize) {
        const float p = __expf(row[c] - best - logSum);
        dst[c] = (p - ((c == label) ? 1.0f : 0.0f)) * scale;
    }
    if (lane == 0) {
        rowLoss[warp] = -(row[label] - best - logSum);
    }
}
// end snippet

// db = column sums of dY. One thread per column, walking the rows.
//
// At n = 10 this launches ten threads and cannot fill one SM, let alone 40.
// It stays because it is honest about where the time goes: the profile shows
// a kernel that is pure launch overhead, and the fix is folding the bias
// gradient into the GEMM's epilogue as a second output, which is exactly the
// change the write-up is supposed to propose.
__global__ void columnSum(const float* __restrict__ y, float* __restrict__ out,
                          int m, int n) {
    const int col = blockIdx.x * blockDim.x + threadIdx.x;
    if (col < n) {
        float s = 0.0f;
        for (int r = 0; r < m; ++r) {
            s += y[static_cast<size_t>(r) * n + col];
        }
        out[col] = s;
    }
}

// ReLU backward, in place on dH. The mask comes from the activation itself,
// because relu(x) > 0 exactly when x > 0, so no pre-activation buffer has to
// be kept alive across the backward pass.
__global__ void reluBackward(float* __restrict__ dh,
                             const float* __restrict__ act, int total) {
    const int i = blockIdx.x * blockDim.x + threadIdx.x;
    if (i < total) {
        dh[i] = (act[i] > 0.0f) ? dh[i] : 0.0f;
    }
}

// SGD with momentum over a contiguous slice of the parameter vector.
//
// The four tensors live in one buffer, so the fused path calls this once per
// step over all 203,530 parameters and the naive path calls it four times
// over the four slices. Same arithmetic, same bytes, three extra launches.
// snippet: step
__global__ void sgdMomentum(float* __restrict__ w, const float* __restrict__ g,
                            float* __restrict__ v, float lr, float momentum,
                            int n) {
    const int stride = gridDim.x * blockDim.x;
    for (int i = blockIdx.x * blockDim.x + threadIdx.x; i < n; i += stride) {
        const float vel = momentum * v[i] - lr * g[i];
        v[i] = vel;
        w[i] += vel;
    }
}
// end snippet

// Host helpers

// splitmix64. A counter-based generator: sample i depends only on i and the
// seed, so the dataset is identical on every machine and reproducible from
// the two constants above without shipping a file.
static uint64_t mix64(uint64_t x) {
    x += 0x9e3779b97f4a7c15ull;
    x = (x ^ (x >> 30)) * 0xbf58476d1ce4e5b9ull;
    x = (x ^ (x >> 27)) * 0x94d049bb133111ebull;
    return x ^ (x >> 31);
}

// Uniform in [-1, 1) from a stream index. 53 bits into a double, then a
// single rounding to float, so the value does not depend on the host's
// floating point settings.
static float uniform(uint64_t stream, uint64_t index) {
    const uint64_t bits = mix64(stream * 0x2545f4914f6cdd1dull + index);
    const double unit = static_cast<double>(bits >> 11) / 9007199254740992.0;
    return static_cast<float>(2.0 * unit - 1.0);
}

static uint32_t readBe32(const unsigned char* p) {
    return (static_cast<uint32_t>(p[0]) << 24) |
           (static_cast<uint32_t>(p[1]) << 16) |
           (static_cast<uint32_t>(p[2]) << 8) | static_cast<uint32_t>(p[3]);
}

static std::vector<unsigned char> readGzip(const std::string& path) {
    gzFile file = gzopen(path.c_str(), "rb");
    if (file == nullptr) {
        std::fprintf(stderr, "cannot open dataset file %s\n", path.c_str());
        std::exit(EXIT_FAILURE);
    }
    std::vector<unsigned char> bytes;
    unsigned char chunk[1 << 15];
    for (;;) {
        const int n = gzread(file, chunk, sizeof(chunk));
        if (n < 0) {
            int error = Z_OK;
            const char* message = gzerror(file, &error);
            std::fprintf(stderr, "cannot read %s: %s\n", path.c_str(),
                         message);
            gzclose(file);
            std::exit(EXIT_FAILURE);
        }
        if (n == 0) {
            break;
        }
        bytes.insert(bytes.end(), chunk, chunk + n);
    }
    if (gzclose(file) != Z_OK) {
        std::fprintf(stderr, "cannot close dataset file %s\n", path.c_str());
        std::exit(EXIT_FAILURE);
    }
    return bytes;
}

static void loadIdxSplit(const std::string& imagePath,
                         const std::string& labelPath, int expected,
                         float* images, int* labels) {
    const std::vector<unsigned char> imageBytes = readGzip(imagePath);
    const std::vector<unsigned char> labelBytes = readGzip(labelPath);
    if (imageBytes.size() < 16 || labelBytes.size() < 8 ||
        readBe32(imageBytes.data()) != 0x00000803u ||
        readBe32(labelBytes.data()) != 0x00000801u) {
        std::fprintf(stderr, "invalid Fashion-MNIST IDX header\n");
        std::exit(EXIT_FAILURE);
    }
    const uint32_t imageCount = readBe32(imageBytes.data() + 4);
    const uint32_t rows = readBe32(imageBytes.data() + 8);
    const uint32_t cols = readBe32(imageBytes.data() + 12);
    const uint32_t labelCount = readBe32(labelBytes.data() + 4);
    const size_t pixels = static_cast<size_t>(expected) * kInput;
    if (imageCount != static_cast<uint32_t>(expected) ||
        labelCount != static_cast<uint32_t>(expected) || rows != 28 ||
        cols != 28 || imageBytes.size() != 16 + pixels ||
        labelBytes.size() != 8 + static_cast<size_t>(expected)) {
        std::fprintf(stderr,
                     "unexpected Fashion-MNIST shape in %s and %s: "
                     "%u images, %u labels, %ux%u, expected %d at 28x28\n",
                     imagePath.c_str(), labelPath.c_str(), imageCount,
                     labelCount, rows, cols, expected);
        std::exit(EXIT_FAILURE);
    }
    for (size_t i = 0; i < pixels; ++i) {
        images[i] = static_cast<float>(imageBytes[16 + i]) / 255.0f;
    }
    for (int i = 0; i < expected; ++i) {
        labels[i] = static_cast<int>(labelBytes[8 + i]);
        if (labels[i] < 0 || labels[i] >= kClasses) {
            std::fprintf(stderr, "invalid label %d at row %d in %s\n",
                         labels[i], i, labelPath.c_str());
            std::exit(EXIT_FAILURE);
        }
    }
}

static int rowsInStep(int step) {
    const int first = step * kBatch;
    const int left = kTrain - first;
    return (left < kBatch) ? left : kBatch;
}

// The relative tolerance a float32 dot product of depth k earns against a
// float64 reference. Worst-case error growth for a naive summation is k
// units in the last place, and the observed growth for uncorrelated terms is
// closer to sqrt(k); this takes sqrt(k) and multiplies by 8 for the two
// layers, the bias add and the accumulation order the tiles impose.
static double relToleranceFor(int k) {
    return 8.0 * std::sqrt(static_cast<double>(k)) *
           static_cast<double>(FLT_EPSILON);
}

static int blocksFor(int n, int threads) {
    return (n + threads - 1) / threads;
}

// Times a launch sequence 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.
template <typename LaunchFn>
static float timeOnce(LaunchFn launch) {
    cudaEvent_t start, stop;
    CUDA_CHECK(cudaEventCreate(&start));
    CUDA_CHECK(cudaEventCreate(&stop));
    CUDA_CHECK(cudaEventRecord(start));
    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;
}

// Plain data. No methods and no hidden destructor: main owns an explicit
// cleanup lambda used by both the normal and profile-only paths.
struct Net {
    float* params;
    float* grads;
    float* vel;
    float* x;
    float* act;
    float* logits;
    float* dlogits;
    float* dact;
    float* rowLoss;
    int* labels;
};

// The float64 CPU reference: forward, analytic backward, and a training loop

// The reference never allocates; the caller owns every buffer. It promotes
// the same float parameters the GPU reads, so it checks the kernels'
// arithmetic and not the author's choice of initial values.
static void forwardCpu(const std::vector<double>& p, const float* x, int m,
                       std::vector<double>& act, std::vector<double>& logits) {
    for (int r = 0; r < m; ++r) {
        for (int h = 0; h < kHidden; ++h) {
            double s = p[kOffB1 + h];
            for (int i = 0; i < kInput; ++i) {
                s += static_cast<double>(
                         x[static_cast<size_t>(r) * kInput + i]) *
                     p[kOffW1 + static_cast<size_t>(i) * kHidden + h];
            }
            act[static_cast<size_t>(r) * kHidden + h] = (s > 0.0) ? s : 0.0;
        }
        for (int c = 0; c < kClasses; ++c) {
            double s = p[kOffB2 + c];
            for (int h = 0; h < kHidden; ++h) {
                s += act[static_cast<size_t>(r) * kHidden + h] *
                     p[kOffW2 + static_cast<size_t>(h) * kClasses + c];
            }
            logits[static_cast<size_t>(r) * kClasses + c] = s;
        }
    }
}

// Mean cross entropy over the batch, computed with the max subtraction so
// the reference is stable everywhere the kernel claims to be.
static double lossCpu(const std::vector<double>& logits, const int* labels,
                      int m) {
    double total = 0.0;
    for (int r = 0; r < m; ++r) {
        const double* row = &logits[static_cast<size_t>(r) * kClasses];
        double best = row[0];
        for (int c = 1; c < kClasses; ++c) {
            best = (row[c] > best) ? row[c] : best;
        }
        double sum = 0.0;
        for (int c = 0; c < kClasses; ++c) {
            sum += std::exp(row[c] - best);
        }
        total += -(row[labels[r]] - best - std::log(sum));
    }
    return total / m;
}

// Analytic backward in float64. dlogits carries the 1/m of the mean already,
// so every gradient below is the gradient of the mean loss.
static void backwardCpu(const std::vector<double>& p, const float* x,
                        const int* labels, int m,
                        const std::vector<double>& act,
                        const std::vector<double>& logits,
                        std::vector<double>& grads) {
    std::vector<double> dlogits(static_cast<size_t>(m) * kClasses, 0.0);
    for (int r = 0; r < m; ++r) {
        const double* row = &logits[static_cast<size_t>(r) * kClasses];
        double best = row[0];
        for (int c = 1; c < kClasses; ++c) {
            best = (row[c] > best) ? row[c] : best;
        }
        double sum = 0.0;
        for (int c = 0; c < kClasses; ++c) {
            sum += std::exp(row[c] - best);
        }
        for (int c = 0; c < kClasses; ++c) {
            const double prob = std::exp(row[c] - best) / sum;
            const double hot = (c == labels[r]) ? 1.0 : 0.0;
            dlogits[static_cast<size_t>(r) * kClasses + c] =
                (prob - hot) / static_cast<double>(m);
        }
    }

    for (size_t i = 0; i < grads.size(); ++i) {
        grads[i] = 0.0;
    }
    std::vector<double> dact(static_cast<size_t>(m) * kHidden, 0.0);
    for (int r = 0; r < m; ++r) {
        for (int c = 0; c < kClasses; ++c) {
            const double d = dlogits[static_cast<size_t>(r) * kClasses + c];
            grads[kOffB2 + c] += d;
            for (int h = 0; h < kHidden; ++h) {
                grads[kOffW2 + static_cast<size_t>(h) * kClasses + c] +=
                    act[static_cast<size_t>(r) * kHidden + h] * d;
                dact[static_cast<size_t>(r) * kHidden + h] +=
                    d * p[kOffW2 + static_cast<size_t>(h) * kClasses + c];
            }
        }
    }
    for (int r = 0; r < m; ++r) {
        for (int h = 0; h < kHidden; ++h) {
            const size_t idxA = static_cast<size_t>(r) * kHidden + h;
            const double d = (act[idxA] > 0.0) ? dact[idxA] : 0.0;
            grads[kOffB1 + h] += d;
            for (int i = 0; i < kInput; ++i) {
                grads[kOffW1 + static_cast<size_t>(i) * kHidden + h] +=
                    static_cast<double>(
                        x[static_cast<size_t>(r) * kInput + i]) *
                    d;
            }
        }
    }
}

// Argmax accuracy in percent, from a float64 forward over the test split.
static double accuracyCpu(const std::vector<double>& p, const float* x,
                          const int* labels, int m, std::vector<double>& act,
                          std::vector<double>& logits) {
    forwardCpu(p, x, m, act, logits);
    int right = 0;
    for (int r = 0; r < m; ++r) {
        const double* row = &logits[static_cast<size_t>(r) * kClasses];
        int best = 0;
        for (int c = 1; c < kClasses; ++c) {
            if (row[c] > row[best]) {
                best = c;
            }
        }
        if (best == labels[r]) {
            ++right;
        }
    }
    return 100.0 * static_cast<double>(right) / static_cast<double>(m);
}

// The GPU training step

static void forwardGpu(const Net& net, const float* x, int m, bool fused) {
    const dim3 block(kTileDim, kTileDim);
    const dim3 grid1(blocksFor(kHidden, kTileDim), blocksFor(m, kTileDim));
    const dim3 grid2(blocksFor(kClasses, kTileDim), blocksFor(m, kTileDim));
    if (fused) {
        gemmTiled<false, false, true, true>
            <<<grid1, block>>>(x, net.params + kOffW1, net.params + kOffB1,
                               net.act, m, kHidden, kInput);
        CUDA_CHECK(cudaGetLastError());
        gemmTiled<false, false, true, false><<<grid2, block>>>(
            net.act, net.params + kOffW2, net.params + kOffB2, net.logits, m,
            kClasses, kHidden);
        CUDA_CHECK(cudaGetLastError());
        return;
    }

    const int hidden = m * kHidden;
    const int out = m * kClasses;
    gemmTiled<false, false, false, false><<<grid1, block>>>(
        x, net.params + kOffW1, nullptr, net.act, m, kHidden, kInput);
    CUDA_CHECK(cudaGetLastError());
    addBias<<<blocksFor(hidden, kThreadsPerBlock), kThreadsPerBlock>>>(
        net.act, net.params + kOffB1, m, kHidden);
    CUDA_CHECK(cudaGetLastError());
    reluForward<<<blocksFor(hidden, kThreadsPerBlock), kThreadsPerBlock>>>(
        net.act, hidden);
    CUDA_CHECK(cudaGetLastError());
    gemmTiled<false, false, false, false>
        <<<grid2, block>>>(net.act, net.params + kOffW2, nullptr, net.logits, m,
                           kClasses, kHidden);
    CUDA_CHECK(cudaGetLastError());
    addBias<<<blocksFor(out, kThreadsPerBlock), kThreadsPerBlock>>>(
        net.logits, net.params + kOffB2, m, kClasses);
    CUDA_CHECK(cudaGetLastError());
}

static void lossAndGradsGpu(const Net& net, const float* x, const int* labels,
                            int m) {
    const dim3 block(kTileDim, kTileDim);
    const int warpsPerBlock = kThreadsPerBlock / kWarpSize;
    softmaxCrossEntropy<<<blocksFor(m, warpsPerBlock), kThreadsPerBlock>>>(
        net.logits, labels, net.dlogits, net.rowLoss, m, kClasses,
        1.0f / static_cast<float>(m));
    CUDA_CHECK(cudaGetLastError());

    // dW2 = act^T dlogits, [kHidden x kClasses].
    const dim3 gridW2(blocksFor(kClasses, kTileDim),
                      blocksFor(kHidden, kTileDim));
    gemmTiled<true, false, false, false>
        <<<gridW2, block>>>(net.act, net.dlogits, nullptr, net.grads + kOffW2,
                            kHidden, kClasses, m);
    CUDA_CHECK(cudaGetLastError());
    columnSum<<<blocksFor(kClasses, kThreadsPerBlock), kThreadsPerBlock>>>(
        net.dlogits, net.grads + kOffB2, m, kClasses);
    CUDA_CHECK(cudaGetLastError());

    // dact = dlogits W2^T, [m x kHidden], then the ReLU mask in place.
    const dim3 gridAct(blocksFor(kHidden, kTileDim), blocksFor(m, kTileDim));
    gemmTiled<false, true, false, false>
        <<<gridAct, block>>>(net.dlogits, net.params + kOffW2, nullptr,
                             net.dact, m, kHidden, kClasses);
    CUDA_CHECK(cudaGetLastError());
    const int hidden = m * kHidden;
    reluBackward<<<blocksFor(hidden, kThreadsPerBlock), kThreadsPerBlock>>>(
        net.dact, net.act, hidden);
    CUDA_CHECK(cudaGetLastError());

    // dW1 = x^T dact, [kInput x kHidden].
    const dim3 gridW1(blocksFor(kHidden, kTileDim),
                      blocksFor(kInput, kTileDim));
    gemmTiled<true, false, false, false><<<gridW1, block>>>(
        x, net.dact, nullptr, net.grads + kOffW1, kInput, kHidden, m);
    CUDA_CHECK(cudaGetLastError());
    columnSum<<<blocksFor(kHidden, kThreadsPerBlock), kThreadsPerBlock>>>(
        net.dact, net.grads + kOffB1, m, kHidden);
    CUDA_CHECK(cudaGetLastError());
}

static void stepGpu(const Net& net, bool fused) {
    const int threads = kThreadsPerBlock;
    if (fused) {
        sgdMomentum<<<blocksFor(kParamCount, threads), threads>>>(
            net.params, net.grads, net.vel, kLearningRate, kMomentum,
            kParamCount);
        CUDA_CHECK(cudaGetLastError());
        return;
    }
    const int offs[4] = {kOffW1, kOffB1, kOffW2, kOffB2};
    const int sizes[4] = {kInput * kHidden, kHidden, kHidden * kClasses,
                          kClasses};
    for (int t = 0; t < 4; ++t) {
        sgdMomentum<<<blocksFor(sizes[t], threads), threads>>>(
            net.params + offs[t], net.grads + offs[t], net.vel + offs[t],
            kLearningRate, kMomentum, sizes[t]);
        CUDA_CHECK(cudaGetLastError());
    }
}

// One training step: forward, loss and dlogits, backward, update. The NVTX
// ranges are what the nsys report groups by, so the write-up can name a
// phase and not only a kernel.
static void trainStep(const Net& net, const float* x, const int* labels, int m,
                      bool fused) {
    nvtxRangePushA("forward");
    forwardGpu(net, x, m, fused);
    nvtxRangePop();
    nvtxRangePushA("loss-backward");
    lossAndGradsGpu(net, x, labels, m);
    nvtxRangePop();
    nvtxRangePushA("step");
    stepGpu(net, fused);
    nvtxRangePop();
}

static double meanLossGpu(const Net& net, int m, std::vector<float>& scratch) {
    CUDA_CHECK(cudaMemcpy(scratch.data(), net.rowLoss,
                          static_cast<size_t>(m) * sizeof(float),
                          cudaMemcpyDeviceToHost));
    double total = 0.0;
    for (int r = 0; r < m; ++r) {
        total += scratch[r];
    }
    return total / m;
}

static double accuracyGpu(const Net& net, const float* dTestX,
                          const std::vector<int>& testLabels,
                          std::vector<float>& logitScratch) {
    forwardGpu(net, dTestX, kTest, true);
    CUDA_CHECK(cudaGetLastError());
    CUDA_CHECK(cudaDeviceSynchronize());
    CUDA_CHECK(cudaMemcpy(logitScratch.data(), net.logits,
                          static_cast<size_t>(kTest) * kClasses * sizeof(float),
                          cudaMemcpyDeviceToHost));
    int right = 0;
    for (int r = 0; r < kTest; ++r) {
        const float* row = &logitScratch[static_cast<size_t>(r) * kClasses];
        int best = 0;
        for (int c = 1; c < kClasses; ++c) {
            if (row[c] > row[best]) {
                best = c;
            }
        }
        if (best == testLabels[r]) {
            ++right;
        }
    }
    return 100.0 * static_cast<double>(right) / static_cast<double>(kTest);
}

int main(int argc, char** argv) {
    bool profileOnly = false;
    std::string dataDir = "../../data/day100";
    for (int i = 1; i < argc; ++i) {
        if (std::strcmp(argv[i], "--profile") == 0) {
            profileOnly = true;
        } else if (std::strcmp(argv[i], "--data") == 0 && i + 1 < argc) {
            dataDir = argv[++i];
        } else {
            std::fprintf(stderr, "usage: %s [--profile] [--data DIR]\n",
                         argv[0]);
            return EXIT_FAILURE;
        }
    }

    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(
        "net %d-%d-%d, %d parameters, batch %d, lr %.2f, momentum "
        "%.2f, %d epochs, seed %llu\n",
        kInput, kHidden, kClasses, kParamCount, kBatch,
        static_cast<double>(kLearningRate), static_cast<double>(kMomentum),
        kEpochs, static_cast<unsigned long long>(kSeed));

    // The vendored Fashion-MNIST subset. SOURCE.md beside these files records
    // the source, MIT license, regeneration command and hashes.
    const int total = kTrain + kTest;
    std::vector<float> hostX(static_cast<size_t>(total) * kInput);
    std::vector<int> hostY(total);
    loadIdxSplit(dataDir + "/train-images-idx3-ubyte.gz",
                 dataDir + "/train-labels-idx1-ubyte.gz", kTrain,
                 hostX.data(), hostY.data());
    loadIdxSplit(dataDir + "/t10k-images-idx3-ubyte.gz",
                 dataDir + "/t10k-labels-idx1-ubyte.gz", kTest,
                 hostX.data() + static_cast<size_t>(kTrain) * kInput,
                 hostY.data() + kTrain);
    std::printf("dataset: Fashion-MNIST, %d train and %d held-out test images\n",
                kTrain, kTest);

    // One seeded shuffle of the training split, done once. Every epoch walks
    // the same order, on the GPU and on the CPU reference alike.
    for (int i = kTrain - 1; i > 0; --i) {
        const int j = static_cast<int>(mix64(kSeed + 5ull + i) % (i + 1));
        for (int c = 0; c < kInput; ++c) {
            const float tmp = hostX[static_cast<size_t>(i) * kInput + c];
            hostX[static_cast<size_t>(i) * kInput + c] =
                hostX[static_cast<size_t>(j) * kInput + c];
            hostX[static_cast<size_t>(j) * kInput + c] = tmp;
        }
        const int tmpY = hostY[i];
        hostY[i] = hostY[j];
        hostY[j] = tmpY;
    }

    // Print the subset's class histogram so the exact training input is part
    // of the transcript rather than an implicit property of the data files.
    {
        int hist[kClasses] = {0};
        for (int s = 0; s < kTrain; ++s) {
            ++hist[hostY[s]];
        }
        std::printf("train class counts:");
        for (int c = 0; c < kClasses; ++c) {
            std::printf(" %d", hist[c]);
        }
        std::printf("\n");
    }

    // Initial parameters, one buffer, shared by both training runs.
    std::vector<float> initParams(kParamCount, 0.0f);
    {
        const double a1 = std::sqrt(6.0 / (kInput + kHidden));
        const double a2 = std::sqrt(6.0 / (kHidden + kClasses));
        for (int i = 0; i < kInput * kHidden; ++i) {
            initParams[kOffW1 + i] =
                static_cast<float>(uniform(kSeed + 7ull, i) * a1);
        }
        for (int i = 0; i < kHidden * kClasses; ++i) {
            initParams[kOffW2 + i] =
                static_cast<float>(uniform(kSeed + 8ull, i) * a2);
        }
    }
    std::vector<double> initParamsD(kParamCount);
    for (int i = 0; i < kParamCount; ++i) {
        initParamsD[i] = initParams[i];
    }

    // Device buffers.
    Net net;
    CUDA_CHECK(cudaMalloc(&net.params, kParamCount * sizeof(float)));
    CUDA_CHECK(cudaMalloc(&net.grads, kParamCount * sizeof(float)));
    CUDA_CHECK(cudaMalloc(&net.vel, kParamCount * sizeof(float)));
    CUDA_CHECK(cudaMalloc(&net.x,
                          static_cast<size_t>(total) * kInput * sizeof(float)));
    CUDA_CHECK(cudaMalloc(
        &net.act, static_cast<size_t>(kActRows) * kHidden * sizeof(float)));
    CUDA_CHECK(cudaMalloc(
        &net.logits, static_cast<size_t>(kActRows) * kClasses * sizeof(float)));
    CUDA_CHECK(cudaMalloc(&net.dlogits, static_cast<size_t>(kActRows) *
                                            kClasses * sizeof(float)));
    CUDA_CHECK(cudaMalloc(
        &net.dact, static_cast<size_t>(kActRows) * kHidden * sizeof(float)));
    CUDA_CHECK(cudaMalloc(&net.rowLoss, kActRows * sizeof(float)));
    CUDA_CHECK(cudaMalloc(&net.labels, total * sizeof(int)));
    CUDA_CHECK(cudaMemcpy(net.x, hostX.data(),
                          static_cast<size_t>(total) * kInput * sizeof(float),
                          cudaMemcpyHostToDevice));
    CUDA_CHECK(cudaMemcpy(net.labels, hostY.data(), total * sizeof(int),
                          cudaMemcpyHostToDevice));

    const auto freeNet = [&] {
        CUDA_CHECK(cudaFree(net.params));
        CUDA_CHECK(cudaFree(net.grads));
        CUDA_CHECK(cudaFree(net.vel));
        CUDA_CHECK(cudaFree(net.x));
        CUDA_CHECK(cudaFree(net.act));
        CUDA_CHECK(cudaFree(net.logits));
        CUDA_CHECK(cudaFree(net.dlogits));
        CUDA_CHECK(cudaFree(net.dact));
        CUDA_CHECK(cudaFree(net.rowLoss));
        CUDA_CHECK(cudaFree(net.labels));
    };

    if (profileOnly) {
        CUDA_CHECK(cudaMemcpy(net.params, initParams.data(),
                              kParamCount * sizeof(float),
                              cudaMemcpyHostToDevice));
        CUDA_CHECK(cudaMemset(net.vel, 0, kParamCount * sizeof(float)));

        const nvtxStringHandle_t profileEpochMessage =
            nvtxDomainRegisterStringA(nullptr, "profile-epoch");
        nvtxEventAttributes_t profileEpoch = {};
        profileEpoch.version = NVTX_VERSION;
        profileEpoch.size = NVTX_EVENT_ATTRIB_STRUCT_SIZE;
        profileEpoch.messageType = NVTX_MESSAGE_TYPE_REGISTERED;
        profileEpoch.message.registered = profileEpochMessage;

        // Warm every kernel before opening the capture range, then expose one
        // representative fused epoch and nothing from the CPU reference or
        // the correctness harness to Nsight Systems.
        for (int s = 0; s < kStepsPerEpoch; ++s) {
            const int rows = rowsInStep(s);
            trainStep(net, net.x + static_cast<size_t>(s) * kBatch * kInput,
                      net.labels + s * kBatch, rows, true);
        }
        CUDA_CHECK(cudaDeviceSynchronize());

        nvtxRangePushEx(&profileEpoch);
        for (int s = 0; s < kStepsPerEpoch; ++s) {
            const int rows = rowsInStep(s);
            trainStep(net, net.x + static_cast<size_t>(s) * kBatch * kInput,
                      net.labels + s * kBatch, rows, true);
        }
        CUDA_CHECK(cudaDeviceSynchronize());
        nvtxRangePop();
        std::printf("profile-only: captured one fused epoch after warm-up\n");
        freeNet();
        return EXIT_SUCCESS;
    }

    float* const dTestX = net.x + static_cast<size_t>(kTrain) * kInput;
    std::vector<int> testLabels(hostY.begin() + kTrain, hostY.end());
    std::vector<float> logitScratch(static_cast<size_t>(kActRows) * kClasses);
    std::vector<float> lossScratch(kActRows);
    std::vector<double> actCpu(static_cast<size_t>(kActRows) * kHidden);
    std::vector<double> logitsCpu(static_cast<size_t>(kActRows) * kClasses);

    int status = EXIT_SUCCESS;

    // Cases 1 and 2: first-light forward at batch 1, then the fixed batch.
    {
        CUDA_CHECK(cudaMemcpy(net.params, initParams.data(),
                              kParamCount * sizeof(float),
                              cudaMemcpyHostToDevice));
        const int batches[2] = {1, kBatch};
        const char* names[2] = {"onebatch-forward", "batch-forward"};
        for (int which = 0; which < 2; ++which) {
            const int m = batches[which];
            forwardGpu(net, net.x, m, true);
            CUDA_CHECK(cudaGetLastError());
            CUDA_CHECK(cudaDeviceSynchronize());
            CUDA_CHECK(cudaMemcpy(
                logitScratch.data(), net.logits,
                static_cast<size_t>(m) * kClasses * sizeof(float),
                cudaMemcpyDeviceToHost));
            forwardCpu(initParamsD, hostX.data(), m, actCpu, logitsCpu);

            const double rtol = relToleranceFor(kInput);
            double maxAbs = 0.0;
            for (int i = 0; i < m * kClasses; ++i) {
                maxAbs = std::fmax(maxAbs, std::fabs(logitsCpu[i]));
            }
            const double atol = rtol * maxAbs;
            double worstRatio = 0.0;
            double worstError = 0.0;
            double worstAllowed = 0.0;
            int worstAt = -1;
            bool finite = true;
            for (int i = 0; i < m * kClasses; ++i) {
                const double want = logitsCpu[i];
                const double got = logitScratch[i];
                if (!std::isfinite(got)) {
                    finite = false;
                    worstAt = i;
                    break;
                }
                const double error = std::fabs(got - want);
                const double allowed = atol + rtol * std::fabs(want);
                const double ratio = error / allowed;
                if (ratio > worstRatio) {
                    worstRatio = ratio;
                    worstError = error;
                    worstAllowed = allowed;
                    worstAt = i;
                }
            }
            std::printf(
                "case %d %s: worst |got-ref|/allowance %.3f at %d, "
                "error %.3e, allowance %.3e\n",
                which + 1, names[which], worstRatio, worstAt, worstError,
                worstAllowed);
            std::printf(
                "  tolerance |got-ref| <= %.3e + %.3e*|ref| "
                "(depth %d, max |ref| %.3e)\n",
                atol, rtol, kInput, maxAbs);
            if (!finite || !(worstRatio <= 1.0)) {
                std::fprintf(stderr, "case %d FAILED\n", which + 1);
                status = EXIT_FAILURE;
                break;
            }
        }
    }

    // Case 3: the softmax at plus and minus 60.
    if (status == EXIT_SUCCESS) {
        const int m = 4;
        std::vector<float> extreme(static_cast<size_t>(m) * kClasses);
        std::vector<int> labels(m);
        for (int r = 0; r < m; ++r) {
            for (int c = 0; c < kClasses; ++c) {
                extreme[static_cast<size_t>(r) * kClasses + c] =
                    ((r + c) % 2 == 0) ? 60.0f : -60.0f;
            }
            labels[r] = r % kClasses;
        }
        CUDA_CHECK(cudaMemcpy(net.logits, extreme.data(),
                              extreme.size() * sizeof(float),
                              cudaMemcpyHostToDevice));
        CUDA_CHECK(cudaMemcpy(net.labels, labels.data(), m * sizeof(int),
                              cudaMemcpyHostToDevice));
        const int warpsPerBlock = kThreadsPerBlock / kWarpSize;
        softmaxCrossEntropy<<<blocksFor(m, warpsPerBlock), kThreadsPerBlock>>>(
            net.logits, net.labels, net.dlogits, net.rowLoss, m, kClasses,
            1.0f / static_cast<float>(m));
        CUDA_CHECK(cudaGetLastError());
        CUDA_CHECK(cudaDeviceSynchronize());

        std::vector<double> extremeD(extreme.size());
        for (size_t i = 0; i < extreme.size(); ++i) {
            extremeD[i] = extreme[i];
        }
        const double want = lossCpu(extremeD, labels.data(), m);
        const double got = meanLossGpu(net, m, lossScratch);
        const bool finite = std::isfinite(got);
        // Depth 10, one exp and one log per term, so the float32 path earns
        // a looser bound than a dot product of the same length.
        const double rtol = 4.0 * relToleranceFor(kClasses);
        const double rel = std::fabs(got - want) / std::fabs(want);
        std::printf(
            "case 3 loss-extreme: gpu %.6f, float64 %.6f, relative "
            "%.3e, tolerance %.3e\n",
            got, want, rel, rtol);
        if (!finite || !(rel <= rtol)) {
            std::fprintf(stderr, "case 3 FAILED\n");
            status = EXIT_FAILURE;
        }
        CUDA_CHECK(cudaMemcpy(net.labels, hostY.data(), total * sizeof(int),
                              cudaMemcpyHostToDevice));
    }

    // Case 4: 200 parameter coordinates, gradients three ways.
    if (status == EXIT_SUCCESS) {
        CUDA_CHECK(cudaMemcpy(net.params, initParams.data(),
                              kParamCount * sizeof(float),
                              cudaMemcpyHostToDevice));
        forwardGpu(net, net.x, kGradBatch, true);
        lossAndGradsGpu(net, net.x, net.labels, kGradBatch);
        CUDA_CHECK(cudaGetLastError());
        CUDA_CHECK(cudaDeviceSynchronize());
        std::vector<float> gradGpu(kParamCount);
        CUDA_CHECK(cudaMemcpy(gradGpu.data(), net.grads,
                              kParamCount * sizeof(float),
                              cudaMemcpyDeviceToHost));

        std::vector<double> gradCpu(kParamCount, 0.0);
        forwardCpu(initParamsD, hostX.data(), kGradBatch, actCpu, logitsCpu);
        backwardCpu(initParamsD, hostX.data(), hostY.data(), kGradBatch, actCpu,
                    logitsCpu, gradCpu);

        // The finite difference is the independent witness: it never touches
        // either backward pass, only the float64 forward and the loss.
        const int starts[4] = {kOffW1, kOffB1, kOffW2, kOffB2};
        const int sizes[4] = {kInput * kHidden, kHidden, kHidden * kClasses,
                              kClasses};
        const char* names[4] = {"dW1", "db1", "dW2", "db2"};
        const double eps = 1e-5;
        double worstFd = 0.0;
        double worstGpu = 0.0;
        std::vector<double> probe(initParamsD);
        for (int t = 0; t < 4; ++t) {
            for (int s = 0; s < kGradCoordsPerTensor; ++s) {
                const int idx = starts[t] + static_cast<int>(
                    mix64(kSeed + 101ull * t + s) % sizes[t]);
                const double keep = probe[idx];
                probe[idx] = keep + eps;
                forwardCpu(probe, hostX.data(), kGradBatch, actCpu, logitsCpu);
                const double up = lossCpu(logitsCpu, hostY.data(), kGradBatch);
                probe[idx] = keep - eps;
                forwardCpu(probe, hostX.data(), kGradBatch, actCpu, logitsCpu);
                const double down =
                    lossCpu(logitsCpu, hostY.data(), kGradBatch);
                probe[idx] = keep;
                const double fd = (up - down) / (2.0 * eps);

                // Denominator: this coordinate's own gradient, floored at a
                // hundredth of the largest gradient in its tensor. A dead
                // ReLU makes a coordinate's true gradient exactly zero, and
                // dividing float32 noise by zero would fail a correct
                // kernel.
                double big = 0.0;
                for (int i = 0; i < sizes[t]; ++i) {
                    const double v = std::fabs(gradCpu[starts[t] + i]);
                    big = (v > big) ? v : big;
                }
                const double floorScale = 0.01 * big;
                const double scale = (std::fabs(gradCpu[idx]) > floorScale)
                                         ? std::fabs(gradCpu[idx])
                                         : floorScale;
                const double relFd = std::fabs(fd - gradCpu[idx]) / scale;
                const double relGpu =
                    std::fabs(gradGpu[idx] - gradCpu[idx]) / scale;
                worstFd = (relFd > worstFd) ? relFd : worstFd;
                worstGpu = (relGpu > worstGpu) ? relGpu : worstGpu;
                if (s == 0) {
                    std::printf(
                        "  %s[%d]: gpu %.6e  float64 %.6e  finite "
                        "diff %.6e\n",
                        names[t], idx - starts[t],
                        static_cast<double>(gradGpu[idx]), gradCpu[idx], fd);
                }
            }
        }
        // The finite difference carries its own truncation error, of order
        // eps squared on a float64 loss, so it gets the looser bound. The
        // GPU gradients are a float32 reduction of depth kGradBatch feeding
        // one of depth kInput, so they get the depth-scaled one.
        const double fdTol = 1e-4;
        const double gpuTol = relToleranceFor(kInput * kGradBatch);
        std::printf(
            "case 4 grad-check: %d coordinates, float64 vs finite "
            "difference %.3e (tol %.3e), gpu vs float64 %.3e "
            "(tol %.3e)\n",
            4 * kGradCoordsPerTensor, worstFd, fdTol, worstGpu, gpuTol);
        if (!(worstFd <= fdTol) || !(worstGpu <= gpuTol)) {
            std::fprintf(stderr, "case 4 FAILED\n");
            status = EXIT_FAILURE;
        }
    }

    // Case 5: one momentum SGD update against a float64 reference.
    if (status == EXIT_SUCCESS) {
        constexpr int n = 1024;
        std::vector<float> p(n), g(n), v(n), got(n);
        std::vector<double> want(n);
        for (int i = 0; i < n; ++i) {
            p[i] = uniform(kSeed + 20ull, i);
            g[i] = uniform(kSeed + 21ull, i);
            v[i] = uniform(kSeed + 22ull, i);
            const double nextV = kMomentum * static_cast<double>(v[i]) -
                                 kLearningRate * static_cast<double>(g[i]);
            want[i] = static_cast<double>(p[i]) + nextV;
        }
        CUDA_CHECK(cudaMemcpy(net.params, p.data(), n * sizeof(float),
                              cudaMemcpyHostToDevice));
        CUDA_CHECK(cudaMemcpy(net.grads, g.data(), n * sizeof(float),
                              cudaMemcpyHostToDevice));
        CUDA_CHECK(cudaMemcpy(net.vel, v.data(), n * sizeof(float),
                              cudaMemcpyHostToDevice));
        sgdMomentum<<<blocksFor(n, kThreadsPerBlock), kThreadsPerBlock>>>(
            net.params, net.grads, net.vel, kLearningRate, kMomentum, n);
        CUDA_CHECK(cudaGetLastError());
        CUDA_CHECK(cudaDeviceSynchronize());
        CUDA_CHECK(cudaMemcpy(got.data(), net.params, n * sizeof(float),
                              cudaMemcpyDeviceToHost));
        constexpr double rtol = 1.0e-5;
        constexpr double atol = 1.0e-6;
        double worstRatio = 0.0;
        double worstError = 0.0;
        double worstAllowed = 0.0;
        int worstAt = -1;
        bool finite = true;
        for (int i = 0; i < n; ++i) {
            if (!std::isfinite(got[i])) {
                finite = false;
                worstAt = i;
                break;
            }
            const double error = std::fabs(got[i] - want[i]);
            const double allowed = atol + rtol * std::fabs(want[i]);
            const double ratio = error / allowed;
            if (ratio > worstRatio) {
                worstRatio = ratio;
                worstError = error;
                worstAllowed = allowed;
                worstAt = i;
            }
        }
        std::printf(
            "case 5 step: worst |got-ref|/allowance %.3f at %d, "
            "error %.3e, allowance %.3e\n",
            worstRatio, worstAt, worstError, worstAllowed);
        std::printf("  tolerance |got-ref| <= %.1e + %.1e*|ref|\n", atol,
                    rtol);
        if (!finite || !(worstRatio <= 1.0)) {
            std::fprintf(stderr, "case 5 FAILED\n");
            status = EXIT_FAILURE;
        }
    }

    // Case 6: overfit 32 images.
    if (status == EXIT_SUCCESS) {
        CUDA_CHECK(cudaMemcpy(net.params, initParams.data(),
                              kParamCount * sizeof(float),
                              cudaMemcpyHostToDevice));
        CUDA_CHECK(cudaMemset(net.vel, 0, kParamCount * sizeof(float)));
        double last = 0.0;
        for (int s = 0; s < kOverfitSteps; ++s) {
            trainStep(net, net.x, net.labels, kOverfitSamples, true);
        }
        CUDA_CHECK(cudaDeviceSynchronize());
        forwardGpu(net, net.x, kOverfitSamples, true);
        lossAndGradsGpu(net, net.x, net.labels, kOverfitSamples);
        CUDA_CHECK(cudaDeviceSynchronize());
        last = meanLossGpu(net, kOverfitSamples, lossScratch);
        std::printf("case 6 overfit-32: loss after %d steps %.6f, gate %.3f\n",
                    kOverfitSteps, last, kOverfitLossGate);
        if (!(last < kOverfitLossGate)) {
            std::fprintf(stderr, "case 6 FAILED\n");
            status = EXIT_FAILURE;
        }
    }

    // The float64 CPU training run the accuracy gate compares against.
    double accCpuPct = 0.0;
    if (status == EXIT_SUCCESS) {
        std::vector<double> p(initParamsD);
        std::vector<double> v(kParamCount, 0.0);
        std::vector<double> g(kParamCount, 0.0);
        for (int e = 0; e < kEpochs; ++e) {
            for (int s = 0; s < kStepsPerEpoch; ++s) {
                const int rows = rowsInStep(s);
                const float* bx =
                    &hostX[static_cast<size_t>(s) * kBatch * kInput];
                const int* by = &hostY[static_cast<size_t>(s) * kBatch];
                forwardCpu(p, bx, rows, actCpu, logitsCpu);
                backwardCpu(p, bx, by, rows, actCpu, logitsCpu, g);
                for (int i = 0; i < kParamCount; ++i) {
                    v[i] = kMomentum * v[i] - kLearningRate * g[i];
                    p[i] += v[i];
                }
            }
        }
        accCpuPct = accuracyCpu(p, &hostX[static_cast<size_t>(kTrain) * kInput],
                                testLabels.data(), kTest, actCpu, logitsCpu);
        std::printf("float64 CPU reference: test accuracy %.2f percent\n",
                    accCpuPct);
    }

    // Cases 7 and 8: held-out accuracy, then fused/unfused epoch timing.
    double epochMean[2] = {0.0, 0.0};
    for (int variant = 0; variant < 2 && status == EXIT_SUCCESS; ++variant) {
        const bool fused = (variant == 0);
        const char* name = fused ? "case 7 train" : "case 8 perf-unfused";
        CUDA_CHECK(cudaMemcpy(net.params, initParams.data(),
                              kParamCount * sizeof(float),
                              cudaMemcpyHostToDevice));
        CUDA_CHECK(cudaMemset(net.vel, 0, kParamCount * sizeof(float)));

        double sum = 0.0;
        double lo = 1e300;
        double hi = 0.0;
        for (int e = 0; e < kEpochs; ++e) {
            const float ms = timeOnce([&]() {
                for (int s = 0; s < kStepsPerEpoch; ++s) {
                    const int rows = rowsInStep(s);
                    trainStep(net,
                              net.x + static_cast<size_t>(s) * kBatch * kInput,
                              net.labels + static_cast<size_t>(s) * kBatch,
                              rows, fused);
                }
            });
            // Epoch 0 launches every kernel the timed epochs launch, so the
            // lazy module load for all of them is paid before the clock
            // starts. It is dropped, not averaged in.
            if (e >= kWarmupEpochs) {
                sum += ms;
                lo = (ms < lo) ? ms : lo;
                hi = (ms > hi) ? ms : hi;
            }
        }
        const int timed = kEpochs - kWarmupEpochs;
        const double mean = sum / timed;
        epochMean[variant] = mean;
        // Evaluation always takes the fused forward path: it is the same
        // arithmetic either way, and it keeps the accuracy comparison about
        // training rather than about which kernels drew the test logits.
        const double accGpu =
            accuracyGpu(net, dTestX, testLabels, logitScratch);
        std::printf(
            "%s: %d timed epochs, mean %.3f ms (min %.3f, max %.3f), "
            "test accuracy %.2f percent\n",
            name, timed, mean, lo, hi, accGpu);
        std::printf("  per step %.4f ms over %d steps per epoch\n",
                    mean / kStepsPerEpoch, kStepsPerEpoch);
        if (!(std::fabs(accGpu - accCpuPct) <= kAccuracyGatePoints)) {
            std::fprintf(stderr,
                         "%s FAILED: accuracy %.2f is more than %.2f points "
                         "from the float64 reference %.2f\n",
                         name, accGpu, kAccuracyGatePoints, accCpuPct);
            status = EXIT_FAILURE;
        }
    }

    if (status == EXIT_SUCCESS) {
        // The bronze comparison. Fusion removes no arithmetic at all: both
        // runs do the same multiplies and adds on the same bytes of input.
        // Whatever this ratio is, it is launches and round trips.
        std::printf(
            "naive / fused epoch time: %.3f (naive %.3f ms over "
            "fused %.3f ms)\n",
            epochMean[1] / epochMean[0], epochMean[1], epochMean[0]);
        std::printf("launches per step: fused 10, naive 17\n");
    }

    // The same explicit cleanup is taken whether the cases passed or failed.
    freeNet();

    if (status != EXIT_SUCCESS) {
        std::fprintf(stderr, "at least one case failed\n");
        return EXIT_FAILURE;
    }
    std::printf("all eight cases passed\n");
    return EXIT_SUCCESS;
}