COURSE / SOURCE

norms.cu

All lessons
Source filecode/day95-norms/norms.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 95: the three row-wise kernels a transformer runs on every token,
// written with one shape. One block owns one row, every thread strides over
// that row, a warp-shuffle reduction collapses the row to one number, and
// that number is broadcast back to every thread before the write pass.
// Softmax needs two of those reductions (max, then sum), two-pass layer norm
// needs two (mean, then variance), Welford's one-pass form needs one merged
// reduction, and RMS norm needs one plain one. That count is the whole
// performance story of the page.
//
// Correctness is a double-precision host reference over the original float
// inputs, with the day 66 K-scaled tolerance where K is the row length,
// because every output depends on a reduction over exactly that many terms.
//
// Two claims are tested, but only the first is a correctness gate:
//   1. softmaxRowNaive must produce a non-finite value on the part 1 row.
//      If it does not, the overflow this day exists to explain did not
//      happen and the page is wrong about expf.
//   2. rmsNormRow is predicted to beat layerNormRowTwoPass because it does
//      fewer read sweeps. The measured result is reported as held or refuted;
//      scheduling and cache effects are not correctness failures.
//
// Build: nvcc -std=c++17 -O3 -arch=sm_75 -o norms norms.cu
// Run:   ./norms

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

// 256 threads, eight warps, the course default. Every kernel here assumes
// exactly this block shape: the block reduction writes one partial per warp
// into an eight-slot shared array and every thread reads all eight back.
constexpr int kThreadsPerBlock = 256;
constexpr int kWarpSize = 32;
constexpr int kWarpsPerBlock = kThreadsPerBlock / kWarpSize;
constexpr unsigned int kFullMask = 0xffffffffu;

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

// torch.nn.LayerNorm's default eps, used here for RMS norm as well. It is
// not torch.nn.RMSNorm's default: that one takes eps=None and then uses the
// machine epsilon of the compute type, 1.1920929e-07 for fp32 input, so a
// learner comparing part 0 against PyTorch has to pass eps=1e-5 explicitly.
// Both defaults checked 2026-09-01 at
// https://docs.pytorch.org/docs/2.13/generated/torch.nn.LayerNorm.html and
// https://docs.pytorch.org/docs/2.13/generated/torch.nn.RMSNorm.html
constexpr float kEps = 1e-5f;

// Two cases with the same element count and the same bytes, so the timing
// table compares reduction depth rather than problem size. 16384 rows of
// 1024 and 4096 rows of 4096 both move 64 MiB in and 64 MiB out, which is
// sixteen times the T4's 4 MiB L2, so neither case hides in cache.
constexpr int kNumCases = 2;
constexpr int kNumPaths = 5;
constexpr size_t kCaseRows[kNumCases] = {16384, 4096};
constexpr size_t kCaseCols[kNumCases] = {1024, 4096};
constexpr size_t kMaxElems = 16777216;
constexpr size_t kMaxCols = 4096;

// The part 1 row: 1024 logits spread over 200, which is what an unshifted
// exponential cannot survive and a shifted one can.
constexpr size_t kDemoCols = 1024;

// The part 0 row, small enough to paste into Python and compare by eye.
constexpr size_t kTorchCols = 8;

// f32 base rtol from the course tolerance table, and the mantissa bit count
// that sets eps = 2^-p.
constexpr double kF32Rtol = 1e-5;
constexpr int kF32MantissaBits = 23;

// The comparative timing gate's slack. Both kernels re-read a row that may
// still be in L2, so their times can converge; 5 percent is the margin
// inside which "no faster" is a tie rather than a refutation.
constexpr double kTimingSlack = 1.05;

static_assert(kWarpsPerBlock * kWarpSize == kThreadsPerBlock,
              "the block reduction stores one partial per warp in an "
              "eight-slot shared array, so the block must be eight warps");
static_assert(kCaseRows[0] * kCaseCols[0] == kMaxElems &&
                  kCaseRows[1] * kCaseCols[1] == kMaxElems,
              "both cases must move identical bytes, or the timing table "
              "compares problem size instead of reduction depth");
static_assert(kCaseCols[0] <= kMaxCols && kCaseCols[1] <= kMaxCols &&
                  kDemoCols <= kMaxCols && kTorchCols <= kMaxCols,
              "gamma and beta are allocated once at kMaxCols");
static_assert(kCaseRows[0] <= 65535 && kCaseRows[1] <= 65535,
              "one block per row, and the demo keeps the grid inside the "
              "65535-block x dimension every compute capability guarantees");

// The reduction every kernel on this page is built from. Day 23's
// __shfl_down_sync over a full warp: five steps, no shared memory, no
// barrier, and after them lane 0 holds the warp's total.
__device__ float warpReduceSum(float v) {
    for (int offset = kWarpSize / 2; offset > 0; offset /= 2) {
        v += __shfl_down_sync(kFullMask, v, offset);
    }
    return v;
}

__device__ float warpReduceMax(float v) {
    for (int offset = kWarpSize / 2; offset > 0; offset /= 2) {
        v = fmaxf(v, __shfl_down_sync(kFullMask, v, offset));
    }
    return v;
}

// Warp partials into shared memory, then every thread reads all eight back
// and adds them. That last loop is the broadcast: a row-wise kernel needs
// the total in every thread, not in lane 0, because every thread has
// elements to scale with it. Eight shared loads per thread is cheaper than
// a second shuffle stage plus a broadcast shuffle, and it is the same shape
// for sum and max.
//
// The trailing barrier exists so the caller may reuse `smem` for a second
// reduction, which softmax and two-pass layer norm both do.
__device__ float blockReduceSum(float v, float* smem) {
    const unsigned int lane = threadIdx.x % kWarpSize;
    const unsigned int warp = threadIdx.x / kWarpSize;
    v = warpReduceSum(v);
    if (lane == 0) {
        smem[warp] = v;
    }
    __syncthreads();
    float total = 0.0f;
    for (int w = 0; w < kWarpsPerBlock; ++w) {
        total += smem[w];
    }
    __syncthreads();
    return total;
}

__device__ float blockReduceMax(float v, float* smem) {
    const unsigned int lane = threadIdx.x % kWarpSize;
    const unsigned int warp = threadIdx.x / kWarpSize;
    v = warpReduceMax(v);
    if (lane == 0) {
        smem[warp] = v;
    }
    __syncthreads();
    float total = -INFINITY;
    for (int w = 0; w < kWarpsPerBlock; ++w) {
        total = fmaxf(total, smem[w]);
    }
    __syncthreads();
    return total;
}

// Softmax without the max subtraction, which is how everyone writes it the
// first time. Correct arithmetic, correct on ordinary data, and it returns
// inf and NaN the moment a logit passes about 88.7, because that is where
// expf overflows a float. Part 1 gates on it failing.
//
// One thread handles cols/256 elements of one row, strided by the block, so
// the 32 lanes of a warp read 32 consecutive floats and every load
// coalesces. Two read sweeps and one write: the sum, then a write loop that
// reads the row a second time.
//
// Launch assumption: one block of exactly kThreadsPerBlock threads per row.
// snippet: softmax-naive
__global__ void softmaxRowNaive(const float* __restrict__ in,
                                float* __restrict__ out, size_t cols) {
    __shared__ float smem[kWarpsPerBlock];
    const float* row = in + blockIdx.x * cols;
    float* outRow = out + blockIdx.x * cols;

    float partial = 0.0f;
    for (size_t c = threadIdx.x; c < cols; c += blockDim.x) {
        partial += expf(row[c]);
    }
    const float total = blockReduceSum(partial, smem);

    for (size_t c = threadIdx.x; c < cols; c += blockDim.x) {
        outRow[c] = expf(row[c]) / total;
    }
}
// end snippet

// The same softmax with the shift that makes it finite. Subtracting the row
// max leaves every exponent at or below zero, so expf cannot overflow; the
// mathematical result is unchanged because the shift cancels between
// numerator and denominator. Three read sweeps and one write: max, sum, and
// the write loop's own read.
//
// Memory and launch assumption: identical to softmaxRowNaive. The extra
// pass is a re-read of a row that may still be in L1 or L2, which is the
// thing the bandwidth column of part 3 is measuring.
// snippet: softmax-safe
__global__ void softmaxRowSafe(const float* __restrict__ in,
                               float* __restrict__ out, size_t cols) {
    __shared__ float smem[kWarpsPerBlock];
    const float* row = in + blockIdx.x * cols;
    float* outRow = out + blockIdx.x * cols;

    float best = -INFINITY;
    for (size_t c = threadIdx.x; c < cols; c += blockDim.x) {
        best = fmaxf(best, row[c]);
    }
    const float rowMax = blockReduceMax(best, smem);

    float partial = 0.0f;
    for (size_t c = threadIdx.x; c < cols; c += blockDim.x) {
        partial += expf(row[c] - rowMax);
    }
    const float total = blockReduceSum(partial, smem);

    for (size_t c = threadIdx.x; c < cols; c += blockDim.x) {
        outRow[c] = expf(row[c] - rowMax) / total;
    }
}
// end snippet

// Layer norm the way the paper reads: mean first, then variance around that
// mean, then the affine write. Three read sweeps and one write.
//
// The variance is the biased one, divided by cols rather than cols - 1,
// which is what torch.nn.LayerNorm computes.
//
// Memory and launch assumption: identical to softmaxRowNaive.
__global__ void layerNormRowTwoPass(const float* __restrict__ in,
                                    const float* __restrict__ gamma,
                                    const float* __restrict__ beta,
                                    float* __restrict__ out, size_t cols) {
    __shared__ float smem[kWarpsPerBlock];
    const float* row = in + blockIdx.x * cols;
    float* outRow = out + blockIdx.x * cols;

    float partial = 0.0f;
    for (size_t c = threadIdx.x; c < cols; c += blockDim.x) {
        partial += row[c];
    }
    const float mean = blockReduceSum(partial, smem) / static_cast<float>(cols);

    float sq = 0.0f;
    for (size_t c = threadIdx.x; c < cols; c += blockDim.x) {
        const float d = row[c] - mean;
        sq += d * d;
    }
    const float var = blockReduceSum(sq, smem) / static_cast<float>(cols);
    const float scale = rsqrtf(var + kEps);

    for (size_t c = threadIdx.x; c < cols; c += blockDim.x) {
        outRow[c] = (row[c] - mean) * scale * gamma[c] + beta[c];
    }
}

// Merges two Welford accumulators. Combining (n, mean, M2) pairs is what
// makes the one-pass form parallel: a thread folds its own elements in one
// at a time, then the tree folds threads into warps and warps into a block.
// The guard keeps a lane with no elements from dividing by zero, which
// happens whenever cols is under 256.
__device__ void welfordMerge(float* count, float* mean, float* m2, float bCount,
                             float bMean, float bM2) {
    const float total = *count + bCount;
    if (total > 0.0f) {
        const float delta = bMean - *mean;
        *mean += delta * (bCount / total);
        *m2 += bM2 + delta * delta * (*count) * (bCount / total);
    }
    *count = total;
}

// Layer norm in one read pass. Welford carries a running mean and a running
// sum of squared deviations from that running mean, so the variance falls
// out of the same sweep that produced the mean. Two read sweeps instead of
// three, at the cost of a divide per element and a three-value reduction
// instead of a one-value one.
//
// Memory and launch assumption: identical to softmaxRowNaive.
__global__ void layerNormRowWelford(const float* __restrict__ in,
                                    const float* __restrict__ gamma,
                                    const float* __restrict__ beta,
                                    float* __restrict__ out, size_t cols) {
    __shared__ float sCount[kWarpsPerBlock];
    __shared__ float sMean[kWarpsPerBlock];
    __shared__ float sM2[kWarpsPerBlock];
    const float* row = in + blockIdx.x * cols;
    float* outRow = out + blockIdx.x * cols;

    // snippet: welford
    float count = 0.0f;
    float mean = 0.0f;
    float m2 = 0.0f;
    for (size_t c = threadIdx.x; c < cols; c += blockDim.x) {
        const float x = row[c];
        count += 1.0f;
        const float delta = x - mean;
        mean += delta / count;
        m2 += delta * (x - mean);
    }

    for (int offset = kWarpSize / 2; offset > 0; offset /= 2) {
        welfordMerge(&count, &mean, &m2,
                     __shfl_down_sync(kFullMask, count, offset),
                     __shfl_down_sync(kFullMask, mean, offset),
                     __shfl_down_sync(kFullMask, m2, offset));
    }
    // end snippet

    const unsigned int lane = threadIdx.x % kWarpSize;
    const unsigned int warp = threadIdx.x / kWarpSize;
    if (lane == 0) {
        sCount[warp] = count;
        sMean[warp] = mean;
        sM2[warp] = m2;
    }
    __syncthreads();

    count = 0.0f;
    mean = 0.0f;
    m2 = 0.0f;
    for (int w = 0; w < kWarpsPerBlock; ++w) {
        welfordMerge(&count, &mean, &m2, sCount[w], sMean[w], sM2[w]);
    }

    const float scale = rsqrtf(m2 / static_cast<float>(cols) + kEps);
    for (size_t c = threadIdx.x; c < cols; c += blockDim.x) {
        outRow[c] = (row[c] - mean) * scale * gamma[c] + beta[c];
    }
}

// RMS norm: the mean is gone, so there is nothing to subtract and nothing to
// centre. One reduction over the squares, one write pass that folds gamma
// in, and the read of gamma is a broadcast every block repeats. Two read
// sweeps and one write, one reduction, no beta.
//
// The formula is torch.nn.RMSNorm's, eps inside the square root:
// out = x / sqrt(mean(x^2) + eps) * gamma.
//
// Memory and launch assumption: identical to softmaxRowNaive.
// snippet: rmsnorm
__global__ void rmsNormRow(const float* __restrict__ in,
                           const float* __restrict__ gamma,
                           float* __restrict__ out, size_t cols) {
    __shared__ float smem[kWarpsPerBlock];
    const float* row = in + blockIdx.x * cols;
    float* outRow = out + blockIdx.x * cols;

    float sq = 0.0f;
    for (size_t c = threadIdx.x; c < cols; c += blockDim.x) {
        const float x = row[c];
        sq += x * x;
    }
    const float meanSq = blockReduceSum(sq, smem) / static_cast<float>(cols);
    const float scale = rsqrtf(meanSq + kEps);

    for (size_t c = threadIdx.x; c < cols; c += blockDim.x) {
        outRow[c] = row[c] * scale * gamma[c];
    }
}
// end snippet

// The floor, measured in this process rather than borrowed from another
// day: read every element once, write it once, no arithmetic and no
// reduction. No row-wise kernel on this page can beat it, and every GB/s
// figure in part 3 is computed from the same eight bytes per element, so
// the rows are comparable.
//
// One thread, one element, grid-stride so the launch shape is free.
__global__ void copyFloor(const float* __restrict__ in, float* __restrict__ out,
                          size_t n) {
    const size_t step = gridDim.x * static_cast<size_t>(blockDim.x);
    for (size_t i = blockIdx.x * static_cast<size_t>(blockDim.x) + threadIdx.x;
         i < n; i += step) {
        out[i] = in[i];
    }
}

// Host references in double over the original float inputs. Plain loops,
// written for obvious correctness. Each one implements the published
// formula rather than mirroring the kernel's loop structure: the softmax
// reference shifts by the row max because that is what softmax is, the
// layer norm reference takes two sweeps because that is the definition, and
// both kernels being compared to it use a different order on purpose.
static void softmaxCpu(const float* in, double* out, size_t rows, size_t cols) {
    for (size_t r = 0; r < rows; ++r) {
        const float* row = in + r * cols;
        double best = -HUGE_VAL;
        for (size_t c = 0; c < cols; ++c) {
            if (static_cast<double>(row[c]) > best) {
                best = static_cast<double>(row[c]);
            }
        }
        double total = 0.0;
        for (size_t c = 0; c < cols; ++c) {
            total += std::exp(static_cast<double>(row[c]) - best);
        }
        for (size_t c = 0; c < cols; ++c) {
            out[r * cols + c] =
                std::exp(static_cast<double>(row[c]) - best) / total;
        }
    }
}

static void layerNormCpu(const float* in, const float* gamma, const float* beta,
                         double* out, size_t rows, size_t cols) {
    for (size_t r = 0; r < rows; ++r) {
        const float* row = in + r * cols;
        double total = 0.0;
        for (size_t c = 0; c < cols; ++c) {
            total += static_cast<double>(row[c]);
        }
        const double mean = total / static_cast<double>(cols);
        double sq = 0.0;
        for (size_t c = 0; c < cols; ++c) {
            const double d = static_cast<double>(row[c]) - mean;
            sq += d * d;
        }
        const double var = sq / static_cast<double>(cols);
        const double scale = 1.0 / std::sqrt(var + static_cast<double>(kEps));
        for (size_t c = 0; c < cols; ++c) {
            out[r * cols + c] = (static_cast<double>(row[c]) - mean) * scale *
                                    static_cast<double>(gamma[c]) +
                                static_cast<double>(beta[c]);
        }
    }
}

static void rmsNormCpu(const float* in, const float* gamma, double* out,
                       size_t rows, size_t cols) {
    for (size_t r = 0; r < rows; ++r) {
        const float* row = in + r * cols;
        double sq = 0.0;
        for (size_t c = 0; c < cols; ++c) {
            sq += static_cast<double>(row[c]) * static_cast<double>(row[c]);
        }
        const double meanSq = sq / static_cast<double>(cols);
        const double scale =
            1.0 / std::sqrt(meanSq + static_cast<double>(kEps));
        for (size_t c = 0; c < cols; ++c) {
            out[r * cols + c] = static_cast<double>(row[c]) * scale *
                                static_cast<double>(gamma[c]);
        }
    }
}

// The K-scaled tolerance from day 66. Every output on this page depends on a
// reduction over one row, so K is the row length and not the element count:
// rounding errors across a row are uncorrelated, so they grow like sqrt(K)
// rather than K, and the factor 4 is slack for a different-but-valid
// summation order. A fixed 1e-5 would fail correct code at 4096 columns.
// snippet: tolerance
static double kScaledRtol(double tableRtol, int mantissaBits, size_t kDim) {
    const double eps = std::ldexp(1.0, -mantissaBits);
    const double grown = 4.0 * eps * std::sqrt(static_cast<double>(kDim));
    return grown > tableRtol ? grown : tableRtol;
}
// end snippet

// Worst error as a fraction of the gate |got - ref| <= atol + rtol * |ref|.
// A fraction above 1 anywhere is a failure and the caller gets the first
// offending index. A non-finite output fails outright: inf here means the
// overflow part 1 demonstrates, and NaN is never close enough.
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;
}

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;
}

static size_t countNonFinite(const float* x, size_t n) {
    size_t bad = 0;
    for (size_t i = 0; i < n; ++i) {
        if (!std::isfinite(x[i])) {
            ++bad;
        }
    }
    return bad;
}

// Classifies a float the way day 47's flush-to-zero census did: normal,
// subnormal, or exactly zero. A shifted softmax puts its smallest terms in
// the subnormal band, and a build with -ftz=true moves them to zero.
static void classify(const float* x, size_t n, size_t* normal,
                     size_t* subnormal, size_t* zero) {
    *normal = 0;
    *subnormal = 0;
    *zero = 0;
    for (size_t i = 0; i < n; ++i) {
        if (x[i] == 0.0f) {
            ++(*zero);
        } else if (std::fabs(x[i]) < 1.17549435e-38f) {
            ++(*subnormal);
        } else {
            ++(*normal);
        }
    }
}

// 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;
}

// Inputs that force rounding and give each row its own maximum, so a kernel
// that computes one row's statistics and applies them to another fails the
// gate. The period 211 divides neither case's row length, and the row
// offset shifts the maximum by up to 0.75, which is enough that a softmax
// reading the wrong row's max lands outside the tolerance.
static void makeInputs(std::vector<float>* h_in, size_t rows, size_t cols) {
    for (size_t r = 0; r < rows; ++r) {
        const float offset = (static_cast<float>(r % 7) - 3.0f) * 0.25f;
        for (size_t c = 0; c < cols; ++c) {
            const size_t e = r * cols + c;
            (*h_in)[e] =
                (static_cast<float>(static_cast<int>(e % 211) - 105)) * 0.017f +
                offset;
        }
    }
}

static void makeAffine(std::vector<float>* h_gamma, std::vector<float>* h_beta,
                       size_t cols) {
    for (size_t c = 0; c < cols; ++c) {
        (*h_gamma)[c] =
            1.0f + (static_cast<float>(static_cast<int>(c % 5) - 2)) * 0.05f;
        (*h_beta)[c] =
            (static_cast<float>(static_cast<int>(c % 3) - 1)) * 0.02f;
    }
}

int main() {
    const int device = 0;
    CUDA_CHECK(cudaSetDevice(device));
    cudaDeviceProp prop;
    CUDA_CHECK(cudaGetDeviceProperties(&prop, device));
    std::printf("GPU: %s (compute capability %d.%d), %d SMs, %zu KiB L2\n\n",
                prop.name, prop.major, prop.minor, prop.multiProcessorCount,
                static_cast<size_t>(prop.l2CacheSize) / 1024);

    std::vector<float> h_in(kMaxElems);
    std::vector<float> h_out(kMaxElems);
    std::vector<double> h_want(kMaxElems);
    std::vector<float> h_gamma(kMaxCols);
    std::vector<float> h_beta(kMaxCols);

    float* d_in = nullptr;
    float* d_out = nullptr;
    float* d_gamma = nullptr;
    float* d_beta = nullptr;
    CUDA_CHECK(cudaMalloc(&d_in, kMaxElems * sizeof(float)));
    CUDA_CHECK(cudaMalloc(&d_out, kMaxElems * sizeof(float)));
    CUDA_CHECK(cudaMalloc(&d_gamma, kMaxCols * sizeof(float)));
    CUDA_CHECK(cudaMalloc(&d_beta, kMaxCols * 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.
    const char* badKernel = nullptr;
    size_t badIndex = 0;
    size_t badRows = 0;
    size_t badCols = 0;
    int overflowGateFailed = 0;

    // Part 0: eight values a learner can paste into PyTorch. gamma is all
    // ones here so the comparison is against torch.nn.RMSNorm's defaults.
    for (size_t c = 0; c < kTorchCols; ++c) {
        h_in[c] = static_cast<float>(static_cast<int>(c) - 3) * 0.5f + 0.25f;
        h_gamma[c] = 1.0f;
    }
    CUDA_CHECK(cudaMemcpy(d_in, h_in.data(), kTorchCols * sizeof(float),
                          cudaMemcpyHostToDevice));
    CUDA_CHECK(cudaMemcpy(d_gamma, h_gamma.data(), kTorchCols * sizeof(float),
                          cudaMemcpyHostToDevice));
    rmsNormRow<<<1, kThreadsPerBlock>>>(d_in, d_gamma, d_out, kTorchCols);
    CUDA_CHECK(cudaGetLastError());
    CUDA_CHECK(cudaDeviceSynchronize());
    CUDA_CHECK(cudaMemcpy(h_out.data(), d_out, kTorchCols * sizeof(float),
                          cudaMemcpyDeviceToHost));
    std::printf("Part 0: one row for the PyTorch comparison, eps = %g\n",
                static_cast<double>(kEps));
    std::printf("  input      ");
    for (size_t c = 0; c < kTorchCols; ++c) {
        std::printf("%9.4f", static_cast<double>(h_in[c]));
    }
    std::printf("\n  rmsNormRow ");
    for (size_t c = 0; c < kTorchCols; ++c) {
        std::printf("%9.6f", static_cast<double>(h_out[c]));
    }
    std::printf("\n\n");

    // Part 1: the overflow. One row of 1024 logits spread from the maximum
    // down 200, which is what an attention score row looks like after a
    // model has learned anything at all.
    for (size_t c = 0; c < kDemoCols; ++c) {
        h_in[c] = 120.0f - static_cast<float>(c % 201);
    }
    CUDA_CHECK(cudaMemcpy(d_in, h_in.data(), kDemoCols * sizeof(float),
                          cudaMemcpyHostToDevice));

    std::printf("Part 1: softmax over 1024 logits from -80 to 120\n");
    softmaxRowNaive<<<1, kThreadsPerBlock>>>(d_in, d_out, kDemoCols);
    CUDA_CHECK(cudaGetLastError());
    CUDA_CHECK(cudaDeviceSynchronize());
    CUDA_CHECK(cudaMemcpy(h_out.data(), d_out, kDemoCols * sizeof(float),
                          cudaMemcpyDeviceToHost));
    const size_t naiveBad = countNonFinite(h_out.data(), kDemoCols);
    std::printf("  softmaxRowNaive non-finite outputs   %zu of %zu\n", naiveBad,
                kDemoCols);

    softmaxRowSafe<<<1, kThreadsPerBlock>>>(d_in, d_out, kDemoCols);
    CUDA_CHECK(cudaGetLastError());
    CUDA_CHECK(cudaDeviceSynchronize());
    CUDA_CHECK(cudaMemcpy(h_out.data(), d_out, kDemoCols * sizeof(float),
                          cudaMemcpyDeviceToHost));
    const size_t safeBad = countNonFinite(h_out.data(), kDemoCols);
    size_t normal = 0;
    size_t subnormal = 0;
    size_t zero = 0;
    classify(h_out.data(), kDemoCols, &normal, &subnormal, &zero);
    softmaxCpu(h_in.data(), h_want.data(), 1, kDemoCols);
    double safeSum = 0.0;
    for (size_t c = 0; c < kDemoCols; ++c) {
        safeSum += static_cast<double>(h_out[c]);
    }
    std::printf("  softmaxRowSafe  non-finite outputs   %zu of %zu\n", safeBad,
                kDemoCols);
    std::printf("  softmaxRowSafe  probabilities sum to %.9f\n", safeSum);
    std::printf(
        "  softmaxRowSafe  census: %zu normal, %zu subnormal, "
        "%zu exactly zero\n\n",
        normal, subnormal, zero);

    // The overflow is this day's first claim, so it is a branch. If the
    // naive kernel survives a row it has no business surviving, the page's
    // explanation of expf is wrong and the run must say so.
    if (naiveBad == 0 || safeBad != 0) {
        overflowGateFailed = 1;
        std::fprintf(stderr,
                     "part 1 gate failed: naive produced %zu non-finite "
                     "outputs (expected more than 0) and safe produced %zu "
                     "(expected 0)\n",
                     naiveBad, safeBad);
    }

    struct Path {
        const char* name;
        int kind;  // 0 softmax, 1 layer norm, 2 rms norm
    };
    const Path paths[] = {
        {"softmaxRowNaive", 0},     {"softmaxRowSafe", 0},
        {"layerNormRowTwoPass", 1}, {"layerNormRowWelford", 1},
        {"rmsNormRow", 2},
    };

    for (int caseIdx = 0;
         caseIdx < kNumCases && badKernel == nullptr && overflowGateFailed == 0;
         ++caseIdx) {
        const size_t rows = kCaseRows[caseIdx];
        const size_t cols = kCaseCols[caseIdx];
        const size_t elems = rows * cols;
        const size_t bytes = elems * sizeof(float);
        const dim3 grid(static_cast<unsigned int>(rows));

        makeInputs(&h_in, rows, cols);
        makeAffine(&h_gamma, &h_beta, cols);
        CUDA_CHECK(
            cudaMemcpy(d_in, h_in.data(), bytes, cudaMemcpyHostToDevice));
        CUDA_CHECK(cudaMemcpy(d_gamma, h_gamma.data(), cols * sizeof(float),
                              cudaMemcpyHostToDevice));
        CUDA_CHECK(cudaMemcpy(d_beta, h_beta.data(), cols * sizeof(float),
                              cudaMemcpyHostToDevice));

        const double rtol = kScaledRtol(kF32Rtol, kF32MantissaBits, cols);
        std::printf("case %zu rows x %zu cols, %zu MiB in and %zu MiB out\n",
                    rows, cols, bytes / (1024 * 1024), bytes / (1024 * 1024));
        std::printf(
            "  gate: |got-ref| <= rtol*max|ref| + rtol*|ref|, "
            "rtol = max(1e-5, 4*2^-23*sqrt(%zu)) = %.3e\n",
            cols, rtol);
        std::printf("  %-22s %13s %10s\n", "kernel", "max abs err",
                    "gate frac");

        for (int pathIdx = 0; pathIdx < kNumPaths && badKernel == nullptr;
             ++pathIdx) {
            // The output is cleared before each correctness launch so a
            // kernel that skips elements cannot inherit a right answer.
            CUDA_CHECK(cudaMemset(d_out, 0, bytes));
            switch (paths[pathIdx].kind) {
                case 0:
                    if (pathIdx == 0) {
                        softmaxRowNaive<<<grid, kThreadsPerBlock>>>(d_in, d_out,
                                                                    cols);
                    } else {
                        softmaxRowSafe<<<grid, kThreadsPerBlock>>>(d_in, d_out,
                                                                   cols);
                    }
                    softmaxCpu(h_in.data(), h_want.data(), rows, cols);
                    break;
                case 1:
                    if (pathIdx == 2) {
                        layerNormRowTwoPass<<<grid, kThreadsPerBlock>>>(
                            d_in, d_gamma, d_beta, d_out, cols);
                    } else {
                        layerNormRowWelford<<<grid, kThreadsPerBlock>>>(
                            d_in, d_gamma, d_beta, d_out, cols);
                    }
                    layerNormCpu(h_in.data(), h_gamma.data(), h_beta.data(),
                                 h_want.data(), rows, cols);
                    break;
                default:
                    rmsNormRow<<<grid, kThreadsPerBlock>>>(d_in, d_gamma, d_out,
                                                           cols);
                    rmsNormCpu(h_in.data(), h_gamma.data(), h_want.data(), rows,
                               cols);
                    break;
            }
            CUDA_CHECK(cudaGetLastError());
            CUDA_CHECK(cudaDeviceSynchronize());
            CUDA_CHECK(
                cudaMemcpy(h_out.data(), d_out, bytes, cudaMemcpyDeviceToHost));

            const double atol = rtol * maxAbs(h_want.data(), elems);
            size_t firstBad = elems;
            const double frac = worstGateFraction(h_out.data(), h_want.data(),
                                                  elems, rtol, atol, &firstBad);
            const double absErr =
                maxAbsError(h_out.data(), h_want.data(), elems);
            std::printf("  %-22s %13.3e %10.4f\n", paths[pathIdx].name, absErr,
                        frac);
            if (firstBad != elems) {
                badKernel = paths[pathIdx].name;
                badIndex = firstBad;
                badRows = rows;
                badCols = cols;
            }
        }
        if (badKernel != nullptr) {
            break;
        }

        // Timing. Every row is priced at the eight bytes the algorithm has
        // to move, one read and one write per element, so a kernel whose
        // extra read sweeps miss cache reports a GB/s under the floor and
        // one whose extra passes hit cache does not. That is a statement
        // about DRAM traffic inferred from a timer, not a counter reading:
        // day 49 is where the counters settle it.
        const int floorBlocks = static_cast<int>(
            (elems + kThreadsPerBlock - 1) / kThreadsPerBlock / 8);
        const float floorMs = timeKernel([&] {
            copyFloor<<<floorBlocks, kThreadsPerBlock>>>(d_in, d_out, elems);
        });
        const float naiveMs = timeKernel([&] {
            softmaxRowNaive<<<grid, kThreadsPerBlock>>>(d_in, d_out, cols);
        });
        const float safeMs = timeKernel([&] {
            softmaxRowSafe<<<grid, kThreadsPerBlock>>>(d_in, d_out, cols);
        });
        const float twoPassMs = timeKernel([&] {
            layerNormRowTwoPass<<<grid, kThreadsPerBlock>>>(
                d_in, d_gamma, d_beta, d_out, cols);
        });
        const float welfordMs = timeKernel([&] {
            layerNormRowWelford<<<grid, kThreadsPerBlock>>>(
                d_in, d_gamma, d_beta, d_out, cols);
        });
        const float rmsMs = timeKernel([&] {
            rmsNormRow<<<grid, kThreadsPerBlock>>>(d_in, d_gamma, d_out, cols);
        });

        const double askBytes = static_cast<double>(elems) * 8.0;
        const float times[6] = {floorMs,   naiveMs,   safeMs,
                                twoPassMs, welfordMs, rmsMs};
        const char* names[6] = {"copyFloor",           "softmaxRowNaive",
                                "softmaxRowSafe",      "layerNormRowTwoPass",
                                "layerNormRowWelford", "rmsNormRow"};
        // Row touches: every read sweep plus the write sweep. Only one read
        // and one write are compulsory; the rest are re-reads that the
        // caches may or may not absorb, which is what the last column
        // reports.
        const int touches[6] = {2, 3, 4, 4, 3, 3};
        std::printf("  %-22s %8s %9s %10s %8s\n", "kernel", "touches", "ms",
                    "ask GB/s", "% floor");
        for (int k = 0; k < 6; ++k) {
            const double gbs = askBytes / (times[k] * 1.0e6);
            const double floorGbs = askBytes / (floorMs * 1.0e6);
            std::printf("  %-22s %8d %9.4f %10.1f %8.1f\n", names[k],
                        touches[k], times[k], gbs, 100.0 * gbs / floorGbs);
        }

        // This is a hardware-dependent performance prediction, not a
        // correctness condition. Keep the five-percent band so noisy ties
        // are not called refutations, but never turn the process red for it.
        const bool timingPredictionHeld =
            rmsMs <= twoPassMs * kTimingSlack;
        std::printf(
            "  RMS-vs-two-pass prediction: %s (%.4f ms vs %.4f ms; "
            "%.0f%% slack)\n",
            timingPredictionHeld ? "held" : "refuted",
            static_cast<double>(rmsMs), static_cast<double>(twoPassMs),
            100.0 * (kTimingSlack - 1.0));
        std::printf("\n");
    }

    CUDA_CHECK(cudaFree(d_in));
    CUDA_CHECK(cudaFree(d_out));
    CUDA_CHECK(cudaFree(d_gamma));
    CUDA_CHECK(cudaFree(d_beta));

    if (badKernel != nullptr) {
        std::fprintf(stderr,
                     "%s failed its gate at %zu (row %zu, col %zu) on case "
                     "%zu x %zu\n",
                     badKernel, badIndex, badIndex / badCols,
                     badIndex % badCols, badRows, badCols);
        return EXIT_FAILURE;
    }
    if (overflowGateFailed != 0) {
        return EXIT_FAILURE;
    }

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