COURSE / SOURCE

quant_matmul.cu

All lessons
Source filecode/day98-quantization/quant_matmul.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 98: quantization and sparsity. An INT8 matmul with scaling, built as a
// three-row ladder over day 16's tiled kernel, plus a host-side census that
// prices FP8 and MX-style block scaling on the same weight matrix, plus the
// accuracy half of 2:4 structured pruning.
//
// The contract this program measures: inside the tile loop an INT8 matmul is
// exact integer arithmetic, so every bit of error is spent at the two
// boundaries, quantizing the inputs and rescaling in the epilogue. That is
// the opposite of day 71's FP16 story, where the accumulator decided.
//
// The three INT8 rows change one thing each:
//   1. one scale for all of B, one symmetric scale for A (the common form)
//   2. a scale per column of B, whose column maxima span 64x here
//   3. and an asymmetric scale with a zero point on A, which is non-negative
//
// Not measured here, and named on the page as such: INT8 tensor cores (the
// T4 has them at 7.5; this kernel runs integer MADs on CUDA cores), FP8
// tensor cores (8.9), and the 2:4 sparse tensor-core path through cuSPARSELt
// (8.0). Every FP8 and 2:4 number below is an accuracy number, never a time.
//
// Build: nvcc -std=c++17 -O3 -arch=sm_75 -o quant_matmul quant_matmul.cu
// Run:   ./quant_matmul

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

#include <cuda_fp8.h>
#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)

// Day 16's tile: 16 x 16 is one 256-thread block, the course default. The
// INT8 tiles cost 256 bytes each against the FP32 pair's 1 KiB each, which
// is a quarter of the shared memory and none of the interest: the access
// pattern is unchanged, so the table compares number formats and nothing
// else.
constexpr int kTileDim = 16;
constexpr int kWarmupRuns = 3;
constexpr int kTimedRuns = 10;

// Square cases, all multiples of the tile. K grows 8x across the ladder so
// error growth with the reduction depth is visible in one table, and the
// footprints straddle the T4's 4 MiB L2: three 256 x 256 float buffers are
// 768 KiB and live in cache, while the 2048 case moves 48 MiB in FP32 and
// 12 MiB in INT8, which is where narrower loads can pay.
constexpr int kNumCases = 3;
constexpr size_t kCaseSize[kNumCases] = {256, 1024, 2048};
constexpr size_t kMaxDim = 2048;

constexpr int kInt8Max = 127;
constexpr int kMxBlock = 32;  // the MX block length, taken along k
constexpr float kFp8E4m3Max = 448.0f;

// Base rtol for the FP32 row from the course tolerance table, with 23
// mantissa bits behind it. The INT8 rows do not use it; their gate is
// absolute and derived in columnGateAtol below.
constexpr double kF32Rtol = 1e-5;
constexpr int kF32MantissaBits = 23;

// The slack factor in every gate here, the same 4 day 66 uses: it is room
// for a valid-but-different summation order, not room for a bug.
constexpr double kGateSlack = 4.0;

// B's column gains run 1, 1/2, ... 1/64 and repeat, so the column maxima of
// the weight matrix span 64x. Real weight tensors do this and per-tensor
// quantization is where it hurts. The outlier every 37th element is the
// other half of the story, the one block scaling exists for.
constexpr int kColumnGainCount = 7;
constexpr int kOutlierPeriod = 37;
constexpr float kOutlierGain = 8.0f;

static_assert(kTileDim * kTileDim == 256,
              "a 16 x 16 tile is one 256-thread block, the course default");
static_assert((kTileDim * kTileDim) % 32 == 0,
              "block size must be a whole number of warps");
static_assert(kCaseSize[0] % kTileDim == 0 && kCaseSize[1] % kTileDim == 0 &&
                  kCaseSize[2] % kTileDim == 0,
              "every case divides by the tile, so all four kernels time the "
              "same clean geometry");
static_assert(kCaseSize[0] < kCaseSize[1] && kCaseSize[1] < kCaseSize[2] &&
                  kCaseSize[2] <= kMaxDim,
              "cases ascend so the error column reads as K grows, and "
              "everything fits the one kMaxDim allocation");
static_assert(kMaxDim % kMxBlock == 0,
              "the MX census walks whole blocks of 32 down each column");
static_assert(2 * kMaxDim * 127 * 128 < 2147483647,
              "the INT32 accumulator must not overflow: |acc| <= K*127*128 "
              "and the zero-point correction is bounded by the same product, "
              "so twice it has to fit a signed 32-bit integer");
static_assert(2 * kTileDim * kTileDim * sizeof(float) <= 48 * 1024,
              "the FP32 tiles must fit the 48 KiB a block gets on a T4 "
              "without the cudaFuncSetAttribute opt-in");

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

// Day 16's tiled kernel, unchanged, in FP32. This is the baseline row of
// every table below and the reference the INT8 rows are timed against.
//
// Memory: threadIdx.x is the fastest index, so a warp's half-row of the B
// tile fill reads 16 consecutive floats, 64 bytes. The INT8 kernel changes
// the element width and nothing else about this pattern.
//
// Launch assumption: exactly kTileDim x kTileDim threads per block and a
// grid that rounds up on both axes. Guards cover the loads and the store,
// never the barrier.
__global__ void matmulTiledF32(const float* __restrict__ a,
                               const float* __restrict__ b,
                               float* __restrict__ c, size_t m, size_t n,
                               size_t k) {
    __shared__ float tileA[kTileDim][kTileDim];
    __shared__ float tileB[kTileDim][kTileDim];

    const size_t col =
        blockIdx.x * static_cast<size_t>(blockDim.x) + threadIdx.x;
    const size_t row =
        blockIdx.y * static_cast<size_t>(blockDim.y) + threadIdx.y;
    const size_t tiles = (k + kTileDim - 1) / kTileDim;

    float acc = 0.0f;
    for (size_t tileIdx = 0; tileIdx < tiles; ++tileIdx) {
        const size_t aCol = tileIdx * kTileDim + threadIdx.x;
        const size_t bRow = tileIdx * kTileDim + threadIdx.y;

        tileA[threadIdx.y][threadIdx.x] =
            (row < m && aCol < k) ? a[row * k + aCol] : 0.0f;
        tileB[threadIdx.y][threadIdx.x] =
            (bRow < k && col < n) ? b[bRow * n + col] : 0.0f;
        __syncthreads();

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

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

// The INT8 ladder as one kernel with its two decisions as template
// parameters, so a row of the table differs from the row above it in exactly
// one of them.
//
// kPerChannelB picks between one scale for the whole weight matrix and one
// per output column. kAsymA adds the zero point that lets a non-negative
// activation use all 256 codes instead of 128; its price is the correction
// term, which needs the column sums of the quantized B and is therefore
// O(N^2) work in the epilogue rather than O(N^3) work in the loop.
//
// The loop itself is exact. Products of two 8-bit integers accumulated in
// INT32 round nothing at all, and the static_assert above says why the sum
// cannot overflow. All of this kernel's error was already spent by the
// quantizer before the launch.
//
// Memory: identical to matmulTiledF32 except that every element is one byte,
// so a warp's B tile fill reads 16 bytes where the FP32 kernel read 64.
//
// Launch assumption: identical to matmulTiledF32.
template <bool kPerChannelB, bool kAsymA>
__global__ void matmulInt8(const int8_t* __restrict__ a,
                           const int8_t* __restrict__ b,
                           const float* __restrict__ bScales,
                           const int* __restrict__ bColSums, float aScale,
                           float bScale, int aZero, float* __restrict__ c,
                           size_t m, size_t n, size_t k) {
    __shared__ int8_t tileA[kTileDim][kTileDim];
    __shared__ int8_t tileB[kTileDim][kTileDim];

    const size_t col =
        blockIdx.x * static_cast<size_t>(blockDim.x) + threadIdx.x;
    const size_t row =
        blockIdx.y * static_cast<size_t>(blockDim.y) + threadIdx.y;
    const size_t tiles = (k + kTileDim - 1) / kTileDim;

    // The A guard fills with the code that means zero, which is the zero
    // point rather than the literal 0 once the scale is asymmetric. Nothing
    // in this program takes the guard, since every case divides by the tile,
    // and a fill that is only right by accident is worse than no fill.
    const int8_t aPad = kAsymA ? static_cast<int8_t>(aZero) : int8_t{0};

    // snippet: int8-inner
    int acc = 0;
    for (size_t tileIdx = 0; tileIdx < tiles; ++tileIdx) {
        const size_t aCol = tileIdx * kTileDim + threadIdx.x;
        const size_t bRow = tileIdx * kTileDim + threadIdx.y;

        tileA[threadIdx.y][threadIdx.x] =
            (row < m && aCol < k) ? a[row * k + aCol] : aPad;
        tileB[threadIdx.y][threadIdx.x] =
            (bRow < k && col < n) ? b[bRow * n + col] : int8_t{0};
        __syncthreads();

        for (int p = 0; p < kTileDim; ++p) {
            acc += static_cast<int>(tileA[threadIdx.y][p]) *
                   static_cast<int>(tileB[p][threadIdx.x]);
        }
        __syncthreads();
    }

    if (row < m && col < n) {
        int corrected = acc;
        if constexpr (kAsymA) {
            corrected -= aZero * bColSums[col];
        }
        const float scale = aScale * (kPerChannelB ? bScales[col] : bScale);
        c[row * n + col] = static_cast<float>(corrected) * scale;
    }
    // end snippet
}

// Round to nearest, then clamp. lrintf follows the current rounding mode,
// which is round-to-nearest-even unless a program changes it, and nothing
// here does. The clamp is not decoration: a value that lands on 128 after
// rounding wraps to -128 in an int8_t, which turns the largest weight in the
// matrix into the most negative one.
static int8_t toInt8(float v) {
    long r = std::lrintf(v);
    if (r > 127) {
        r = 127;
    }
    if (r < -128) {
        r = -128;
    }
    return static_cast<int8_t>(r);
}

// Symmetric per-tensor quantization. One scale, no zero point, the code 0
// means the value 0. Dequantizing is q * scale.
static float quantizeSymmetric(const float* x, int8_t* q, size_t n) {
    float amax = 0.0f;
    for (size_t i = 0; i < n; ++i) {
        if (std::fabs(x[i]) > amax) {
            amax = std::fabs(x[i]);
        }
    }
    const float scale =
        (amax > 0.0f) ? amax / static_cast<float>(kInt8Max) : 1.0f;
    for (size_t i = 0; i < n; ++i) {
        q[i] = toInt8(x[i] / scale);
    }
    return scale;
}

// Asymmetric per-tensor quantization: 256 codes spread over the observed
// range, with the zero point being the code that dequantizes to zero.
// Dequantizing is (q - zero) * scale. On a non-negative tensor this is worth
// exactly one bit against the symmetric form, because symmetric spends half
// its codes on negatives that never occur.
static float quantizeAsymmetric(const float* x, int8_t* q, size_t n,
                                int* zero) {
    float lo = x[0];
    float hi = x[0];
    for (size_t i = 1; i < n; ++i) {
        if (x[i] < lo) {
            lo = x[i];
        }
        if (x[i] > hi) {
            hi = x[i];
        }
    }
    const float span = hi - lo;
    const float scale = (span > 0.0f) ? span / 255.0f : 1.0f;
    const int z = static_cast<int>(std::lrintf(-128.0f - lo / scale));
    *zero = z;
    for (size_t i = 0; i < n; ++i) {
        q[i] = toInt8(x[i] / scale + static_cast<float>(z));
    }
    return scale;
}

// snippet: quantize-columns
// Symmetric quantization with one scale per column of B, which is one scale
// per output channel of the layer. The scale a column gets is set by that
// column's own largest magnitude, so a column whose values are 64x smaller
// than the matrix maximum keeps its full 8 bits instead of being rounded
// into two or three codes.
static void quantizeColumns(const float* b, int8_t* q, float* scales,
                            float* colMax, size_t k, size_t n) {
    for (size_t col = 0; col < n; ++col) {
        float amax = 0.0f;
        for (size_t p = 0; p < k; ++p) {
            const float v = std::fabs(b[p * n + col]);
            if (v > amax) {
                amax = v;
            }
        }
        colMax[col] = amax;
        scales[col] =
            (amax > 0.0f) ? amax / static_cast<float>(kInt8Max) : 1.0f;
        for (size_t p = 0; p < k; ++p) {
            q[p * n + col] = toInt8(b[p * n + col] / scales[col]);
        }
    }
}
// end snippet

// Column sums of the quantized weights, which is the whole cost of the zero
// point: the asymmetric epilogue needs sum_k b_q[k][col] to undo the offset,
// and that sum does not depend on the activations, so a real inference stack
// computes it once when the weights are quantized rather than per call.
static void columnSums(const int8_t* q, int* sums, size_t k, size_t n) {
    for (size_t col = 0; col < n; ++col) {
        int total = 0;
        for (size_t p = 0; p < k; ++p) {
            total += static_cast<int>(q[p * n + col]);
        }
        sums[col] = total;
    }
}

// 2:4 structured pruning along k: in every group of four consecutive weights
// down a column, keep the two largest by magnitude and zero the other two.
// That is the pattern the sparse tensor cores can skip, and forcing it is
// free of any hardware requirement. Using it for speed is not; see the
// README.
static void prune24(const float* b, float* out, size_t k, size_t n) {
    for (size_t col = 0; col < n; ++col) {
        for (size_t base = 0; base < k; base += 4) {
            const size_t span = (k - base < 4) ? (k - base) : 4;
            size_t keep0 = base;
            size_t keep1 = base;
            float best = -1.0f;
            float second = -1.0f;
            for (size_t j = 0; j < span; ++j) {
                const float v = std::fabs(b[(base + j) * n + col]);
                if (v > best) {
                    second = best;
                    keep1 = keep0;
                    best = v;
                    keep0 = base + j;
                } else if (v > second) {
                    second = v;
                    keep1 = base + j;
                }
            }
            for (size_t j = 0; j < span; ++j) {
                const size_t p = base + j;
                out[p * n + col] =
                    (p == keep0 || p == keep1) ? b[p * n + col] : 0.0f;
            }
        }
    }
}

// CPU reference in double over the ORIGINAL float inputs, never the
// quantized copies. A reference built from the dequantized values would
// forgive the quantization error, which is the only error this day has.
// Written for obvious correctness: plain loops, no blocking.
static void matmulCpu(const float* a, const float* b, double* c, size_t m,
                      size_t n, size_t k) {
    for (size_t row = 0; row < m; ++row) {
        for (size_t col = 0; col < n; ++col) {
            double total = 0.0;
            for (size_t p = 0; p < k; ++p) {
                total += static_cast<double>(a[row * k + p]) *
                         static_cast<double>(b[p * n + col]);
            }
            c[row * n + col] = total;
        }
    }
}

// snippet: column-gate
// The absolute tolerance for one output column, derived rather than tuned.
// Write a = a_q_dequantized + ea with |ea| <= sa/2 and likewise for b. One
// term of the dot product then carries at most |a|*sb/2 + |b|*sa/2, and with
// sa = amax/127 and sb = bColMax/127 that is amax*bColMax/127. A
// correctness gate cannot assume those K errors are independent, so its
// deterministic bound adds their magnitudes and scales with K. The sqrt(K)
// estimate remains useful as a statistical diagnostic, but it does not decide
// whether a kernel is correct. kGateSlack is the same factor 4 day 66 uses
// for a valid-but-different summation order and the omitted second-order term.
//
// The tolerance is per column because the columns of B differ by 64x here.
// One tolerance built from the matrix maximum would be 64x too loose on the
// smallest column, which is exactly the mistake the per-tensor row makes.
static double columnGateAtol(double amax, double bColMax, size_t kDim) {
    return kGateSlack * static_cast<double>(kDim) * amax * bColMax /
           static_cast<double>(kInt8Max);
}

static double columnStatAtol(double amax, double bColMax, size_t kDim) {
    return kGateSlack * std::sqrt(static_cast<double>(kDim)) * amax * bColMax /
           static_cast<double>(kInt8Max);
}
// end snippet

// Worst |got - ref| / atol over the matrix, judged per column, plus the
// first element that exceeds its own column's tolerance. A non-finite output
// fails outright: inf here means an overflow bug and NaN is never close.
static double worstColumnFraction(const float* got, const double* want,
                                  const double* colAtol, size_t m, size_t n,
                                  size_t* firstBad) {
    double worst = 0.0;
    *firstBad = m * n;
    for (size_t row = 0; row < m; ++row) {
        for (size_t col = 0; col < n; ++col) {
            const size_t i = row * n + col;
            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 / colAtol[col];
            if (frac > worst) {
                worst = frac;
            }
            if (frac > 1.0 && *firstBad == m * n) {
                *firstBad = i;
            }
        }
    }
    return worst;
}

// The root mean square error over the whole output, which is the comparator
// the accumulator-free rows are judged on. A max over four million elements
// moves with whichever single element happened to round worst; an RMS over
// the same elements moves with the format.
static double rmsError(const float* got, const double* want, size_t n) {
    double sumSq = 0.0;
    for (size_t i = 0; i < n; ++i) {
        const double err = std::fabs(static_cast<double>(got[i]) - want[i]);
        sumSq += err * err;
    }
    return std::sqrt(sumSq / static_cast<double>(n));
}

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

// The K-scaled relative tolerance day 66 established, used only for the FP32
// row: rtol = max(table value, 4 * eps * sqrt(K)) with eps = 2^-p.
static double kScaledRtol(double tableRtol, int mantissaBits, size_t kDim) {
    const double eps = std::ldexp(1.0, -mantissaBits);
    const double grown =
        kGateSlack * eps * std::sqrt(static_cast<double>(kDim));
    return grown > tableRtol ? grown : tableRtol;
}

static double worstRelFraction(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;
}

// A 32-bit linear congruential generator, so the inputs are the same on
// every machine and every run. The constants are Numerical Recipes'. Why not
// day 71's periodic pattern: quantization errors of a periodic sequence are
// themselves periodic, so they would correlate down a dot product and grow
// like K instead of sqrt(K). The deterministic gate allows that worst case,
// but the statistical diagnostic would then measure the pattern rather than
// the format.
static uint32_t lcgNext(uint32_t* state) {
    *state = *state * 1664525u + 1013904223u;
    return *state;
}

static float unitFloat(uint32_t* state) {
    return static_cast<float>(lcgNext(state) >> 8) / 16777216.0f;
}

// A is activations after a ReLU, so non-negative in [0, 1). That is what
// makes the zero point worth a bit. B is weights, signed and zero-centered,
// with two structures a real weight tensor has and a uniform random matrix
// does not: a per-column gain spanning 64x, and an outlier every 37th
// element that is 8x the rest of its column.
static void makeInputs(std::vector<float>* h_a, std::vector<float>* h_b,
                       size_t m, size_t n, size_t k) {
    uint32_t state = 20260901u;
    for (size_t i = 0; i < m * k; ++i) {
        (*h_a)[i] = unitFloat(&state);
    }
    for (size_t p = 0; p < k; ++p) {
        for (size_t col = 0; col < n; ++col) {
            const float gain =
                std::ldexp(1.0f, -static_cast<int>(col % kColumnGainCount));
            const float outlier =
                ((p * n + col) % kOutlierPeriod == 0) ? kOutlierGain : 1.0f;
            (*h_b)[p * n + col] =
                (2.0f * unitFloat(&state) - 1.0f) * gain * outlier;
        }
    }
}

// What this card's tensor cores can execute, from the compute capability
// read at run time. INT8 mma is supported from 7.5, FP8 (e4m3, e5m2) from
// 8.9, and the 2:4 sparse tensor-core path from 8.0. Sources are in the
// README. This program's kernels use no tensor cores at all; the report
// exists so the transcript states what the card could and could not have
// been doing.
// snippet: capability-report
static void printCapabilityReport(int major, int minor) {
    const int cc = major * 10 + minor;
    struct Floor {
        const char* name;
        int minCc;
    };
    const Floor floors[] = {
        {"INT8 tensor cores", 75},
        {"2:4 sparse tensor cores (cuSPARSELt)", 80},
        {"FP8 e4m3/e5m2 tensor cores", 89},
    };
    std::printf("tensor-core support at compute capability %d.%d:\n", major,
                minor);
    for (const Floor& f : floors) {
        std::printf("  %-38s %s\n", f.name,
                    cc >= f.minCc ? "yes" : "no on this card");
    }
    std::printf(
        "  this program uses none of them: its INT8 loop is integer MADs\n"
        "  on CUDA cores, and its FP8 and 2:4 rows are host arithmetic\n");
}
// end snippet

// snippet: activation-census
// What the zero point is worth, on its own, before any matmul runs. A is
// non-negative, so symmetric quantization spends codes -127 to -1 on values
// that never occur and its step is amax/127. The asymmetric form spreads 256
// codes over the observed range instead, so its step is span/255 and its
// round-trip error is half. This is arithmetic, not hardware: it needs no
// GPU and it comes out the same on every machine.
static void printActivationCensus(const float* a, size_t n, double amax,
                                  double* symRms, double* asymRms) {
    std::vector<int8_t> qSym(n);
    std::vector<int8_t> qAsym(n);
    int zero = 0;
    const float symScale = quantizeSymmetric(a, qSym.data(), n);
    const float asymScale = quantizeAsymmetric(a, qAsym.data(), n, &zero);

    double symMax = 0.0;
    double asymMax = 0.0;
    double symSq = 0.0;
    double asymSq = 0.0;
    for (size_t i = 0; i < n; ++i) {
        const double eSym =
            std::fabs(static_cast<double>(qSym[i]) * symScale - a[i]);
        const double eAsym =
            std::fabs(static_cast<double>(qAsym[i] - zero) * asymScale - a[i]);
        if (eSym > symMax) {
            symMax = eSym;
        }
        if (eAsym > asymMax) {
            asymMax = eAsym;
        }
        symSq += eSym * eSym;
        asymSq += eAsym * eAsym;
    }
    *symRms = std::sqrt(symSq / static_cast<double>(n));
    *asymRms = std::sqrt(asymSq / static_cast<double>(n));

    std::printf(
        "activation census on A (%zu values in [0, 1)), errors as a "
        "fraction of max|A| = %.4f\n",
        n, amax);
    std::printf("  %-30s %14s %14s\n", "format", "max err", "rms err");
    std::printf("  %-30s %14.3e %14.3e\n", "INT8 symmetric, one scale",
                symMax / amax, *symRms / amax);
    std::printf("  %-30s %14.3e %14.3e\n", "INT8 asymmetric, zero point",
                asymMax / amax, *asymRms / amax);
    std::printf("\n");
}
// end snippet

// What each storage format costs before any matmul runs: the round-trip
// error of quantizing B and dequantizing it, reported against B's own
// largest magnitude so the four rows are comparable. This needs no GPU and
// no compute capability, which is the point: it is the format's arithmetic,
// not the hardware's.
static void printStorageCensus(const float* b, const float* colMax, size_t k,
                               size_t n, double bmax) {
    const size_t elems = k * n;
    double maxErr[4] = {0.0, 0.0, 0.0, 0.0};
    double sumSq[4] = {0.0, 0.0, 0.0, 0.0};

    const float tensorScaleInt8 =
        static_cast<float>(bmax) / static_cast<float>(kInt8Max);
    const float tensorScaleFp8 = static_cast<float>(bmax) / kFp8E4m3Max;

    for (size_t col = 0; col < n; ++col) {
        const float colScaleInt8 =
            (colMax[col] > 0.0f) ? colMax[col] / static_cast<float>(kInt8Max)
                                 : 1.0f;
        for (size_t base = 0; base < k; base += kMxBlock) {
            // One shared scale per 32 values down a column, which is the
            // shape MXFP8 and NVFP4 standardise in hardware. Here it is done
            // in software on FP8 values, so it measures the idea and not the
            // instruction.
            float blockMax = 0.0f;
            for (int j = 0; j < kMxBlock; ++j) {
                const float v = std::fabs(b[(base + j) * n + col]);
                if (v > blockMax) {
                    blockMax = v;
                }
            }
            const float blockScaleFp8 =
                (blockMax > 0.0f) ? blockMax / kFp8E4m3Max : 1.0f;

            for (int j = 0; j < kMxBlock; ++j) {
                const float x = b[(base + j) * n + col];
                const double back[4] = {
                    static_cast<double>(toInt8(x / tensorScaleInt8)) *
                        tensorScaleInt8,
                    static_cast<double>(toInt8(x / colScaleInt8)) *
                        colScaleInt8,
                    static_cast<double>(
                        static_cast<float>(__nv_fp8_e4m3(x / tensorScaleFp8))) *
                        tensorScaleFp8,
                    static_cast<double>(
                        static_cast<float>(__nv_fp8_e4m3(x / blockScaleFp8))) *
                        blockScaleFp8,
                };
                for (int r = 0; r < 4; ++r) {
                    const double err = std::fabs(back[r] - x);
                    if (err > maxErr[r]) {
                        maxErr[r] = err;
                    }
                    sumSq[r] += err * err;
                }
            }
        }
    }

    const char* names[4] = {
        "INT8, one scale for B",
        "INT8, one scale per column",
        "FP8 e4m3, one scale for B",
        "FP8 e4m3, one scale per 32",
    };
    std::printf(
        "storage census on B (%zux%zu), errors as a fraction of "
        "max|B| = %.4f\n",
        k, n, bmax);
    std::printf("  %-30s %14s %14s\n", "format", "max err", "rms err");
    for (int r = 0; r < 4; ++r) {
        std::printf("  %-30s %14.3e %14.3e\n", names[r], maxErr[r] / bmax,
                    std::sqrt(sumSq[r] / static_cast<double>(elems)) / bmax);
    }
    std::printf("\n");
}

// Times a launch with CUDA events and returns the mean milliseconds per run.
//
// Copy it verbatim; the alternative is four copies 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;
}

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\n", prop.name, prop.major,
                prop.minor);
    printCapabilityReport(prop.major, prop.minor);
    std::printf("\n");

    const size_t maxElems = kMaxDim * kMaxDim;

    std::vector<float> h_a(maxElems);
    std::vector<float> h_b(maxElems);
    std::vector<float> h_bPruned(maxElems);
    std::vector<int8_t> h_aqSym(maxElems);
    std::vector<int8_t> h_aqAsym(maxElems);
    std::vector<int8_t> h_bqTensor(maxElems);
    std::vector<int8_t> h_bqColumn(maxElems);
    std::vector<int8_t> h_bqPruned(maxElems);
    std::vector<float> h_bScales(kMaxDim);
    std::vector<float> h_bColMax(kMaxDim);
    std::vector<float> h_bPrunedScales(kMaxDim);
    std::vector<float> h_bPrunedColMax(kMaxDim);
    std::vector<int> h_bColSums(kMaxDim);
    std::vector<double> h_colGateAtol(kMaxDim);
    std::vector<double> h_colStatAtol(kMaxDim);
    std::vector<float> h_c(maxElems);
    std::vector<double> h_want(maxElems);

    float* d_a = nullptr;
    float* d_b = nullptr;
    int8_t* d_aqSym = nullptr;
    int8_t* d_aqAsym = nullptr;
    int8_t* d_bqTensor = nullptr;
    int8_t* d_bqColumn = nullptr;
    int8_t* d_bqPruned = nullptr;
    float* d_bScales = nullptr;
    float* d_bPrunedScales = nullptr;
    int* d_bColSums = nullptr;
    float* d_c = nullptr;
    CUDA_CHECK(cudaMalloc(&d_a, maxElems * sizeof(float)));
    CUDA_CHECK(cudaMalloc(&d_b, maxElems * sizeof(float)));
    CUDA_CHECK(cudaMalloc(&d_aqSym, maxElems));
    CUDA_CHECK(cudaMalloc(&d_aqAsym, maxElems));
    CUDA_CHECK(cudaMalloc(&d_bqTensor, maxElems));
    CUDA_CHECK(cudaMalloc(&d_bqColumn, maxElems));
    CUDA_CHECK(cudaMalloc(&d_bqPruned, maxElems));
    CUDA_CHECK(cudaMalloc(&d_bScales, kMaxDim * sizeof(float)));
    CUDA_CHECK(cudaMalloc(&d_bPrunedScales, kMaxDim * sizeof(float)));
    CUDA_CHECK(cudaMalloc(&d_bColSums, kMaxDim * sizeof(int)));
    CUDA_CHECK(cudaMalloc(&d_c, maxElems * sizeof(float)));

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

    for (int caseIdx = 0; caseIdx < kNumCases; ++caseIdx) {
        const size_t size = kCaseSize[caseIdx];
        const size_t elems = size * size;
        const size_t outBytes = elems * sizeof(float);

        makeInputs(&h_a, &h_b, size, size, size);
        prune24(h_b.data(), h_bPruned.data(), size, size);

        const float aScaleSym =
            quantizeSymmetric(h_a.data(), h_aqSym.data(), elems);
        int aZero = 0;
        const float aScaleAsym =
            quantizeAsymmetric(h_a.data(), h_aqAsym.data(), elems, &aZero);
        const float bScaleTensor =
            quantizeSymmetric(h_b.data(), h_bqTensor.data(), elems);
        quantizeColumns(h_b.data(), h_bqColumn.data(), h_bScales.data(),
                        h_bColMax.data(), size, size);
        quantizeColumns(h_bPruned.data(), h_bqPruned.data(),
                        h_bPrunedScales.data(), h_bPrunedColMax.data(), size,
                        size);
        columnSums(h_bqColumn.data(), h_bColSums.data(), size, size);

        CUDA_CHECK(cudaMemcpy(d_a, h_a.data(), elems * sizeof(float),
                              cudaMemcpyHostToDevice));
        CUDA_CHECK(cudaMemcpy(d_b, h_b.data(), elems * sizeof(float),
                              cudaMemcpyHostToDevice));
        CUDA_CHECK(
            cudaMemcpy(d_aqSym, h_aqSym.data(), elems, cudaMemcpyHostToDevice));
        CUDA_CHECK(cudaMemcpy(d_aqAsym, h_aqAsym.data(), elems,
                              cudaMemcpyHostToDevice));
        CUDA_CHECK(cudaMemcpy(d_bqTensor, h_bqTensor.data(), elems,
                              cudaMemcpyHostToDevice));
        CUDA_CHECK(cudaMemcpy(d_bqColumn, h_bqColumn.data(), elems,
                              cudaMemcpyHostToDevice));
        CUDA_CHECK(cudaMemcpy(d_bqPruned, h_bqPruned.data(), elems,
                              cudaMemcpyHostToDevice));
        CUDA_CHECK(cudaMemcpy(d_bScales, h_bScales.data(), size * sizeof(float),
                              cudaMemcpyHostToDevice));
        CUDA_CHECK(cudaMemcpy(d_bPrunedScales, h_bPrunedScales.data(),
                              size * sizeof(float), cudaMemcpyHostToDevice));
        CUDA_CHECK(cudaMemcpy(d_bColSums, h_bColSums.data(), size * sizeof(int),
                              cudaMemcpyHostToDevice));

        double amax = 0.0;
        for (size_t i = 0; i < elems; ++i) {
            if (h_a[i] > amax) {
                amax = h_a[i];
            }
        }
        double bmax = 0.0;
        for (size_t col = 0; col < size; ++col) {
            if (h_bColMax[col] > bmax) {
                bmax = h_bColMax[col];
            }
            h_colGateAtol[col] =
                columnGateAtol(amax, h_bColMax[col], size);
            h_colStatAtol[col] =
                columnStatAtol(amax, h_bColMax[col], size);
        }

        // The censuses run once, on the largest case, and before the double
        // reference so a reader watching the transcript sees them without
        // waiting for a few billion host multiply-adds.
        if (caseIdx == kNumCases - 1) {
            double symRms = 0.0;
            double asymRms = 0.0;
            printActivationCensus(h_a.data(), elems, amax, &symRms, &asymRms);
            printStorageCensus(h_b.data(), h_bColMax.data(), size, size, bmax);
            // The zero point's claim as a branch, made where it is a
            // property of the format rather than of this weight matrix: on a
            // tensor that never goes negative, spending 256 codes instead of
            // 128 has to halve the round-trip error. 1.5 is the margin, well
            // under the 2 the arithmetic gives and well above noise, since
            // both numbers are deterministic host arithmetic.
            if (symRms < 1.5 * asymRms) {
                comparativeGateFailed = 1;
                std::fprintf(stderr,
                             "the zero point failed to buy a bit on a "
                             "non-negative tensor: symmetric rms %.3e "
                             "against asymmetric %.3e\n",
                             symRms, asymRms);
                break;
            }
        }

        matmulCpu(h_a.data(), h_b.data(), h_want.data(), size, size, size);

        const double refMax = maxAbs(h_want.data(), elems);
        const double f32Rtol = kScaledRtol(kF32Rtol, kF32MantissaBits, size);
        const double f32Atol = f32Rtol * refMax;

        std::printf("case %zux%zux%zu, max |reference| = %.4f\n", size, size,
                    size, refMax);
        std::printf("  a scale sym %.3e, a scale asym %.3e, zero point %d\n",
                    aScaleSym, aScaleAsym, aZero);
        std::printf(
            "  b scale (one for all) %.3e, per-column scales span "
            "%.3e to %.3e\n",
            bScaleTensor, h_bScales[0], h_bScales[kColumnGainCount - 1]);
        std::printf("  f32 gate: |got-ref| <= %.3e + %.3e*|ref|\n", f32Atol,
                    f32Rtol);
        std::printf(
            "  int8 gate is per column: |got-ref| <= 4*K*amax*"
            "max|B[:,c]|/127\n");
        std::printf(
            "  sqrt(K) fraction is reported as a statistical diagnostic, "
            "not gated\n");

        const dim3 block(kTileDim, kTileDim);
        const dim3 grid(static_cast<unsigned int>(
                            ceilDiv(size, static_cast<size_t>(kTileDim))),
                        static_cast<unsigned int>(
                            ceilDiv(size, static_cast<size_t>(kTileDim))));

        struct Path {
            const char* name;
            int gated;
        };
        // The per-tensor row is reported and never gated. It is predicted to
        // fail the per-column gate, and a row whose failure is the lesson
        // must not stop the program; day 62 does the same with its racy
        // kernels.
        const Path paths[4] = {
            {"matmulTiledF32", 1},
            {"int8 one scale", 0},
            {"int8 per-column", 1},
            {"int8 per-col + zp", 1},
        };
        double absErr[4] = {0.0, 0.0, 0.0, 0.0};
        double rmsErr[4] = {0.0, 0.0, 0.0, 0.0};
        double gateFrac[4] = {0.0, 0.0, 0.0, 0.0};
        double statFrac[4] = {0.0, 0.0, 0.0, 0.0};
        float meanMs[4] = {0.0f, 0.0f, 0.0f, 0.0f};

        std::printf("  %-18s %11s %11s %9s %9s %9s %7s\n", "kernel",
                    "max err", "rms err", "gate frac", "sqrtK frac",
                    "time (ms)", "vs f32");
        for (int pathIdx = 0; pathIdx < 4 && badKernel == nullptr; ++pathIdx) {
            // Cleared before each correctness launch so a kernel that skips
            // elements cannot inherit a right answer.
            CUDA_CHECK(cudaMemset(d_c, 0, outBytes));
            switch (pathIdx) {
                case 0:
                    matmulTiledF32<<<grid, block>>>(d_a, d_b, d_c, size, size,
                                                    size);
                    break;
                case 1:
                    matmulInt8<false, false><<<grid, block>>>(
                        d_aqSym, d_bqTensor, d_bScales, d_bColSums, aScaleSym,
                        bScaleTensor, 0, d_c, size, size, size);
                    break;
                case 2:
                    matmulInt8<true, false><<<grid, block>>>(
                        d_aqSym, d_bqColumn, d_bScales, d_bColSums, aScaleSym,
                        bScaleTensor, 0, d_c, size, size, size);
                    break;
                default:
                    matmulInt8<true, true><<<grid, block>>>(
                        d_aqAsym, d_bqColumn, d_bScales, d_bColSums, aScaleAsym,
                        bScaleTensor, aZero, d_c, size, size, size);
                    break;
            }
            CUDA_CHECK(cudaGetLastError());
            CUDA_CHECK(cudaDeviceSynchronize());
            CUDA_CHECK(
                cudaMemcpy(h_c.data(), d_c, outBytes, cudaMemcpyDeviceToHost));

            size_t firstBad = elems;
            absErr[pathIdx] = maxAbsError(h_c.data(), h_want.data(), elems);
            rmsErr[pathIdx] = rmsError(h_c.data(), h_want.data(), elems);
            gateFrac[pathIdx] = worstColumnFraction(
                h_c.data(), h_want.data(), h_colGateAtol.data(), size, size,
                &firstBad);
            size_t ignoredStatBad = elems;
            statFrac[pathIdx] = worstColumnFraction(
                h_c.data(), h_want.data(), h_colStatAtol.data(), size, size,
                &ignoredStatBad);
            if (pathIdx == 0) {
                size_t firstRelBad = elems;
                const double relFrac =
                    worstRelFraction(h_c.data(), h_want.data(), elems, f32Rtol,
                                     f32Atol, &firstRelBad);
                if (firstRelBad != elems || !(relFrac == relFrac)) {
                    badCase = caseIdx;
                    badKernel = paths[0].name;
                    badIndex = firstRelBad;
                    break;
                }
            } else if (paths[pathIdx].gated != 0 && firstBad != elems) {
                badCase = caseIdx;
                badKernel = paths[pathIdx].name;
                badIndex = firstBad;
                break;
            }

            switch (pathIdx) {
                case 0:
                    meanMs[0] = timeKernel([&] {
                        matmulTiledF32<<<grid, block>>>(d_a, d_b, d_c, size,
                                                        size, size);
                    });
                    break;
                case 1:
                    meanMs[1] = timeKernel([&] {
                        matmulInt8<false, false><<<grid, block>>>(
                            d_aqSym, d_bqTensor, d_bScales, d_bColSums,
                            aScaleSym, bScaleTensor, 0, d_c, size, size, size);
                    });
                    break;
                case 2:
                    meanMs[2] = timeKernel([&] {
                        matmulInt8<true, false><<<grid, block>>>(
                            d_aqSym, d_bqColumn, d_bScales, d_bColSums,
                            aScaleSym, bScaleTensor, 0, d_c, size, size, size);
                    });
                    break;
                default:
                    meanMs[3] = timeKernel([&] {
                        matmulInt8<true, true><<<grid, block>>>(
                            d_aqAsym, d_bqColumn, d_bScales, d_bColSums,
                            aScaleAsym, bScaleTensor, aZero, d_c, size, size,
                            size);
                    });
                    break;
            }
            std::printf("  %-18s %11.3e %11.3e %9.3f %9.3f %9.4f %7.2f\n",
                        paths[pathIdx].name, absErr[pathIdx], rmsErr[pathIdx],
                        gateFrac[pathIdx], statFrac[pathIdx], meanMs[pathIdx],
                        static_cast<double>(meanMs[0]) / meanMs[pathIdx]);
        }
        if (badKernel != nullptr) {
            break;
        }

        // The 2:4 row: the same per-column INT8 kernel over a B that has had
        // half its weights zeroed in the fixed pattern, judged against the
        // DENSE reference. It answers what the pattern costs in accuracy.
        // What it does not answer is what the pattern buys in time, because
        // this kernel reads every zero it was handed; see the README.
        CUDA_CHECK(cudaMemset(d_c, 0, outBytes));
        matmulInt8<true, false><<<grid, block>>>(
            d_aqSym, d_bqPruned, d_bPrunedScales, d_bColSums, aScaleSym,
            bScaleTensor, 0, d_c, size, size, size);
        CUDA_CHECK(cudaGetLastError());
        CUDA_CHECK(cudaDeviceSynchronize());
        CUDA_CHECK(
            cudaMemcpy(h_c.data(), d_c, outBytes, cudaMemcpyDeviceToHost));
        const double sparseErr = maxAbsError(h_c.data(), h_want.data(), elems);
        const double sparseRms = rmsError(h_c.data(), h_want.data(), elems);
        std::printf("  %-18s %11.3e %11.3e %9s %9s %7s\n", "int8 2:4 pruned",
                    sparseErr, sparseRms, "n/a", "not timed", "n/a");

        // The day's two claims as branches. Per-column scaling must beat one
        // scale for the whole matrix on the per-column gate, and the zero
        // point must beat symmetric on a non-negative activation. If either
        // is false on this hardware, the run says so instead of the page.
        if (gateFrac[2] >= gateFrac[1]) {
            comparativeGateFailed = 1;
            std::fprintf(stderr,
                         "case %zu: per-column worst gate fraction %.3f is "
                         "not better than per-tensor %.3f\n",
                         size, gateFrac[2], gateFrac[1]);
            break;
        }
        // The zero point is reported end to end and not gated here. It is
        // worth a factor of two on A, which the activation census proves,
        // and A is not what dominates this product: B carries a 64x column
        // range and an outlier every 37th element, so most of the error in
        // the rows above was spent on the weights. A page that gated this
        // ratio would be gating a property of the test matrix.
        std::printf(
            "  per-column cuts the worst gate fraction by %.2fx; its "
            "sqrt(K) diagnostic improves by %.2fx; the zero point moves "
            "end-to-end rms by %.3fx\n\n",
            gateFrac[1] / gateFrac[2], statFrac[1] / statFrac[2],
            rmsErr[2] / rmsErr[3]);
    }

    CUDA_CHECK(cudaFree(d_a));
    CUDA_CHECK(cudaFree(d_b));
    CUDA_CHECK(cudaFree(d_aqSym));
    CUDA_CHECK(cudaFree(d_aqAsym));
    CUDA_CHECK(cudaFree(d_bqTensor));
    CUDA_CHECK(cudaFree(d_bqColumn));
    CUDA_CHECK(cudaFree(d_bqPruned));
    CUDA_CHECK(cudaFree(d_bScales));
    CUDA_CHECK(cudaFree(d_bPrunedScales));
    CUDA_CHECK(cudaFree(d_bColSums));
    CUDA_CHECK(cudaFree(d_c));

    if (badKernel != nullptr) {
        const size_t n = kCaseSize[badCase];
        std::fprintf(stderr,
                     "%s failed its gate at %zu (row %zu, col %zu) on case "
                     "%zux%zux%zu\n",
                     badKernel, badIndex, badIndex / n, badIndex % n, n, n, n);
        return EXIT_FAILURE;
    }
    if (comparativeGateFailed != 0) {
        return EXIT_FAILURE;
    }

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