Source/Packages

Realization.Nvidia.SM86.StreamingAttentionSM86

packages/realizations/cooperative/nvidia-sm86/src/Realization/Nvidia/SM86/StreamingAttentionSM86.alpha

1,129 lines205 declarations66.1 KiBSHA-256 aefcbc3b00ca

def · lines 43–43

saR

Full file
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.