Day 97Module 10
in-technical-review

LLM kernels 3: attention and online softmax

One attention head over 4,096 tokens takes 3 MiB of input and produces 1 MiB of output. The obvious way to compute it writes 64 MiB in between, and touches that buffer five times. Push the context to 131,072 tokens, which several shipping models accept, and the same buffer is 64 GiB for one head of one layer. No GPU has that, and those models run anyway.

So the matrix is not being stored. The kernel that skips it looks impossible at first, because softmax needs the largest element of a row before it can exponentiate anything, and you do not know the largest element until you have seen the row. This page derives the identity that gets around that, builds the kernel on top of it, and measures what the skipped traffic is actually worth in milliseconds, which is a different question with a less flattering answer.

What the matrix costs, and how often you touch it

One head is three matrices of shape N by d, the queries, keys and values. Here d is 64 and N is the sequence length. Attention is softmax(Q K^T / sqrt(d)) V: score every query against every key, turn each row of scores into weights that sum to one, then mix the value vectors with those weights.

The straightforward version does that in three kernels, and the middle one forces the first to write its answer down. Q K^T produces an N by N matrix of scores, S. The softmax reads S to find each row's maximum, reads it again to write the exponentials, and the third kernel reads those back to weight V. That is five passes over N^2 floats: 20 N^2 bytes, against 16 N d bytes for the queries, keys, values and output put together. At N = 4096 those are 320 MiB and 4 MiB. The quadratic term is not a detail of the implementation, it is the implementation.

Both versions do the same arithmetic, about 4 N^2 d floating-point operations. Divide that by the bytes and the naive path sits at 12.8 FLOP per byte, which day 49 measured to be under this T4's ridge point of 29.34, so it is memory bound and its arithmetic intensity does not depend on N at all. A kernel that never writes S has an intensity of N/4 on the same roofline, which climbs with the sequence length instead of sitting still.

Softmax needs the whole row, and does not

The reason people believe attention has to materialise S is a fact about softmax, and the fact is true. Exponentiating raw scores overflows: a score of 100 gives exp(100), which is inf in FP32, so every implementation subtracts the row maximum first. Subtracting the maximum needs the maximum, which needs the row.

The way out is not to avoid the maximum, it is to fix it up. Suppose you have processed part of a row and you hold two numbers: m, the largest score seen so far, and d, the sum of exp(x - m) over those same scores. A new tile arrives whose largest score is bigger. Every term you have already accumulated was scaled by the old maximum, and converting it to the new one is one multiply, because the exponent is a sum:

// exp(x - mNew) == exp(x - mOld) * exp(mOld - mNew)
mNew  = max(mOld, tileMax);
scale = exp(mOld - mNew);         // one factor for the whole history
dNew  = dOld * scale + sum over the tile of exp(x - mNew);
oNew  = oOld * scale + (tile weights) . V_tile;

One factor rescales the running denominator and the running output together, so the state a query row carries is m, d and its d-wide output, and never a row of scores. Milakov and Gimelshein published this as the online normalizer calculation (https://arxiv.org/abs/1805.02867 , checked 2026-09-01); FlashAttention applies it tile by tile with the value mix carried along (https://arxiv.org/abs/2205.14135 , checked 2026-09-01). The widget above steps a sixteen-element row in four tiles with the largest value arriving last, which is the case intuition gets wrong.

Two paths, one head, one reference

Full program in attention.cu, under code/day97-attention/. It runs both paths at N of 256, 1024 and 4096 and prints the bytes before it prints any time.

The baseline is not a straw man. Its two matmuls use day 16's 16 by 16 tile, with the key tile held transposed in shared memory so the global reads stay coalesced. The only thing the naive path is charged for on this page is the matrix it writes to global memory.

The reference is double precision on the host, holding one row of scores rather than the matrix. The two GPU paths sum in different orders, so an FP32 reference could not arbitrate between them. The gate is day 66's K-scaled rule with the depth set to N, since every output element is a sum over N keys, and the program prints the arithmetic beside it.

One key is deliberately loud. Key N-3 sits in the final tile of every sweep and scores far above the rest, so every query row's running maximum is still moving on the last tile. Without it the rescale could fire once and never again, leaving this day's branch untested.

In the online kernel each block owns 64 query rows, each warp owns eight of them, and each lane owns one key of the current 32-key tile and two of the 64 output dimensions. The state that replaces the score matrix is 32 floats per lane:

    float acc[kRowsPerWarp][kDimsPerLane];
    float runMax[kRowsPerWarp];
    float runSum[kRowsPerWarp];
    for (int r = 0; r < kRowsPerWarp; ++r) {
        runMax[r] = -INFINITY;
        runSum[r] = 0.0f;
        for (int c = 0; c < kDimsPerLane; ++c) {
            acc[r][c] = 0.0f;
        }
    }

And here is one tile's worth of the identity, for one query row. The two butterfly reductions are day 23's shuffles, which leave the tile maximum and the tile sum in every lane without a barrier or a scratch buffer:

            float tileMax = score;
            for (int off = kWarpSize / 2; off > 0; off >>= 1) {
                tileMax =
                    fmaxf(tileMax, __shfl_xor_sync(0xffffffffu, tileMax, off));
            }
            const float newMax = fmaxf(runMax[r], tileMax);
            // exp(-inf - newMax) is 0, so the first tile's rescale wipes the
            // initialised state instead of special-casing it.
            const float rescale = expf(runMax[r] - newMax);
            const float weight = expf(score - newMax);

            float tileSum = weight;
            for (int off = kWarpSize / 2; off > 0; off >>= 1) {
                tileSum += __shfl_xor_sync(0xffffffffu, tileSum, off);
            }
            runSum[r] = runSum[r] * rescale + tileSum;
            runMax[r] = newMax;

            for (int c = 0; c < kDimsPerLane; ++c) {
                acc[r][c] *= rescale;
            }
            for (int j = 0; j < kKeyTile; ++j) {
                const float w = __shfl_sync(0xffffffffu, weight, j);
                for (int c = 0; c < kDimsPerLane; ++c) {
                    acc[r][c] += w * tileV[j][c * kWarpSize + lane];
                }
            }

The caveat is the query tile. Sixty-four rows per block means keys and values are re-read once per block, so the kernel's request count carries an N^2 term even though its allocation does not. That ask is printed next to the floor, and whether it reaches DRAM or stops at L2 is a cache question this program cannot answer. A production kernel raises the query tile and hands both matmuls to the tensor cores, which day 72 introduces and this page does not use.

Results

Re-verified on the same Tesla T4 with driver 580.173.02 and CUDA 13.0 (V13.0.88). Errors, resource rows and traffic arithmetic reproduced exactly; all cases passed and the process exited 0. At 4096, online improved from 6.7860 to 5.2910 ms while naive moved from 7.4483 to 8.3724 ms, widening the online win. Both transcripts are listed in front matter; this table is CUDA 13.0.

n path max abs err gate frac time (ms) naive/online
256 naive, 3 kernels 5.941e-09 0.0614 0.0745 0.211
256 online, 1 kernel 5.678e-09 0.0614 0.3540 0.211
1024 naive, 3 kernels 3.023e-09 0.0553 0.8413 0.576
1024 online, 1 kernel 3.372e-09 0.0553 1.4604 0.576
4096 naive, 3 kernels 1.453e-09 0.0578 8.3724 1.582
4096 online, 1 kernel 1.310e-09 0.0550 5.2910 1.582

At n=4096 the arithmetic traffic ratio was 81.0x. The five predictions all held:

  1. Held. Every gate passed, while the online-versus-naive worst gate fractions remained nonzero: 0.0669, 0.0596 and 0.0614.
  2. Held. Online maximum error stayed within 1.12x of naive at every size and was smaller at 256 and 4096.
  3. Held. attentionOnline used 32,896 shared bytes, 61 registers, one block per SM and 25 percent occupancy. Each naive kernel reported four blocks and 100 percent.
  4. Held. At 256, online took 0.3540 ms against naive's 0.0745 ms; naive/online was 0.211.
  5. Held. The ratio rose strictly, 0.211 to 0.576 to 1.582, crossing 1 only at 4096.

The byte win existed at every size; the time win appeared only at 4096. That is the measured distinction between deleting storage and accelerating the finite test case.

Run it yourself

Anything at sm_75 or above builds and runs this, including a free Colab session:

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

No profiler, no root: the occupancy table comes from cudaFuncGetAttributes and cudaOccupancyMaxActiveBlocksPerMultiprocessor, neither of which needs performance counters. No Compiler Explorer embed, because the double-precision reference at 4,096 tokens is billions of multiply-adds on one core and blows the 20 second execution cap by itself. Device memory peaks near 140 MB, the naive path's two square buffers.

Exercise

Write the online kernel yourself. Keep the head width at 64 and the key tile at 32, and produce one kernel that takes Q, K and V and writes O with no N by N allocation anywhere.

Time: 60 to 90 minutes. Submit: your kernel, its worst gate fraction against the shipped naive path at n = 1024, and the size at which it starts beating that path in wall time on your card.

Check: your kernel must pass the two gates the shipped program prints, against the double-precision host reference and against the naive path's output, both with rtol = max(1e-5, 4*2^-23*sqrt(n)). A failure names the query row and head dimension of the first offending element. Two mistakes cause almost all of them: forgetting to rescale the accumulated output when the maximum moves, which is right only when the largest score lands in the first tile, and normalising per tile instead of dividing once at the end, which gives rows that do not sum to one.

Hint 1

You cannot store a row of scores, so ask what a query row carries between tiles instead. Three things, one of them as wide as the head. What has to happen to all three when a later tile holds a bigger score than anything before it?

Hint 2

Write exp(x - mNew) in terms of exp(x - mOld). The exponent is a difference, so the conversion is a single multiplicative factor that does not depend on x. Which of your carried values does that factor apply to, and which one does not need it because it is being replaced?

Solution

attentionOnline in attention.cu is the worked version. Per tile: reduce the tile's scores to a maximum, take newMax = max(runMax, tileMax), compute rescale = exp(runMax - newMax), multiply the running denominator and every element of the running output by it, then add this tile's contribution against newMax. Divide once at the end. Initialising the maximum to negative infinity makes the first rescale exactly zero, so the opening tile needs no special case.

Softmax's combining rule is associative: two partial results, each holding a maximum and a sum taken against it, merge by rescaling the smaller-maximum side. That is why attention tiles at all, and it is the property day 24 leaned on to turn a sum into a tree.

Pitfalls

Your kernel is right in testing and wrong in production. Test data whose largest score lands early never exercises the rescale, so a kernel that forgets it passes. Put the outlier in the last tile, the way this program does, and check that the failure appears.

You subtracted the maximum and still got inf. Then it was not the maximum of everything you exponentiated: the usual cause is updating the running sum before the running maximum, or weighting a tile against its own maximum instead of the running one.

Your denominator drifts at long sequence lengths. The rescale factor is below one and applies to the whole history, so the running sum is a product of many such factors. In FP32 with a large spread of scores it can underflow to zero, and a denominator of zero produces NaN rather than a wrong number. Gate on finiteness, which this program does, rather than on error alone. Day 47 is where the exponential's own accuracy gets priced.

You benchmarked against a naive path nobody would ship. An untiled Q K^T makes any comparison look good and measures the matmul rather than the materialisation. This page gives the baseline day 16's tile for that reason, and it still loses on bytes by the factor the arithmetic predicts.

You fused attention and expected the launch count to pay. Three kernels became one, and three launches are microseconds (day 48 priced an empty one at 2.668 us). What fusion bought here is the buffer, not the launches. Count the bytes before you attribute the result.

You padded every shared tile out of habit. The key tile here is padded by one float because lane j walks down key j, which without the pad puts all 32 lanes on one bank. The value tile is not, because lanes there read consecutive dimensions. Day 15 has the arithmetic that tells the two cases apart, and a bank conflict is worth ruling out before an occupancy change.

Go deeper

Next

Day 98 drops the precision this page kept in FP32, with INT8 on the T4's tensor cores and FP8 and 2:4 sparsity as stretch tiers. A tiny transformer forward pass is this day's stretch: combine the attention kernel with day 95's norms and day 96's elementwise chain, then check every stage before timing the whole block. Day 99 changes direction and teaches the backward shapes that day 100 needs. All three sit in module 10; time them with CUDA events and read occupancy from the runtime rather than guessing it.