===========================================================================
The backward pass, deterministic and in two launches (no atomics). With
P = 2^(S c - L) recomputed from the forward's base-2 log-sum-exp L, the
probability gradient dP = dO V^T and the row term D = rowsum(dO * O)
(= rowsum(P * dP)), the score gradient is dS = P * (dP - D); then
dQ = s dS K, dK = s dS^T Q and dV = P^T dO, s the score scale.
The query launch (streamingAttentionQuerySM86): a block per 64 queries of
a query head, grid (T / 64, K/V heads, groups). It forms D for its rows from dO's
fragments and the forward's output O (and writes it for the key launch),
then walks the key tiles up to the diagonal as the forward does, keeping
dQ. Operands: Q, K, V, dO ([head][T][64], half), K^T ([head][64][T],
half), O ([T][heads x 64], half, the forward's), L ([head][T]); results
dQ ([head][T][64], binary32) and D ([head][T], binary32). Parameters:
Q, K, V, K^T, dO, O, L, dQ, D, then c and s.
514def sbStore64 = (lambda unrestricted address : Nat . (lambda unrestricted value : Nat . (lambda unrestricted offset : Nat .
515 (lambda unrestricted wait : Nat .
516 (constructor SM86InstructionBody SM86StoreGlobal64 (saR address) (saR value) (saU offset) (saAfter wait))))))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.