COURSE / SOURCE

elementwise.cu

All lessons
Source filecode/day96-elementwise/elementwise.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 96: RoPE indexing, the two GELUs, and bias + GELU + residual fused.
//
// Three parts, one program, all of it elementwise and all of it bound by
// memory rather than by arithmetic.
//
//   Part 1  RoPE. The same rotation written three ways, differing only in
//           which two dimensions of a head are called a pair and how the
//           two values are loaded. Adjacent pairs (2j, 2j+1) make every
//           lane read two floats one apart; half-split pairs (j, j+D/2)
//           make the same warp read two contiguous runs. A float2 load
//           fixes the first without changing its arithmetic.
//   Part 2  GELU. The exact form, 0.5x(1+erf(x/sqrt2)), against the tanh
//           approximation, over a sweep of x, with the largest gap and
//           where it happens reported rather than asserted.
//   Part 3  bias -> GELU -> residual, as three kernels and as one. The
//           staged chain moves 7 floats per element, the fused chain 3.
//           The launch ratio is 3 and the byte ratio is 7 to 3, and those
//           are different numbers, so the measurement can say which one
//           the card was charging for.
//
// The byte counts are computed from the constants and printed above the
// times, so the model this program is judged against cannot drift away from
// the program.
//
// Build: nvcc -std=c++17 -O3 -arch=sm_75 -o elementwise elementwise.cu
// Run:   ./elementwise
//
// VERIFIED: Tesla T4, driver 580.173.02, CUDA 12.6, 2026-09-02. The README's
// Nsight Systems report and exported kernel-summary CSV are present.

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

// The timed shapes. Part 1 rotates a [batch, seq, heads, headDim] tensor
// and part 3 runs a [rows, cols] activation through the chain; both come to
// 16,777,216 floats, 64 MiB, so the same five device buffers serve both.
constexpr int kBatch = 8;
constexpr int kSeq = 2048;
constexpr int kHeads = 8;
constexpr int kHeadDim = 128;
constexpr int kRows = 16384;
constexpr int kCols = 1024;
constexpr size_t kElems =
    static_cast<size_t>(kBatch) * kSeq * kHeads * static_cast<size_t>(kHeadDim);

// The correctness shapes, deliberately ragged so that every bounds check
// runs on every launch. 8,160 pairs is not a multiple of 256, 9,797
// elements is not either, and 97 columns leaves 159 idle lanes in the one
// block that covers a row. The timed shapes above are powers of two, so no
// guard fires inside a timed kernel and no timing includes tail handling.
constexpr int kChkBatch = 3;
constexpr int kChkSeq = 17;
constexpr int kChkHeads = 5;
constexpr int kChkDim = 64;
constexpr int kChkRows = 101;
constexpr int kChkCols = 97;
constexpr size_t kChkRopeElems = static_cast<size_t>(kChkBatch) * kChkSeq *
                                 kChkHeads * static_cast<size_t>(kChkDim);
constexpr size_t kChkFuseElems =
    static_cast<size_t>(kChkRows) * static_cast<size_t>(kChkCols);

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

// RoPE's frequency base. theta_i = base^(-2(i-1)/d) for i in [1, d/2], the
// rotary matrix of RoFormer equation 15
// (https://arxiv.org/abs/2104.09864 , checked 2026-09-01).
constexpr double kRopeBase = 10000.0;

// Part 2's sweep. 1 Mi samples spread over [-8, 8] puts about 65,000 of
// them inside the interval where the two GELU forms disagree most, and the
// whole buffer is 4 MiB, so the host can find the maximum itself.
constexpr size_t kGeluSamples = 1048576;
constexpr float kGeluLo = -8.0f;
constexpr float kGeluHi = 8.0f;

// GELU's two forms. kInvSqrt2 is 1/sqrt(2) and kGeluA is sqrt(2/pi), both
// rounded to float; the CPU reference promotes these same rounded floats to
// double rather than using the exact constants, so it checks the kernel's
// arithmetic and not the author's choice of literal. kGeluB is the cubic
// coefficient from the tanh form as PyTorch writes it
// (https://docs.pytorch.org/docs/stable/generated/torch.nn.GELU.html ,
// checked 2026-09-01).
constexpr float kInvSqrt2 = 0.70710678f;
constexpr float kGeluA = 0.7978845608f;
constexpr float kGeluB = 0.044715f;

// The chain's floats per element, which is the whole prediction for part 3.
// Staged: bias reads x and writes t1, GELU reads t1 and writes t2, the
// residual add reads t2 and r and writes out. That is 2 + 2 + 3 = 7. Fused:
// x, r, out. The bias vector itself is kCols floats, 4 KiB against 64 MiB
// per buffer, and it is read by every row, so it is left out of both counts
// and named in the README instead of being smuggled into a ratio.
constexpr double kStagedFloatsPerElem = 7.0;
constexpr double kFusedFloatsPerElem = 3.0;
constexpr double kCopyFloatsPerElem = 2.0;

// RoPE moves two floats in and two out per pair, which is two floats per
// element either way, the same traffic as a copy. The cos and sin tables
// are kSeq * kHeadDim / 2 floats each, 1 MiB for the pair, read by every
// batch and every head; they fit this card's 4 MiB L2 and are excluded for
// the same reason the bias vector is.
constexpr double kRopeFloatsPerElem = 2.0;

// The chain is three dependent elementwise operations with no reduction in
// it, so the tolerance table's f32 row applies unscaled: rtol 1e-5 is about
// 84 eps, which covers three roundings plus one erff, whose documented
// maximum error is 2 ULP (CUDA Programming Guide, mathematical functions
// appendix, checked 2026-09-01). Nothing here accumulates over K terms, so
// there is no sqrt(K) factor to add.
constexpr double kRelTolerance = 1e-5;
constexpr double kAbsTolerance = 1e-6;

// Part 2's one gate. If the exact and the tanh GELU agreed to within this,
// they would be the same function to a float and this page's argument would
// be wrong. 1e-5 is two orders above the float spacing near 1.0 and two
// orders below the gap the closed forms predict, so it separates "these are
// different functions" from "these are the same function rounded".
constexpr double kGeluFormGap = 1e-5;

static_assert(kThreadsPerBlock % 32 == 0,
              "the block is a whole number of warps");
static_assert(kThreadsPerBlock <= 1024,
              "1024 threads is the per-block maximum on every compute "
              "capability this course targets");
static_assert(kHeadDim % 2 == 0 && kChkDim % 2 == 0,
              "RoPE rotates pairs of dimensions, so the head dimension is "
              "even in both shapes");
static_assert(kElems == static_cast<size_t>(kRows) * kCols,
              "part 1 and part 3 share five device buffers, so their timed "
              "element counts must be equal");
static_assert(kElems % (2 * kThreadsPerBlock) == 0,
              "the timed pair count is a whole number of blocks, so no "
              "bounds check runs inside a timed RoPE kernel");
static_assert(kCols % kThreadsPerBlock == 0,
              "a timed row is a whole number of blocks in x, so no guard "
              "fires inside a timed chain kernel");
static_assert((kChkRopeElems / 2) % kThreadsPerBlock != 0,
              "the RoPE check shape must be ragged or its guard never runs");
static_assert(kChkFuseElems % kThreadsPerBlock != 0,
              "the chain check shape must be ragged for the same reason");
static_assert(kChkCols < kThreadsPerBlock,
              "the check shape leaves idle lanes in the row block, which is "
              "the case the 2D guard exists for");
static_assert(kRows <= 65535,
              "rows become gridDim.y, capped at 65535 on every compute "
              "capability this course targets");
static_assert(kGeluSamples <= kElems, "the sweep reuses the big buffers");

// ---------------------------------------------------------------------
// Device-side elementwise pieces
// ---------------------------------------------------------------------

// snippet: gelu-forms
// The exact GELU: x times the standard normal CDF, which is what PyTorch
// computes for approximate='none'.
__device__ __forceinline__ float geluErf(float x) {
    return 0.5f * x * (1.0f + erff(x * kInvSqrt2));
}

// The tanh GELU, the approximation from the original paper and what GPT-2
// shipped. Same shape, different function: the cubic inside the tanh is not
// a rounding of the erf form.
__device__ __forceinline__ float geluTanh(float x) {
    const float inner = kGeluA * (x + kGeluB * x * x * x);
    return 0.5f * x * (1.0f + tanhf(inner));
}
// end snippet

// ---------------------------------------------------------------------
// Part 1: three RoPE kernels
// ---------------------------------------------------------------------

// One thread rotates one pair of dimensions, in the adjacent-pair
// convention of the RoFormer paper: (0,1), (2,3), and so on.
//
// Memory: lane L of a warp owns pair p = base + L, so its two loads are at
// element 2p and 2p + 1. Each load on its own is day 11's stride-2 pattern,
// 32 lanes spread over 256 bytes, so it asks for eight 32-byte sectors and
// uses half of each. The second load asks for the same eight sectors and
// uses the other half, so the DRAM bytes are the same as a coalesced read
// and the cost is the extra request, not extra traffic.
//
// Launch assumption: gridDim.x * blockDim.x >= pairs, halfDim = headDim / 2,
// and the cos and sin tables hold seq * halfDim floats each.
// snippet: rope-interleaved
__global__ void ropeInterleaved(const float* __restrict__ in,
                                const float* __restrict__ cosTab,
                                const float* __restrict__ sinTab,
                                float* __restrict__ out, size_t pairs,
                                int halfDim, int heads, int seq) {
    const size_t p = blockIdx.x * static_cast<size_t>(blockDim.x) + threadIdx.x;
    if (p < pairs) {
        const size_t half = static_cast<size_t>(halfDim);
        const size_t j = p % half;
        const size_t pos = (p / (half * static_cast<size_t>(heads))) %
                           static_cast<size_t>(seq);
        const float c = cosTab[pos * half + j];
        const float s = sinTab[pos * half + j];
        const float x0 = in[2 * p];
        const float x1 = in[2 * p + 1];
        out[2 * p] = x0 * c - x1 * s;
        out[2 * p + 1] = x0 * s + x1 * c;
    }
}
// end snippet

// The same arithmetic in the same order, with both halves of the pair
// fetched as one 8-byte float2 and stored the same way.
//
// Memory: one load instruction per lane, 32 lanes covering the same 256
// contiguous bytes in one request instead of two. The bytes do not change;
// the number of requests halves.
// cudaMalloc returns 256-byte aligned memory and a pair sits at an 8-byte
// offset, so the reinterpret_cast is aligned by construction.
//
// Launch assumption: identical to ropeInterleaved. The results of the two
// kernels are compared for bit equality, not for tolerance.
__global__ void ropeInterleavedVec2(const float* __restrict__ in,
                                    const float* __restrict__ cosTab,
                                    const float* __restrict__ sinTab,
                                    float* __restrict__ out, size_t pairs,
                                    int halfDim, int heads, int seq) {
    const size_t p = blockIdx.x * static_cast<size_t>(blockDim.x) + threadIdx.x;
    if (p < pairs) {
        const size_t half = static_cast<size_t>(halfDim);
        const size_t j = p % half;
        const size_t pos = (p / (half * static_cast<size_t>(heads))) %
                           static_cast<size_t>(seq);
        const float c = cosTab[pos * half + j];
        const float s = sinTab[pos * half + j];
        const float2 v = reinterpret_cast<const float2*>(in)[p];
        float2 r;
        r.x = v.x * c - v.y * s;
        r.y = v.x * s + v.y * c;
        reinterpret_cast<float2*>(out)[p] = r;
    }
}

// The half-split convention: dimension j pairs with j + headDim/2, which is
// what a Llama-style implementation's rotate_half does. Same rotation, a
// different assignment of dimensions to angles, so the output is a
// different tensor unless the projection weights were permuted to match.
//
// Memory: lane L owns pair p = base + L within one head, so the 32 lanes
// read 32 consecutive floats from the low half and 32 consecutive floats
// from the high half. Two coalesced runs, four sectors each, nothing
// wasted, and no float2 needed to get there.
//
// Launch assumption: identical to ropeInterleaved, plus heads * headDim
// dividing the tensor, which the shape constants guarantee.
// snippet: rope-half-split
__global__ void ropeHalfSplit(const float* __restrict__ in,
                              const float* __restrict__ cosTab,
                              const float* __restrict__ sinTab,
                              float* __restrict__ out, size_t pairs,
                              int halfDim, int heads, int seq) {
    const size_t p = blockIdx.x * static_cast<size_t>(blockDim.x) + threadIdx.x;
    if (p < pairs) {
        const size_t half = static_cast<size_t>(halfDim);
        const size_t j = p % half;
        const size_t headBase = (p / half) * 2 * half;
        const size_t pos = (p / (half * static_cast<size_t>(heads))) %
                           static_cast<size_t>(seq);
        const float c = cosTab[pos * half + j];
        const float s = sinTab[pos * half + j];
        const float x0 = in[headBase + j];
        const float x1 = in[headBase + j + half];
        out[headBase + j] = x0 * c - x1 * s;
        out[headBase + j + half] = x0 * s + x1 * c;
    }
}
// end snippet

// ---------------------------------------------------------------------
// Part 2: both GELUs, once, over the same sweep
// ---------------------------------------------------------------------

// Writes the exact GELU and the absolute gap to the tanh form. One thread,
// one sample.
//
// Memory: consecutive threads take consecutive samples, so one warp reads
// 128 contiguous bytes and writes two runs of the same width.
//
// Launch assumption: gridDim.x * blockDim.x >= n.
__global__ void geluBothForms(const float* __restrict__ x,
                              float* __restrict__ exact,
                              float* __restrict__ gap, size_t n) {
    const size_t i = blockIdx.x * static_cast<size_t>(blockDim.x) + threadIdx.x;
    if (i < n) {
        const float v = x[i];
        const float e = geluErf(v);
        exact[i] = e;
        gap[i] = fabsf(e - geluTanh(v));
    }
}

// ---------------------------------------------------------------------
// Part 3: the chain, staged and fused
// ---------------------------------------------------------------------

// out[row][col] = x[row][col] + bias[col]. One thread owns one element and
// blockIdx.y is the row, so the column index needs no division and the bias
// index is the column index.
//
// Memory: a warp covers 32 consecutive columns of one row, 128 contiguous
// bytes, and all 32 lanes read the same 128 bytes of bias that the previous
// warp read, which is an L2 hit after the first row.
//
// Launch assumption: gridDim.y = rows, gridDim.x * blockDim.x >= cols.
__global__ void addBiasRow(const float* __restrict__ x,
                           const float* __restrict__ bias,
                           float* __restrict__ out, size_t cols) {
    const size_t col =
        blockIdx.x * static_cast<size_t>(blockDim.x) + threadIdx.x;
    if (col < cols) {
        const size_t i = blockIdx.y * cols + col;
        out[i] = x[i] + bias[col];
    }
}

// out[i] = gelu(in[i]), exact form. One thread owns one element.
//
// Memory: one warp reads 128 contiguous bytes and writes 128 contiguous
// bytes. Launch assumption: gridDim.x * blockDim.x >= n.
__global__ void applyGeluErf(const float* __restrict__ in,
                             float* __restrict__ out, size_t n) {
    const size_t i = blockIdx.x * static_cast<size_t>(blockDim.x) + threadIdx.x;
    if (i < n) {
        out[i] = geluErf(in[i]);
    }
}

// out[i] = a[i] + r[i]. One thread owns one element, and this is the stage
// that reads two buffers instead of one.
//
// Memory: two coalesced 128-byte reads per warp and one write. Launch
// assumption: gridDim.x * blockDim.x >= n.
__global__ void addResidual(const float* __restrict__ a,
                            const float* __restrict__ r,
                            float* __restrict__ out, size_t n) {
    const size_t i = blockIdx.x * static_cast<size_t>(blockDim.x) + threadIdx.x;
    if (i < n) {
        out[i] = a[i] + r[i];
    }
}

// The same three operations in one kernel. The two intermediates never
// leave a register, so the chain reads x and r and writes out, and nothing
// else crosses the bus.
//
// Memory: identical to addBiasRow's for x and out, plus one more coalesced
// 128-byte read per warp for the residual.
//
// Launch assumption: gridDim.y = rows, gridDim.x * blockDim.x >= cols.
// snippet: chain-fused
__global__ void biasGeluResidualFused(const float* __restrict__ x,
                                      const float* __restrict__ bias,
                                      const float* __restrict__ r,
                                      float* __restrict__ out, size_t cols) {
    const size_t col =
        blockIdx.x * static_cast<size_t>(blockDim.x) + threadIdx.x;
    if (col < cols) {
        const size_t i = blockIdx.y * cols + col;
        out[i] = geluErf(x[i] + bias[col]) + r[i];
    }
}
// end snippet

// Reads the buffer and writes it back, with no arithmetic. No elementwise
// kernel over this data can beat it, and measuring it in this process keeps
// this program's block shape, buffer size and clock state out of the
// comparison.
//
// Memory: one coalesced 128-byte read and one 128-byte write per warp.
// Launch assumption: gridDim.x * blockDim.x >= n.
__global__ void copyFloor(const float* __restrict__ in, float* __restrict__ out,
                          size_t n) {
    const size_t i = blockIdx.x * static_cast<size_t>(blockDim.x) + threadIdx.x;
    if (i < n) {
        out[i] = in[i];
    }
}

// ---------------------------------------------------------------------
// Host references and helpers
// ---------------------------------------------------------------------

// theta_i for pair index j of a head of width headDim, RoFormer equation 15
// with i = j + 1.
static double ropeTheta(int j, int headDim) {
    return std::pow(kRopeBase, -2.0 * static_cast<double>(j) /
                                   static_cast<double>(headDim));
}

// Builds the cos and sin tables the kernels read, one row per position and
// one column per pair. The tables are float because the kernels read
// floats; the references below promote these same rounded values to double,
// so a comparison prices the kernel's arithmetic and not the table's.
static void buildTables(std::vector<float>& cosTab, std::vector<float>& sinTab,
                        int seq, int headDim) {
    const int half = headDim / 2;
    for (int pos = 0; pos < seq; ++pos) {
        for (int j = 0; j < half; ++j) {
            const double angle =
                static_cast<double>(pos) * ropeTheta(j, headDim);
            cosTab[static_cast<size_t>(pos) * half + j] =
                static_cast<float>(std::cos(angle));
            sinTab[static_cast<size_t>(pos) * half + j] =
                static_cast<float>(std::sin(angle));
        }
    }
}

// The adjacent-pair reference, in double. Written for obvious correctness:
// four plain loops and no attempt to share work with the kernel's index
// arithmetic, so an indexing bug in the kernel cannot hide in the reference.
static void ropeInterleavedCpu(const float* in, const float* cosTab,
                               const float* sinTab, double* out, int batch,
                               int seq, int heads, int headDim) {
    const int half = headDim / 2;
    for (int b = 0; b < batch; ++b) {
        for (int pos = 0; pos < seq; ++pos) {
            for (int h = 0; h < heads; ++h) {
                const size_t base =
                    ((static_cast<size_t>(b) * seq + pos) * heads + h) *
                    headDim;
                for (int j = 0; j < half; ++j) {
                    const double c =
                        cosTab[static_cast<size_t>(pos) * half + j];
                    const double s =
                        sinTab[static_cast<size_t>(pos) * half + j];
                    const double x0 = in[base + 2 * j];
                    const double x1 = in[base + 2 * j + 1];
                    out[base + 2 * j] = x0 * c - x1 * s;
                    out[base + 2 * j + 1] = x0 * s + x1 * c;
                }
            }
        }
    }
}

// The half-split reference. Same rotation, dimension j paired with
// j + headDim/2.
static void ropeHalfSplitCpu(const float* in, const float* cosTab,
                             const float* sinTab, double* out, int batch,
                             int seq, int heads, int headDim) {
    const int half = headDim / 2;
    for (int b = 0; b < batch; ++b) {
        for (int pos = 0; pos < seq; ++pos) {
            for (int h = 0; h < heads; ++h) {
                const size_t base =
                    ((static_cast<size_t>(b) * seq + pos) * heads + h) *
                    headDim;
                for (int j = 0; j < half; ++j) {
                    const double c =
                        cosTab[static_cast<size_t>(pos) * half + j];
                    const double s =
                        sinTab[static_cast<size_t>(pos) * half + j];
                    const double x0 = in[base + j];
                    const double x1 = in[base + j + half];
                    out[base + j] = x0 * c - x1 * s;
                    out[base + j + half] = x0 * s + x1 * c;
                }
            }
        }
    }
}

// The exact GELU in double, from the same rounded 1/sqrt(2) the kernel uses.
static double geluErfCpu(double x) {
    return 0.5 * x * (1.0 + std::erf(x * static_cast<double>(kInvSqrt2)));
}

// The chain in double: bias, exact GELU, residual. The reference computes
// erf from the standard library rather than repeating the kernel's
// intrinsic, so this is a check and not a restatement.
static void chainCpu(const float* x, const float* bias, const float* r,
                     double* out, size_t rows, size_t cols) {
    for (size_t row = 0; row < rows; ++row) {
        for (size_t col = 0; col < cols; ++col) {
            const size_t i = row * cols + col;
            const double v = static_cast<double>(x[i]) + bias[col];
            out[i] = geluErfCpu(v) + static_cast<double>(r[i]);
        }
    }
}

// Returns the first index where got and want differ by more than the
// tolerance, or n if they agree everywhere.
static size_t firstMismatch(const float* got, const double* want, size_t n,
                            double rtol, double atol) {
    for (size_t i = 0; i < n; ++i) {
        const double diff = std::fabs(static_cast<double>(got[i]) - want[i]);
        if (diff > atol + rtol * std::fabs(want[i])) {
            return i;
        }
    }
    return n;
}

// Returns the first index where two float buffers are not bit equal, or n.
static size_t firstDifference(const float* a, const float* b, size_t n) {
    for (size_t i = 0; i < n; ++i) {
        if (a[i] != b[i]) {
            return i;
        }
    }
    return n;
}

// Times a launch with CUDA events and returns the mean milliseconds per run.
// Warms up inside itself, because lazy module loading makes the first launch
// of each kernel pay its own load.
template <typename LaunchFn>
static float timeKernel(LaunchFn launch) {
    cudaEvent_t start, stop;
    CUDA_CHECK(cudaEventCreate(&start));
    CUDA_CHECK(cudaEventCreate(&stop));

    for (int i = 0; i < kWarmupRuns; ++i) {
        launch();
    }
    CUDA_CHECK(cudaDeviceSynchronize());
    CUDA_CHECK(cudaGetLastError());

    CUDA_CHECK(cudaEventRecord(start));
    for (int i = 0; i < kTimedRuns; ++i) {
        launch();
    }
    CUDA_CHECK(cudaEventRecord(stop));
    CUDA_CHECK(cudaEventSynchronize(stop));
    CUDA_CHECK(cudaGetLastError());

    float ms = 0.0f;
    CUDA_CHECK(cudaEventElapsedTime(&ms, start, stop));
    CUDA_CHECK(cudaEventDestroy(start));
    CUDA_CHECK(cudaEventDestroy(stop));
    return ms / kTimedRuns;
}

static double gbPerSecond(double floatsPerElem, size_t n, float ms) {
    const double bytes = floatsPerElem * static_cast<double>(n) * sizeof(float);
    return bytes / (static_cast<double>(ms) * 1.0e-3) / 1.0e9;
}

int main() {
    const int device = 0;
    CUDA_CHECK(cudaSetDevice(device));
    cudaDeviceProp prop;
    CUDA_CHECK(cudaGetDeviceProperties(&prop, device));
    std::printf("GPU: %s (compute capability %d.%d)\n", prop.name, prop.major,
                prop.minor);

    int status = EXIT_SUCCESS;
    const size_t bytes = kElems * sizeof(float);
    const size_t pairs = kElems / 2;
    const int halfDim = kHeadDim / 2;
    const size_t tableElems =
        static_cast<size_t>(kSeq) * static_cast<size_t>(halfDim);

    std::printf("\nfloats moved per element\n");
    std::printf("  RoPE, any variant   %.0f\n", kRopeFloatsPerElem);
    std::printf("  staged chain        %.0f  (2 + 2 + 3)\n",
                kStagedFloatsPerElem);
    std::printf("  fused chain         %.0f  (x, residual, out)\n",
                kFusedFloatsPerElem);
    std::printf("  copy floor          %.0f\n", kCopyFloatsPerElem);
    std::printf("  staged / fused      %.3f, against a launch ratio of 3\n",
                kStagedFloatsPerElem / kFusedFloatsPerElem);

    float* d_in = nullptr;
    float* d_r = nullptr;
    float* d_t1 = nullptr;
    float* d_t2 = nullptr;
    float* d_out = nullptr;
    float* d_cos = nullptr;
    float* d_sin = nullptr;
    float* d_bias = nullptr;
    CUDA_CHECK(cudaMalloc(&d_in, bytes));
    CUDA_CHECK(cudaMalloc(&d_r, bytes));
    CUDA_CHECK(cudaMalloc(&d_t1, bytes));
    CUDA_CHECK(cudaMalloc(&d_t2, bytes));
    CUDA_CHECK(cudaMalloc(&d_out, bytes));
    CUDA_CHECK(cudaMalloc(&d_cos, tableElems * sizeof(float)));
    CUDA_CHECK(cudaMalloc(&d_sin, tableElems * sizeof(float)));
    CUDA_CHECK(cudaMalloc(&d_bias, kCols * sizeof(float)));

    // -----------------------------------------------------------------
    // Part 1a: RoPE correctness, on the ragged shape
    // -----------------------------------------------------------------
    {
        const int chkHalf = kChkDim / 2;
        const size_t chkPairs = kChkRopeElems / 2;
        const size_t chkTable =
            static_cast<size_t>(kChkSeq) * static_cast<size_t>(chkHalf);

        std::vector<float> h_q(kChkRopeElems);
        std::vector<float> h_got(kChkRopeElems);
        std::vector<float> h_got2(kChkRopeElems);
        std::vector<float> h_split(kChkRopeElems);
        std::vector<double> h_wantPair(kChkRopeElems);
        std::vector<double> h_wantSplit(kChkRopeElems);
        std::vector<float> h_cos(chkTable);
        std::vector<float> h_sin(chkTable);
        for (size_t i = 0; i < kChkRopeElems; ++i) {
            h_q[i] = static_cast<float>((i % 61) - 30) * 0.125f;
        }
        buildTables(h_cos, h_sin, kChkSeq, kChkDim);

        CUDA_CHECK(cudaMemcpy(d_in, h_q.data(), kChkRopeElems * sizeof(float),
                              cudaMemcpyHostToDevice));
        CUDA_CHECK(cudaMemcpy(d_cos, h_cos.data(), chkTable * sizeof(float),
                              cudaMemcpyHostToDevice));
        CUDA_CHECK(cudaMemcpy(d_sin, h_sin.data(), chkTable * sizeof(float),
                              cudaMemcpyHostToDevice));

        const int chkBlocks = static_cast<int>(
            (chkPairs + kThreadsPerBlock - 1) / kThreadsPerBlock);
        ropeInterleaved<<<chkBlocks, kThreadsPerBlock>>>(
            d_in, d_cos, d_sin, d_t1, chkPairs, chkHalf, kChkHeads, kChkSeq);
        CUDA_CHECK(cudaGetLastError());
        ropeInterleavedVec2<<<chkBlocks, kThreadsPerBlock>>>(
            d_in, d_cos, d_sin, d_t2, chkPairs, chkHalf, kChkHeads, kChkSeq);
        CUDA_CHECK(cudaGetLastError());
        ropeHalfSplit<<<chkBlocks, kThreadsPerBlock>>>(
            d_in, d_cos, d_sin, d_out, chkPairs, chkHalf, kChkHeads, kChkSeq);
        CUDA_CHECK(cudaGetLastError());
        CUDA_CHECK(cudaDeviceSynchronize());

        CUDA_CHECK(cudaMemcpy(h_got.data(), d_t1, kChkRopeElems * sizeof(float),
                              cudaMemcpyDeviceToHost));
        CUDA_CHECK(cudaMemcpy(h_got2.data(), d_t2,
                              kChkRopeElems * sizeof(float),
                              cudaMemcpyDeviceToHost));
        CUDA_CHECK(cudaMemcpy(h_split.data(), d_out,
                              kChkRopeElems * sizeof(float),
                              cudaMemcpyDeviceToHost));

        ropeInterleavedCpu(h_q.data(), h_cos.data(), h_sin.data(),
                           h_wantPair.data(), kChkBatch, kChkSeq, kChkHeads,
                           kChkDim);
        ropeHalfSplitCpu(h_q.data(), h_cos.data(), h_sin.data(),
                         h_wantSplit.data(), kChkBatch, kChkSeq, kChkHeads,
                         kChkDim);

        const size_t badPair =
            firstMismatch(h_got.data(), h_wantPair.data(), kChkRopeElems,
                          kRelTolerance, kAbsTolerance);
        if (badPair != kChkRopeElems) {
            std::fprintf(stderr,
                         "ropeInterleaved wrong at %zu: got %.9g, want %.9g\n",
                         badPair, static_cast<double>(h_got[badPair]),
                         h_wantPair[badPair]);
            status = EXIT_FAILURE;
        }

        const size_t badSplit =
            firstMismatch(h_split.data(), h_wantSplit.data(), kChkRopeElems,
                          kRelTolerance, kAbsTolerance);
        if (badSplit != kChkRopeElems) {
            std::fprintf(stderr,
                         "ropeHalfSplit wrong at %zu: got %.9g, want %.9g\n",
                         badSplit, static_cast<double>(h_split[badSplit]),
                         h_wantSplit[badSplit]);
            status = EXIT_FAILURE;
        }

        // The float2 kernel runs the same operations in the same order on
        // the same values, so it has no licence to differ in any bit. A
        // tolerance here would hide exactly the kind of bug a vectorised
        // rewrite introduces.
        const size_t badVec =
            firstDifference(h_got.data(), h_got2.data(), kChkRopeElems);
        if (badVec != kChkRopeElems) {
            std::fprintf(stderr,
                         "ropeInterleavedVec2 differs from ropeInterleaved at "
                         "%zu: %.9g against %.9g\n",
                         badVec, static_cast<double>(h_got2[badVec]),
                         static_cast<double>(h_got[badVec]));
            status = EXIT_FAILURE;
        }

        // The two conventions must not agree. If they do on this input, the
        // program is not testing what it claims to and every sentence about
        // pairing on the page is unsupported.
        const size_t sameAt =
            firstDifference(h_got.data(), h_split.data(), kChkRopeElems);
        if (sameAt == kChkRopeElems) {
            std::fprintf(stderr,
                         "the two RoPE conventions produced identical "
                         "tensors, which the input was chosen to prevent\n");
            status = EXIT_FAILURE;
        }
        std::printf(
            "\nPart 1: RoPE on %zu elements checked against a double "
            "reference\n  first index where the two conventions differ: %zu\n",
            kChkRopeElems, sameAt);
    }

    // -----------------------------------------------------------------
    // Part 1b: RoPE timing, on the full shape
    // -----------------------------------------------------------------
    if (status == EXIT_SUCCESS) {
        std::vector<float> h_q(kElems);
        std::vector<float> h_cos(tableElems);
        std::vector<float> h_sin(tableElems);
        for (size_t i = 0; i < kElems; ++i) {
            h_q[i] = static_cast<float>((i % 61) - 30) * 0.125f;
        }
        buildTables(h_cos, h_sin, kSeq, kHeadDim);
        CUDA_CHECK(cudaMemcpy(d_in, h_q.data(), bytes, cudaMemcpyHostToDevice));
        CUDA_CHECK(cudaMemcpy(d_cos, h_cos.data(), tableElems * sizeof(float),
                              cudaMemcpyHostToDevice));
        CUDA_CHECK(cudaMemcpy(d_sin, h_sin.data(), tableElems * sizeof(float),
                              cudaMemcpyHostToDevice));

        const int pairBlocks =
            static_cast<int>((pairs + kThreadsPerBlock - 1) / kThreadsPerBlock);
        const int elemBlocks = static_cast<int>(
            (kElems + kThreadsPerBlock - 1) / kThreadsPerBlock);

        const float msPair = timeKernel([&] {
            ropeInterleaved<<<pairBlocks, kThreadsPerBlock>>>(
                d_in, d_cos, d_sin, d_t1, pairs, halfDim, kHeads, kSeq);
        });
        const float msVec = timeKernel([&] {
            ropeInterleavedVec2<<<pairBlocks, kThreadsPerBlock>>>(
                d_in, d_cos, d_sin, d_t2, pairs, halfDim, kHeads, kSeq);
        });
        const float msSplit = timeKernel([&] {
            ropeHalfSplit<<<pairBlocks, kThreadsPerBlock>>>(
                d_in, d_cos, d_sin, d_out, pairs, halfDim, kHeads, kSeq);
        });
        const float msCopy = timeKernel([&] {
            copyFloor<<<elemBlocks, kThreadsPerBlock>>>(d_in, d_out, kElems);
        });

        if (msPair <= 0.0f || msVec <= 0.0f || msSplit <= 0.0f ||
            msCopy <= 0.0f) {
            std::fprintf(stderr,
                         "a part 1 timing was not positive, so the event "
                         "pair never separated\n");
            status = EXIT_FAILURE;
        }

        std::printf(
            "\n  %zu elements, %d heads of %d, sequence %d\n"
            "  %-24s %10s %10s %8s\n",
            kElems, kHeads, kHeadDim, kSeq, "kernel", "ms", "GB/s", "x copy");
        const float rows[] = {msPair, msVec, msSplit, msCopy};
        const char* names[] = {"ropeInterleaved", "ropeInterleavedVec2",
                               "ropeHalfSplit", "copyFloor"};
        const double perElem[] = {kRopeFloatsPerElem, kRopeFloatsPerElem,
                                  kRopeFloatsPerElem, kCopyFloatsPerElem};
        for (int k = 0; k < 4; ++k) {
            std::printf("  %-24s %10.4f %10.1f %8.2f\n", names[k],
                        static_cast<double>(rows[k]),
                        gbPerSecond(perElem[k], kElems, rows[k]),
                        static_cast<double>(msCopy / rows[k]));
        }
    }

    // -----------------------------------------------------------------
    // Part 2: the two GELUs
    // -----------------------------------------------------------------
    if (status == EXIT_SUCCESS) {
        std::vector<float> h_x(kGeluSamples);
        std::vector<float> h_exact(kGeluSamples);
        std::vector<float> h_gap(kGeluSamples);
        std::vector<double> h_want(kGeluSamples);
        const double step = (static_cast<double>(kGeluHi) - kGeluLo) /
                            static_cast<double>(kGeluSamples - 1);
        for (size_t i = 0; i < kGeluSamples; ++i) {
            h_x[i] =
                static_cast<float>(kGeluLo + step * static_cast<double>(i));
        }
        CUDA_CHECK(cudaMemcpy(d_in, h_x.data(), kGeluSamples * sizeof(float),
                              cudaMemcpyHostToDevice));

        const int blocks = static_cast<int>(
            (kGeluSamples + kThreadsPerBlock - 1) / kThreadsPerBlock);
        geluBothForms<<<blocks, kThreadsPerBlock>>>(d_in, d_t1, d_t2,
                                                    kGeluSamples);
        CUDA_CHECK(cudaGetLastError());
        CUDA_CHECK(cudaDeviceSynchronize());
        CUDA_CHECK(cudaMemcpy(h_exact.data(), d_t1,
                              kGeluSamples * sizeof(float),
                              cudaMemcpyDeviceToHost));
        CUDA_CHECK(cudaMemcpy(h_gap.data(), d_t2, kGeluSamples * sizeof(float),
                              cudaMemcpyDeviceToHost));

        for (size_t i = 0; i < kGeluSamples; ++i) {
            h_want[i] = geluErfCpu(static_cast<double>(h_x[i]));
        }
        const size_t bad =
            firstMismatch(h_exact.data(), h_want.data(), kGeluSamples,
                          kRelTolerance, kAbsTolerance);
        if (bad != kGeluSamples) {
            std::fprintf(stderr,
                         "geluErf wrong at x = %.9g: got %.9g, want %.9g\n",
                         static_cast<double>(h_x[bad]),
                         static_cast<double>(h_exact[bad]), h_want[bad]);
            status = EXIT_FAILURE;
        }

        size_t worst = 0;
        for (size_t i = 1; i < kGeluSamples; ++i) {
            if (h_gap[i] > h_gap[worst]) {
                worst = i;
            }
        }
        const double worstGap = static_cast<double>(h_gap[worst]);
        std::printf(
            "\nPart 2: exact GELU against the tanh approximation, %zu "
            "samples over [%.1f, %.1f]\n"
            "  largest absolute gap   %.6e at x = %.6f\n"
            "  exact GELU there       %.6f\n"
            "  relative to that value %.3e\n",
            kGeluSamples, static_cast<double>(kGeluLo),
            static_cast<double>(kGeluHi), worstGap,
            static_cast<double>(h_x[worst]),
            static_cast<double>(h_exact[worst]),
            worstGap / std::fabs(static_cast<double>(h_exact[worst])));

        if (worstGap <= kGeluFormGap) {
            std::fprintf(stderr,
                         "the two GELU forms agreed to %.3e, so this page's "
                         "claim that they are different functions is wrong\n",
                         worstGap);
            status = EXIT_FAILURE;
        }
    }

    // -----------------------------------------------------------------
    // Part 3a: chain correctness, on the ragged shape
    // -----------------------------------------------------------------
    if (status == EXIT_SUCCESS) {
        std::vector<float> h_x(kChkFuseElems);
        std::vector<float> h_res(kChkFuseElems);
        std::vector<float> h_bias(kChkCols);
        std::vector<float> h_staged(kChkFuseElems);
        std::vector<float> h_fused(kChkFuseElems);
        std::vector<double> h_want(kChkFuseElems);
        for (size_t i = 0; i < kChkFuseElems; ++i) {
            h_x[i] = static_cast<float>((i % 53) - 26) * 0.1875f;
            h_res[i] = static_cast<float>((i % 29) - 14) * 0.0625f;
        }
        for (int c = 0; c < kChkCols; ++c) {
            h_bias[c] = static_cast<float>((c % 11) - 5) * 0.25f;
        }
        CUDA_CHECK(cudaMemcpy(d_in, h_x.data(), kChkFuseElems * sizeof(float),
                              cudaMemcpyHostToDevice));
        CUDA_CHECK(cudaMemcpy(d_r, h_res.data(), kChkFuseElems * sizeof(float),
                              cudaMemcpyHostToDevice));
        CUDA_CHECK(cudaMemcpy(d_bias, h_bias.data(), kChkCols * sizeof(float),
                              cudaMemcpyHostToDevice));

        const dim3 chkGrid(
            static_cast<unsigned int>((kChkCols + kThreadsPerBlock - 1) /
                                      kThreadsPerBlock),
            static_cast<unsigned int>(kChkRows));
        const int chkBlocks = static_cast<int>(
            (kChkFuseElems + kThreadsPerBlock - 1) / kThreadsPerBlock);

        addBiasRow<<<chkGrid, kThreadsPerBlock>>>(d_in, d_bias, d_t1, kChkCols);
        CUDA_CHECK(cudaGetLastError());
        applyGeluErf<<<chkBlocks, kThreadsPerBlock>>>(d_t1, d_t2,
                                                      kChkFuseElems);
        CUDA_CHECK(cudaGetLastError());
        addResidual<<<chkBlocks, kThreadsPerBlock>>>(d_t2, d_r, d_out,
                                                     kChkFuseElems);
        CUDA_CHECK(cudaGetLastError());
        CUDA_CHECK(cudaDeviceSynchronize());
        CUDA_CHECK(cudaMemcpy(h_staged.data(), d_out,
                              kChkFuseElems * sizeof(float),
                              cudaMemcpyDeviceToHost));

        biasGeluResidualFused<<<chkGrid, kThreadsPerBlock>>>(d_in, d_bias, d_r,
                                                             d_out, kChkCols);
        CUDA_CHECK(cudaGetLastError());
        CUDA_CHECK(cudaDeviceSynchronize());
        CUDA_CHECK(cudaMemcpy(h_fused.data(), d_out,
                              kChkFuseElems * sizeof(float),
                              cudaMemcpyDeviceToHost));

        chainCpu(h_x.data(), h_bias.data(), h_res.data(), h_want.data(),
                 kChkRows, kChkCols);

        const size_t badStaged =
            firstMismatch(h_staged.data(), h_want.data(), kChkFuseElems,
                          kRelTolerance, kAbsTolerance);
        if (badStaged != kChkFuseElems) {
            std::fprintf(stderr,
                         "staged chain wrong at %zu: got %.9g, want %.9g\n",
                         badStaged, static_cast<double>(h_staged[badStaged]),
                         h_want[badStaged]);
            status = EXIT_FAILURE;
        }
        const size_t badFused =
            firstMismatch(h_fused.data(), h_want.data(), kChkFuseElems,
                          kRelTolerance, kAbsTolerance);
        if (badFused != kChkFuseElems) {
            std::fprintf(stderr,
                         "fused chain wrong at %zu: got %.9g, want %.9g\n",
                         badFused, static_cast<double>(h_fused[badFused]),
                         h_want[badFused]);
            status = EXIT_FAILURE;
        }

        // Counted, never gated. Both versions round to float at the same
        // three points, so they should agree, but the compiler may contract
        // a multiply and an add across a stage boundary in the fused body
        // where the staged one has a store in the way. A difference is a
        // fact about fusion, not a failure.
        size_t differing = 0;
        double widest = 0.0;
        for (size_t i = 0; i < kChkFuseElems; ++i) {
            const double gap = std::fabs(static_cast<double>(h_staged[i]) -
                                         static_cast<double>(h_fused[i]));
            if (gap > 0.0) {
                ++differing;
                if (gap > widest) {
                    widest = gap;
                }
            }
        }
        std::printf(
            "\nPart 3: chain on %zu elements (%d rows of %d) checked "
            "against a double reference\n"
            "  staged and fused differ on %zu of %zu, by at most %.3e\n",
            kChkFuseElems, kChkRows, kChkCols, differing, kChkFuseElems,
            widest);
    }

    // -----------------------------------------------------------------
    // Part 3b: chain timing, on the full shape
    // -----------------------------------------------------------------
    if (status == EXIT_SUCCESS) {
        std::vector<float> h_x(kElems);
        std::vector<float> h_bias(kCols);
        for (size_t i = 0; i < kElems; ++i) {
            h_x[i] = static_cast<float>((i % 53) - 26) * 0.1875f;
        }
        CUDA_CHECK(cudaMemcpy(d_in, h_x.data(), bytes, cudaMemcpyHostToDevice));
        for (size_t i = 0; i < kElems; ++i) {
            h_x[i] = static_cast<float>((i % 29) - 14) * 0.0625f;
        }
        CUDA_CHECK(cudaMemcpy(d_r, h_x.data(), bytes, cudaMemcpyHostToDevice));
        for (int c = 0; c < kCols; ++c) {
            h_bias[c] = static_cast<float>((c % 11) - 5) * 0.25f;
        }
        CUDA_CHECK(cudaMemcpy(d_bias, h_bias.data(), kCols * sizeof(float),
                              cudaMemcpyHostToDevice));

        const dim3 grid(static_cast<unsigned int>(kCols / kThreadsPerBlock),
                        static_cast<unsigned int>(kRows));
        const int elemBlocks = static_cast<int>(
            (kElems + kThreadsPerBlock - 1) / kThreadsPerBlock);

        const float msBias = timeKernel([&] {
            addBiasRow<<<grid, kThreadsPerBlock>>>(d_in, d_bias, d_t1, kCols);
        });
        const float msGelu = timeKernel([&] {
            applyGeluErf<<<elemBlocks, kThreadsPerBlock>>>(d_t1, d_t2, kElems);
        });
        const float msRes = timeKernel([&] {
            addResidual<<<elemBlocks, kThreadsPerBlock>>>(d_t2, d_r, d_out,
                                                          kElems);
        });
        const float msStaged = timeKernel([&] {
            addBiasRow<<<grid, kThreadsPerBlock>>>(d_in, d_bias, d_t1, kCols);
            applyGeluErf<<<elemBlocks, kThreadsPerBlock>>>(d_t1, d_t2, kElems);
            addResidual<<<elemBlocks, kThreadsPerBlock>>>(d_t2, d_r, d_out,
                                                          kElems);
        });
        const float msFused = timeKernel([&] {
            biasGeluResidualFused<<<grid, kThreadsPerBlock>>>(d_in, d_bias, d_r,
                                                              d_out, kCols);
        });
        const float msCopy = timeKernel([&] {
            copyFloor<<<elemBlocks, kThreadsPerBlock>>>(d_in, d_out, kElems);
        });

        if (msStaged <= 0.0f || msFused <= 0.0f || msCopy <= 0.0f) {
            std::fprintf(stderr,
                         "a part 3 timing was not positive, so the event "
                         "pair never separated\n");
            status = EXIT_FAILURE;
        }

        std::printf(
            "\n  %zu elements, %d rows of %d, bias of %d floats\n"
            "  %-24s %10s %8s %10s %8s\n",
            kElems, kRows, kCols, kCols, "version", "ms", "floats", "GB/s",
            "x copy");
        const float times[] = {msBias,   msGelu,  msRes,
                               msStaged, msFused, msCopy};
        const char* names[] = {"addBiasRow",  "applyGeluErf",
                               "addResidual", "staged chain, one region",
                               "fused chain", "copyFloor"};
        const double perElem[] = {2.0,
                                  2.0,
                                  3.0,
                                  kStagedFloatsPerElem,
                                  kFusedFloatsPerElem,
                                  kCopyFloatsPerElem};
        for (int k = 0; k < 6; ++k) {
            std::printf("  %-24s %10.4f %8.0f %10.1f %8.2f\n", names[k],
                        static_cast<double>(times[k]), perElem[k],
                        gbPerSecond(perElem[k], kElems, times[k]),
                        static_cast<double>(msCopy / times[k]));
        }
        std::printf(
            "  three stages summed      %10.4f\n"
            "  staged / fused           %10.3f  (byte ratio %.3f, launch "
            "ratio 3)\n",
            static_cast<double>(msBias + msGelu + msRes),
            static_cast<double>(msStaged / msFused),
            kStagedFloatsPerElem / kFusedFloatsPerElem);

        // More traffic predicts a slower staged chain, but cache residency,
        // launch overhead, and measurement noise are hardware-dependent.
        // Report the outcome without conflating it with correctness.
        std::printf("  fusion performance prediction: %s\n",
                    (msStaged > msFused) ? "held" : "refuted");
    }

    CUDA_CHECK(cudaFree(d_in));
    CUDA_CHECK(cudaFree(d_r));
    CUDA_CHECK(cudaFree(d_t1));
    CUDA_CHECK(cudaFree(d_t2));
    CUDA_CHECK(cudaFree(d_out));
    CUDA_CHECK(cudaFree(d_cos));
    CUDA_CHECK(cudaFree(d_sin));
    CUDA_CHECK(cudaFree(d_bias));
    return status;
}