53 · r50 — FA_TKV 32→16 occupancy experiment (REVERTED)
Result:
fa_prefill_f16kv's KV tile shrinks from 32 columns to 16, smem drops from 34.8 KB to ~33 KB, resident blocks 2 → 3 blocks/SM — parity green, suite 166/0/3, but the whole-prefill wall clock is neutral-to-negative (two rounds of interleaved A/B: 3-pair median −0.5%, 5-pair median −0.01%); more fatally, greedy-32 is no longer byte-identical (an argmax boundary flipped by ULP-level float regrouping). The occupancy gain is cancelled by the doubled per-tile sync/softmax overhead — tile shrinking on FA is a dead lever, and it inherently breaks byte identity. Commit:9128468(docs-only — the code change was reverted after measurement and never landed as a code commit). Date: 2026-09-06.
1. Background — where things stood
r48 (FAP2, register-resident softmax) was the heaviest blow the FA line has taken so far: S/P no longer pass through shared memory, softmax runs directly on the QK^T wmma accumulator fragments, and P becomes the A operand of P·V in registers, in place. The FA kernel dropped from 5.16 ms to 2.12 ms (2.43×), whole prefill 2603.5 → 2749.9 tok/s, and FA's share of the whole wall shrank from 10.2% to ~4.7%. r49 (A-quantize prepass shared-A dedup) took another +2.32%, pushing the baseline to 2797.5 tok/s and vs-llama to 1.18×.
r48's record left the FA line one explicit "next lever": after FAP2 the kernel's top stall became global K/V loads, and r48's ncu device report showed the resource limiting resident blocks is smem (34.82 KB per block → 2 blocks/SM; 124 registers only limits to 4 blocks). Cutting FA_TKV from 32 to 16 halves the K/V staging smem demand linearly, so resident blocks could theoretically reach 3 — less per-byte latency exposure, the classic occupancy logic.
This thread had appeared twice before, and both times it ended as "the kernel moved, the wall didn't":
| Round | Change | Kernel | Wall | Outcome |
|---|---|---|---|---|
| r23 (f16 era) | FA_TKV raised | occ 16.7 → 32.68%, kernel −6.7% | −0.3% | REVERTED |
| r46 (FAP1) | FA_TKV 64→32 + S/P row padding | 5.16 → 4.58 ms (−11%) | +0.27% | REVERTED |
| r50 (this doc) | FA_TKV 32→16 | — (not timed separately) | −0.5% / −0.01% | REVERTED |
Three rounds of same-direction evidence were already enough to constitute a mechanism judgment, but r50 was still worth doing, for two reasons: first, it is a one-#define experiment (r48 symbolized every geometric quantity, so the marginal cost was near zero); second, r46's veto reason was "FA is not yet the wall's critical path," whereas after r48 FA is down to 4.7% — "does occupancy finally pay off at a smaller share" had never been directly answered. The answer is no, and the way it was answered is worth more than the numbers: this round hit two boundaries at once — occupancy is ineffective on FA (the mechanism boundary), and the strict greedy byte-identity gate is unsatisfiable by any FA_TKV change (the verification boundary, which r57 would hit again).
2. Principle — the GPU mechanism
The kernel's shape after FAP2. fa_prefill_f16kv is a 4-warp (128-thread) full-row warp-tile: each warp owns 16 query rows × all FA_TKV KV columns. Each KV tile's lifecycle is:
fa_stage_kv_async: cp.async moves the K/V tile into smem (row stridesstr = hd+8, zero-filled out of range);cp.async.wait_group 0+__syncthreads();- QK^T: 8 wmma steps over the hd dimension, producing
fc[FA_TKV/16]16×16 f32 accumulator fragments; - second
__syncthreads(); - fragment-resident online softmax: per tile, a max reduction (2
__shfl_xor),__expf, a sum reduction (2__shfl_xor), the runningm/lupdate, and 8 O accumulator fragments × alpha rescale; - P converted to f16 in place as the A operand, P·V runs wmma, and a third
__syncthreads()closes the tile.
Why halving the tile does not speed it up. The key is splitting the cost per KV column into two parts:
- wmma arithmetic: the number of QK^T and P·V mmas per column is independent of FA_TKV (the total is fixed);
- per-tile fixed costs: 3
__syncthreads, cp.async commit/wait, the softmax shfl reduction chains, the O-fragment rescale — these are billed per tile.
FA_TKV 32→16 doubles the tile count, so the second class of cost doubles exactly. And the occupancy gain? Occupancy solves "latency hiding": more resident warps let DRAM latency be filled by other warps' issue. But after r48 the kernel's K/V traffic is not that large (7B @3325 tok: per layer KV = 2 × 3325 × 128 × 2 B × 4 kv-head ≈ 6.8 MB, 28 layers ≈ 191 MB), and 2 blocks/SM already hides the latency; a 3rd block only sends more warps to contend for the same L1TEX/DRAM bandwidth and slices each block's L2 working set smaller. Gain ≈ 0, cost = doubled fixed costs → net effect inside the noise band. This is exactly r46's lesson repeating verbatim: occupancy gain cancelled by doubled per-tile sync/softmax overhead.
The occupancy account (derived from r48's ncu device report). GB10 (CC 12.1) has 100 KB of smem available per SM:
| FA_TKV | dynamic smem per block | smem limit | register limit (124 regs) | actual residency |
|---|---|---|---|---|
| 32 (current tree) | (64+64)×136×2 = 34.8 KB | ⌊100/34.8⌋ = 2 | 4 | 2 |
| 16 (experiment) | max(formula, 32 KB) ≈ 33.0 KB | ⌊100/33.0⌋ = 3 | 4 | 3 |
The record notes: this 3 blocks/SM was derived from r48's device configuration — ncu/nsys could not be re-run in that window — but it is self-consistent (34.8→2 and 33.0→3 straddle the integer-division boundary exactly), and the wall-clock result proves that even with 3 blocks truly in hand it would not have been worth it.
Why byte identity is necessarily lost. Online softmax's math is independent of tile partitioning at infinite precision, but f32 is not infinite precision. The running state carries across tiles:
l0 = l0 * a0 + sum0; // once per tile: both l and O are recomputed under the current tile's grouping
acc[ob].x[j] *= aa0; // O accumulators rescaled by alpha, 8 fragments × 8 lanes per tile
Doubling the tile count regroups both the f32 multiply-add chain of l and the rescale chain of acc: the summation order of Σexp and the associativity boundaries of l·a + s all change → each number drifts by ~1 ULP. The vast majority of tokens have top-2 logit gaps far larger than a ULP and argmax is unaffected; but greedy sampling is pure argmax, so if any single token happens to stand on a knife edge where the top-2 gap is ~1e-7, the generation stream forks there and never looks back. The parity gate cannot see it: parity fixtures compare numeric error within tolerance (measured max err ~1e-4, the numbers themselves correct) and never do a whole-stream byte comparison. The only gate that can see it is the greedy byte-identity gate.
3. Implementation
Archival status (flagged up front): r50's code change was a one-line experiment in the working tree that disappeared with the revert once measurement was done — there is no code commit; what survives is only the docs commit 9128468 (the record + all numbers). This doc's "before" excerpts therefore come from the current tree (i.e., the post-revert FA_TKV=32 state, the code still running today); the "after" side has only the formulas and a one-line diff narration from the record — no code to excerpt.
3.1 Design choices (why this shape and not another)
The experiment was deliberately kept minimal: not one character changed except #define FA_TKV 32 → 16. r48's rewrite had already symbolized every geometric quantity in the kernel — the QK^T fragment array fc[FA_TKV/16], P·V's k-loop bound kk0 < FA_TKV, cp.async staging's count c < FA_TKV*hd/8, the softmax fragment arrays [FA_TKV/16*4], Vs = Ks + FA_TKV*sstr, grid/launch — so a one-line define is a complete re-parameterization. This is the implicit asset r48 left behind: symbolized geometry makes tile size an enumerable one-dimensional experiment.
The change immediately flushed out one non-symbolized constant that had slipped through: the tail tile's O write-back reuses smem as a 64×128 f32 stage buffer — FA_TQ * hd * 4 = 32768 B = 32 KB, independent of FA_TKV. At FA_TKV=32 the staging formula yields 34.8 KB ≥ 32 KB and the constant is masked; at 16 the formula yields 25.5 KB < 32 KB and the tail block overruns directly. So the launcher must take the max of the two:
smem = max((FA_TQ + 2*FA_TKV)*(hd+8)*2, // Qs + Ks + Vs staging
FA_TQ * hd * 4) // tail tile's O write-back stage (f32)
= 32 KB (at FA_TKV=16)
3.2 Key code
The symbolized geometry itself (current tree src/cuda_kernels.cu, r48's legacy — the r50 experiment only worked as a one-line change because of it):
// S and P live entirely in wmma accumulator fragments (FAP2 register-resident
// softmax — NO Sf/Pf shared round trip, no m/l/alpha shared arrays). Each warp
// owns a full 16-query-row block x all FA_TKV KV columns, so the online softmax
// (per-row max/sum on the fragments) and the P·V contraction (build the f16
// A-operand from the scaled fragments in place) are both warp-local.
#define FA_TQ 64
#define FA_TKV 32 // r50 experiment: 32 → 16 (one-line change, later reverted)
The tile loop's symbolized consumption of FA_TKV (current tree lines 4045–4058) — tile count = kv_end / FA_TKV, so halving the tile doubles the iteration count:
for (int kt = 0; kt < kv_end; kt += FA_TKV) {
// stage K/V tile (padded stride, zero-filled beyond kv_end)
fa_stage_kv_async(k, v, Ks, Vs, kt, kv_end, hk, hd, stride_kv, sstr, tid, 128);
...
__syncthreads(); // ← per-tile fixed cost #1
// S = Q · K^T via wmma (reduction over hd), this warp's full row block.
// fc[0] = kv cols [0,16), fc[1] = [16, FA_TKV) of this tile.
wmma::fragment<wmma::accumulator, 16, 16, 16, float> fc[FA_TKV / 16];
...
__syncthreads(); // ← per-tile fixed cost #2
The online softmax's per-tile reduction chains (current tree lines 4099–4123) — 2 max-shfl + 2 sum-shfl run once per tile, the bulk of the "billed per tile" fixed costs:
#pragma unroll
for (int off = 1; off <= 2; off <<= 1) {
mnew0 = fmaxf(mnew0, __shfl_xor_sync(0xffffffffu, mnew0, off)); // max tree
mnew1 = fmaxf(mnew1, __shfl_xor_sync(0xffffffffu, mnew1, off));
}
...
#pragma unroll
for (int off = 1; off <= 2; off <<= 1) {
sum0 += __shfl_xor_sync(0xffffffffu, sum0, off); // sum tree
sum1 += __shfl_xor_sync(0xffffffffu, sum1, off);
}
if (mnew0 != -INFINITY) m0 = mnew0;
l0 = l0 * a0 + sum0; l1 = l1 * a1 + sum1; // ← the once-per-tile f32 regrouping
The O-fragment rescale (current tree lines 4132–4138) — this is the statement that breaks byte identity, forming the cross-tile f32 state chain together with l0 = l0*a0 + sum0:
#pragma unroll
for (int ob = 0; ob < 8; ob++) {
acc[ob].x[0] *= aa0; acc[ob].x[1] *= aa0; // x[0,1,4,5] → row r0
acc[ob].x[2] *= aa1; acc[ob].x[3] *= aa1; // x[2,3,6,7] → row r1
acc[ob].x[4] *= aa0; acc[ob].x[5] *= aa0;
acc[ob].x[6] *= aa1; acc[ob].x[7] *= aa1;
}
The tail tile's O write-back reusing smem as the 32 KB stage (current tree lines 4186–4189) — the non-symbolized constant r50 flushed out:
} else {
// Qs/Ks/Vs regions are free after the KV loop: contiguous staging for
// the 64x128 f32 O tile (32 KB < the 34.8 KB smem budget).
float* stage = reinterpret_cast<float*>(smem);
The launcher (current tree lines 4216–4218) — after the revert only the first term is needed; during the r50 experiment it was max(first term, FA_TQ*hd*4):
// Qs + Ks + Vs only (S/P no longer go through shared memory). sstr = hd+8
// padding; Ks/Vs are FA_TKV rows (the r46 launcher's 3*FA_TQ bug is gone).
size_t smem = ((size_t)FA_TQ + 2 * FA_TKV) * (hd + 8) * 2;
3.3 Pitfalls
- The masked non-symbolized constant. The 32 KB O write-back stage is independent of FA_TKV but was masked by the staging formula's headroom (34.8 > 32). The correct procedure for a symbolized change is to first grep every
smemallocation/reuse point for "the largest requirement" rather than trusting the derived formula alone — this time the launcher came within a hair of under-allocating and the tail block overran. - "Parity green" ≠ "safe to land". The parity fixtures were all green (max err 1e-4), the suite 166/0/3 all green, and only greedy-32 blew the fork open in the token stream. Tolerance gates and byte gates test two different fault classes — the former tests "are the numbers right," the latter tests "is the path exactly on the baseline's rails."
- The occupancy numbers could not be re-verified on the spot. ncu/nsys could not run in that window, and 3 blocks/SM was derived from r48's device report. The derivation is self-consistent and the wall-clock conclusion is unambiguous, but the record honestly labels the evidence grade — a derived value never poses as a measured one.
4. Verification
- parity suite (numeric tolerance ≤1e-4, multiple shapes): proves the math is correct — guards against "the change broke the numbers."
- suite 166/0/3 (166 passed / 0 failed / 3 ignored): the regression surface — guards against "touching FA broke something else."
- greedy-32 byte identity (
-n 32 --greedy --seed 42, 2K prompt; whole-stream diff of dump/tmp/g32_r50_base.txtvs/tmp/g32_r50_new.txt): guards against "ULP regrouping quietly changing the sampling path" — this doc's protagonist gate, and the only one that caught the problem. - interleaved A/B measurement (same window, same binary): the 3-pair series base 2834.1 / 2832.9 / 2805.7 (med 2832.9) against the new med, −0.5%; a further 5-pair series at −0.01% — guards against machine drift being read as a gain.
- (occupancy: derived from r48's ncu device report, see §2 — the evidence grade is noted in the record.)
5. Results
- Correctness: parity green; suite 166/0/3; greedy-32 NOT byte-identical — the fork point is one argmax knife-edge token, after which the generated text diverges from "…95% DRAM-bound." to "…95% smem-bound." and each remains coherent. This is benign ULP-level float-regrouping noise (the numbers themselves are right), but against the strict byte gate it is a fail.
- Wall clock (7B @3325 tok, same-window interleaved A/B medians): 3-pair series base med 2832.9 → −0.5%; 5-pair series −0.01%. Both rounds sit on the negative side of the ±2% noise band — the occupancy gain (2→3 blocks/SM) never cashed into anything visible on the wall.
- Veto mechanism: after r48 FA is only ~4.7% of the whole wall, and its top stall (global K/V loads) is already hidden by 2 blocks/SM; halving FA_TKV doubles the per-tile fixed costs (3×
__syncthreads, cp.async commit/wait, the softmax shfl chains, the O-fragment rescale), exactly cancelling the theoretical gain. Three rounds of same-direction evidence (r23 + r46 + r50): tile shrinking / occupancy on FA is a dead lever. - Under what future conditions a retry is worthwhile: only when FA becomes the wall's critical path again AND the campaign accepts downgrading FA-class changes from "greedy byte-identical" to a tolerance-grade gate package (hard argmax gate + penalty-free greedy + control groups — this package was later formalized in the decode campaign's D3a). Under the strict byte-gate policy, no FA_TKV/FA_TQ change can satisfy that gate — this is not an implementation defect but an inherent property of online softmax's regrouping math (r57's FA_TQ experiment hit the same wall again).
6. Lessons
- Occupancy is not a universal lever: first ask "is the top stall still un-hidden," then look at resident block count — when latency is already hidden, adding blocks only adds contention, and the per-tile-billed fixed costs double exactly as tiles shrink.
- Changing tile size = changing float accumulation order: any change that regroups f32 reductions/rescales under new boundaries produces ULP drift, and the greedy-32 byte gate will catch the one knife-edge token parity can never see.
- Symbolized geometry is an investment in "one-dimensional experiments": r48 wrote every geometric quantity as a function of
FA_TKV, which is the only reason r50 could run a complete experiment with one define — and it incidentally exposed the single non-symbolized constant (the 32 KB O write-back stage). - Archive dead levers too: three same-direction negatives (r23+r46+r50) mean the "FA tile-size" direction never needs to be tried again; a negative result's value is closing the whole question, not just this one diff.
← 52-r49-a-quantize-shared-dedup · Index · 54-r51-producer-fused-a-quantize →