Causal attention as a streaming normalised sum: the
forward pass of one head's softmax(Q K^T scale) V without the T x T
scores ever leaving a thread's registers. A block of four warps takes 64
queries of one head (grid x the query tiles, grid y the compact K/V heads,
grid z each K/V head's query group); each warp 16 of them. For every
64-key tile up to and including the diagonal one,
the warp forms its 16 x 64 scores with the tensor cores, folds them into
a running row maximum m, a running row sum l and a running weighted sum O
by the merge rule for normalised weighted sums -- a new maximum m' scales
the old l and O by 2^(m - m') -- and adds the tile's P V. At the end
O / l is the attention output and m + log2 l each row's log-sum-exp (in
base 2), which the backward pass recomputes the probabilities from.
Scores are in base-2 units: t = S c with c the binary32 nearest
scale / ln 2, given at run time; P = 2^(t - m). P is rounded to binary16
for its product with V (as the materialised path rounds its
probabilities); the sums stay binary32. The reduction order differs from
the materialised path's, a numerical change the plan's contract admits
(the gradient check and the loss gate decide it).
Operands come straight from global memory (cached in L1): Q as the A
fragments of HMMA.16816 (row-major, [head][T][64] binary16), K's rows as
the B fragments of S = Q K^T ([head][T][64]), V^T's rows as the B
fragments of P V ([head][64][T]). No shared memory and no barrier: each
warp works alone. The output is binary16, token-major across heads
([T][heads x 64], the heads merged), and the log-sum-exp binary32
([head][T]).
Parameters (constant bank 0 from the SM86 parameter base): pointers to
Q, K, V^T, O and L, then c.
43def saR = (lambda unrestricted index : Nat . (sm86Register (nat-to-byte index)))The compiler supplied declaration spans and resolved links from this source snapshot. This page does not assert that this file belongs to a checked closure.