module Realization.Nvidia.SM86.StreamingAttentionSM86 import Accelerator.SM86.Control import Accelerator.SM86.Immediate import Accelerator.SM86.Instruction import Accelerator.SM86.NumericSemantics import Accelerator.SM86.Program import Accelerator.SM86.Types import Std.Natural import Std.Foundation -- 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. def saR = (lambda unrestricted index : Nat . (sm86Register (nat-to-byte index))) def saU = sm86Unsigned32FromNaturalTruncated def saNone = (constructor SM86Barrier SM86BarrierNone) def saSB0 = (constructor SM86Barrier SM86Barrier0) def saSB1 = (constructor SM86Barrier SM86Barrier1) def saSB2 = (constructor SM86Barrier SM86Barrier2) def saSB3 = (constructor SM86Barrier SM86Barrier3) def saSB4 = (constructor SM86Barrier SM86Barrier4) def saSB5 = (constructor SM86Barrier SM86Barrier5) -- the wait mask naming scoreboards: bit n waits SBn def saWaitNone : Nat = 0 def saWait0 : Nat = 1 def saWait1 : Nat = 2 def saWait2 : Nat = 4 def saWait3 : Nat = 8 def saWait4 : Nat = 16 def saWait5 : Nat = 32 def saWaitAll : Nat = 63 -- Generated with the largest stall (15): Realization.Nvidia.SM86. -- StallCompaction shrinks each to what the fixed latencies need, and -- SM86.Scoreboard checks the result; variable-latency results are ordered -- by the scoreboards named here. def saCtl = (lambda unrestricted write : (family SM86Barrier) . (lambda unrestricted wait : Nat . (constructor SM86Control SM86ControlValue (byte 15) (constructor SM86YieldMode SM86Continue) write saNone (nat-to-byte wait) (byte 0)))) def saPlain = (saCtl saNone saWaitNone) def saAfter = (lambda unrestricted wait : Nat . (saCtl saNone wait)) def saOp = (lambda unrestricted body : (family SM86InstructionBody) . (lambda unrestricted tail : (family SM86Program) . (constructor SM86Program SM86ProgramNext (sm86Instruction body) tail))) def saUnless = (lambda unrestricted predicate : (family SM86Predicate) . (lambda unrestricted body : (family SM86InstructionBody) . (lambda unrestricted tail : (family SM86Program) . (constructor SM86Program SM86ProgramNext (sm86NegatedPredicatedInstruction predicate body) tail)))) def saWhen = (lambda unrestricted predicate : (family SM86Predicate) . (lambda unrestricted body : (family SM86InstructionBody) . (lambda unrestricted tail : (family SM86Program) . (constructor SM86Program SM86ProgramNext (sm86PredicatedInstruction predicate body) tail)))) def saP0 = (constructor SM86Predicate SM86Predicate0) def saP1 = (constructor SM86Predicate SM86Predicate1) -- body 0, body 1, ..., body (count - 1), then the tail def saFor = (lambda unrestricted count : Nat . (lambda unrestricted body : (pi unrestricted index : Nat . (pi unrestricted tail : (family SM86Program) . (family SM86Program))) . (lambda unrestricted tail : (family SM86Program) . (app (nat-eliminate (lambda unrestricted n : Nat . (pi unrestricted start : Nat . (family SM86Program))) (lambda unrestricted start : Nat . tail) (lambda unrestricted predecessor : Nat . (lambda unrestricted rest : (pi unrestricted start : Nat . (family SM86Program)) . (lambda unrestricted start : Nat . (body start (rest (succ start)))))) count) 0)))) -- ---- instructions ---- def saMovImm = (lambda unrestricted d : Nat . (lambda unrestricted value : Nat . (saOp (constructor SM86InstructionBody SM86MoveImmediate (saR d) (saU value) saPlain)))) def saMovImmAfter = (lambda unrestricted d : Nat . (lambda unrestricted value : Nat . (lambda unrestricted wait : Nat . (saOp (constructor SM86InstructionBody SM86MoveImmediate (saR d) (saU value) (saAfter wait)))))) def saMovConst = (lambda unrestricted d : Nat . (lambda unrestricted offset : Nat . (saOp (constructor SM86InstructionBody SM86MoveConstant (saR d) (byte 0) (saU offset) saPlain)))) def saS2R = (lambda unrestricted d : Nat . (lambda unrestricted special : (family SM86SpecialRegister) . (saOp (constructor SM86InstructionBody SM86SpecialToRegister (saR d) special (saCtl saSB5 saWaitNone))))) -- d = a * immediate + c def saImad = (lambda unrestricted d : Nat . (lambda unrestricted a : Nat . (lambda unrestricted immediate : Nat . (lambda unrestricted c : Nat . (saOp (constructor SM86InstructionBody SM86IntegerMultiplyAddImmediate (saR d) (saR a) (saU immediate) (saR c) saPlain)))))) -- d = a + immediate def saAddImm = (lambda unrestricted d : Nat . (lambda unrestricted a : Nat . (lambda unrestricted immediate : Nat . (saOp (constructor SM86InstructionBody SM86IntegerAddThreeImmediate (saR d) (saR a) (saU immediate) saPlain))))) def saShr = (lambda unrestricted d : Nat . (lambda unrestricted a : Nat . (lambda unrestricted amount : Nat . (saOp (constructor SM86InstructionBody SM86ShiftRightImmediate (saR d) (saR a) (nat-to-byte amount) saPlain))))) -- d = a & b def saAnd = (lambda unrestricted d : Nat . (lambda unrestricted a : Nat . (lambda unrestricted b : Nat . (saOp (constructor SM86InstructionBody SM86LogicThreeInputTruthTable (saR d) (saR a) (saR b) (byte 192) saPlain))))) -- d (pair) = a * b + c[0][offset] (64-bit) def saWide = (lambda unrestricted d : Nat . (lambda unrestricted a : Nat . (lambda unrestricted b : Nat . (lambda unrestricted offset : Nat . (saOp (constructor SM86InstructionBody SM86IntegerMultiplyAddWideConstant (saR d) (saR a) (saR b) (byte 0) (saU offset) saPlain)))))) def saLoad = (lambda unrestricted d : Nat . (lambda unrestricted address : Nat . (lambda unrestricted offset : Nat . (lambda unrestricted write : (family SM86Barrier) . (lambda unrestricted wait : Nat . (saOp (constructor SM86InstructionBody SM86LoadGlobal (saR d) (saR address) (saU offset) (saCtl write wait)))))))) def saStore = (lambda unrestricted address : Nat . (lambda unrestricted value : Nat . (lambda unrestricted offset : Nat . (lambda unrestricted wait : Nat . (constructor SM86InstructionBody SM86StoreGlobal (saR address) (saR value) (saU offset) (saAfter wait)))))) def saHmma = (lambda unrestricted d : Nat . (lambda unrestricted a : Nat . (lambda unrestricted b : Nat . (lambda unrestricted write : (family SM86Barrier) . (lambda unrestricted wait : Nat . (saOp (constructor SM86InstructionBody SM86TensorCoreHalfMatrixMultiplyAccumulate16x8x16Float32 (saR d) (saR a) (saR b) (saR d) (saCtl write wait)))))))) def saFmul = (lambda unrestricted d : Nat . (lambda unrestricted a : Nat . (lambda unrestricted b : Nat . (lambda unrestricted wait : Nat . (saOp (constructor SM86InstructionBody SM86FloatMultiply (saR d) (saR a) (saR b) (saAfter wait))))))) def saFadd = (lambda unrestricted d : Nat . (lambda unrestricted a : Nat . (lambda unrestricted b : Nat . (lambda unrestricted wait : Nat . (saOp (constructor SM86InstructionBody SM86FloatAdd (saR d) (saR a) (saR b) (saAfter wait))))))) def saFfma = (lambda unrestricted d : Nat . (lambda unrestricted a : Nat . (lambda unrestricted b : Nat . (lambda unrestricted c : Nat . (saOp (constructor SM86InstructionBody SM86FloatFusedMultiplyAdd (saR d) (saR a) (saR b) (saR c) saPlain)))))) def saFmax = (lambda unrestricted d : Nat . (lambda unrestricted a : Nat . (lambda unrestricted b : Nat . (lambda unrestricted wait : Nat . (saOp (constructor SM86InstructionBody SM86FloatMinimumOrMaximum (saR d) (saR a) (saR b) (constructor SM86FloatExtremum SM86FloatMaximum) (saAfter wait))))))) def saFneg = (lambda unrestricted d : Nat . (lambda unrestricted a : Nat . (saOp (constructor SM86InstructionBody SM86FloatNegate (saR d) (saR a) saPlain)))) def saMufu = (lambda unrestricted d : Nat . (lambda unrestricted a : Nat . (lambda unrestricted operation : (family SM86MultiFunction) . (lambda unrestricted wait : Nat . (saOp (constructor SM86InstructionBody SM86MultiFunctionUnitApproximation (saR d) (saR a) operation (saCtl saSB4 wait))))))) def saEx2 = (lambda unrestricted d : Nat . (lambda unrestricted a : Nat . (lambda unrestricted wait : Nat . (saMufu d a (constructor SM86MultiFunction SM86ExponentialBase2) wait)))) -- butterfly across the lanes of a quad def saShfl = (lambda unrestricted d : Nat . (lambda unrestricted a : Nat . (lambda unrestricted lane : Nat . (lambda unrestricted wait : Nat . (saOp (constructor SM86InstructionBody SM86WarpShuffle (saR d) (saR a) (nat-to-byte lane) (saU 31) (constructor SM86ShuffleMode SM86ShuffleButterfly) (saCtl saSB5 wait))))))) -- d = (binary16 high, binary16 low) def saPack = (lambda unrestricted d : Nat . (lambda unrestricted high : Nat . (lambda unrestricted low : Nat . (saOp (constructor SM86InstructionBody SM86FloatPairToPackedHalfPair (saR d) (saR high) (saR low) saPlain))))) -- P = a > immediate (unsigned) def saGreater = (lambda unrestricted predicate : (family SM86Predicate) . (lambda unrestricted a : Nat . (lambda unrestricted immediate : Nat . (saOp (constructor SM86InstructionBody SM86PredicateGreaterThanImmediate predicate (saR a) (saU immediate) saPlain))))) -- ---- the geometry ---- def saHeadWidth : Nat = 64 def saTile : Nat = 64 def saThreads : Nat = 128 def saHalfBytes : Nat = 2 -- binary32 -infinity def saMinusInfinity : Nat = 0xff800000 def saParameterBase : Nat = 0x160 def saArgument = (lambda unrestricted index : Nat . (naturalAdd saParameterBase (naturalMultiply index 8))) def saConstantBytes : Nat = (saArgument 6) -- ---- registers ---- def saTid : Nat = 0 def saLane : Nat = 1 def saWarp : Nat = 2 def saQuad : Nat = 3 def saQuadLane : Nat = 4 def saOne : Nat = 5 def saKeyOffset : Nat = 6 def saValueOffset : Nat = 7 def saCount : Nat = 8 def saMask : Nat = 9 def saKeyPointer : Nat = 10 def saValuePointer : Nat = 12 def saPointer : Nat = 14 def saMaskBase : Nat = 16 def saScratch : Nat = 17 def saScale : Nat = 18 def saTileIndex : Nat = 19 def saHead : Nat = 192 def saRow : Nat = 193 -- Query backward occupies R194..R243, so the grouped K/V selector lives -- above that entire allocation rather than aliasing its transpose pointer. def saKeyValueHead : Nat = 244 -- the A fragments of Q: k-slice kt's four registers def saQ = (lambda unrestricted kt : Nat . (naturalAdd 20 (naturalMultiply kt 4))) def saQReg = (lambda unrestricted kt : Nat . (lambda unrestricted k : Nat . (naturalAdd (saQ kt) k))) -- the score accumulators: n-tile nt's four binary32 registers def saS = (lambda unrestricted nt : Nat . (naturalAdd 36 (naturalMultiply nt 4))) def saO = (lambda unrestricted nt : Nat . (naturalAdd 68 (naturalMultiply nt 4))) -- K's (then V^T's) B fragments: slice k, n-tile nt, half h def saB = (lambda unrestricted k : Nat . (lambda unrestricted nt : Nat . (naturalAdd 100 (naturalMultiply (naturalAdd (naturalMultiply k 8) nt) 2)))) -- P's A fragments: k-slice kk def saPA = (lambda unrestricted kk : Nat . (naturalAdd 164 (naturalMultiply kk 4))) def saM = (lambda unrestricted h : Nat . (naturalAdd 180 h)) def saL = (lambda unrestricted h : Nat . (naturalAdd 182 h)) def saTmp = (lambda unrestricted n : Nat . (naturalAdd 184 n)) def saRegisters : Nat = 248 -- the element (row half h, column b) of an accumulator tile: c0 c1 the -- first row, c2 c3 the row 8 below def saElem = (lambda unrestricted tile : Nat . (lambda unrestricted h : Nat . (lambda unrestricted b : Nat . (naturalAdd tile (naturalAdd (naturalMultiply 2 h) b))))) -- ---- one key tile ---- -- the tile's K fragments (all 32 slices x tiles x halves), S zeroed def saLoadKeys = (lambda unrestricted tail : (family SM86Program) . (saFor 32 (lambda unrestricted i : Nat . (let unrestricted k = (naturalDivideUnchecked i 8) in (let unrestricted nt = (naturalModuloUnchecked i 8) in (lambda unrestricted rest : (family SM86Program) . (saLoad (saB k nt) saKeyPointer (naturalAdd (naturalMultiply nt 1024) (naturalMultiply k 32)) saSB0 (naturalSelect (naturalIsZero i) saWait3 saWaitNone) (saLoad (succ (saB k nt)) saKeyPointer (naturalAdd (naturalMultiply nt 1024) (naturalAdd (naturalMultiply k 32) 16)) saSB0 saWaitNone rest)))))) (saFor 32 (lambda unrestricted i : Nat . (saMovImm (naturalAdd 36 i) 0)) tail))) -- S = Q K^T: slice by slice, the eight n-tiles of a slice in flight together def saScores = (lambda unrestricted tail : (family SM86Program) . (saFor 4 (lambda unrestricted k : Nat . (saFor 8 (lambda unrestricted nt : Nat . (saHmma (saS nt) (saQ k) (saB k nt) saSB2 (naturalSelect (naturalIsZero nt) (naturalSelect (naturalIsZero k) (naturalAdd saWait0 saWait1) saWait2) saWaitNone))))) tail)) -- the tile's V^T fragments into the registers K's had (the products that -- read them have finished: the first load waits for them) def saLoadValues = (lambda unrestricted seq : Nat . (lambda unrestricted tail : (family SM86Program) . (saFor 32 (lambda unrestricted i : Nat . (let unrestricted k = (naturalDivideUnchecked i 8) in (let unrestricted nt = (naturalModuloUnchecked i 8) in (let unrestricted rowOffset = (naturalMultiply nt (naturalMultiply 8 (naturalMultiply seq saHalfBytes))) in (lambda unrestricted rest : (family SM86Program) . (saLoad (saB k nt) saValuePointer (naturalAdd rowOffset (naturalMultiply k 32)) saSB0 (naturalSelect (naturalIsZero i) saWait2 saWaitNone) (saLoad (succ (saB k nt)) saValuePointer (naturalAdd rowOffset (naturalAdd (naturalMultiply k 32) 16)) saSB0 saWaitNone rest))))))) tail))) -- the diagonal tile's mask: element (row r + 8h, column 8 nt + 2 (lane % 4) + b) -- is masked when its column exceeds its row: when 8 nt + b - 8 h exceeds -- 16 warp + lane / 4 - 2 (lane % 4). saMask holds that plus 64, plus 4096 -- for every tile after this one, so no element of an earlier tile is masked def saMaskTile = (lambda unrestricted tail : (family SM86Program) . (saFor 32 (lambda unrestricted i : Nat . (let unrestricted nt = (naturalDivideUnchecked i 4) in (let unrestricted h = (naturalDivideUnchecked (naturalModuloUnchecked i 4) 2) in (let unrestricted b = (naturalModuloUnchecked i 2) in (let unrestricted threshold = (naturalSaturatingSubtract (naturalAdd (naturalAdd (naturalMultiply 8 nt) b) 63) (naturalMultiply 8 h)) in (lambda unrestricted rest : (family SM86Program) . (saGreater saP0 saMask threshold (saUnless saP0 (constructor SM86InstructionBody SM86MoveImmediate (saR (saElem (saS nt) h b)) (saU saMinusInfinity) saPlain) rest)))))))) tail)) -- one row half h: its sixteen scores, scaled to base 2 (the products have -- finished: the first multiply waits for them) def saScaleRow = (lambda unrestricted h : Nat . (lambda unrestricted tail : (family SM86Program) . (saFor 16 (lambda unrestricted i : Nat . (let unrestricted e = (saElem (saS (naturalDivideUnchecked i 2)) h (naturalModuloUnchecked i 2)) in (saFmul e e saScale saWaitNone))) tail))) -- the row's maximum over the quad's 64 columns, in tmp h def saRowMax = (lambda unrestricted h : Nat . (lambda unrestricted tail : (family SM86Program) . (let unrestricted into = (saTmp h) in (saFmax into (saElem (saS 0) h 0) (saElem (saS 0) h 1) saWaitNone (saFor 14 (lambda unrestricted j : Nat . (let unrestricted i = (naturalAdd j 2) in (saFmax into into (saElem (saS (naturalDivideUnchecked i 2)) h (naturalModuloUnchecked i 2)) saWaitNone))) (saShfl (saTmp 2) into 1 saWaitNone (saFmax into into (saTmp 2) saWait5 (saShfl (saTmp 2) into 2 saWaitNone (saFmax into into (saTmp 2) saWait5 tail))))))))) -- m' = max(m, the row's maximum) in tmp h; alpha = 2^(m - m') in tmp 4 + h; -- m = m'; -m' in tmp 6 + h def saRescaleFactor = (lambda unrestricted h : Nat . (lambda unrestricted tail : (family SM86Program) . (saFmax (saTmp h) (saTmp h) (saM h) saWaitNone (saFneg (saTmp (naturalAdd 6 h)) (saTmp h) (saFadd (saTmp (naturalAdd 4 h)) (saM h) (saTmp (naturalAdd 6 h)) saWaitNone (saEx2 (saTmp (naturalAdd 4 h)) (saTmp (naturalAdd 4 h)) saWaitNone (saOp (constructor SM86InstructionBody SM86IntegerAddThreeImmediate (saR (saM h)) (saR (saTmp h)) (saU 0) saPlain) tail))))))) -- P = 2^(t - m') in place of t; the row's sum of P into l, rescaled def saExponentials = (lambda unrestricted h : Nat . (lambda unrestricted tail : (family SM86Program) . (saFor 16 (lambda unrestricted i : Nat . (let unrestricted e = (saElem (saS (naturalDivideUnchecked i 2)) h (naturalModuloUnchecked i 2)) in (lambda unrestricted rest : (family SM86Program) . (saFadd e e (saTmp (naturalAdd 6 h)) saWaitNone (saEx2 e e saWaitNone rest))))) -- l = l alpha + (the sum of the sixteen P), once the exponentials are in (saFmul (saL h) (saL h) (saTmp (naturalAdd 4 h)) saWait4 (saFor 16 (lambda unrestricted i : Nat . (saFadd (saL h) (saL h) (saElem (saS (naturalDivideUnchecked i 2)) h (naturalModuloUnchecked i 2)) saWaitNone)) tail))))) -- O's row half h rescaled by alpha (the previous tile's products have -- finished: the first multiply waits for them) def saRescaleOutput = (lambda unrestricted h : Nat . (lambda unrestricted tail : (family SM86Program) . (saFor 16 (lambda unrestricted i : Nat . (let unrestricted e = (saElem (saO (naturalDivideUnchecked i 2)) h (naturalModuloUnchecked i 2)) in (saFmul e e (saTmp (naturalAdd 4 h)) (naturalSelect (naturalIsZero i) saWait3 saWaitNone)))) tail))) -- P's A fragments for P V: k-slice kk is n-tiles 2 kk and 2 kk + 1 def saPackP = (lambda unrestricted tail : (family SM86Program) . (saFor 4 (lambda unrestricted kk : Nat . (let unrestricted left = (saS (naturalMultiply 2 kk)) in (let unrestricted right = (saS (succ (naturalMultiply 2 kk))) in (lambda unrestricted rest : (family SM86Program) . (saPack (saPA kk) (succ left) left (saPack (succ (saPA kk)) (naturalAdd left 3) (naturalAdd left 2) (saPack (naturalAdd (saPA kk) 2) (succ right) right (saPack (naturalAdd (saPA kk) 3) (naturalAdd right 3) (naturalAdd right 2) rest)))))))) tail)) -- O += P V: the first product waits for V^T's loads def saValues = (lambda unrestricted tail : (family SM86Program) . (saFor 4 (lambda unrestricted k : Nat . (saFor 8 (lambda unrestricted nt : Nat . (saHmma (saO nt) (saPA k) (saB k nt) saSB3 (naturalSelect (naturalIsZero nt) (naturalSelect (naturalIsZero k) saWait0 saWait3) saWaitNone))))) tail)) def saTileBody = (lambda unrestricted seq : Nat . (lambda unrestricted tail : (family SM86Program) . (saWide saKeyPointer saKeyOffset saOne (saArgument 1) (saWide saValuePointer saValueOffset saOne (saArgument 2) (saImad saMask saCount 4096 saMaskBase (saLoadKeys (saScores (saLoadValues seq (saMaskTile (saScaleRow 0 (saScaleRow 1 (saRowMax 0 (saRowMax 1 (saRescaleFactor 0 (saRescaleFactor 1 (saExponentials 0 (saExponentials 1 (saRescaleOutput 0 (saRescaleOutput 1 (saPackP (saValues tail))))))))))))))))))))) -- ---- the loop's tail: the next key tile, one fewer left ---- def saLoopTail = (lambda unrestricted seq : Nat . (lambda unrestricted bodyInstructions : Nat . (saAddImm saKeyOffset saKeyOffset (naturalMultiply saTile (naturalMultiply saHeadWidth saHalfBytes)) (saAddImm saValueOffset saValueOffset (naturalMultiply saTile saHalfBytes) (saAddImm saCount saCount 4294967295 (saGreater saP1 saCount 0 (saWhen saP1 (constructor SM86InstructionBody SM86Branch (saU (naturalSaturatingSubtract 4294967296 (naturalMultiply 16 (naturalAdd bodyInstructions 5)))) (sm86Unsigned32 (byte 255) (byte 255) (byte 131) (byte 3)) saPlain) sm86ProgramEmpty))))))) -- ---- the epilogue: O / l to the merged plane, m + log2 l ---- def saEpilogue = (lambda unrestricted seq : Nat . (lambda unrestricted heads : Nat . (lambda unrestricted tail : (family SM86Program) . (let unrestricted rowBytes = (naturalMultiply heads (naturalMultiply saHeadWidth saHalfBytes)) in (saFor 2 (lambda unrestricted h : Nat . (lambda unrestricted rest : (family SM86Program) . (saShfl (saTmp 2) (saL h) 1 saWait3 (saFadd (saL h) (saL h) (saTmp 2) saWait5 (saShfl (saTmp 2) (saL h) 2 saWaitNone (saFadd (saL h) (saL h) (saTmp 2) saWait5 (saMufu (saTmp (naturalAdd 4 h)) (saL h) (constructor SM86MultiFunction SM86Reciprocal) saWaitNone (saMufu (saTmp (naturalAdd 6 h)) (saL h) (constructor SM86MultiFunction SM86LogarithmBase2) saWaitNone (saFor 16 (lambda unrestricted i : Nat . (let unrestricted e = (saElem (saO (naturalDivideUnchecked i 2)) h (naturalModuloUnchecked i 2)) in (saFmul e e (saTmp (naturalAdd 4 h)) (naturalSelect (naturalIsZero i) saWait4 saWaitNone)))) (saFadd (saM h) (saM h) (saTmp (naturalAdd 6 h)) saWaitNone rest)))))))))) -- the output: this thread's rows r and r + 8, columns 8 nt + 2 (lane % 4) (saWide saPointer saScratch saOne (saArgument 3) (saFor 16 (lambda unrestricted i : Nat . (let unrestricted nt = (naturalDivideUnchecked i 2) in (let unrestricted h = (naturalModuloUnchecked i 2) in (lambda unrestricted rest : (family SM86Program) . (saPack (saTmp 2) (succ (saElem (saO nt) h 0)) (saElem (saO nt) h 0) (saOp (saStore saPointer (saTmp 2) (naturalAdd (naturalMultiply nt 16) (naturalMultiply h (naturalMultiply 8 rowBytes))) saWaitNone) rest)))))) -- the log-sum-exp, by the quad's first lane (saWide saPointer saRow saOne (saArgument 4) (saGreater saP0 saQuadLane 0 (saUnless saP0 (saStore saPointer (saM 0) 0 saWaitNone) (saUnless saP0 (saStore saPointer (saM 1) 32 saWaitNone) tail))))))))))) -- ---- the program ---- -- seq a multiple of 64, heads the number of 64-wide heads -- Grid Y selects the compact K/V head and grid Z selects one of its query -- heads. This avoids a device integer divide and keeps K/V planes compact. -- The equal-head case uses one group and a unit Z dimension. def streamingAttentionGroupedSM86Admitted = (lambda unrestricted seq : Nat . (lambda unrestricted heads : Nat . (lambda unrestricted keyValueHeads : Nat . (naturalAnd (naturalNonzero keyValueHeads) (naturalAnd (naturalNonzero heads) (naturalAnd (naturalNonzero seq) (naturalAnd (naturalEqual (naturalModuloUnchecked seq saTile) 0) (naturalAnd (naturalEqual (naturalModuloUnchecked heads (naturalSelect (naturalNonzero keyValueHeads) keyValueHeads 1)) 0) (naturalAnd (naturalLessOrEqual heads 65535) (naturalLessOrEqual (naturalMultiply heads (naturalMultiply saHeadWidth (naturalMultiply seq 4))) 4294967295)))))))))) def streamingAttentionForwardGroupedUncheckedSM86 = (lambda unrestricted seq : Nat . (lambda unrestricted heads : Nat . (lambda unrestricted keyValueHeads : Nat . (let unrestricted groups = (naturalDivideUnchecked heads keyValueHeads) in (let unrestricted headBytes = (naturalMultiply seq (naturalMultiply saHeadWidth saHalfBytes)) in (let unrestricted rowBytes = (naturalMultiply heads (naturalMultiply saHeadWidth saHalfBytes)) in (let unrestricted body = (saTileBody seq sm86ProgramEmpty) in (saS2R saTid (constructor SM86SpecialRegister SM86ThreadIdX) (saS2R saTileIndex (constructor SM86SpecialRegister SM86CooperativeThreadArrayIdX) (saS2R saKeyValueHead (constructor SM86SpecialRegister SM86CooperativeThreadArrayIdY) (saS2R saHead (constructor SM86SpecialRegister SM86CooperativeThreadArrayIdZ) (saMovImm saOne 1 -- the special registers are in once this waits (saMovImmAfter saScratch 31 saWait5 (saImad saHead saKeyValueHead groups saHead (saAnd saLane saTid saScratch (saShr saWarp saTid 5 (saShr saQuad saLane 2 (saMovImm saScratch 3 (saAnd saQuadLane saLane saScratch -- the thread's query row: 64 tile + 16 warp + lane / 4 (saImad saRow saWarp 16 saQuad (saImad saRow saTileIndex saTile saRow -- Q: head h's plane, the row, the quad lane's pair of columns (saMovImm saScratch 0 (saImad saScratch saRow (naturalMultiply saHeadWidth saHalfBytes) saScratch (saImad saScratch saQuadLane 4 saScratch (saImad saScratch saHead headBytes saScratch (saWide saPointer saScratch saOne (saArgument 0) (saFor 16 (lambda unrestricted i : Nat . (let unrestricted kt = (naturalDivideUnchecked i 4) in (let unrestricted j = (naturalModuloUnchecked i 4) in (saLoad (saQReg kt j) saPointer (naturalAdd (naturalMultiply kt 32) (naturalAdd (naturalMultiply (naturalModuloUnchecked j 2) (naturalMultiply 8 (naturalMultiply saHeadWidth saHalfBytes))) (naturalMultiply (naturalDivideUnchecked j 2) 16))) saSB1 saWaitNone)))) -- the output's offset in the merged plane (kept in saScratch): the row, -- head h's 64 columns, the quad lane's pair (saMovImm saScratch 0 (saImad saScratch saRow rowBytes saScratch (saImad saScratch saHead (naturalMultiply saHeadWidth saHalfBytes) saScratch (saImad saScratch saQuadLane 4 saScratch -- the log-sum-exp's (kept in saRow): 4 (h seq + row) (saImad saRow saHead seq saRow (saMovImm saKeyOffset 0 (saImad saRow saRow 4 saKeyOffset -- K from key 0: the quad's row lane / 4, its pair of columns (saImad saKeyOffset saQuad (naturalMultiply saHeadWidth saHalfBytes) saKeyOffset (saImad saKeyOffset saQuadLane 4 saKeyOffset (saImad saKeyOffset saKeyValueHead headBytes saKeyOffset -- V^T from key 0: row lane / 4 of head h's 64 x seq plane (saMovImm saValueOffset 0 (saImad saValueOffset saQuad (naturalMultiply seq saHalfBytes) saValueOffset (saImad saValueOffset saQuadLane 4 saValueOffset (saImad saValueOffset saKeyValueHead headBytes saValueOffset -- the tiles up to the diagonal one; the mask's base, 16 warp + lane / 4 -- - 2 (lane % 4) + 64 - 4096 (saAddImm saCount saTileIndex 1 (saImad saMaskBase saWarp 16 saQuad (saImad saMaskBase saQuadLane 4294967294 saMaskBase (saAddImm saMaskBase saMaskBase (naturalSaturatingSubtract 4294967296 (naturalSaturatingSubtract 4096 64)) (saMovConst saScale (saArgument 5) (saMovImm (saM 0) saMinusInfinity (saMovImm (saM 1) saMinusInfinity (saMovImm (saL 0) 0 (saMovImm (saL 1) 0 (saFor 32 (lambda unrestricted i : Nat . (saMovImm (naturalAdd 68 i) 0)) (sm86ProgramAppend body (sm86ProgramAppend (saLoopTail seq (sm86ProgramCount body)) (saEpilogue seq heads (saOp (constructor SM86InstructionBody SM86Exit (saAfter saWaitAll)) sm86ProgramEmpty))))))))))))))))))))))))))))))))))))))))))))))))))))))) def streamingAttentionForwardGroupedSM86 = (lambda unrestricted seq : Nat . (lambda unrestricted heads : Nat . (lambda unrestricted keyValueHeads : Nat . (lambda erased admitted : (equal Nat (streamingAttentionGroupedSM86Admitted seq heads keyValueHeads) 1) . (streamingAttentionForwardGroupedUncheckedSM86 seq heads keyValueHeads))))) def streamingAttentionForwardSM86 = (lambda unrestricted seq : Nat . (lambda unrestricted heads : Nat . (streamingAttentionForwardGroupedUncheckedSM86 seq heads heads))) def streamingAttentionForwardSM86Registers : Nat = saRegisters def streamingAttentionForwardSM86Threads : Nat = saThreads def streamingAttentionForwardSM86ConstantBytes : Nat = saConstantBytes def streamingAttentionForwardQueryArgument : Nat = 0 def streamingAttentionForwardKeyArgument : Nat = 1 def streamingAttentionForwardValueArgument : Nat = 2 def streamingAttentionForwardOutputArgument : Nat = 3 def streamingAttentionForwardLogSumArgument : Nat = 4 def streamingAttentionForwardScaleArgument : Nat = 5 -- =========================================================================== -- 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. def sbStore64 = (lambda unrestricted address : Nat . (lambda unrestricted value : Nat . (lambda unrestricted offset : Nat . (lambda unrestricted wait : Nat . (constructor SM86InstructionBody SM86StoreGlobal64 (saR address) (saR value) (saU offset) (saAfter wait)))))) def sbWiden = (lambda unrestricted d : Nat . (lambda unrestricted a : Nat . (lambda unrestricted high : Nat . (lambda unrestricted wait : Nat . (saOp (constructor SM86InstructionBody SM86HalfToFloat (saR d) (saR a) (eliminate StdBool (lambda unrestricted current : (family StdBool) . (family SM86HalfSelector)) (stdBoolFromNatural high) (branch StdTrue . (constructor SM86HalfSelector SM86HighHalf)) (branch StdFalse . (constructor SM86HalfSelector SM86LowHalf))) (saAfter wait))))))) def sbNegate = (lambda unrestricted d : Nat . (lambda unrestricted a : Nat . (lambda unrestricted wait : Nat . (saOp (constructor SM86InstructionBody SM86FloatNegate (saR d) (saR a) (saAfter wait)))))) -- registers beyond the forward's def sqKTOffset : Nat = 194 def sqDOut = (lambda unrestricted kt : Nat . (naturalAdd 196 (naturalMultiply kt 4))) def sqDP = (lambda unrestricted nt : Nat . (naturalAdd 212 (naturalMultiply nt 4))) -- the row terms: -L (184 + h) and -D (186 + h) for the rows r and r + 8 def sqNegL = (lambda unrestricted h : Nat . (naturalAdd 184 h)) def sqNegD = (lambda unrestricted h : Nat . (naturalAdd 186 h)) def sqTmp = (lambda unrestricted n : Nat . (naturalAdd 188 n)) def sqRegisters : Nat = 248 -- the dQ accumulators are the forward's O's; dS's A fragments P's -- K^T's rows as the B fragments of dQ += dS K: slice kk (keys 16 kk ..), -- n-tile nt (columns 8 nt ..) of head h's 64 x seq plane def sqLoadKeysT = (lambda unrestricted seq : Nat . (lambda unrestricted tail : (family SM86Program) . (saFor 32 (lambda unrestricted i : Nat . (let unrestricted k = (naturalDivideUnchecked i 8) in (let unrestricted nt = (naturalModuloUnchecked i 8) in (let unrestricted rowOffset = (naturalMultiply nt (naturalMultiply 8 (naturalMultiply seq saHalfBytes))) in (lambda unrestricted rest : (family SM86Program) . (saLoad (saB k nt) saValuePointer (naturalAdd rowOffset (naturalMultiply k 32)) saSB0 (naturalSelect (naturalIsZero i) saWait2 saWaitNone) (saLoad (succ (saB k nt)) saValuePointer (naturalAdd rowOffset (naturalAdd (naturalMultiply k 32) 16)) saSB0 saWaitNone rest))))))) tail))) -- a tile's rows (K's or V's) as B fragments, the first load waiting on `wait` def sqLoadRows = (lambda unrestricted pointer : Nat . (lambda unrestricted wait : Nat . (lambda unrestricted tail : (family SM86Program) . (saFor 32 (lambda unrestricted i : Nat . (let unrestricted k = (naturalDivideUnchecked i 8) in (let unrestricted nt = (naturalModuloUnchecked i 8) in (lambda unrestricted rest : (family SM86Program) . (saLoad (saB k nt) pointer (naturalAdd (naturalMultiply nt 1024) (naturalMultiply k 32)) saSB0 (naturalSelect (naturalIsZero i) wait saWaitNone) (saLoad (succ (saB k nt)) pointer (naturalAdd (naturalMultiply nt 1024) (naturalAdd (naturalMultiply k 32) 16)) saSB0 saWaitNone rest)))))) tail)))) -- accumulators = A (4-register fragments from `a`) x the B fragments, slice -- by slice; the first product waits on `first`, each slice after on SB2 def sqProducts = (lambda unrestricted accumulator : (pi unrestricted nt : Nat . Nat) . (lambda unrestricted a : Nat . (lambda unrestricted first : Nat . (lambda unrestricted tail : (family SM86Program) . (saFor 4 (lambda unrestricted k : Nat . (saFor 8 (lambda unrestricted nt : Nat . (saHmma (accumulator nt) (naturalAdd a (naturalMultiply k 4)) (saB k nt) saSB2 (naturalSelect (naturalIsZero nt) (naturalSelect (naturalIsZero k) first saWait2) saWaitNone))))) tail))))) -- P = 2^(t - L) and dS = P (dP - D), into S's registers (the products have -- finished: the first operation waits for them) def sqScoreGradient = (lambda unrestricted tail : (family SM86Program) . (saFor 32 (lambda unrestricted i : Nat . (let unrestricted nt = (naturalDivideUnchecked i 4) in (let unrestricted h = (naturalDivideUnchecked (naturalModuloUnchecked i 4) 2) in (let unrestricted b = (naturalModuloUnchecked i 2) in (let unrestricted e = (saElem (saS nt) h b) in (lambda unrestricted rest : (family SM86Program) . (saFfma e e saScale (sqNegL h) (saEx2 e e saWaitNone rest)))))))) (saFor 32 (lambda unrestricted i : Nat . (let unrestricted nt = (naturalDivideUnchecked i 4) in (let unrestricted h = (naturalDivideUnchecked (naturalModuloUnchecked i 4) 2) in (let unrestricted b = (naturalModuloUnchecked i 2) in (let unrestricted e = (saElem (saS nt) h b) in (let unrestricted g = (saElem (sqDP nt) h b) in (lambda unrestricted rest : (family SM86Program) . (saFadd g g (sqNegD h) saWaitNone (saFmul e e g (naturalSelect (naturalIsZero i) saWait4 saWaitNone) rest))))))))) tail))) def sqTileBody = (lambda unrestricted seq : Nat . (lambda unrestricted tail : (family SM86Program) . (saImad saMask saCount 4096 saMaskBase -- S = Q K^T (saWide saKeyPointer saKeyOffset saOne (saArgument 1) (sqLoadRows saKeyPointer saWait3 (saFor 32 (lambda unrestricted i : Nat . (saMovImm (naturalAdd 36 i) 0)) (saFor 32 (lambda unrestricted i : Nat . (saMovImm (naturalAdd 212 i) 0)) (sqProducts saS (saQ 0) (naturalAdd saWait0 saWait1) -- dP = dO V^T (V's rows at K's offset) (saWide saKeyPointer saKeyOffset saOne (saArgument 2) (sqLoadRows saKeyPointer saWait2 (sqProducts sqDP (sqDOut 0) saWait0 -- K^T's fragments for dQ, once dP's products are done (saWide saValuePointer sqKTOffset saOne (saArgument 3) (sqLoadKeysT seq (saMaskTile (sqScoreGradient (saPackP (saValues tail))))))))))))))))) def sqLoopTail = (lambda unrestricted bodyInstructions : Nat . (saAddImm saKeyOffset saKeyOffset (naturalMultiply saTile (naturalMultiply saHeadWidth saHalfBytes)) (saAddImm sqKTOffset sqKTOffset (naturalMultiply saTile saHalfBytes) (saAddImm saCount saCount 4294967295 (saGreater saP1 saCount 0 (saWhen saP1 (constructor SM86InstructionBody SM86Branch (saU (naturalSaturatingSubtract 4294967296 (naturalMultiply 16 (naturalAdd bodyInstructions 5)))) (sm86Unsigned32 (byte 255) (byte 255) (byte 131) (byte 3)) saPlain) sm86ProgramEmpty)))))) -- D for the thread's rows: its sixteen dO values of row r (and of r + 8) -- against O's, summed over the quad; -D and -L kept, D written by the -- quad's first lane def sqRowTerms = (lambda unrestricted seq : Nat . (lambda unrestricted heads : Nat . (lambda unrestricted tail : (family SM86Program) . (let unrestricted rowBytes = (naturalMultiply heads (naturalMultiply saHeadWidth saHalfBytes)) in -- O in dO's fragment pattern, into the B registers (free before the loop) (saWide saPointer saScratch saOne (saArgument 5) (saFor 16 (lambda unrestricted i : Nat . (let unrestricted kt = (naturalDivideUnchecked i 4) in (let unrestricted j = (naturalModuloUnchecked i 4) in (saLoad (naturalAdd 100 i) saPointer (naturalAdd (naturalMultiply kt 32) (naturalAdd (naturalMultiply (naturalModuloUnchecked j 2) (naturalMultiply 8 rowBytes)) (naturalMultiply (naturalDivideUnchecked j 2) 16))) saSB0 saWaitNone)))) (saMovImm (sqTmp 0) 0 (saMovImm (sqTmp 1) 0 -- register j of slice kt is row r (j even) or r + 8 (j odd); both halves (saFor 32 (lambda unrestricted i : Nat . (let unrestricted r = (naturalDivideUnchecked i 2) in (let unrestricted high = (naturalModuloUnchecked i 2) in (let unrestricted h = (naturalModuloUnchecked (naturalModuloUnchecked r 4) 2) in (lambda unrestricted rest : (family SM86Program) . (sbWiden (sqTmp 2) (naturalAdd 100 r) high (naturalSelect (naturalIsZero i) (naturalAdd saWait0 saWait1) saWaitNone) (sbWiden (sqTmp 3) (naturalAdd (sqDOut 0) r) high saWaitNone (saFfma (sqTmp h) (sqTmp 2) (sqTmp 3) (sqTmp h) rest)))))))) (saFor 2 (lambda unrestricted h : Nat . (lambda unrestricted rest : (family SM86Program) . (saShfl (sqTmp 2) (sqTmp h) 1 saWaitNone (saFadd (sqTmp h) (sqTmp h) (sqTmp 2) saWait5 (saShfl (sqTmp 2) (sqTmp h) 2 saWaitNone (saFadd (sqTmp h) (sqTmp h) (sqTmp 2) saWait5 (saFneg (sqNegD h) (sqTmp h) rest))))))) -- D and L at 4 (h seq + row): D by the quad's first lane; -L (saWide saPointer saRow saOne (saArgument 8) (saGreater saP0 saQuadLane 0 (saUnless saP0 (saStore saPointer (sqTmp 0) 0 saWaitNone) (saUnless saP0 (saStore saPointer (sqTmp 1) 32 saWaitNone) (saWide saPointer saRow saOne (saArgument 6) (saLoad (sqTmp 2) saPointer 0 saSB1 saWaitNone (saLoad (sqTmp 3) saPointer 32 saSB1 saWaitNone (sbNegate (sqNegL 0) (sqTmp 2) saWait1 (sbNegate (sqNegL 1) (sqTmp 3) saWaitNone tail))))))))))))))))))) -- dQ = s (the sum): the thread's rows r and r + 8, pairs of binary32 columns def sqEpilogue = (lambda unrestricted tail : (family SM86Program) . (saMovConst (sqTmp 4) (saArgument 10) (saFor 32 (lambda unrestricted i : Nat . (saFmul (naturalAdd 68 i) (naturalAdd 68 i) (sqTmp 4) (naturalSelect (naturalIsZero i) saWait3 saWaitNone))) (saWide saPointer saScratch saOne (saArgument 7) (saFor 16 (lambda unrestricted i : Nat . (let unrestricted nt = (naturalDivideUnchecked i 2) in (let unrestricted h = (naturalModuloUnchecked i 2) in (saOp (sbStore64 saPointer (saElem (saO nt) h 0) (naturalAdd (naturalMultiply nt 32) (naturalMultiply h (naturalMultiply 8 (naturalMultiply saHeadWidth 4)))) saWaitNone))))) tail))))) def streamingAttentionQueryGroupedUncheckedSM86 = (lambda unrestricted seq : Nat . (lambda unrestricted heads : Nat . (lambda unrestricted keyValueHeads : Nat . (let unrestricted groups = (naturalDivideUnchecked heads keyValueHeads) in (let unrestricted headBytes = (naturalMultiply seq (naturalMultiply saHeadWidth saHalfBytes)) in (let unrestricted rowBytes = (naturalMultiply heads (naturalMultiply saHeadWidth saHalfBytes)) in (let unrestricted body = (sqTileBody seq sm86ProgramEmpty) in (saS2R saTid (constructor SM86SpecialRegister SM86ThreadIdX) (saS2R saTileIndex (constructor SM86SpecialRegister SM86CooperativeThreadArrayIdX) (saS2R saKeyValueHead (constructor SM86SpecialRegister SM86CooperativeThreadArrayIdY) (saS2R saHead (constructor SM86SpecialRegister SM86CooperativeThreadArrayIdZ) (saMovImm saOne 1 (saMovImmAfter saScratch 31 saWait5 (saImad saHead saKeyValueHead groups saHead (saAnd saLane saTid saScratch (saShr saWarp saTid 5 (saShr saQuad saLane 2 (saMovImm saScratch 3 (saAnd saQuadLane saLane saScratch (saImad saRow saWarp 16 saQuad (saImad saRow saTileIndex saTile saRow -- Q's and dO's A fragments: head h's plane, the row, the quad lane's pair (saMovImm saScratch 0 (saImad saScratch saRow (naturalMultiply saHeadWidth saHalfBytes) saScratch (saImad saScratch saQuadLane 4 saScratch (saImad saScratch saHead headBytes saScratch (saWide saPointer saScratch saOne (saArgument 0) (saFor 16 (lambda unrestricted i : Nat . (let unrestricted kt = (naturalDivideUnchecked i 4) in (let unrestricted j = (naturalModuloUnchecked i 4) in (saLoad (saQReg kt j) saPointer (naturalAdd (naturalMultiply kt 32) (naturalAdd (naturalMultiply (naturalModuloUnchecked j 2) (naturalMultiply 8 (naturalMultiply saHeadWidth saHalfBytes))) (naturalMultiply (naturalDivideUnchecked j 2) 16))) saSB1 saWaitNone)))) (saWide saPointer saScratch saOne (saArgument 4) (saFor 16 (lambda unrestricted i : Nat . (let unrestricted kt = (naturalDivideUnchecked i 4) in (let unrestricted j = (naturalModuloUnchecked i 4) in (saLoad (naturalAdd (sqDOut kt) j) saPointer (naturalAdd (naturalMultiply kt 32) (naturalAdd (naturalMultiply (naturalModuloUnchecked j 2) (naturalMultiply 8 (naturalMultiply saHeadWidth saHalfBytes))) (naturalMultiply (naturalDivideUnchecked j 2) 16))) saSB1 saWaitNone)))) -- dQ's offset (binary32 rows), kept in saKeyPointer's partner until the -- epilogue: saScratch becomes O's (merged plane) for the row terms (saMovImm saKeyOffset 0 (saImad saKeyOffset saRow (naturalMultiply saHeadWidth 4) saKeyOffset (saImad saKeyOffset saQuadLane 8 saKeyOffset (saImad saKeyOffset saHead (naturalMultiply seq (naturalMultiply saHeadWidth 4)) saKeyOffset (saMovImm saScratch 0 (saImad saScratch saRow rowBytes saScratch (saImad saScratch saHead (naturalMultiply saHeadWidth saHalfBytes) saScratch (saImad saScratch saQuadLane 4 saScratch -- the row terms' offset: 4 (h seq + row) (saImad saRow saHead seq saRow (saMovImm saValueOffset 0 (saImad saRow saRow 4 saValueOffset (sqRowTerms seq heads -- dQ's offset into saScratch; K's and V's from key 0, K^T's (saAddImm saScratch saKeyOffset 0 (saMovImm saKeyOffset 0 (saImad saKeyOffset saQuad (naturalMultiply saHeadWidth saHalfBytes) saKeyOffset (saImad saKeyOffset saQuadLane 4 saKeyOffset (saImad saKeyOffset saKeyValueHead headBytes saKeyOffset (saMovImm sqKTOffset 0 (saImad sqKTOffset saQuad (naturalMultiply seq saHalfBytes) sqKTOffset (saImad sqKTOffset saQuadLane 4 sqKTOffset (saImad sqKTOffset saKeyValueHead headBytes sqKTOffset (saAddImm saCount saTileIndex 1 (saImad saMaskBase saWarp 16 saQuad (saImad saMaskBase saQuadLane 4294967294 saMaskBase (saAddImm saMaskBase saMaskBase (naturalSaturatingSubtract 4294967296 (naturalSaturatingSubtract 4096 64)) (saMovConst saScale (saArgument 9) (saFor 32 (lambda unrestricted i : Nat . (saMovImm (naturalAdd 68 i) 0)) (sm86ProgramAppend body (sm86ProgramAppend (sqLoopTail (sm86ProgramCount body)) (sqEpilogue (saOp (constructor SM86InstructionBody SM86Exit (saAfter saWaitAll)) sm86ProgramEmpty)))))))))))))))))))))))))))))))))))))))))))))))))))))))))))) def streamingAttentionQueryGroupedSM86 = (lambda unrestricted seq : Nat . (lambda unrestricted heads : Nat . (lambda unrestricted keyValueHeads : Nat . (lambda erased admitted : (equal Nat (streamingAttentionGroupedSM86Admitted seq heads keyValueHeads) 1) . (streamingAttentionQueryGroupedUncheckedSM86 seq heads keyValueHeads))))) def streamingAttentionQuerySM86 = (lambda unrestricted seq : Nat . (lambda unrestricted heads : Nat . (streamingAttentionQueryGroupedUncheckedSM86 seq heads heads))) def streamingAttentionQuerySM86Registers : Nat = sqRegisters def streamingAttentionQuerySM86Threads : Nat = saThreads def streamingAttentionQuerySM86ConstantBytes : Nat = (saArgument 11) def streamingAttentionQueryQueryArgument : Nat = 0 def streamingAttentionQueryKeyArgument : Nat = 1 def streamingAttentionQueryValueArgument : Nat = 2 def streamingAttentionQueryTransposedKeyArgument : Nat = 3 def streamingAttentionQueryOutputGradientArgument : Nat = 4 def streamingAttentionQueryAttentionOutputArgument : Nat = 5 def streamingAttentionQueryLogSumArgument : Nat = 6 def streamingAttentionQueryQueryGradientArgument : Nat = 7 def streamingAttentionQueryRowTermArgument : Nat = 8 def streamingAttentionQueryBaseTwoScaleArgument : Nat = 9 def streamingAttentionQueryNaturalScaleArgument : Nat = 10 -- The key launch (streamingAttentionKeySM86): a block of two warps per 32 -- keys of a K/V head, grid (T / 32, K/V heads, groups). A warp holds its 16 keys' K and V -- rows as A fragments and walks the 32-query sub-tiles from the last one -- down to its diagonal (the last, masked): S^T = K Q^T, dP^T = V dO^T, -- P^T = 2^(S^T c - L), dS^T = P^T (dP^T - D) with L and D the queries', -- dV += P^T dO and dK += dS^T Q. Operands: Q, K, V, dO ([head][T][64], -- half), Q^T, dO^T ([head][64][T], half), L and D ([head][T]); results dK -- ([head][T][64]) and dV^T ([head][64][T]), binary32. Parameters: Q, K, -- V, dO, Q^T, dO^T, L, D, dK, dV^T, then c and s. def skK = (lambda unrestricted kt : Nat . (naturalAdd 20 (naturalMultiply kt 4))) def skV = (lambda unrestricted kt : Nat . (naturalAdd 36 (naturalMultiply kt 4))) def skS = (lambda unrestricted nt : Nat . (naturalAdd 52 (naturalMultiply nt 4))) def skDP = (lambda unrestricted nt : Nat . (naturalAdd 68 (naturalMultiply nt 4))) def skDV = (lambda unrestricted nt : Nat . (naturalAdd 84 (naturalMultiply nt 4))) def skDK = (lambda unrestricted nt : Nat . (naturalAdd 116 (naturalMultiply nt 4))) -- the B fragments of slice k, n-tile nt among `tiles` n-tiles def skB = (lambda unrestricted tiles : Nat . (lambda unrestricted k : Nat . (lambda unrestricted nt : Nat . (naturalAdd 148 (naturalMultiply (naturalAdd (naturalMultiply k tiles) nt) 2))))) def skPA = (lambda unrestricted kk : Nat . (naturalAdd 180 (naturalMultiply kk 4))) def skSA = (lambda unrestricted kk : Nat . (naturalAdd 188 (naturalMultiply kk 4))) -- -L and -D of the thread's columns: n-tile nt, column b def skNegL = (lambda unrestricted nt : Nat . (lambda unrestricted b : Nat . (naturalAdd 196 (naturalAdd (naturalMultiply nt 2) b)))) def skNegD = (lambda unrestricted nt : Nat . (lambda unrestricted b : Nat . (naturalAdd 204 (naturalAdd (naturalMultiply nt 2) b)))) def skTmp = (lambda unrestricted n : Nat . (naturalAdd 212 n)) def skHead : Nat = 220 def skRow : Nat = 221 def skLogOffset : Nat = 222 def skKeyValueHead : Nat = 223 def skRegisters : Nat = 232 def skSubTile : Nat = 32 def skThreads : Nat = 64 -- rows ([T][64] half) as the B fragments of a product over the 64 columns: -- slices k 0..3, n-tiles nt 0..3 (the sub-tile's 32 rows) def skLoadRows = (lambda unrestricted pointer : Nat . (lambda unrestricted wait : Nat . (lambda unrestricted tail : (family SM86Program) . (saFor 16 (lambda unrestricted i : Nat . (let unrestricted k = (naturalDivideUnchecked i 4) in (let unrestricted nt = (naturalModuloUnchecked i 4) in (lambda unrestricted rest : (family SM86Program) . (saLoad (skB 4 k nt) pointer (naturalAdd (naturalMultiply nt 1024) (naturalMultiply k 32)) saSB0 (naturalSelect (naturalIsZero i) wait saWaitNone) (saLoad (succ (skB 4 k nt)) pointer (naturalAdd (naturalMultiply nt 1024) (naturalAdd (naturalMultiply k 32) 16)) saSB0 saWaitNone rest)))))) tail)))) -- transposed rows ([64][T] half) as the B fragments of a product over the -- sub-tile's 32 positions: slices kk 0..1, n-tiles nt 0..7 (the 64 columns) def skLoadColumns = (lambda unrestricted seq : Nat . (lambda unrestricted pointer : Nat . (lambda unrestricted wait : Nat . (lambda unrestricted tail : (family SM86Program) . (saFor 16 (lambda unrestricted i : Nat . (let unrestricted k = (naturalDivideUnchecked i 8) in (let unrestricted nt = (naturalModuloUnchecked i 8) in (let unrestricted rowOffset = (naturalMultiply nt (naturalMultiply 8 (naturalMultiply seq saHalfBytes))) in (lambda unrestricted rest : (family SM86Program) . (saLoad (skB 8 k nt) pointer (naturalAdd rowOffset (naturalMultiply k 32)) saSB0 (naturalSelect (naturalIsZero i) wait saWaitNone) (saLoad (succ (skB 8 k nt)) pointer (naturalAdd rowOffset (naturalAdd (naturalMultiply k 32) 16)) saSB0 saWaitNone rest))))))) tail))))) def skProducts = (lambda unrestricted accumulator : (pi unrestricted nt : Nat . Nat) . (lambda unrestricted a : (pi unrestricted k : Nat . Nat) . (lambda unrestricted slices : Nat . (lambda unrestricted tiles : Nat . (lambda unrestricted barrier : (family SM86Barrier) . (lambda unrestricted wait : Nat . (lambda unrestricted first : Nat . (lambda unrestricted tail : (family SM86Program) . (saFor slices (lambda unrestricted k : Nat . (saFor tiles (lambda unrestricted nt : Nat . (saHmma (accumulator nt) (a k) (skB tiles k nt) barrier (naturalSelect (naturalIsZero nt) (naturalSelect (naturalIsZero k) first wait) saWaitNone))))) tail))))))))) -- the diagonal sub-tile's mask: element (key r + 8h, query 8 nt + 2 (lane % 4) -- + b) is kept when the query is not before the key: saMask holds -- 2 (lane % 4) - 16 warp - lane / 4 + 128, plus 4096 for every sub-tile -- before the diagonal one def skMaskTile = (lambda unrestricted tail : (family SM86Program) . (saFor 16 (lambda unrestricted i : Nat . (let unrestricted nt = (naturalDivideUnchecked i 4) in (let unrestricted h = (naturalDivideUnchecked (naturalModuloUnchecked i 4) 2) in (let unrestricted b = (naturalModuloUnchecked i 2) in (let unrestricted threshold = (naturalSaturatingSubtract (naturalAdd 127 (naturalMultiply 8 h)) (naturalAdd (naturalMultiply 8 nt) b)) in (lambda unrestricted rest : (family SM86Program) . (saGreater saP0 saMask threshold (saUnless saP0 (constructor SM86InstructionBody SM86MoveImmediate (saR (saElem (skS nt) h b)) (saU saMinusInfinity) saPlain) rest)))))))) tail)) -- P^T = 2^(t - L), dS^T = P^T (dP^T - D), in S^T's registers def skScoreGradient = (lambda unrestricted tail : (family SM86Program) . (saFor 16 (lambda unrestricted i : Nat . (let unrestricted nt = (naturalDivideUnchecked i 4) in (let unrestricted h = (naturalDivideUnchecked (naturalModuloUnchecked i 4) 2) in (let unrestricted b = (naturalModuloUnchecked i 2) in (let unrestricted e = (saElem (skS nt) h b) in (lambda unrestricted rest : (family SM86Program) . (saFfma e e saScale (skNegL nt b) (saEx2 e e saWaitNone rest)))))))) -- P^T's fragments, before dS^T replaces it (saFor 2 (lambda unrestricted kk : Nat . (let unrestricted left = (skS (naturalMultiply 2 kk)) in (let unrestricted right = (skS (succ (naturalMultiply 2 kk))) in (lambda unrestricted rest : (family SM86Program) . (saOp (constructor SM86InstructionBody SM86FloatPairToPackedHalfPair (saR (skPA kk)) (saR (succ left)) (saR left) (saAfter (naturalSelect (naturalIsZero kk) saWait4 saWaitNone))) (saPack (succ (skPA kk)) (naturalAdd left 3) (naturalAdd left 2) (saPack (naturalAdd (skPA kk) 2) (succ right) right (saPack (naturalAdd (skPA kk) 3) (naturalAdd right 3) (naturalAdd right 2) rest)))))))) (saFor 16 (lambda unrestricted i : Nat . (let unrestricted nt = (naturalDivideUnchecked i 4) in (let unrestricted h = (naturalDivideUnchecked (naturalModuloUnchecked i 4) 2) in (let unrestricted b = (naturalModuloUnchecked i 2) in (let unrestricted e = (saElem (skS nt) h b) in (let unrestricted g = (saElem (skDP nt) h b) in (lambda unrestricted rest : (family SM86Program) . (saFadd g g (skNegD nt b) saWaitNone (saFmul e e g saWaitNone rest))))))))) (saFor 2 (lambda unrestricted kk : Nat . (let unrestricted left = (skS (naturalMultiply 2 kk)) in (let unrestricted right = (skS (succ (naturalMultiply 2 kk))) in (lambda unrestricted rest : (family SM86Program) . (saPack (skSA kk) (succ left) left (saPack (succ (skSA kk)) (naturalAdd left 3) (naturalAdd left 2) (saPack (naturalAdd (skSA kk) 2) (succ right) right (saPack (naturalAdd (skSA kk) 3) (naturalAdd right 3) (naturalAdd right 2) rest)))))))) tail))))) def skTileBody = (lambda unrestricted seq : Nat . (lambda unrestricted tail : (family SM86Program) . (saImad saMask saCount 4096 saMaskBase -- the sub-tile's L and D (their negatives) (saWide saPointer skLogOffset saOne (saArgument 6) (saFor 8 (lambda unrestricted i : Nat . (saLoad (skTmp i) saPointer (naturalAdd (naturalMultiply (naturalDivideUnchecked i 2) 32) (naturalMultiply (naturalModuloUnchecked i 2) 4)) saSB1 (naturalSelect (naturalIsZero i) saWait3 saWaitNone))) (saFor 8 (lambda unrestricted i : Nat . (sbNegate (skNegL (naturalDivideUnchecked i 2) (naturalModuloUnchecked i 2)) (skTmp i) (naturalSelect (naturalIsZero i) saWait1 saWaitNone))) (saWide saPointer skLogOffset saOne (saArgument 7) (saFor 8 (lambda unrestricted i : Nat . (saLoad (skTmp i) saPointer (naturalAdd (naturalMultiply (naturalDivideUnchecked i 2) 32) (naturalMultiply (naturalModuloUnchecked i 2) 4)) saSB1 saWaitNone)) (saFor 8 (lambda unrestricted i : Nat . (sbNegate (skNegD (naturalDivideUnchecked i 2) (naturalModuloUnchecked i 2)) (skTmp i) (naturalSelect (naturalIsZero i) saWait1 saWaitNone))) -- S^T = K Q^T (saWide saKeyPointer saKeyOffset saOne (saArgument 0) (skLoadRows saKeyPointer saWait3 (saFor 16 (lambda unrestricted i : Nat . (saMovImm (naturalAdd 52 i) 0)) (saFor 16 (lambda unrestricted i : Nat . (saMovImm (naturalAdd 68 i) 0)) (skProducts skS skK 4 4 saSB2 saWait2 (naturalAdd saWait0 saWait1) -- dP^T = V dO^T (saWide saKeyPointer saKeyOffset saOne (saArgument 3) (skLoadRows saKeyPointer saWait2 (skProducts skDP skV 4 4 saSB2 saWait2 saWait0 -- dO^T's fragments for dV, once dP^T's products are done (saWide saValuePointer saValueOffset saOne (saArgument 5) (skLoadColumns seq saValuePointer saWait2 (skMaskTile (skScoreGradient (skProducts skDV skPA 2 8 saSB3 saWait3 saWait0 -- Q^T's fragments for dK, once dV's products have read dO^T's (saWide saValuePointer saValueOffset saOne (saArgument 4) (skLoadColumns seq saValuePointer saWait3 (skProducts skDK skSA 2 8 saSB3 saWait3 saWait0 tail))))))))))))))))))))))))) def skLoopTail = (lambda unrestricted bodyInstructions : Nat . (saAddImm saKeyOffset saKeyOffset (naturalSaturatingSubtract 4294967296 (naturalMultiply skSubTile (naturalMultiply saHeadWidth saHalfBytes))) (saAddImm saValueOffset saValueOffset (naturalSaturatingSubtract 4294967296 (naturalMultiply skSubTile saHalfBytes)) (saAddImm skLogOffset skLogOffset (naturalSaturatingSubtract 4294967296 (naturalMultiply skSubTile 4)) (saAddImm saCount saCount 4294967295 (saGreater saP1 saCount 0 (saWhen saP1 (constructor SM86InstructionBody SM86Branch (saU (naturalSaturatingSubtract 4294967296 (naturalMultiply 16 (naturalAdd bodyInstructions 6)))) (sm86Unsigned32 (byte 255) (byte 255) (byte 131) (byte 3)) saPlain) sm86ProgramEmpty))))))) -- dK = s (the sum), rows; dV^T, columns: element (key r + 8h, column -- 8 nt + 2 (lane % 4) + b) at (8 nt + b) seq + 8 h words past the thread's def skEpilogue = (lambda unrestricted seq : Nat . (lambda unrestricted tail : (family SM86Program) . (saMovConst (skTmp 0) (saArgument 11) (saFor 32 (lambda unrestricted i : Nat . (saFmul (naturalAdd 116 i) (naturalAdd 116 i) (skTmp 0) (naturalSelect (naturalIsZero i) saWait3 saWaitNone))) (saWide saPointer saScratch saOne (saArgument 8) (saFor 16 (lambda unrestricted i : Nat . (let unrestricted nt = (naturalDivideUnchecked i 2) in (let unrestricted h = (naturalModuloUnchecked i 2) in (saOp (sbStore64 saPointer (saElem (skDK nt) h 0) (naturalAdd (naturalMultiply nt 32) (naturalMultiply h (naturalMultiply 8 (naturalMultiply saHeadWidth 4)))) saWaitNone))))) -- dV^T's offset: head h's 64 x seq plane, the quad lane's first column, -- the key (saMovImm (skTmp 1) 0 (saImad (skTmp 1) saQuadLane (naturalMultiply 8 seq) (skTmp 1) (saImad (skTmp 1) skHead (naturalMultiply saHeadWidth (naturalMultiply seq 4)) (skTmp 1) (saImad (skTmp 1) skRow 4 (skTmp 1) (saWide saPointer (skTmp 1) saOne (saArgument 9) (saFor 32 (lambda unrestricted i : Nat . (let unrestricted nt = (naturalDivideUnchecked i 4) in (let unrestricted h = (naturalDivideUnchecked (naturalModuloUnchecked i 4) 2) in (let unrestricted b = (naturalModuloUnchecked i 2) in (saOp (saStore saPointer (saElem (skDV nt) h b) (naturalAdd (naturalMultiply (naturalAdd (naturalMultiply 8 nt) b) (naturalMultiply seq 4)) (naturalMultiply h 32)) saWaitNone)))))) tail)))))))))))) -- Each query head writes a separate dK/dV plane. A later reduction sums -- those planes into the compact K/V gradient; concurrent CTAs never race on -- one compact gradient address. def streamingAttentionKeyGroupedUncheckedSM86 = (lambda unrestricted seq : Nat . (lambda unrestricted heads : Nat . (lambda unrestricted keyValueHeads : Nat . (let unrestricted groups = (naturalDivideUnchecked heads keyValueHeads) in (let unrestricted headBytes = (naturalMultiply seq (naturalMultiply saHeadWidth saHalfBytes)) in (let unrestricted subTiles = (naturalDivideUnchecked seq skSubTile) in (let unrestricted body = (skTileBody seq sm86ProgramEmpty) in (saS2R saTid (constructor SM86SpecialRegister SM86ThreadIdX) (saS2R saTileIndex (constructor SM86SpecialRegister SM86CooperativeThreadArrayIdX) (saS2R skKeyValueHead (constructor SM86SpecialRegister SM86CooperativeThreadArrayIdY) (saS2R skHead (constructor SM86SpecialRegister SM86CooperativeThreadArrayIdZ) (saMovImm saOne 1 (saMovImmAfter saScratch 31 saWait5 (saImad skHead skKeyValueHead groups skHead (saAnd saLane saTid saScratch (saShr saWarp saTid 5 (saShr saQuad saLane 2 (saMovImm saScratch 3 (saAnd saQuadLane saLane saScratch -- the thread's key row: 32 tile + 16 warp + lane / 4 (saImad skRow saWarp 16 saQuad (saImad skRow saTileIndex skSubTile skRow -- K's and V's A fragments (saMovImm saScratch 0 (saImad saScratch skRow (naturalMultiply saHeadWidth saHalfBytes) saScratch (saImad saScratch saQuadLane 4 saScratch (saImad saScratch skKeyValueHead headBytes saScratch (saWide saPointer saScratch saOne (saArgument 1) (saFor 16 (lambda unrestricted i : Nat . (let unrestricted kt = (naturalDivideUnchecked i 4) in (let unrestricted j = (naturalModuloUnchecked i 4) in (saLoad (naturalAdd (skK kt) j) saPointer (naturalAdd (naturalMultiply kt 32) (naturalAdd (naturalMultiply (naturalModuloUnchecked j 2) (naturalMultiply 8 (naturalMultiply saHeadWidth saHalfBytes))) (naturalMultiply (naturalDivideUnchecked j 2) 16))) saSB1 saWaitNone)))) (saWide saPointer saScratch saOne (saArgument 2) (saFor 16 (lambda unrestricted i : Nat . (let unrestricted kt = (naturalDivideUnchecked i 4) in (let unrestricted j = (naturalModuloUnchecked i 4) in (saLoad (naturalAdd (skV kt) j) saPointer (naturalAdd (naturalMultiply kt 32) (naturalAdd (naturalMultiply (naturalModuloUnchecked j 2) (naturalMultiply 8 (naturalMultiply saHeadWidth saHalfBytes))) (naturalMultiply (naturalDivideUnchecked j 2) 16))) saSB1 saWaitNone)))) -- dK's offset (binary32 rows) in saScratch (saMovImm saScratch 0 (saImad saScratch skRow (naturalMultiply saHeadWidth 4) saScratch (saImad saScratch saQuadLane 8 saScratch (saImad saScratch skHead (naturalMultiply seq (naturalMultiply saHeadWidth 4)) saScratch -- the queries start at the last sub-tile: Q's and dO's rows (row lane / 4, -- the quad lane's pair), Q^T's and dO^T's columns (row lane / 4), L's and D's (saMovImm saKeyOffset 0 (saImad saKeyOffset saQuad (naturalMultiply saHeadWidth saHalfBytes) saKeyOffset (saImad saKeyOffset saQuadLane 4 saKeyOffset (saImad saKeyOffset skHead headBytes saKeyOffset (saAddImm saKeyOffset saKeyOffset (naturalMultiply (naturalSaturatingSubtract subTiles 1) (naturalMultiply skSubTile (naturalMultiply saHeadWidth saHalfBytes))) (saMovImm saValueOffset 0 (saImad saValueOffset saQuad (naturalMultiply seq saHalfBytes) saValueOffset (saImad saValueOffset saQuadLane 4 saValueOffset (saImad saValueOffset skHead headBytes saValueOffset (saAddImm saValueOffset saValueOffset (naturalMultiply (naturalSaturatingSubtract subTiles 1) (naturalMultiply skSubTile saHalfBytes)) (saMovImm skLogOffset 0 (saImad skLogOffset saQuadLane 8 skLogOffset (saImad skLogOffset skHead (naturalMultiply seq 4) skLogOffset (saAddImm skLogOffset skLogOffset (naturalMultiply (naturalSaturatingSubtract subTiles 1) (naturalMultiply skSubTile 4)) -- the sub-tiles from the last down to the diagonal; the mask's base, -- 2 (lane % 4) - 16 warp - lane / 4 + 128 - 4096 (saMovImm saCount subTiles (saImad saCount saTileIndex 4294967295 saCount (saMovImm saMaskBase 0 (saImad saMaskBase saQuadLane 2 saMaskBase (saImad saMaskBase saWarp 4294967280 saMaskBase (saImad saMaskBase saQuad 4294967295 saMaskBase (saAddImm saMaskBase saMaskBase (naturalSaturatingSubtract 4294967296 (naturalSaturatingSubtract 4096 128)) (saMovConst saScale (saArgument 10) (saFor 64 (lambda unrestricted i : Nat . (saMovImm (naturalAdd 84 i) 0)) (sm86ProgramAppend body (sm86ProgramAppend (skLoopTail (sm86ProgramCount body)) (skEpilogue seq (saOp (constructor SM86InstructionBody SM86Exit (saAfter saWaitAll)) sm86ProgramEmpty)))))))))))))))))))))))))))))))))))))))))))))))))))))))))))) def streamingAttentionKeyGroupedSM86 = (lambda unrestricted seq : Nat . (lambda unrestricted heads : Nat . (lambda unrestricted keyValueHeads : Nat . (lambda erased admitted : (equal Nat (streamingAttentionGroupedSM86Admitted seq heads keyValueHeads) 1) . (streamingAttentionKeyGroupedUncheckedSM86 seq heads keyValueHeads))))) def streamingAttentionKeySM86 = (lambda unrestricted seq : Nat . (lambda unrestricted heads : Nat . (streamingAttentionKeyGroupedUncheckedSM86 seq heads heads))) def streamingAttentionKeySM86Registers : Nat = skRegisters def streamingAttentionKeySM86Threads : Nat = skThreads def streamingAttentionKeySM86ConstantBytes : Nat = (saArgument 12) def streamingAttentionKeyQueryArgument : Nat = 0 def streamingAttentionKeyKeyArgument : Nat = 1 def streamingAttentionKeyValueArgument : Nat = 2 def streamingAttentionKeyOutputGradientArgument : Nat = 3 def streamingAttentionKeyTransposedQueryArgument : Nat = 4 def streamingAttentionKeyTransposedOutputGradientArgument : Nat = 5 def streamingAttentionKeyLogSumArgument : Nat = 6 def streamingAttentionKeyRowTermArgument : Nat = 7 def streamingAttentionKeyKeyGradientArgument : Nat = 8 def streamingAttentionKeyValueGradientArgument : Nat = 9 def streamingAttentionKeyBaseTwoScaleArgument : Nat = 10 def streamingAttentionKeyNaturalScaleArgument : Nat = 11 -- Sum the query-head dK or dV planes of one K/V group into its compact -- gradient plane. Both backward outputs have 64*T binary32 elements per -- head, though dV is transposed within that plane. A CTA owns 128 elements -- of one K/V head; no atomics or cross-CTA ordering are needed. The caller -- supplies disjoint source/destination regions and launches grid -- (64*T/128, keyValueHeads, 1) for T divisible by 64. def streamingAttentionReduceGroupedGradientUncheckedSM86 = (lambda unrestricted seq : Nat . (lambda unrestricted heads : Nat . (lambda unrestricted keyValueHeads : Nat . (let unrestricted groups = (naturalDivideUnchecked heads keyValueHeads) in (let unrestricted planeWords = (naturalMultiply saHeadWidth seq) in (let unrestricted planeBytes = (naturalMultiply planeWords 4) in (saS2R saTid (constructor SM86SpecialRegister SM86ThreadIdX) (saS2R 2 (constructor SM86SpecialRegister SM86CooperativeThreadArrayIdX) (saS2R 3 (constructor SM86SpecialRegister SM86CooperativeThreadArrayIdY) (saMovImmAfter 7 0 saWait5 (saImad 4 2 saThreads saTid (saImad 5 3 planeWords 4 (saImad 6 3 (naturalMultiply groups planeWords) 4 (saMovImm 11 4 (saWide 12 6 11 (saArgument 0) (saWide 14 5 11 (saArgument 1) (saMovImm 9 0 (saFor groups (lambda unrestricted group : Nat . (lambda unrestricted rest : (family SM86Program) . (saLoad 8 12 (naturalMultiply group planeBytes) saSB0 saWaitNone (saFadd 9 9 8 saWait0 rest)))) (saOp (saStore 14 9 0 saWaitNone) (saOp (constructor SM86InstructionBody SM86Exit (saAfter saWaitAll)) sm86ProgramEmpty)))))))))))))))))))) def streamingAttentionReduceGroupedGradientSM86 = (lambda unrestricted seq : Nat . (lambda unrestricted heads : Nat . (lambda unrestricted keyValueHeads : Nat . (lambda erased admitted : (equal Nat (streamingAttentionGroupedSM86Admitted seq heads keyValueHeads) 1) . (streamingAttentionReduceGroupedGradientUncheckedSM86 seq heads keyValueHeads))))) def streamingAttentionReduceGroupedGradientSM86Registers : Nat = 24 def streamingAttentionReduceGroupedGradientSM86Threads : Nat = saThreads def streamingAttentionReduceGroupedGradientSM86ConstantBytes : Nat = (saArgument 2) def streamingAttentionReduceSourceArgument : Nat = 0 def streamingAttentionReduceDestinationArgument : Nat = 1 def streamingAttentionReduceGroupedGradientSM86Blocks = (lambda unrestricted seq : Nat . (naturalDivideUnchecked (naturalMultiply saHeadWidth seq) saThreads)) -- These dimensions are part of the realization's contract, not facts for a -- pairing to restate. The admission above guards the divisor and geometry. def streamingAttentionGroupedSM86QueryBlocks = (lambda unrestricted seq : Nat . (naturalDivideUnchecked seq saTile)) def streamingAttentionGroupedSM86KeyBlocks = (lambda unrestricted seq : Nat . (naturalDivideUnchecked seq skSubTile)) def streamingAttentionGroupedSM86Groups = (lambda unrestricted heads : Nat . (lambda unrestricted keyValueHeads : Nat . (naturalDivideUnchecked heads keyValueHeads)))