Day 95Module 10
in-technical-review

LLM kernels 1: softmax, layer norm, RMS norm

A row of attention scores arrives from a matmul with its largest entry near 120. The softmax over it is the formula from the paper, written the obvious way:

partial += expf(row[c]);          // sum the exponentials
outRow[c] = expf(row[c]) / total; // then divide

expf(120.0f) is inf. The sum is inf. Every probability in the row comes back nan, including the ones whose logits were perfectly ordinary. The arithmetic is right and the kernel is unusable.

Ninety-four days have built the pieces that fix it. Today they assemble: the three normalization kernels that run on every token of every transformer, all three built from day 23's shuffle and day 24's reduction, all three limited by the same thing, and one of them cheaper for a reason you can count before you measure.

One block, one row, one number

All three take a row of activations, reduce it to one or two scalars, and scale every element of that row by them. The row is the unit of work, so the block is the unit of parallelism: one block owns one row start to finish, and no block talks to another.

Inside the block it is day 24's ladder with the top two rungs kept. Each of the 256 threads strides across the row accumulating its own partial, so at 1024 columns a thread holds four elements and the 32 lanes of a warp read 32 consecutive floats, which coalesces. Five __shfl_down_sync steps collapse each warp to one value, the eight warp totals go through shared memory, and every thread reads all eight back and folds them.

That last step is the one a reduction lesson does not need. Day 24 finished when lane 0 of block 0 held the answer; here all 256 threads have elements waiting to be divided by it. The __syncthreads() before the read makes the shared array safe to write, and the one after it makes the array safe to reuse for a second reduction, which softmax and two-pass layer norm both need. Drop that second barrier and the max pass and the sum pass share eight slots without agreeing on when.

Diagram: 256 threads to one number, and back. One row drawn as a horizontal strip on top, 256 thread boxes beneath it, then a reduction tree, then an arrow fanning back out to all 256 boxes. Band 1, the strided read: each thread box points at 4 cells of a 1024-column row, its lane's four cells spaced 256 apart. Caption "4 elements per thread, 32 consecutive floats per warp." Band 2, the shuffle: 32 lanes of one warp collapsing over 5 levels to a single value, with no shared memory drawn at all. Caption "32 lanes to 1 in 5 steps, no barrier." Band 3, shared memory and the broadcast: 8 slots written by 8 lanes, then 256 arrows reading all 8 back. Caption "8 partials in, 1 total out to every thread." Alt text: "A row reduced and broadcast back. Four elements per thread, then five shuffle steps take 32 lanes to one, then eight warp partials become one number that all 256 threads read."

The subtraction that is not an optimisation

Subtracting the row maximum looks like a trick and courses present it as one. It changes no mathematics: the factor exp(-max) cancels exactly between numerator and denominator, so on a row that would have worked anyway it buys nothing. What it buys is the right to hold a logit above 88.7, which is where expf overflows a float and turns the denominator into inf.

After the shift every exponent is at or below zero, so the top of the range is safe by construction. The bottom is where the fix stops being free, and day 47 measured that edge: expf returns a subnormal below about -87.34 and exactly zero below about -103.28, which is 149 times the natural log of 2. Its census of one softmax found 17 normal results, 832 subnormal and 175 exactly zero, and -ftz=true moved all 832 subnormals to zero. The shift stops the top of the row from overflowing and lets the bottom underflow, quietly, in terms that were going to contribute nothing. Part 1 of this program prints the same census for a row spanning 200 in logit space.

The naive kernel stays in the program, gated, and it passes on the ordinary inputs of part 2 because its arithmetic is correct. Day 25 shipped a racy kernel that produced the right answer on all 100 runs and was still wrong. A test that never feeds a kernel the input that breaks it has not tested the kernel.

Counting passes instead of flops

Full program in norms.cu, under code/day95-norms/. Six kernels, two cases, and three properties that decide what the timing table means.

Every kernel is priced at eight bytes per element. One read and one write per element is all any of them owes the memory system. The extra sweeps, the max pass and the variance pass, re-read a row that may still be in L1 or L2 rather than out in global memory, so a kernel whose re-reads hit cache reports a figure near the floor and one whose re-reads miss reports a figure well under it. The floor is copyFloor, measured in this process so the block shape, the buffer size and the clock state are shared.

Both cases move the same bytes. 16,384 rows of 1,024 and 4,096 rows of 4,096 are each 64 MiB in and 64 MiB out, sixteen times the T4's 4 MiB L2. Only the row length changes, so the table compares reduction depth rather than problem size.

The reference is double precision on the host, and the tolerance grows with the row rather than being fixed, because every output depends on a sum over the whole row:

static double kScaledRtol(double tableRtol, int mantissaBits, size_t kDim) {
    const double eps = std::ldexp(1.0, -mantissaBits);
    const double grown = 4.0 * eps * std::sqrt(static_cast<double>(kDim));
    return grown > tableRtol ? grown : tableRtol;
}

Here is the safe softmax, and the passes are countable in it: max, sum, write, with the row read in all three.

__global__ void softmaxRowSafe(const float* __restrict__ in,
                               float* __restrict__ out, size_t cols) {
    __shared__ float smem[kWarpsPerBlock];
    const float* row = in + blockIdx.x * cols;
    float* outRow = out + blockIdx.x * cols;

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

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

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

Layer norm reads the row three times for a different reason: the definition is two sweeps, mean and then variance around that mean. Welford's method carries a running mean and a running sum of squared deviations from it, so the variance falls out of the sweep that produced the mean. The per-thread loop and the warp merge are the whole idea:

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

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

RMS norm gets there by deleting the mean instead of computing it faster. Nothing to centre, so one reduction over the squares is the whole statistic and gamma folds into the write: three row touches, one reduction, no beta. That is a bandwidth argument for RMS norm rather than a modelling one.

Note. The program will not let that claim slide. rmsNormRow must not be more than 5 percent slower than layerNormRowTwoPass or the run returns EXIT_FAILURE. If deleting a read sweep buys nothing on this card, the reason given here is not the real one.

Results

Re-verified on the same Tesla T4 with driver 580.173.02 and CUDA 13.0 (V13.0.88). Correctness, error values and performance conclusions reproduced; timing drift stayed tiny. Both transcripts are listed in front matter; the table below is CUDA 13.0. The pinned PyTorch documentation elsewhere on this page remains a framework reference, not a toolkit stamp.

case kernel touches ms ask GB/s % floor
16384 x 1024 copyFloor 2 0.5560 241.4 100.0
16384 x 1024 softmaxRowSafe 4 0.9336 143.8 59.5
16384 x 1024 layerNormRowTwoPass 4 0.8811 152.3 63.1
16384 x 1024 layerNormRowWelford 3 1.1351 118.2 49.0
16384 x 1024 rmsNormRow 3 0.7475 179.6 74.4
4096 x 4096 copyFloor 2 0.5562 241.3 100.0
4096 x 4096 softmaxRowSafe 4 1.1664 115.1 47.7
4096 x 4096 layerNormRowTwoPass 4 1.0320 130.1 53.9
4096 x 4096 layerNormRowWelford 3 0.9196 145.9 60.5
4096 x 4096 rmsNormRow 3 0.8244 162.8 67.5

softmaxRowNaive measured 159.8 and 146.4 ask GB/s at the two shapes. The four predictions resolved as follows.

  1. Refuted. At 1024 columns only RMS norm cleared 70 percent of the floor; at 4096 none did. Safe softmax again reached 59.5 and 47.7 percent. The simple memory-bound model did not predict how much the extra reduction sweeps cost.
  2. Partly held, refuted overall. RMS norm beat two-pass layer norm at both shapes. Welford lost at 1024 columns, 1.1351 against 0.8811 ms, then beat it at 4096. RMS and Welford did not tie.
  3. Held. Naive softmax produced 179 non-finite outputs on the deliberate overflow row, safe softmax produced zero and summed to 0.999999982. All normal-case correctness rows passed.
  4. Held. The 241.3 to 241.4 GB/s copy floor remained close to day 49's 244.7 GB/s.

The RMS-versus-two-pass performance prediction printed held at both shapes, but the wider touch-count model did not explain Welford or softmax.

Run it yourself

A free Colab session is enough. Nothing on this page is above sm_75, there is no library to link and no profiler to install:

nvcc -std=c++17 -O3 -arch=sm_75 -o norms norms.cu

No Compiler Explorer embed for this one. The double-precision host references sweep 16.7 million elements five times per case, which is tens of seconds of CPU work, and shrinking the case to fit the 20 second cap would put both cases inside L2 and delete the measurement.

Exercise

Write the fused RMS norm yourself. Start from layerNormRowTwoPass, cut it to one reduction over the squares, fold gamma into the write, and check it against PyTorch as well as the shipped double reference.

Time: 30 to 45 minutes. Submit: your kernel, its gate fraction, its ask GB/s beside layerNormRowTwoPass's, and the largest difference between your part 0 row and PyTorch's.

Check: two gates. The harness gate is the one the program already prints, |got-ref| <= atol + rtol*|ref| with rtol = max(1e-5, 4 x 2^-23 x sqrt(C)) over a double reference, and a non-finite output fails outright with its row and column named. The measurement gate is a ratio inside your own run: your kernel's ask GB/s must be at least layerNormRowTwoPass's, because it touches the row one fewer time and a version that is not faster kept a sweep it did not need. For the PyTorch half, run this on the eight values part 0 prints:

import torch
x = torch.tensor([-1.25, -0.75, -0.25, 0.25, 0.75, 1.25, 1.75, 2.25])
eps = 1e-5
print(x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + eps))
Hint 1

Layer norm needs two statistics and RMS norm needs one. Which of the two sweeps in layerNormRowTwoPass computes something RMS norm never refers to, and what happens to the other sweep once you stop subtracting anything?

Hint 2

Look at where eps goes. torch.nn.LayerNorm defaults to 1e-5, but torch.nn.RMSNorm defaults to eps=None, meaning the machine epsilon of the compute type, 1.1920929e-07 for fp32 input (checked 2026-09-01 at https://docs.pytorch.org/docs/2.13/generated/torch.nn.RMSNorm.html ). If you and PyTorch disagree in the fourth digit and nowhere else, that is two different epsilons.

Solution

rmsNormRow in norms.cu is the reference version: one strided sweep accumulating x * x, one blockReduceSum, rsqrtf(meanSq + kEps), and a write loop multiplying by the scale and by gamma[c]. The mean is gone, so beta goes with it and the row is touched three times instead of four.

Count row touches before writing the kernel and you can say which of two normalization layers is cheaper without running either. Sweeps over the row are what these kernels spend, not arithmetic.

Pitfalls

Your softmax returns nan and only on some rows. The row's maximum crossed 88.7 and expf overflowed. Subtract the row max before the exponential, always, not only when a test fails. It costs one extra sweep, which the caches usually absorb.

Your one-pass layer norm computes the variance as E[x^2] - E[x]^2. That form is one reduction and it is the wrong one. When the mean is large relative to the spread the two terms are nearly equal, the subtraction cancels most of the significant digits, and the variance comes out too small or negative with rsqrtf returning a nonsense scale. Welford costs a divide per element and does not cancel.

You reused the shared array for the second reduction without a barrier. The max pass and the sum pass in a safe softmax both write eight slots. Without a __syncthreads() between the last read of the first reduction and the first write of the second, some warps read the new value and some the old, and the row is scaled by a number that never existed. Day 62's racecheck names this class of bug in one command.

You fused the whole block into one kernel and it got slower. Day 48 measured what fusion pays: a four-kernel chain won 3.08x on a 3.00x byte cut, so the win was the intermediate buffers. A normalization layer has none left to delete.

You compiled with -use_fast_math and the tolerance stopped holding. The flag brings -ftz=true with it, so the subnormal tail of a shifted softmax becomes exactly zero and the row sums differently. Fast math is a per-file decision and day 47 prices it one sub-flag at a time. Time the result with CUDA events and a warm-up per kernel, never a host clock: day 9 measured a round trip at 45.822 ms against a 0.786 ms kernel.

Go deeper

Next

Day 96 keeps the row and drops the reduction: GELU, RoPE and the bias add are elementwise, with no barrier and no broadcast, so the only question is how many fit in one sweep. Day 97 puts both halves together in attention, where the row is too long to hold and the online softmax finds the maximum and the sum in the single sweep this page said needed two.