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.
rmsNormRowmust not be more than 5 percent slower thanlayerNormRowTwoPassor the run returnsEXIT_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.
- 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.
- 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.
- 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.
- 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
- CUDA C++ Programming Guide, "Warp Shuffle Functions": https://docs.nvidia.com/cuda/cuda-programming-guide/05-appendices/cpp-language-extensions.html (checked 2026-09-01)
torch.nn.LayerNorm, the biased variance and the 1e-5 default: https://docs.pytorch.org/docs/2.13/generated/torch.nn.LayerNorm.html (checked 2026-09-01)torch.nn.RMSNorm, eps inside the square root and theeps=Nonedefault: https://docs.pytorch.org/docs/2.13/generated/torch.nn.RMSNorm.html (checked 2026-09-01)- Zhang and Sennrich, "Root Mean Square Layer Normalization", the paper that dropped the mean: https://arxiv.org/abs/1910.07467
- PMPP 4th edition, chapter 10, on reduction
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.