Source/Packages

Realization.Nvidia.SM86.StreamingAttentionSM86

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

1,203 lines221 declarations70.9 KiBSHA-256 23d4e5e2aa2a

Complete file · line 198

StreamingAttentionSM86.alpha

Definition view
1module Realization.Nvidia.SM86.StreamingAttentionSM86
2
3import Accelerator.SM86.Control
4import Accelerator.SM86.Immediate
5import Accelerator.SM86.Instruction
6import Accelerator.SM86.NumericSemantics
7import Accelerator.SM86.Program
8import Accelerator.SM86.Types
9import Std.Natural
10import Std.Foundation
11
12-- Causal attention as a streaming normalised sum: the
13-- forward pass of one head's softmax(Q K^T scale) V without the T x T
14-- scores ever leaving a thread's registers.  A block of four warps takes 64
15-- queries of one head (grid x the query tiles, grid y the compact K/V heads,
16-- grid z each K/V head's query group); each warp 16 of them. For every
17-- 64-key tile up to and including the diagonal one,
18-- the warp forms its 16 x 64 scores with the tensor cores, folds them into
19-- a running row maximum m, a running row sum l and a running weighted sum O
20-- by the merge rule for normalised weighted sums -- a new maximum m' scales
21-- the old l and O by 2^(m - m') -- and adds the tile's P V.  At the end
22-- O / l is the attention output and m + log2 l each row's log-sum-exp (in
23-- base 2), which the backward pass recomputes the probabilities from.
24--
25-- Scores are in base-2 units: t = S c with c the binary32 nearest
26-- scale / ln 2, given at run time; P = 2^(t - m).  P is rounded to binary16
27-- for its product with V (as the materialised path rounds its
28-- probabilities); the sums stay binary32.  The reduction order differs from
29-- the materialised path's, a numerical change the plan's contract admits
30-- (the gradient check and the loss gate decide it).
31--
32-- Operands come straight from global memory (cached in L1): Q as the A
33-- fragments of HMMA.16816 (row-major, [head][T][64] binary16), K's rows as
34-- the B fragments of S = Q K^T ([head][T][64]), V^T's rows as the B
35-- fragments of P V ([head][64][T]).  No shared memory and no barrier: each
36-- warp works alone.  The output is binary16, token-major across heads
37-- ([T][heads x 64], the heads merged), and the log-sum-exp binary32
38-- ([head][T]).
39--
40-- Parameters (constant bank 0 from the SM86 parameter base): pointers to
41-- Q, K, V^T, O and L, then c.
42
43def saR = (lambda unrestricted index : Nat . (sm86Register (nat-to-byte index)))
44def saU = sm86Unsigned32FromNaturalTruncated
45def saNone = (constructor SM86Barrier SM86BarrierNone)
46def saSB0 = (constructor SM86Barrier SM86Barrier0)
47def saSB1 = (constructor SM86Barrier SM86Barrier1)
48def saSB2 = (constructor SM86Barrier SM86Barrier2)
49def saSB3 = (constructor SM86Barrier SM86Barrier3)
50def saSB4 = (constructor SM86Barrier SM86Barrier4)
51def saSB5 = (constructor SM86Barrier SM86Barrier5)
52-- the wait mask naming scoreboards: bit n waits SBn
53def saWaitNone : Nat = 0
54def saWait0 : Nat = 1
55def saWait1 : Nat = 2
56def saWait2 : Nat = 4
57def saWait3 : Nat = 8
58def saWait4 : Nat = 16
59def saWait5 : Nat = 32
60def saWaitAll : Nat = 63
61
62-- Generated with the largest stall (15): Realization.Nvidia.SM86.
63-- StallCompaction shrinks each to what the fixed latencies need, and
64-- SM86.Scoreboard checks the result; variable-latency results are ordered
65-- by the scoreboards named here.
66def saCtl = (lambda unrestricted write : (family SM86Barrier) . (lambda unrestricted wait : Nat .
67  (constructor SM86Control SM86ControlValue (byte 15) (constructor SM86YieldMode SM86Continue)
68    write saNone (nat-to-byte wait) (byte 0))))
69def saPlain = (saCtl saNone saWaitNone)
70def saAfter = (lambda unrestricted wait : Nat . (saCtl saNone wait))
71
72def saOp = (lambda unrestricted body : (family SM86InstructionBody) . (lambda unrestricted tail : (family SM86Program) .
73  (constructor SM86Program SM86ProgramNext (sm86Instruction body) tail)))
74def saUnless = (lambda unrestricted predicate : (family SM86Predicate) . (lambda unrestricted body : (family SM86InstructionBody) .
75  (lambda unrestricted tail : (family SM86Program) .
76    (constructor SM86Program SM86ProgramNext (sm86NegatedPredicatedInstruction predicate body) tail))))
77def saWhen = (lambda unrestricted predicate : (family SM86Predicate) . (lambda unrestricted body : (family SM86InstructionBody) .
78  (lambda unrestricted tail : (family SM86Program) .
79    (constructor SM86Program SM86ProgramNext (sm86PredicatedInstruction predicate body) tail))))
80def saP0 = (constructor SM86Predicate SM86Predicate0)
81def saP1 = (constructor SM86Predicate SM86Predicate1)
82
83-- body 0, body 1, ..., body (count - 1), then the tail
84def saFor = (lambda unrestricted count : Nat .
85  (lambda unrestricted body : (pi unrestricted index : Nat . (pi unrestricted tail : (family SM86Program) . (family SM86Program))) .
86    (lambda unrestricted tail : (family SM86Program) .
87      (app
88        (nat-eliminate (lambda unrestricted n : Nat . (pi unrestricted start : Nat . (family SM86Program)))
89          (lambda unrestricted start : Nat . tail)
90          (lambda unrestricted predecessor : Nat .
91            (lambda unrestricted rest : (pi unrestricted start : Nat . (family SM86Program)) .
92              (lambda unrestricted start : Nat . (body start (rest (succ start))))))
93          count)
94        0))))
95
96-- ---- instructions ----
97def saMovImm = (lambda unrestricted d : Nat . (lambda unrestricted value : Nat .
98  (saOp (constructor SM86InstructionBody SM86MoveImmediate (saR d) (saU value) saPlain))))
99def saMovImmAfter = (lambda unrestricted d : Nat . (lambda unrestricted value : Nat . (lambda unrestricted wait : Nat .
100  (saOp (constructor SM86InstructionBody SM86MoveImmediate (saR d) (saU value) (saAfter wait))))))
101def saMovConst = (lambda unrestricted d : Nat . (lambda unrestricted offset : Nat .
102  (saOp (constructor SM86InstructionBody SM86MoveConstant (saR d) (byte 0) (saU offset) saPlain))))
103def saS2R = (lambda unrestricted d : Nat . (lambda unrestricted special : (family SM86SpecialRegister) .
104  (saOp (constructor SM86InstructionBody SM86SpecialToRegister (saR d) special (saCtl saSB5 saWaitNone)))))
105-- d = a * immediate + c
106def saImad = (lambda unrestricted d : Nat . (lambda unrestricted a : Nat . (lambda unrestricted immediate : Nat . (lambda unrestricted c : Nat .
107  (saOp (constructor SM86InstructionBody SM86IntegerMultiplyAddImmediate (saR d) (saR a) (saU immediate) (saR c) saPlain))))))
108-- d = a + immediate
109def saAddImm = (lambda unrestricted d : Nat . (lambda unrestricted a : Nat . (lambda unrestricted immediate : Nat .
110  (saOp (constructor SM86InstructionBody SM86IntegerAddThreeImmediate (saR d) (saR a) (saU immediate) saPlain)))))
111def saShr = (lambda unrestricted d : Nat . (lambda unrestricted a : Nat . (lambda unrestricted amount : Nat .
112  (saOp (constructor SM86InstructionBody SM86ShiftRightImmediate (saR d) (saR a) (nat-to-byte amount) saPlain)))))
113-- d = a & b
114def saAnd = (lambda unrestricted d : Nat . (lambda unrestricted a : Nat . (lambda unrestricted b : Nat .
115  (saOp (constructor SM86InstructionBody SM86LogicThreeInputTruthTable (saR d) (saR a) (saR b) (byte 192) saPlain)))))
116-- d (pair) = a * b + c[0][offset] (64-bit)
117def saWide = (lambda unrestricted d : Nat . (lambda unrestricted a : Nat . (lambda unrestricted b : Nat . (lambda unrestricted offset : Nat .
118  (saOp (constructor SM86InstructionBody SM86IntegerMultiplyAddWideConstant (saR d) (saR a) (saR b) (byte 0) (saU offset) saPlain))))))
119def saLoad = (lambda unrestricted d : Nat . (lambda unrestricted address : Nat . (lambda unrestricted offset : Nat .
120  (lambda unrestricted write : (family SM86Barrier) . (lambda unrestricted wait : Nat .
121    (saOp (constructor SM86InstructionBody SM86LoadGlobal (saR d) (saR address) (saU offset) (saCtl write wait))))))))
122def saStore = (lambda unrestricted address : Nat . (lambda unrestricted value : Nat . (lambda unrestricted offset : Nat .
123  (lambda unrestricted wait : Nat .
124    (constructor SM86InstructionBody SM86StoreGlobal (saR address) (saR value) (saU offset) (saAfter wait))))))
125def saHmma = (lambda unrestricted d : Nat . (lambda unrestricted a : Nat . (lambda unrestricted b : Nat .
126  (lambda unrestricted write : (family SM86Barrier) . (lambda unrestricted wait : Nat .
127    (saOp (constructor SM86InstructionBody SM86TensorCoreHalfMatrixMultiplyAccumulate16x8x16Float32
128      (saR d) (saR a) (saR b) (saR d) (saCtl write wait))))))))
129def saFmul = (lambda unrestricted d : Nat . (lambda unrestricted a : Nat . (lambda unrestricted b : Nat . (lambda unrestricted wait : Nat .
130  (saOp (constructor SM86InstructionBody SM86FloatMultiply (saR d) (saR a) (saR b) (saAfter wait)))))))
131def saFadd = (lambda unrestricted d : Nat . (lambda unrestricted a : Nat . (lambda unrestricted b : Nat . (lambda unrestricted wait : Nat .
132  (saOp (constructor SM86InstructionBody SM86FloatAdd (saR d) (saR a) (saR b) (saAfter wait)))))))
133def saFfma = (lambda unrestricted d : Nat . (lambda unrestricted a : Nat . (lambda unrestricted b : Nat . (lambda unrestricted c : Nat .
134  (saOp (constructor SM86InstructionBody SM86FloatFusedMultiplyAdd (saR d) (saR a) (saR b) (saR c) saPlain))))))
135def saFmax = (lambda unrestricted d : Nat . (lambda unrestricted a : Nat . (lambda unrestricted b : Nat . (lambda unrestricted wait : Nat .
136  (saOp (constructor SM86InstructionBody SM86FloatMinimumOrMaximum (saR d) (saR a) (saR b)
137    (constructor SM86FloatExtremum SM86FloatMaximum) (saAfter wait)))))))
138def saFneg = (lambda unrestricted d : Nat . (lambda unrestricted a : Nat .
139  (saOp (constructor SM86InstructionBody SM86FloatNegate (saR d) (saR a) saPlain))))
140def saMufu = (lambda unrestricted d : Nat . (lambda unrestricted a : Nat . (lambda unrestricted operation : (family SM86MultiFunction) .
141  (lambda unrestricted wait : Nat .
142    (saOp (constructor SM86InstructionBody SM86MultiFunctionUnitApproximation (saR d) (saR a) operation (saCtl saSB4 wait)))))))
143def saEx2 = (lambda unrestricted d : Nat . (lambda unrestricted a : Nat . (lambda unrestricted wait : Nat .
144  (saMufu d a (constructor SM86MultiFunction SM86ExponentialBase2) wait))))
145-- butterfly across the lanes of a quad
146def saShfl = (lambda unrestricted d : Nat . (lambda unrestricted a : Nat . (lambda unrestricted lane : Nat . (lambda unrestricted wait : Nat .
147  (saOp (constructor SM86InstructionBody SM86WarpShuffle (saR d) (saR a) (nat-to-byte lane) (saU 31)
148    (constructor SM86ShuffleMode SM86ShuffleButterfly) (saCtl saSB5 wait)))))))
149-- d = (binary16 high, binary16 low)
150def saPack = (lambda unrestricted d : Nat . (lambda unrestricted high : Nat . (lambda unrestricted low : Nat .
151  (saOp (constructor SM86InstructionBody SM86FloatPairToPackedHalfPair (saR d) (saR high) (saR low) saPlain)))))
152-- P = a > immediate (unsigned)
153def saGreater = (lambda unrestricted predicate : (family SM86Predicate) . (lambda unrestricted a : Nat . (lambda unrestricted immediate : Nat .
154  (saOp (constructor SM86InstructionBody SM86PredicateGreaterThanImmediate predicate (saR a) (saU immediate) saPlain)))))
155
156-- ---- the geometry ----
157def saHeadWidth : Nat = 64
158def saTile : Nat = 64
159def saThreads : Nat = 128
160def saHalfBytes : Nat = 2
161-- binary32 -infinity
162def saMinusInfinity : Nat = 0xff800000
163def saParameterBase : Nat = 0x160
164def saArgument = (lambda unrestricted index : Nat . (naturalAdd saParameterBase (naturalMultiply index 8)))
165def saConstantBytes : Nat = (saArgument 6)
166
167-- ---- registers ----
168def saTid : Nat = 0
169def saLane : Nat = 1
170def saWarp : Nat = 2
171def saQuad : Nat = 3
172def saQuadLane : Nat = 4
173def saOne : Nat = 5
174def saKeyOffset : Nat = 6
175def saValueOffset : Nat = 7
176def saCount : Nat = 8
177def saMask : Nat = 9
178def saKeyPointer : Nat = 10
179def saValuePointer : Nat = 12
180def saPointer : Nat = 14
181-- Output stores can still be in flight when the log-sum pointer is formed.
182-- Keep their address pair live until the stores retire instead of reusing it.
183def saForwardLogSumPointer : Nat = 194
184-- A global store may still consume its source register after the next pack
185-- issues. Give each of the sixteen output stores a distinct packed word;
186-- physical row-zero readback otherwise varied despite identical Q/K/V and
187-- log-sum inputs. These registers follow the two-word log-sum address pair.
188def saForwardOutputPackBase : Nat =
189  (naturalAdd saForwardLogSumPointer 2)
190def saForwardOutputPackRegister =
191  (lambda unrestricted index : Nat .
192    (naturalAdd saForwardOutputPackBase index))
193def saMaskBase : Nat = 16
194def saScratch : Nat = 17
195def saScale : Nat = 18
196def saTileIndex : Nat = 19
197def saHead : Nat = 192
198def saRow : Nat = 193
199-- Query backward occupies R194..R243, so the grouped K/V selector lives
200-- above that entire allocation rather than aliasing its transpose pointer.
201def saKeyValueHead : Nat = 244
202-- the A fragments of Q: k-slice kt's four registers
203def saQ = (lambda unrestricted kt : Nat . (naturalAdd 20 (naturalMultiply kt 4)))
204def saQReg = (lambda unrestricted kt : Nat . (lambda unrestricted k : Nat . (naturalAdd (saQ kt) k)))
205-- the score accumulators: n-tile nt's four binary32 registers
206def saS = (lambda unrestricted nt : Nat . (naturalAdd 36 (naturalMultiply nt 4)))
207def saO = (lambda unrestricted nt : Nat . (naturalAdd 68 (naturalMultiply nt 4)))
208-- K's (then V^T's) B fragments: slice k, n-tile nt, half h
209def saB = (lambda unrestricted k : Nat . (lambda unrestricted nt : Nat .
210  (naturalAdd 100 (naturalMultiply (naturalAdd (naturalMultiply k 8) nt) 2))))
211-- P's A fragments: k-slice kk
212def saPA = (lambda unrestricted kk : Nat . (naturalAdd 164 (naturalMultiply kk 4)))
213def saM = (lambda unrestricted h : Nat . (naturalAdd 180 h))
214def saL = (lambda unrestricted h : Nat . (naturalAdd 182 h))
215def saTmp = (lambda unrestricted n : Nat . (naturalAdd 184 n))
216def saRegisters : Nat = 248
217
218-- the element (row half h, column b) of an accumulator tile: c0 c1 the
219-- first row, c2 c3 the row 8 below
220def saElem = (lambda unrestricted tile : Nat . (lambda unrestricted h : Nat . (lambda unrestricted b : Nat .
221  (naturalAdd tile (naturalAdd (naturalMultiply 2 h) b)))))
222
223-- ---- one key tile ----
224-- the tile's K fragments (all 32 slices x tiles x halves), S zeroed
225def saLoadKeys = (lambda unrestricted tail : (family SM86Program) .
226  (saFor 32 (lambda unrestricted i : Nat .
227    (let unrestricted k = (naturalDivideUnchecked i 8) in (let unrestricted nt = (naturalModuloUnchecked i 8) in
228      (lambda unrestricted rest : (family SM86Program) .
229        (saLoad (saB k nt) saKeyPointer (naturalAdd (naturalMultiply nt 1024) (naturalMultiply k 32)) saSB0 (naturalSelect (naturalIsZero i) saWait3 saWaitNone)
230        (saLoad (succ (saB k nt)) saKeyPointer (naturalAdd (naturalMultiply nt 1024) (naturalAdd (naturalMultiply k 32) 16)) saSB0 saWaitNone
231          rest))))))
232  (saFor 32 (lambda unrestricted i : Nat . (saMovImm (naturalAdd 36 i) 0))
233    tail)))
234
235-- S = Q K^T: slice by slice, the eight n-tiles of a slice in flight together
236def saScores = (lambda unrestricted tail : (family SM86Program) .
237  (saFor 4 (lambda unrestricted k : Nat .
238    (saFor 8 (lambda unrestricted nt : Nat .
239      (saHmma (saS nt) (saQ k) (saB k nt) saSB2
240        (naturalSelect (naturalIsZero nt) (naturalSelect (naturalIsZero k) (naturalAdd saWait0 saWait1) saWait2) saWaitNone)))))
241    tail))
242
243-- the tile's V^T fragments into the registers K's had (the products that
244-- read them have finished: the first load waits for them)
245def saLoadValues = (lambda unrestricted seq : Nat . (lambda unrestricted tail : (family SM86Program) .
246  (saFor 32 (lambda unrestricted i : Nat .
247    (let unrestricted k = (naturalDivideUnchecked i 8) in (let unrestricted nt = (naturalModuloUnchecked i 8) in
248      (let unrestricted rowOffset = (naturalMultiply nt (naturalMultiply 8 (naturalMultiply seq saHalfBytes))) in
249      (lambda unrestricted rest : (family SM86Program) .
250        (saLoad (saB k nt) saValuePointer (naturalAdd rowOffset (naturalMultiply k 32)) saSB0 (naturalSelect (naturalIsZero i) saWait2 saWaitNone)
251        (saLoad (succ (saB k nt)) saValuePointer (naturalAdd rowOffset (naturalAdd (naturalMultiply k 32) 16)) saSB0 saWaitNone
252          rest)))))))
253    tail)))
254
255-- the diagonal tile's mask: element (row r + 8h, column 8 nt + 2 (lane % 4) + b)
256-- is masked when its column exceeds its row: when 8 nt + b - 8 h exceeds
257-- 16 warp + lane / 4 - 2 (lane % 4).  saMask holds that plus 64, plus 4096
258-- for every tile after this one, so no element of an earlier tile is masked
259def saMaskTile = (lambda unrestricted tail : (family SM86Program) .
260  (saFor 32 (lambda unrestricted i : Nat .
261    (let unrestricted nt = (naturalDivideUnchecked i 4) in (let unrestricted h = (naturalDivideUnchecked (naturalModuloUnchecked i 4) 2) in
262      (let unrestricted b = (naturalModuloUnchecked i 2) in
263      (let unrestricted threshold = (naturalSaturatingSubtract (naturalAdd (naturalAdd (naturalMultiply 8 nt) b) 63) (naturalMultiply 8 h)) in
264      (lambda unrestricted rest : (family SM86Program) .
265        (saGreater saP0 saMask threshold
266          (saUnless saP0 (constructor SM86InstructionBody SM86MoveImmediate (saR (saElem (saS nt) h b)) (saU saMinusInfinity) saPlain)
267            rest))))))))
268    tail))
269
270-- one row half h: its sixteen scores, scaled to base 2 (the products have
271-- finished: the first multiply waits for them)
272def saScaleRow = (lambda unrestricted h : Nat . (lambda unrestricted tail : (family SM86Program) .
273  (saFor 16 (lambda unrestricted i : Nat .
274    (let unrestricted e = (saElem (saS (naturalDivideUnchecked i 2)) h (naturalModuloUnchecked i 2)) in
275      (saFmul e e saScale saWaitNone)))
276    tail)))
277
278-- the row's maximum over the quad's 64 columns, in tmp h
279def saRowMax = (lambda unrestricted h : Nat . (lambda unrestricted tail : (family SM86Program) .
280  (let unrestricted into = (saTmp h) in
281  (saFmax into (saElem (saS 0) h 0) (saElem (saS 0) h 1) saWaitNone
282    (saFor 14 (lambda unrestricted j : Nat .
283      (let unrestricted i = (naturalAdd j 2) in
284        (saFmax into into (saElem (saS (naturalDivideUnchecked i 2)) h (naturalModuloUnchecked i 2)) saWaitNone)))
285      (saShfl (saTmp 2) into 1 saWaitNone
286      (saFmax into into (saTmp 2) saWait5
287      (saShfl (saTmp 2) into 2 saWaitNone
288      (saFmax into into (saTmp 2) saWait5
289        tail)))))))))
290
291-- m' = max(m, the row's maximum) in tmp h; alpha = 2^(m - m') in tmp 4 + h;
292-- m = m'; -m' in tmp 6 + h
293def saRescaleFactor = (lambda unrestricted h : Nat . (lambda unrestricted tail : (family SM86Program) .
294  (saFmax (saTmp h) (saTmp h) (saM h) saWaitNone
295  (saFneg (saTmp (naturalAdd 6 h)) (saTmp h)
296  (saFadd (saTmp (naturalAdd 4 h)) (saM h) (saTmp (naturalAdd 6 h)) saWaitNone
297  (saEx2 (saTmp (naturalAdd 4 h)) (saTmp (naturalAdd 4 h)) saWaitNone
298  (saOp (constructor SM86InstructionBody SM86IntegerAddThreeImmediate (saR (saM h)) (saR (saTmp h)) (saU 0) saPlain)
299    tail)))))))
300
301-- P = 2^(t - m') in place of t; the row's sum of P into l, rescaled
302def saExponentials = (lambda unrestricted h : Nat . (lambda unrestricted tail : (family SM86Program) .
303  (saFor 16 (lambda unrestricted i : Nat .
304    (let unrestricted e = (saElem (saS (naturalDivideUnchecked i 2)) h (naturalModuloUnchecked i 2)) in
305      (lambda unrestricted rest : (family SM86Program) .
306        (saFadd e e (saTmp (naturalAdd 6 h)) saWaitNone
307          (saEx2 e e saWaitNone rest)))))
308    -- l = l alpha + (the sum of the sixteen P), once the exponentials are in
309    (saFmul (saL h) (saL h) (saTmp (naturalAdd 4 h)) saWait4
310    (saFor 16 (lambda unrestricted i : Nat .
311      (saFadd (saL h) (saL h) (saElem (saS (naturalDivideUnchecked i 2)) h (naturalModuloUnchecked i 2)) saWaitNone))
312      tail)))))
313
314-- O's row half h rescaled by alpha (the previous tile's products have
315-- finished: the first multiply waits for them)
316def saRescaleOutput = (lambda unrestricted h : Nat . (lambda unrestricted tail : (family SM86Program) .
317  (saFor 16 (lambda unrestricted i : Nat .
318    (let unrestricted e = (saElem (saO (naturalDivideUnchecked i 2)) h (naturalModuloUnchecked i 2)) in
319      (saFmul e e (saTmp (naturalAdd 4 h)) (naturalSelect (naturalIsZero i) saWait3 saWaitNone))))
320    tail)))
321
322-- P's A fragments for P V: k-slice kk is n-tiles 2 kk and 2 kk + 1
323def saPackP = (lambda unrestricted tail : (family SM86Program) .
324  (saFor 4 (lambda unrestricted kk : Nat .
325    (let unrestricted left = (saS (naturalMultiply 2 kk)) in (let unrestricted right = (saS (succ (naturalMultiply 2 kk))) in
326      (lambda unrestricted rest : (family SM86Program) .
327        (saPack (saPA kk) (succ left) left
328        (saPack (succ (saPA kk)) (naturalAdd left 3) (naturalAdd left 2)
329        (saPack (naturalAdd (saPA kk) 2) (succ right) right
330        (saPack (naturalAdd (saPA kk) 3) (naturalAdd right 3) (naturalAdd right 2)
331          rest))))))))
332    tail))
333
334-- O += P V: the first product waits for V^T's loads
335def saValues = (lambda unrestricted tail : (family SM86Program) .
336  (saFor 4 (lambda unrestricted k : Nat .
337    (saFor 8 (lambda unrestricted nt : Nat .
338      (saHmma (saO nt) (saPA k) (saB k nt) saSB3
339        (naturalSelect (naturalIsZero nt) (naturalSelect (naturalIsZero k) saWait0 saWait3) saWaitNone)))))
340    tail))
341
342def saTileBody = (lambda unrestricted seq : Nat . (lambda unrestricted tail : (family SM86Program) .
343  (saWide saKeyPointer saKeyOffset saOne (saArgument 1)
344  (saWide saValuePointer saValueOffset saOne (saArgument 2)
345  (saImad saMask saCount 4096 saMaskBase
346  (saLoadKeys
347  (saScores
348  (saLoadValues seq
349  (saMaskTile
350  (saScaleRow 0 (saScaleRow 1
351  (saRowMax 0 (saRowMax 1
352  (saRescaleFactor 0 (saRescaleFactor 1
353  (saExponentials 0 (saExponentials 1
354  (saRescaleOutput 0 (saRescaleOutput 1
355  (saPackP
356  (saValues
357    tail)))))))))))))))))))))
358
359-- ---- the loop's tail: the next key tile, one fewer left ----
360def saLoopTail = (lambda unrestricted seq : Nat . (lambda unrestricted bodyInstructions : Nat .
361  (saAddImm saKeyOffset saKeyOffset (naturalMultiply saTile (naturalMultiply saHeadWidth saHalfBytes))
362  (saAddImm saValueOffset saValueOffset (naturalMultiply saTile saHalfBytes)
363  (saAddImm saCount saCount 4294967295
364  (saGreater saP1 saCount 0
365  (saWhen saP1
366    (constructor SM86InstructionBody SM86Branch
367      (saU (naturalSaturatingSubtract 4294967296 (naturalMultiply 16 (naturalAdd bodyInstructions 5))))
368      (sm86Unsigned32 (byte 255) (byte 255) (byte 131) (byte 3))
369      saPlain)
370    sm86ProgramEmpty)))))))
371
372-- ---- the epilogue: O / l to the merged plane, m + log2 l ----
373def saEpilogue = (lambda unrestricted seq : Nat . (lambda unrestricted heads : Nat . (lambda unrestricted tail : (family SM86Program) .
374  (let unrestricted rowBytes = (naturalMultiply heads (naturalMultiply saHeadWidth saHalfBytes)) in
375  (saFor 2 (lambda unrestricted h : Nat . (lambda unrestricted rest : (family SM86Program) .
376    (saShfl (saTmp 2) (saL h) 1 saWait3
377    (saFadd (saL h) (saL h) (saTmp 2) saWait5
378    (saShfl (saTmp 2) (saL h) 2 saWaitNone
379    (saFadd (saL h) (saL h) (saTmp 2) saWait5
380    (saMufu (saTmp (naturalAdd 4 h)) (saL h) (constructor SM86MultiFunction SM86Reciprocal) saWaitNone
381    (saMufu (saTmp (naturalAdd 6 h)) (saL h) (constructor SM86MultiFunction SM86LogarithmBase2) saWaitNone
382    (saFor 16 (lambda unrestricted i : Nat .
383      (let unrestricted e = (saElem (saO (naturalDivideUnchecked i 2)) h (naturalModuloUnchecked i 2)) in
384        (saFmul e e (saTmp (naturalAdd 4 h)) (naturalSelect (naturalIsZero i) saWait4 saWaitNone))))
385    (saFadd (saM h) (saM h) (saTmp (naturalAdd 6 h)) saWaitNone
386      rest))))))))))
387  -- the output: this thread's rows r and r + 8, columns 8 nt + 2 (lane % 4)
388  (saWide saPointer saScratch saOne (saArgument 3)
389  (saFor 16 (lambda unrestricted i : Nat .
390    (let unrestricted nt = (naturalDivideUnchecked i 2) in (let unrestricted h = (naturalModuloUnchecked i 2) in
391      (lambda unrestricted rest : (family SM86Program) .
392        (saPack (saForwardOutputPackRegister i) (succ (saElem (saO nt) h 0)) (saElem (saO nt) h 0)
393        (saOp (saStore saPointer (saForwardOutputPackRegister i) (naturalAdd (naturalMultiply nt 16) (naturalMultiply h (naturalMultiply 8 rowBytes))) saWaitNone)
394          rest))))))
395  -- the log-sum-exp, by the quad's first lane
396  (saWide saForwardLogSumPointer saRow saOne (saArgument 4)
397  (saGreater saP0 saQuadLane 0
398  (saUnless saP0 (saStore saForwardLogSumPointer (saM 0) 0 saWaitNone)
399  (saUnless saP0 (saStore saForwardLogSumPointer (saM 1) 32 saWaitNone)
400    tail)))))))))))
401
402-- ---- the program ----
403-- seq a multiple of 64, heads the number of 64-wide heads
404-- Grid Y selects the compact K/V head and grid Z selects one of its query
405-- heads.  This avoids a device integer divide and keeps K/V planes compact.
406-- The equal-head case uses one group and a unit Z dimension.
407def streamingAttentionGroupedSM86Admitted =
408  (lambda unrestricted seq : Nat .
409    (lambda unrestricted heads : Nat .
410      (lambda unrestricted keyValueHeads : Nat .
411        (naturalAnd (naturalNonzero keyValueHeads)
412          (naturalAnd (naturalNonzero heads)
413            (naturalAnd (naturalNonzero seq)
414              (naturalAnd (naturalEqual (naturalModuloUnchecked seq saTile) 0)
415                (naturalAnd
416                  (naturalEqual
417                    (naturalModuloUnchecked heads
418                      (naturalSelect (naturalNonzero keyValueHeads) keyValueHeads 1)) 0)
419                  (naturalAnd (naturalLessOrEqual heads 65535)
420                    (naturalLessOrEqual
421                      (naturalMultiply heads (naturalMultiply saHeadWidth (naturalMultiply seq 4)))
422                      4294967295))))))))))
423
424def streamingAttentionForwardGroupedUncheckedSM86 = (lambda unrestricted seq : Nat . (lambda unrestricted heads : Nat . (lambda unrestricted keyValueHeads : Nat .
425  (let unrestricted groups = (naturalDivideUnchecked heads keyValueHeads) in
426  (let unrestricted headBytes = (naturalMultiply seq (naturalMultiply saHeadWidth saHalfBytes)) in
427  (let unrestricted rowBytes = (naturalMultiply heads (naturalMultiply saHeadWidth saHalfBytes)) in
428  (let unrestricted body = (saTileBody seq sm86ProgramEmpty) in
429  (saS2R saTid (constructor SM86SpecialRegister SM86ThreadIdX)
430  (saS2R saTileIndex (constructor SM86SpecialRegister SM86CooperativeThreadArrayIdX)
431  (saS2R saKeyValueHead (constructor SM86SpecialRegister SM86CooperativeThreadArrayIdY)
432  (saS2R saHead (constructor SM86SpecialRegister SM86CooperativeThreadArrayIdZ)
433  (saMovImm saOne 1
434  -- the special registers are in once this waits
435  (saMovImmAfter saScratch 31 saWait5
436  (saImad saHead saKeyValueHead groups saHead
437  (saAnd saLane saTid saScratch
438  (saShr saWarp saTid 5
439  (saShr saQuad saLane 2
440  (saMovImm saScratch 3
441  (saAnd saQuadLane saLane saScratch
442  -- the thread's query row: 64 tile + 16 warp + lane / 4
443  (saImad saRow saWarp 16 saQuad
444  (saImad saRow saTileIndex saTile saRow
445  -- Q: head h's plane, the row, the quad lane's pair of columns
446  (saMovImm saScratch 0
447  (saImad saScratch saRow (naturalMultiply saHeadWidth saHalfBytes) saScratch
448  (saImad saScratch saQuadLane 4 saScratch
449  (saImad saScratch saHead headBytes saScratch
450  (saWide saPointer saScratch saOne (saArgument 0)
451  (saFor 16 (lambda unrestricted i : Nat .
452    (let unrestricted kt = (naturalDivideUnchecked i 4) in (let unrestricted j = (naturalModuloUnchecked i 4) in
453      (saLoad (saQReg kt j) saPointer
454        (naturalAdd (naturalMultiply kt 32)
455          (naturalAdd (naturalMultiply (naturalModuloUnchecked j 2) (naturalMultiply 8 (naturalMultiply saHeadWidth saHalfBytes)))
456            (naturalMultiply (naturalDivideUnchecked j 2) 16)))
457        saSB1 saWaitNone))))
458  -- the output's offset in the merged plane (kept in saScratch): the row,
459  -- head h's 64 columns, the quad lane's pair
460  (saMovImm saScratch 0
461  (saImad saScratch saRow rowBytes saScratch
462  (saImad saScratch saHead (naturalMultiply saHeadWidth saHalfBytes) saScratch
463  (saImad saScratch saQuadLane 4 saScratch
464  -- the log-sum-exp's (kept in saRow): 4 (h seq + row)
465  (saImad saRow saHead seq saRow
466  (saMovImm saKeyOffset 0
467  (saImad saRow saRow 4 saKeyOffset
468  -- K from key 0: the quad's row lane / 4, its pair of columns
469  (saImad saKeyOffset saQuad (naturalMultiply saHeadWidth saHalfBytes) saKeyOffset
470  (saImad saKeyOffset saQuadLane 4 saKeyOffset
471  (saImad saKeyOffset saKeyValueHead headBytes saKeyOffset
472  -- V^T from key 0: row lane / 4 of head h's 64 x seq plane
473  (saMovImm saValueOffset 0
474  (saImad saValueOffset saQuad (naturalMultiply seq saHalfBytes) saValueOffset
475  (saImad saValueOffset saQuadLane 4 saValueOffset
476  (saImad saValueOffset saKeyValueHead headBytes saValueOffset
477  -- the tiles up to the diagonal one; the mask's base, 16 warp + lane / 4
478  -- - 2 (lane % 4) + 64 - 4096
479  (saAddImm saCount saTileIndex 1
480  (saImad saMaskBase saWarp 16 saQuad
481  (saImad saMaskBase saQuadLane 4294967294 saMaskBase
482  (saAddImm saMaskBase saMaskBase (naturalSaturatingSubtract 4294967296 (naturalSaturatingSubtract 4096 64))
483  (saMovConst saScale (saArgument 5)
484  (saMovImm (saM 0) saMinusInfinity (saMovImm (saM 1) saMinusInfinity
485  (saMovImm (saL 0) 0 (saMovImm (saL 1) 0
486  (saFor 32 (lambda unrestricted i : Nat . (saMovImm (naturalAdd 68 i) 0))
487  (sm86ProgramAppend body
488    (sm86ProgramAppend (saLoopTail seq (sm86ProgramCount body))
489      (saEpilogue seq heads
490        (saOp (constructor SM86InstructionBody SM86Exit (saAfter saWaitAll)) sm86ProgramEmpty)))))))))))))))))))))))))))))))))))))))))))))))))))))))
491
492def streamingAttentionForwardGroupedSM86 =
493  (lambda unrestricted seq : Nat . (lambda unrestricted heads : Nat .
494    (lambda unrestricted keyValueHeads : Nat .
495      (lambda erased admitted : (equal Nat (streamingAttentionGroupedSM86Admitted seq heads keyValueHeads) 1) .
496        (streamingAttentionForwardGroupedUncheckedSM86 seq heads keyValueHeads)))))
497def streamingAttentionForwardSM86 = (lambda unrestricted seq : Nat . (lambda unrestricted heads : Nat .
498  (streamingAttentionForwardGroupedUncheckedSM86 seq heads heads)))
499
500def streamingAttentionForwardSM86Registers : Nat = saRegisters
501def streamingAttentionForwardSM86Threads : Nat = saThreads
502def streamingAttentionForwardSM86ConstantBytes : Nat = saConstantBytes
503def streamingAttentionForwardQueryArgument : Nat = 0
504def streamingAttentionForwardKeyArgument : Nat = 1
505def streamingAttentionForwardValueArgument : Nat = 2
506def streamingAttentionForwardOutputArgument : Nat = 3
507def streamingAttentionForwardLogSumArgument : Nat = 4
508def streamingAttentionForwardScaleArgument : Nat = 5
509
510-- ===========================================================================
511-- The backward pass, deterministic and in two launches (no atomics).  With
512-- P = 2^(S c - L) recomputed from the forward's base-2 log-sum-exp L, the
513-- probability gradient dP = dO V^T and the row term D = rowsum(dO * O)
514-- (= rowsum(P * dP)), the score gradient is dS = P * (dP - D); then
515-- dQ = s dS K, dK = s dS^T Q and dV = P^T dO, s the score scale.
516--
517-- The query launch (streamingAttentionQuerySM86): a block per 64 queries of
518-- a query head, grid (T / 64, K/V heads, groups). It forms D for its rows from dO's
519-- fragments and the forward's output O (and writes it for the key launch),
520-- then walks the key tiles up to the diagonal as the forward does, keeping
521-- dQ.  Operands: Q, K, V, dO ([head][T][64], half), K^T ([head][64][T],
522-- half), O ([T][heads x 64], half, the forward's), L ([head][T]); results
523-- dQ ([head][T][64], binary32) and D ([head][T], binary32).  Parameters:
524-- Q, K, V, K^T, dO, O, L, dQ, D, then c and s.
525
526def sbStore64 = (lambda unrestricted address : Nat . (lambda unrestricted value : Nat . (lambda unrestricted offset : Nat .
527  (lambda unrestricted wait : Nat .
528    (constructor SM86InstructionBody SM86StoreGlobal64 (saR address) (saR value) (saU offset) (saAfter wait))))))
529def sbWiden = (lambda unrestricted d : Nat . (lambda unrestricted a : Nat . (lambda unrestricted high : Nat . (lambda unrestricted wait : Nat .
530  (saOp (constructor SM86InstructionBody SM86HalfToFloat (saR d) (saR a)
531    (eliminate StdBool (lambda unrestricted current : (family StdBool) . (family SM86HalfSelector)) (stdBoolFromNatural high)
532      (branch StdTrue . (constructor SM86HalfSelector SM86HighHalf))
533      (branch StdFalse . (constructor SM86HalfSelector SM86LowHalf)))
534    (saAfter wait)))))))
535def sbNegate = (lambda unrestricted d : Nat . (lambda unrestricted a : Nat . (lambda unrestricted wait : Nat .
536  (saOp (constructor SM86InstructionBody SM86FloatNegate (saR d) (saR a) (saAfter wait))))))
537
538-- registers beyond the forward's
539def sqKTOffset : Nat = 194
540def sqDOut = (lambda unrestricted kt : Nat . (naturalAdd 196 (naturalMultiply kt 4)))
541def sqDP = (lambda unrestricted nt : Nat . (naturalAdd 212 (naturalMultiply nt 4)))
542-- the row terms: -L (184 + h) and -D (186 + h) for the rows r and r + 8
543def sqNegL = (lambda unrestricted h : Nat . (naturalAdd 184 h))
544def sqNegD = (lambda unrestricted h : Nat . (naturalAdd 186 h))
545def sqTmp = (lambda unrestricted n : Nat . (naturalAdd 188 n))
546-- Q loads may still use R14/R15 when the output-gradient address is formed.
547-- Give dO its own pair; otherwise some row terms read the wrong operand and
548-- vary across identical physical runs. The dO loads finish before the row-
549-- term stores, which can reuse this pair. Both finish before the dP loop
550-- assigns these words, keeping register demand below 256.
551def sqOutputGradientLoadPointer : Nat = 240
552def sqRowTermStorePointer : Nat = sqOutputGradientLoadPointer
553def sqRegisters : Nat = 248
554-- the dQ accumulators are the forward's O's; dS's A fragments P's
555
556-- K^T's rows as the B fragments of dQ += dS K: slice kk (keys 16 kk ..),
557-- n-tile nt (columns 8 nt ..) of head h's 64 x seq plane
558def sqLoadKeysT = (lambda unrestricted seq : Nat . (lambda unrestricted tail : (family SM86Program) .
559  (saFor 32 (lambda unrestricted i : Nat .
560    (let unrestricted k = (naturalDivideUnchecked i 8) in (let unrestricted nt = (naturalModuloUnchecked i 8) in
561      (let unrestricted rowOffset = (naturalMultiply nt (naturalMultiply 8 (naturalMultiply seq saHalfBytes))) in
562      (lambda unrestricted rest : (family SM86Program) .
563        (saLoad (saB k nt) saValuePointer (naturalAdd rowOffset (naturalMultiply k 32)) saSB0 (naturalSelect (naturalIsZero i) saWait2 saWaitNone)
564        (saLoad (succ (saB k nt)) saValuePointer (naturalAdd rowOffset (naturalAdd (naturalMultiply k 32) 16)) saSB0 saWaitNone
565          rest)))))))
566    tail)))
567
568-- a tile's rows (K's or V's) as B fragments, the first load waiting on `wait`
569def sqLoadRows = (lambda unrestricted pointer : Nat . (lambda unrestricted wait : Nat . (lambda unrestricted tail : (family SM86Program) .
570  (saFor 32 (lambda unrestricted i : Nat .
571    (let unrestricted k = (naturalDivideUnchecked i 8) in (let unrestricted nt = (naturalModuloUnchecked i 8) in
572      (lambda unrestricted rest : (family SM86Program) .
573        (saLoad (saB k nt) pointer (naturalAdd (naturalMultiply nt 1024) (naturalMultiply k 32)) saSB0 (naturalSelect (naturalIsZero i) wait saWaitNone)
574        (saLoad (succ (saB k nt)) pointer (naturalAdd (naturalMultiply nt 1024) (naturalAdd (naturalMultiply k 32) 16)) saSB0 saWaitNone
575          rest))))))
576    tail))))
577
578-- accumulators = A (4-register fragments from `a`) x the B fragments, slice
579-- by slice; the first product waits on `first`, each slice after on SB2
580def sqProducts = (lambda unrestricted accumulator : (pi unrestricted nt : Nat . Nat) . (lambda unrestricted a : Nat .
581  (lambda unrestricted first : Nat . (lambda unrestricted tail : (family SM86Program) .
582    (saFor 4 (lambda unrestricted k : Nat .
583      (saFor 8 (lambda unrestricted nt : Nat .
584        (saHmma (accumulator nt) (naturalAdd a (naturalMultiply k 4)) (saB k nt) saSB2
585          (naturalSelect (naturalIsZero nt) (naturalSelect (naturalIsZero k) first saWait2) saWaitNone)))))
586      tail)))))
587
588-- P = 2^(t - L) and dS = P (dP - D), into S's registers (the products have
589-- finished: the first operation waits for them)
590def sqScoreGradient = (lambda unrestricted tail : (family SM86Program) .
591  (saFor 32 (lambda unrestricted i : Nat .
592    (let unrestricted nt = (naturalDivideUnchecked i 4) in (let unrestricted h = (naturalDivideUnchecked (naturalModuloUnchecked i 4) 2) in
593      (let unrestricted b = (naturalModuloUnchecked i 2) in
594      (let unrestricted e = (saElem (saS nt) h b) in
595      (lambda unrestricted rest : (family SM86Program) .
596        (saFfma e e saScale (sqNegL h)
597          (saEx2 e e saWaitNone rest))))))))
598  (saFor 32 (lambda unrestricted i : Nat .
599    (let unrestricted nt = (naturalDivideUnchecked i 4) in (let unrestricted h = (naturalDivideUnchecked (naturalModuloUnchecked i 4) 2) in
600      (let unrestricted b = (naturalModuloUnchecked i 2) in
601      (let unrestricted e = (saElem (saS nt) h b) in (let unrestricted g = (saElem (sqDP nt) h b) in
602      (lambda unrestricted rest : (family SM86Program) .
603        (saFadd g g (sqNegD h) saWaitNone
604          (saFmul e e g (naturalSelect (naturalIsZero i) saWait4 saWaitNone) rest)))))))))
605    tail)))
606
607def sqTileBody = (lambda unrestricted seq : Nat . (lambda unrestricted tail : (family SM86Program) .
608  (saImad saMask saCount 4096 saMaskBase
609  -- S = Q K^T
610  (saWide saKeyPointer saKeyOffset saOne (saArgument 1)
611  (sqLoadRows saKeyPointer saWait3
612  (saFor 32 (lambda unrestricted i : Nat . (saMovImm (naturalAdd 36 i) 0))
613  (saFor 32 (lambda unrestricted i : Nat . (saMovImm (naturalAdd 212 i) 0))
614  (sqProducts saS (saQ 0) (naturalAdd saWait0 saWait1)
615  -- dP = dO V^T (V's rows at K's offset)
616  (saWide saKeyPointer saKeyOffset saOne (saArgument 2)
617  (sqLoadRows saKeyPointer saWait2
618  (sqProducts sqDP (sqDOut 0) saWait0
619  -- K^T's fragments for dQ, once dP's products are done
620  (saWide saValuePointer sqKTOffset saOne (saArgument 3)
621  (sqLoadKeysT seq
622  (saMaskTile
623  (sqScoreGradient
624  (saPackP
625  (saValues
626    tail)))))))))))))))))
627
628def sqLoopTail = (lambda unrestricted bodyInstructions : Nat .
629  (saAddImm saKeyOffset saKeyOffset (naturalMultiply saTile (naturalMultiply saHeadWidth saHalfBytes))
630  (saAddImm sqKTOffset sqKTOffset (naturalMultiply saTile saHalfBytes)
631  (saAddImm saCount saCount 4294967295
632  (saGreater saP1 saCount 0
633  (saWhen saP1
634    (constructor SM86InstructionBody SM86Branch
635      (saU (naturalSaturatingSubtract 4294967296 (naturalMultiply 16 (naturalAdd bodyInstructions 5))))
636      (sm86Unsigned32 (byte 255) (byte 255) (byte 131) (byte 3))
637      saPlain)
638    sm86ProgramEmpty))))))
639
640-- D for the thread's rows: its sixteen dO values of row r (and of r + 8)
641-- against O's, summed over the quad; -D and -L kept, D written by the
642-- quad's first lane
643def sqRowTerms = (lambda unrestricted seq : Nat . (lambda unrestricted heads : Nat . (lambda unrestricted tail : (family SM86Program) .
644  (let unrestricted rowBytes = (naturalMultiply heads (naturalMultiply saHeadWidth saHalfBytes)) in
645  -- O in dO's fragment pattern, into the B registers (free before the loop)
646  (saWide saPointer saScratch saOne (saArgument 5)
647  (saFor 16 (lambda unrestricted i : Nat .
648    (let unrestricted kt = (naturalDivideUnchecked i 4) in (let unrestricted j = (naturalModuloUnchecked i 4) in
649      (saLoad (naturalAdd 100 i) saPointer
650        (naturalAdd (naturalMultiply kt 32)
651          (naturalAdd (naturalMultiply (naturalModuloUnchecked j 2) (naturalMultiply 8 rowBytes))
652            (naturalMultiply (naturalDivideUnchecked j 2) 16)))
653        saSB0 saWaitNone))))
654  (saMovImm (sqTmp 0) 0 (saMovImm (sqTmp 1) 0
655  -- register j of slice kt is row r (j even) or r + 8 (j odd); both halves
656  (saFor 32 (lambda unrestricted i : Nat .
657    (let unrestricted r = (naturalDivideUnchecked i 2) in (let unrestricted high = (naturalModuloUnchecked i 2) in
658      (let unrestricted h = (naturalModuloUnchecked (naturalModuloUnchecked r 4) 2) in
659      (lambda unrestricted rest : (family SM86Program) .
660        (sbWiden (sqTmp 2) (naturalAdd 100 r) high (naturalSelect (naturalIsZero i) (naturalAdd saWait0 saWait1) saWaitNone)
661        (sbWiden (sqTmp 3) (naturalAdd (sqDOut 0) r) high saWaitNone
662        (saFfma (sqTmp h) (sqTmp 2) (sqTmp 3) (sqTmp h)
663          rest))))))))
664  (saFor 2 (lambda unrestricted h : Nat . (lambda unrestricted rest : (family SM86Program) .
665    (saShfl (sqTmp 2) (sqTmp h) 1 saWaitNone
666    (saFadd (sqTmp h) (sqTmp h) (sqTmp 2) saWait5
667    (saShfl (sqTmp 2) (sqTmp h) 2 saWaitNone
668    (saFadd (sqTmp h) (sqTmp h) (sqTmp 2) saWait5
669    (saFneg (sqNegD h) (sqTmp h)
670      rest)))))))
671  -- D and L at 4 (h seq + row): D by the quad's first lane; -L
672  (saWide sqRowTermStorePointer saRow saOne (saArgument 8)
673  (saGreater saP0 saQuadLane 0
674  (saUnless saP0 (saStore sqRowTermStorePointer (sqTmp 0) 0 saWaitNone)
675  (saUnless saP0 (saStore sqRowTermStorePointer (sqTmp 1) 32 saWaitNone)
676  (saWide saPointer saRow saOne (saArgument 6)
677  (saLoad (sqTmp 2) saPointer 0 saSB1 saWaitNone
678  (saLoad (sqTmp 3) saPointer 32 saSB1 saWaitNone
679  (sbNegate (sqNegL 0) (sqTmp 2) saWait1
680  (sbNegate (sqNegL 1) (sqTmp 3) saWaitNone
681    tail)))))))))))))))))))
682
683-- dQ = s (the sum): the thread's rows r and r + 8, pairs of binary32 columns
684def sqEpilogue = (lambda unrestricted tail : (family SM86Program) .
685  (saMovConst (sqTmp 4) (saArgument 10)
686  (saFor 32 (lambda unrestricted i : Nat .
687    (saFmul (naturalAdd 68 i) (naturalAdd 68 i) (sqTmp 4) (naturalSelect (naturalIsZero i) saWait3 saWaitNone)))
688  (saWide saPointer saScratch saOne (saArgument 7)
689  (saFor 16 (lambda unrestricted i : Nat .
690    (let unrestricted nt = (naturalDivideUnchecked i 2) in (let unrestricted h = (naturalModuloUnchecked i 2) in
691      (saOp (sbStore64 saPointer (saElem (saO nt) h 0)
692        (naturalAdd (naturalMultiply nt 32) (naturalMultiply h (naturalMultiply 8 (naturalMultiply saHeadWidth 4))))
693        saWaitNone)))))
694    tail)))))
695
696def streamingAttentionQueryGroupedUncheckedSM86 = (lambda unrestricted seq : Nat . (lambda unrestricted heads : Nat . (lambda unrestricted keyValueHeads : Nat .
697  (let unrestricted groups = (naturalDivideUnchecked heads keyValueHeads) in
698  (let unrestricted headBytes = (naturalMultiply seq (naturalMultiply saHeadWidth saHalfBytes)) in
699  (let unrestricted rowBytes = (naturalMultiply heads (naturalMultiply saHeadWidth saHalfBytes)) in
700  (let unrestricted body = (sqTileBody seq sm86ProgramEmpty) in
701  (saS2R saTid (constructor SM86SpecialRegister SM86ThreadIdX)
702  (saS2R saTileIndex (constructor SM86SpecialRegister SM86CooperativeThreadArrayIdX)
703  (saS2R saKeyValueHead (constructor SM86SpecialRegister SM86CooperativeThreadArrayIdY)
704  (saS2R saHead (constructor SM86SpecialRegister SM86CooperativeThreadArrayIdZ)
705  (saMovImm saOne 1
706  (saMovImmAfter saScratch 31 saWait5
707  (saImad saHead saKeyValueHead groups saHead
708  (saAnd saLane saTid saScratch
709  (saShr saWarp saTid 5
710  (saShr saQuad saLane 2
711  (saMovImm saScratch 3
712  (saAnd saQuadLane saLane saScratch
713  (saImad saRow saWarp 16 saQuad
714  (saImad saRow saTileIndex saTile saRow
715  -- Q's and dO's A fragments: head h's plane, the row, the quad lane's pair
716  (saMovImm saScratch 0
717  (saImad saScratch saRow (naturalMultiply saHeadWidth saHalfBytes) saScratch
718  (saImad saScratch saQuadLane 4 saScratch
719  (saImad saScratch saHead headBytes saScratch
720  (saWide saPointer saScratch saOne (saArgument 0)
721  (saFor 16 (lambda unrestricted i : Nat .
722    (let unrestricted kt = (naturalDivideUnchecked i 4) in (let unrestricted j = (naturalModuloUnchecked i 4) in
723      (saLoad (saQReg kt j) saPointer
724        (naturalAdd (naturalMultiply kt 32)
725          (naturalAdd (naturalMultiply (naturalModuloUnchecked j 2) (naturalMultiply 8 (naturalMultiply saHeadWidth saHalfBytes)))
726            (naturalMultiply (naturalDivideUnchecked j 2) 16)))
727        saSB1 saWaitNone))))
728  (saWide sqOutputGradientLoadPointer saScratch saOne (saArgument 4)
729  (saFor 16 (lambda unrestricted i : Nat .
730    (let unrestricted kt = (naturalDivideUnchecked i 4) in (let unrestricted j = (naturalModuloUnchecked i 4) in
731      (saLoad (naturalAdd (sqDOut kt) j) sqOutputGradientLoadPointer
732        (naturalAdd (naturalMultiply kt 32)
733          (naturalAdd (naturalMultiply (naturalModuloUnchecked j 2) (naturalMultiply 8 (naturalMultiply saHeadWidth saHalfBytes)))
734            (naturalMultiply (naturalDivideUnchecked j 2) 16)))
735        saSB1 saWaitNone))))
736  -- dQ's offset (binary32 rows), kept in saKeyPointer's partner until the
737  -- epilogue: saScratch becomes O's (merged plane) for the row terms
738  (saMovImm saKeyOffset 0
739  (saImad saKeyOffset saRow (naturalMultiply saHeadWidth 4) saKeyOffset
740  (saImad saKeyOffset saQuadLane 8 saKeyOffset
741  (saImad saKeyOffset saHead (naturalMultiply seq (naturalMultiply saHeadWidth 4)) saKeyOffset
742  (saMovImm saScratch 0
743  (saImad saScratch saRow rowBytes saScratch
744  (saImad saScratch saHead (naturalMultiply saHeadWidth saHalfBytes) saScratch
745  (saImad saScratch saQuadLane 4 saScratch
746  -- the row terms' offset: 4 (h seq + row)
747  (saImad saRow saHead seq saRow
748  (saMovImm saValueOffset 0
749  (saImad saRow saRow 4 saValueOffset
750  (sqRowTerms seq heads
751  -- dQ's offset into saScratch; K's and V's from key 0, K^T's
752  (saAddImm saScratch saKeyOffset 0
753  (saMovImm saKeyOffset 0
754  (saImad saKeyOffset saQuad (naturalMultiply saHeadWidth saHalfBytes) saKeyOffset
755  (saImad saKeyOffset saQuadLane 4 saKeyOffset
756  (saImad saKeyOffset saKeyValueHead headBytes saKeyOffset
757  (saMovImm sqKTOffset 0
758  (saImad sqKTOffset saQuad (naturalMultiply seq saHalfBytes) sqKTOffset
759  (saImad sqKTOffset saQuadLane 4 sqKTOffset
760  (saImad sqKTOffset saKeyValueHead headBytes sqKTOffset
761  (saAddImm saCount saTileIndex 1
762  (saImad saMaskBase saWarp 16 saQuad
763  (saImad saMaskBase saQuadLane 4294967294 saMaskBase
764  (saAddImm saMaskBase saMaskBase (naturalSaturatingSubtract 4294967296 (naturalSaturatingSubtract 4096 64))
765  (saMovConst saScale (saArgument 9)
766  (saFor 32 (lambda unrestricted i : Nat . (saMovImm (naturalAdd 68 i) 0))
767  (sm86ProgramAppend body
768    (sm86ProgramAppend (sqLoopTail (sm86ProgramCount body))
769      (sqEpilogue
770        (saOp (constructor SM86InstructionBody SM86Exit (saAfter saWaitAll)) sm86ProgramEmpty))))))))))))))))))))))))))))))))))))))))))))))))))))))))))))
771
772def streamingAttentionQueryGroupedSM86 =
773  (lambda unrestricted seq : Nat . (lambda unrestricted heads : Nat .
774    (lambda unrestricted keyValueHeads : Nat .
775      (lambda erased admitted : (equal Nat (streamingAttentionGroupedSM86Admitted seq heads keyValueHeads) 1) .
776        (streamingAttentionQueryGroupedUncheckedSM86 seq heads keyValueHeads)))))
777def streamingAttentionQuerySM86 = (lambda unrestricted seq : Nat . (lambda unrestricted heads : Nat .
778  (streamingAttentionQueryGroupedUncheckedSM86 seq heads heads)))
779
780def streamingAttentionQuerySM86Registers : Nat = sqRegisters
781def streamingAttentionQuerySM86Threads : Nat = saThreads
782def streamingAttentionQuerySM86ConstantBytes : Nat = (saArgument 11)
783def streamingAttentionQueryQueryArgument : Nat = 0
784def streamingAttentionQueryKeyArgument : Nat = 1
785def streamingAttentionQueryValueArgument : Nat = 2
786def streamingAttentionQueryTransposedKeyArgument : Nat = 3
787def streamingAttentionQueryOutputGradientArgument : Nat = 4
788def streamingAttentionQueryAttentionOutputArgument : Nat = 5
789def streamingAttentionQueryLogSumArgument : Nat = 6
790def streamingAttentionQueryQueryGradientArgument : Nat = 7
791def streamingAttentionQueryRowTermArgument : Nat = 8
792def streamingAttentionQueryBaseTwoScaleArgument : Nat = 9
793def streamingAttentionQueryNaturalScaleArgument : Nat = 10
794
795-- The key launch (streamingAttentionKeySM86): a block of two warps per 32
796-- keys of a K/V head, grid (T / 32, K/V heads, groups). A warp holds its 16 keys' K and V
797-- rows as A fragments and walks the 32-query sub-tiles from the last one
798-- down to its diagonal (the last, masked): S^T = K Q^T, dP^T = V dO^T,
799-- P^T = 2^(S^T c - L), dS^T = P^T (dP^T - D) with L and D the queries',
800-- dV += P^T dO and dK += dS^T Q.  Operands: Q, K, V, dO ([head][T][64],
801-- half), Q^T, dO^T ([head][64][T], half), L and D ([head][T]); results dK
802-- ([head][T][64]) and dV^T ([head][64][T]), binary32.  Parameters: Q, K,
803-- V, dO, Q^T, dO^T, L, D, dK, dV^T, then c and s.
804
805def skK = (lambda unrestricted kt : Nat . (naturalAdd 20 (naturalMultiply kt 4)))
806def skV = (lambda unrestricted kt : Nat . (naturalAdd 36 (naturalMultiply kt 4)))
807def skS = (lambda unrestricted nt : Nat . (naturalAdd 52 (naturalMultiply nt 4)))
808def skDP = (lambda unrestricted nt : Nat . (naturalAdd 68 (naturalMultiply nt 4)))
809def skDV = (lambda unrestricted nt : Nat . (naturalAdd 84 (naturalMultiply nt 4)))
810def skDK = (lambda unrestricted nt : Nat . (naturalAdd 116 (naturalMultiply nt 4)))
811-- the B fragments of slice k, n-tile nt among `tiles` n-tiles
812def skB = (lambda unrestricted tiles : Nat . (lambda unrestricted k : Nat . (lambda unrestricted nt : Nat .
813  (naturalAdd 148 (naturalMultiply (naturalAdd (naturalMultiply k tiles) nt) 2)))))
814def skPA = (lambda unrestricted kk : Nat . (naturalAdd 180 (naturalMultiply kk 4)))
815def skSA = (lambda unrestricted kk : Nat . (naturalAdd 188 (naturalMultiply kk 4)))
816-- -L and -D of the thread's columns: n-tile nt, column b
817def skNegL = (lambda unrestricted nt : Nat . (lambda unrestricted b : Nat . (naturalAdd 196 (naturalAdd (naturalMultiply nt 2) b))))
818def skNegD = (lambda unrestricted nt : Nat . (lambda unrestricted b : Nat . (naturalAdd 204 (naturalAdd (naturalMultiply nt 2) b))))
819def skTmp = (lambda unrestricted n : Nat . (naturalAdd 212 n))
820def skHead : Nat = 220
821def skRow : Nat = 221
822def skLogOffset : Nat = 222
823def skKeyValueHead : Nat = 223
824-- dK stores remain in flight while the dV address is formed. Their pair
825-- cannot share the address registers used by the later dV stores.
826def skKeyGradientStorePointer : Nat = 224
827def skRegisters : Nat = 232
828def skSubTile : Nat = 32
829def skThreads : Nat = 64
830
831-- The wide grouped launch can leave the 3090's channel waiting with no
832-- retired fence. A per-head launch uses the same body and output planes but
833-- reads its K/V head and within-group query head from one CB0 slot. Keeping
834-- this selection in the realization avoids copying the backward kernel into
835-- a system; the launch schedule determines the head traversal.
836def skLoadHeadIndices = (lambda unrestricted fromConstant : Nat .
837  (lambda unrestricted tail : (family SM86Program) .
838    (nat-eliminate (lambda unrestricted mode : Nat . (family SM86Program))
839      (saS2R skKeyValueHead (constructor SM86SpecialRegister SM86CooperativeThreadArrayIdY)
840        (saS2R skHead (constructor SM86SpecialRegister SM86CooperativeThreadArrayIdZ) tail))
841      (lambda unrestricted predecessor : Nat . (lambda unrestricted induction : (family SM86Program) .
842        (saMovConst skKeyValueHead (naturalAdd (saArgument 12) 4)
843          (saMovConst skHead (saArgument 12) tail))))
844      fromConstant)))
845
846-- A key launch may own a contiguous subset of key tiles. The caller derives
847-- the subset's base from its launch schedule; the ordinary unsplit kernel
848-- keeps using CTA x directly. This keeps all query-loop and output-address
849-- semantics in one implementation while bounding the work of one launch.
850def skLoadTileBase = (lambda unrestricted fromConstant : Nat .
851  (lambda unrestricted tail : (family SM86Program) .
852    (nat-eliminate (lambda unrestricted mode : Nat . (family SM86Program))
853      tail
854      (lambda unrestricted predecessor : Nat . (lambda unrestricted induction : (family SM86Program) .
855        (saMovConst saScratch (saArgument 13)
856          (saImad saTileIndex saScratch 1 saTileIndex tail))))
857      fromConstant)))
858
859-- rows ([T][64] half) as the B fragments of a product over the 64 columns:
860-- slices k 0..3, n-tiles nt 0..3 (the sub-tile's 32 rows)
861def skLoadRows = (lambda unrestricted pointer : Nat . (lambda unrestricted wait : Nat . (lambda unrestricted tail : (family SM86Program) .
862  (saFor 16 (lambda unrestricted i : Nat .
863    (let unrestricted k = (naturalDivideUnchecked i 4) in (let unrestricted nt = (naturalModuloUnchecked i 4) in
864      (lambda unrestricted rest : (family SM86Program) .
865        (saLoad (skB 4 k nt) pointer (naturalAdd (naturalMultiply nt 1024) (naturalMultiply k 32)) saSB0 (naturalSelect (naturalIsZero i) wait saWaitNone)
866        (saLoad (succ (skB 4 k nt)) pointer (naturalAdd (naturalMultiply nt 1024) (naturalAdd (naturalMultiply k 32) 16)) saSB0 saWaitNone
867          rest))))))
868    tail))))
869
870-- transposed rows ([64][T] half) as the B fragments of a product over the
871-- sub-tile's 32 positions: slices kk 0..1, n-tiles nt 0..7 (the 64 columns)
872def skLoadColumns = (lambda unrestricted seq : Nat . (lambda unrestricted pointer : Nat . (lambda unrestricted wait : Nat .
873  (lambda unrestricted tail : (family SM86Program) .
874    (saFor 16 (lambda unrestricted i : Nat .
875      (let unrestricted k = (naturalDivideUnchecked i 8) in (let unrestricted nt = (naturalModuloUnchecked i 8) in
876        (let unrestricted rowOffset = (naturalMultiply nt (naturalMultiply 8 (naturalMultiply seq saHalfBytes))) in
877        (lambda unrestricted rest : (family SM86Program) .
878          (saLoad (skB 8 k nt) pointer (naturalAdd rowOffset (naturalMultiply k 32)) saSB0 (naturalSelect (naturalIsZero i) wait saWaitNone)
879          (saLoad (succ (skB 8 k nt)) pointer (naturalAdd rowOffset (naturalAdd (naturalMultiply k 32) 16)) saSB0 saWaitNone
880            rest)))))))
881      tail)))))
882
883def skProducts = (lambda unrestricted accumulator : (pi unrestricted nt : Nat . Nat) . (lambda unrestricted a : (pi unrestricted k : Nat . Nat) .
884  (lambda unrestricted slices : Nat . (lambda unrestricted tiles : Nat . (lambda unrestricted barrier : (family SM86Barrier) .
885  (lambda unrestricted wait : Nat . (lambda unrestricted first : Nat . (lambda unrestricted tail : (family SM86Program) .
886    (saFor slices (lambda unrestricted k : Nat .
887      (saFor tiles (lambda unrestricted nt : Nat .
888        (saHmma (accumulator nt) (a k) (skB tiles k nt) barrier
889          (naturalSelect (naturalIsZero nt) (naturalSelect (naturalIsZero k) first wait) saWaitNone)))))
890      tail)))))))))
891
892-- the diagonal sub-tile's mask: element (key r + 8h, query 8 nt + 2 (lane % 4)
893-- + b) is kept when the query is not before the key: saMask holds
894-- 2 (lane % 4) - 16 warp - lane / 4 + 128, plus 4096 for every sub-tile
895-- before the diagonal one
896def skMaskTile = (lambda unrestricted tail : (family SM86Program) .
897  (saFor 16 (lambda unrestricted i : Nat .
898    (let unrestricted nt = (naturalDivideUnchecked i 4) in (let unrestricted h = (naturalDivideUnchecked (naturalModuloUnchecked i 4) 2) in
899      (let unrestricted b = (naturalModuloUnchecked i 2) in
900      (let unrestricted threshold = (naturalSaturatingSubtract (naturalAdd 127 (naturalMultiply 8 h)) (naturalAdd (naturalMultiply 8 nt) b)) in
901      (lambda unrestricted rest : (family SM86Program) .
902        (saGreater saP0 saMask threshold
903          (saUnless saP0 (constructor SM86InstructionBody SM86MoveImmediate (saR (saElem (skS nt) h b)) (saU saMinusInfinity) saPlain)
904            rest))))))))
905    tail))
906
907-- P^T = 2^(t - L), dS^T = P^T (dP^T - D), in S^T's registers
908def skScoreGradient = (lambda unrestricted tail : (family SM86Program) .
909  (saFor 16 (lambda unrestricted i : Nat .
910    (let unrestricted nt = (naturalDivideUnchecked i 4) in (let unrestricted h = (naturalDivideUnchecked (naturalModuloUnchecked i 4) 2) in
911      (let unrestricted b = (naturalModuloUnchecked i 2) in
912      (let unrestricted e = (saElem (skS nt) h b) in
913      (lambda unrestricted rest : (family SM86Program) .
914        (saFfma e e saScale (skNegL nt b)
915          (saEx2 e e saWaitNone rest))))))))
916  -- P^T's fragments, before dS^T replaces it
917  (saFor 2 (lambda unrestricted kk : Nat .
918    (let unrestricted left = (skS (naturalMultiply 2 kk)) in (let unrestricted right = (skS (succ (naturalMultiply 2 kk))) in
919      (lambda unrestricted rest : (family SM86Program) .
920        (saOp (constructor SM86InstructionBody SM86FloatPairToPackedHalfPair (saR (skPA kk)) (saR (succ left)) (saR left)
921          (saAfter (naturalSelect (naturalIsZero kk) saWait4 saWaitNone)))
922        (saPack (succ (skPA kk)) (naturalAdd left 3) (naturalAdd left 2)
923        (saPack (naturalAdd (skPA kk) 2) (succ right) right
924        (saPack (naturalAdd (skPA kk) 3) (naturalAdd right 3) (naturalAdd right 2)
925          rest))))))))
926  (saFor 16 (lambda unrestricted i : Nat .
927    (let unrestricted nt = (naturalDivideUnchecked i 4) in (let unrestricted h = (naturalDivideUnchecked (naturalModuloUnchecked i 4) 2) in
928      (let unrestricted b = (naturalModuloUnchecked i 2) in
929      (let unrestricted e = (saElem (skS nt) h b) in (let unrestricted g = (saElem (skDP nt) h b) in
930      (lambda unrestricted rest : (family SM86Program) .
931        (saFadd g g (skNegD nt b) saWaitNone
932          (saFmul e e g saWaitNone rest)))))))))
933  (saFor 2 (lambda unrestricted kk : Nat .
934    (let unrestricted left = (skS (naturalMultiply 2 kk)) in (let unrestricted right = (skS (succ (naturalMultiply 2 kk))) in
935      (lambda unrestricted rest : (family SM86Program) .
936        (saPack (skSA kk) (succ left) left
937        (saPack (succ (skSA kk)) (naturalAdd left 3) (naturalAdd left 2)
938        (saPack (naturalAdd (skSA kk) 2) (succ right) right
939        (saPack (naturalAdd (skSA kk) 3) (naturalAdd right 3) (naturalAdd right 2)
940          rest))))))))
941    tail)))))
942
943def skTileBody = (lambda unrestricted seq : Nat . (lambda unrestricted tail : (family SM86Program) .
944  (saImad saMask saCount 4096 saMaskBase
945  -- the sub-tile's L and D (their negatives)
946  (saWide saPointer skLogOffset saOne (saArgument 6)
947  (saFor 8 (lambda unrestricted i : Nat .
948    (saLoad (skTmp i) saPointer (naturalAdd (naturalMultiply (naturalDivideUnchecked i 2) 32) (naturalMultiply (naturalModuloUnchecked i 2) 4))
949      saSB1 (naturalSelect (naturalIsZero i) saWait3 saWaitNone)))
950  (saFor 8 (lambda unrestricted i : Nat .
951    (sbNegate (skNegL (naturalDivideUnchecked i 2) (naturalModuloUnchecked i 2)) (skTmp i) (naturalSelect (naturalIsZero i) saWait1 saWaitNone)))
952  (saWide saPointer skLogOffset saOne (saArgument 7)
953  (saFor 8 (lambda unrestricted i : Nat .
954    (saLoad (skTmp i) saPointer (naturalAdd (naturalMultiply (naturalDivideUnchecked i 2) 32) (naturalMultiply (naturalModuloUnchecked i 2) 4))
955      saSB1 saWaitNone))
956  (saFor 8 (lambda unrestricted i : Nat .
957    (sbNegate (skNegD (naturalDivideUnchecked i 2) (naturalModuloUnchecked i 2)) (skTmp i) (naturalSelect (naturalIsZero i) saWait1 saWaitNone)))
958  -- S^T = K Q^T
959  (saWide saKeyPointer saKeyOffset saOne (saArgument 0)
960  (skLoadRows saKeyPointer saWait3
961  (saFor 16 (lambda unrestricted i : Nat . (saMovImm (naturalAdd 52 i) 0))
962  (saFor 16 (lambda unrestricted i : Nat . (saMovImm (naturalAdd 68 i) 0))
963  (skProducts skS skK 4 4 saSB2 saWait2 (naturalAdd saWait0 saWait1)
964  -- dP^T = V dO^T
965  (saWide saKeyPointer saKeyOffset saOne (saArgument 3)
966  (skLoadRows saKeyPointer saWait2
967  (skProducts skDP skV 4 4 saSB2 saWait2 saWait0
968  -- dO^T's fragments for dV, once dP^T's products are done
969  (saWide saValuePointer saValueOffset saOne (saArgument 5)
970  (skLoadColumns seq saValuePointer saWait2
971  (skMaskTile
972  (skScoreGradient
973  (skProducts skDV skPA 2 8 saSB3 saWait3 saWait0
974  -- Q^T's fragments for dK, once dV's products have read dO^T's
975  (saWide saValuePointer saValueOffset saOne (saArgument 4)
976  (skLoadColumns seq saValuePointer saWait3
977  (skProducts skDK skSA 2 8 saSB3 saWait3 saWait0
978    tail)))))))))))))))))))))))))
979
980def skLoopTail = (lambda unrestricted bodyInstructions : Nat .
981  (saAddImm saKeyOffset saKeyOffset (naturalSaturatingSubtract 4294967296 (naturalMultiply skSubTile (naturalMultiply saHeadWidth saHalfBytes)))
982  (saAddImm saValueOffset saValueOffset (naturalSaturatingSubtract 4294967296 (naturalMultiply skSubTile saHalfBytes))
983  (saAddImm skLogOffset skLogOffset (naturalSaturatingSubtract 4294967296 (naturalMultiply skSubTile 4))
984  (saAddImm saCount saCount 4294967295
985  (saGreater saP1 saCount 0
986  (saWhen saP1
987    (constructor SM86InstructionBody SM86Branch
988      (saU (naturalSaturatingSubtract 4294967296 (naturalMultiply 16 (naturalAdd bodyInstructions 6))))
989      (sm86Unsigned32 (byte 255) (byte 255) (byte 131) (byte 3))
990      saPlain)
991    sm86ProgramEmpty)))))))
992
993-- dK = s (the sum), rows; dV^T, columns: element (key r + 8h, column
994-- 8 nt + 2 (lane % 4) + b) at (8 nt + b) seq + 8 h words past the thread's
995def skEpilogue = (lambda unrestricted seq : Nat . (lambda unrestricted tail : (family SM86Program) .
996  (saMovConst (skTmp 0) (saArgument 11)
997  (saFor 32 (lambda unrestricted i : Nat .
998    (saFmul (naturalAdd 116 i) (naturalAdd 116 i) (skTmp 0) (naturalSelect (naturalIsZero i) saWait3 saWaitNone)))
999  (saWide skKeyGradientStorePointer saScratch saOne (saArgument 8)
1000  (saFor 16 (lambda unrestricted i : Nat .
1001    (let unrestricted nt = (naturalDivideUnchecked i 2) in (let unrestricted h = (naturalModuloUnchecked i 2) in
1002      (saOp (sbStore64 skKeyGradientStorePointer (saElem (skDK nt) h 0)
1003        (naturalAdd (naturalMultiply nt 32) (naturalMultiply h (naturalMultiply 8 (naturalMultiply saHeadWidth 4))))
1004        saWaitNone)))))
1005  -- dV^T's offset: head h's 64 x seq plane, the quad lane's first column,
1006  -- the key
1007  (saMovImm (skTmp 1) 0
1008  (saImad (skTmp 1) saQuadLane (naturalMultiply 8 seq) (skTmp 1)
1009  (saImad (skTmp 1) skHead (naturalMultiply saHeadWidth (naturalMultiply seq 4)) (skTmp 1)
1010  (saImad (skTmp 1) skRow 4 (skTmp 1)
1011  (saWide saPointer (skTmp 1) saOne (saArgument 9)
1012  (saFor 32 (lambda unrestricted i : Nat .
1013    (let unrestricted nt = (naturalDivideUnchecked i 4) in (let unrestricted h = (naturalDivideUnchecked (naturalModuloUnchecked i 4) 2) in
1014      (let unrestricted b = (naturalModuloUnchecked i 2) in
1015      (saOp (saStore saPointer (saElem (skDV nt) h b)
1016        (naturalAdd (naturalMultiply (naturalAdd (naturalMultiply 8 nt) b) (naturalMultiply seq 4)) (naturalMultiply h 32))
1017        saWaitNone))))))
1018    tail))))))))))))
1019
1020-- Each query head writes a separate dK/dV plane.  A later reduction sums
1021-- those planes into the compact K/V gradient; concurrent CTAs never race on
1022-- one compact gradient address.
1023def streamingAttentionKeyGroupedWithHeadSourceUncheckedSM86 = (lambda unrestricted seq : Nat . (lambda unrestricted heads : Nat . (lambda unrestricted keyValueHeads : Nat . (lambda unrestricted fromConstant : Nat . (lambda unrestricted fromTileBase : Nat .
1024  (let unrestricted groups = (naturalDivideUnchecked heads keyValueHeads) in
1025  (let unrestricted headBytes = (naturalMultiply seq (naturalMultiply saHeadWidth saHalfBytes)) in
1026  (let unrestricted subTiles = (naturalDivideUnchecked seq skSubTile) in
1027  (let unrestricted body = (skTileBody seq sm86ProgramEmpty) in
1028  (saS2R saTid (constructor SM86SpecialRegister SM86ThreadIdX)
1029  (saS2R saTileIndex (constructor SM86SpecialRegister SM86CooperativeThreadArrayIdX)
1030  (skLoadTileBase fromTileBase
1031  (skLoadHeadIndices fromConstant
1032  (saMovImm saOne 1
1033  (saMovImmAfter saScratch 31 saWait5
1034  (saImad skHead skKeyValueHead groups skHead
1035  (saAnd saLane saTid saScratch
1036  (saShr saWarp saTid 5
1037  (saShr saQuad saLane 2
1038  (saMovImm saScratch 3
1039  (saAnd saQuadLane saLane saScratch
1040  -- the thread's key row: 32 tile + 16 warp + lane / 4
1041  (saImad skRow saWarp 16 saQuad
1042  (saImad skRow saTileIndex skSubTile skRow
1043  -- K's and V's A fragments
1044  (saMovImm saScratch 0
1045  (saImad saScratch skRow (naturalMultiply saHeadWidth saHalfBytes) saScratch
1046  (saImad saScratch saQuadLane 4 saScratch
1047  (saImad saScratch skKeyValueHead headBytes saScratch
1048  (saWide saPointer saScratch saOne (saArgument 1)
1049  (saFor 16 (lambda unrestricted i : Nat .
1050    (let unrestricted kt = (naturalDivideUnchecked i 4) in (let unrestricted j = (naturalModuloUnchecked i 4) in
1051      (saLoad (naturalAdd (skK kt) j) saPointer
1052        (naturalAdd (naturalMultiply kt 32)
1053          (naturalAdd (naturalMultiply (naturalModuloUnchecked j 2) (naturalMultiply 8 (naturalMultiply saHeadWidth saHalfBytes)))
1054            (naturalMultiply (naturalDivideUnchecked j 2) 16)))
1055        saSB1 saWaitNone))))
1056  (saWide saPointer saScratch saOne (saArgument 2)
1057  (saFor 16 (lambda unrestricted i : Nat .
1058    (let unrestricted kt = (naturalDivideUnchecked i 4) in (let unrestricted j = (naturalModuloUnchecked i 4) in
1059      (saLoad (naturalAdd (skV kt) j) saPointer
1060        (naturalAdd (naturalMultiply kt 32)
1061          (naturalAdd (naturalMultiply (naturalModuloUnchecked j 2) (naturalMultiply 8 (naturalMultiply saHeadWidth saHalfBytes)))
1062            (naturalMultiply (naturalDivideUnchecked j 2) 16)))
1063        saSB1 saWaitNone))))
1064  -- dK's offset (binary32 rows) in saScratch
1065  (saMovImm saScratch 0
1066  (saImad saScratch skRow (naturalMultiply saHeadWidth 4) saScratch
1067  (saImad saScratch saQuadLane 8 saScratch
1068  (saImad saScratch skHead (naturalMultiply seq (naturalMultiply saHeadWidth 4)) saScratch
1069  -- the queries start at the last sub-tile: Q's and dO's rows (row lane / 4,
1070  -- the quad lane's pair), Q^T's and dO^T's columns (row lane / 4), L's and D's
1071  (saMovImm saKeyOffset 0
1072  (saImad saKeyOffset saQuad (naturalMultiply saHeadWidth saHalfBytes) saKeyOffset
1073  (saImad saKeyOffset saQuadLane 4 saKeyOffset
1074  (saImad saKeyOffset skHead headBytes saKeyOffset
1075  (saAddImm saKeyOffset saKeyOffset (naturalMultiply (naturalSaturatingSubtract subTiles 1) (naturalMultiply skSubTile (naturalMultiply saHeadWidth saHalfBytes)))
1076  (saMovImm saValueOffset 0
1077  (saImad saValueOffset saQuad (naturalMultiply seq saHalfBytes) saValueOffset
1078  (saImad saValueOffset saQuadLane 4 saValueOffset
1079  (saImad saValueOffset skHead headBytes saValueOffset
1080  (saAddImm saValueOffset saValueOffset (naturalMultiply (naturalSaturatingSubtract subTiles 1) (naturalMultiply skSubTile saHalfBytes))
1081  (saMovImm skLogOffset 0
1082  (saImad skLogOffset saQuadLane 8 skLogOffset
1083  (saImad skLogOffset skHead (naturalMultiply seq 4) skLogOffset
1084  (saAddImm skLogOffset skLogOffset (naturalMultiply (naturalSaturatingSubtract subTiles 1) (naturalMultiply skSubTile 4))
1085  -- the sub-tiles from the last down to the diagonal; the mask's base,
1086  -- 2 (lane % 4) - 16 warp - lane / 4 + 128 - 4096
1087  (saMovImm saCount subTiles
1088  (saImad saCount saTileIndex 4294967295 saCount
1089  (saMovImm saMaskBase 0
1090  (saImad saMaskBase saQuadLane 2 saMaskBase
1091  (saImad saMaskBase saWarp 4294967280 saMaskBase
1092  (saImad saMaskBase saQuad 4294967295 saMaskBase
1093  (saAddImm saMaskBase saMaskBase (naturalSaturatingSubtract 4294967296 (naturalSaturatingSubtract 4096 128))
1094  (saMovConst saScale (saArgument 10)
1095  (saFor 64 (lambda unrestricted i : Nat . (saMovImm (naturalAdd 84 i) 0))
1096  (sm86ProgramAppend body
1097    (sm86ProgramAppend (skLoopTail (sm86ProgramCount body))
1098      (skEpilogue seq
1099        (saOp (constructor SM86InstructionBody SM86Exit (saAfter saWaitAll)) sm86ProgramEmpty))))))))))))))))))))))))))))))))))))))))))))))))))))))))))))))
1100
1101def streamingAttentionKeyGroupedUncheckedSM86 = (lambda unrestricted seq : Nat .
1102  (lambda unrestricted heads : Nat . (lambda unrestricted keyValueHeads : Nat .
1103    (streamingAttentionKeyGroupedWithHeadSourceUncheckedSM86 seq heads keyValueHeads 0 0))))
1104
1105def streamingAttentionKeyGroupedSM86 =
1106  (lambda unrestricted seq : Nat . (lambda unrestricted heads : Nat .
1107    (lambda unrestricted keyValueHeads : Nat .
1108      (lambda erased admitted : (equal Nat (streamingAttentionGroupedSM86Admitted seq heads keyValueHeads) 1) .
1109        (streamingAttentionKeyGroupedUncheckedSM86 seq heads keyValueHeads)))))
1110def streamingAttentionKeyGroupedByConstantSM86 =
1111  (lambda unrestricted seq : Nat . (lambda unrestricted heads : Nat .
1112    (lambda unrestricted keyValueHeads : Nat .
1113      (lambda erased admitted : (equal Nat (streamingAttentionGroupedSM86Admitted seq heads keyValueHeads) 1) .
1114        (streamingAttentionKeyGroupedWithHeadSourceUncheckedSM86 seq heads keyValueHeads 1 0)))))
1115def streamingAttentionKeyGroupedByConstantAndTileSM86 =
1116  (lambda unrestricted seq : Nat . (lambda unrestricted heads : Nat .
1117    (lambda unrestricted keyValueHeads : Nat .
1118      (lambda erased admitted : (equal Nat (streamingAttentionGroupedSM86Admitted seq heads keyValueHeads) 1) .
1119        (streamingAttentionKeyGroupedWithHeadSourceUncheckedSM86 seq heads keyValueHeads 1 1)))))
1120def streamingAttentionKeySM86 = (lambda unrestricted seq : Nat . (lambda unrestricted heads : Nat .
1121  (streamingAttentionKeyGroupedUncheckedSM86 seq heads heads)))
1122
1123def streamingAttentionKeySM86Registers : Nat = skRegisters
1124def streamingAttentionKeySM86Threads : Nat = skThreads
1125def streamingAttentionKeySM86ConstantBytes : Nat = (saArgument 12)
1126def streamingAttentionKeyHeadConstantBytes : Nat = (saArgument 13)
1127def streamingAttentionKeyTileBaseConstantBytes : Nat = (saArgument 14)
1128-- The 3090 Baguette full-grid launch did not retire, while the same update
1129-- with half as many key CTAs per launch retired all its update fences on
1130-- 2026-09-28. This bound is provisional until the complete two-part update
1131-- and speed comparison qualify; it belongs to this realization, not to the
1132-- architecture-neutral grouped decoder.
1133def streamingAttentionKeyContiguousBlocksPerLaunch : Nat = 64
1134def streamingAttentionKeyTileBaseArgument : Nat = 13
1135def streamingAttentionKeyGroupArgument : Nat = 12
1136def streamingAttentionKeyQueryArgument : Nat = 0
1137def streamingAttentionKeyKeyArgument : Nat = 1
1138def streamingAttentionKeyValueArgument : Nat = 2
1139def streamingAttentionKeyOutputGradientArgument : Nat = 3
1140def streamingAttentionKeyTransposedQueryArgument : Nat = 4
1141def streamingAttentionKeyTransposedOutputGradientArgument : Nat = 5
1142def streamingAttentionKeyLogSumArgument : Nat = 6
1143def streamingAttentionKeyRowTermArgument : Nat = 7
1144def streamingAttentionKeyKeyGradientArgument : Nat = 8
1145def streamingAttentionKeyValueGradientArgument : Nat = 9
1146def streamingAttentionKeyBaseTwoScaleArgument : Nat = 10
1147def streamingAttentionKeyNaturalScaleArgument : Nat = 11
1148
1149-- Sum the query-head dK or dV planes of one K/V group into its compact
1150-- gradient plane. Both backward outputs have 64*T binary32 elements per
1151-- head, though dV is transposed within that plane. A CTA owns 128 elements
1152-- of one K/V head; no atomics or cross-CTA ordering are needed. The caller
1153-- supplies disjoint source/destination regions and launches grid
1154-- (64*T/128, keyValueHeads, 1) for T divisible by 64.
1155def streamingAttentionReduceGroupedGradientUncheckedSM86 =
1156  (lambda unrestricted seq : Nat .
1157    (lambda unrestricted heads : Nat .
1158      (lambda unrestricted keyValueHeads : Nat .
1159        (let unrestricted groups = (naturalDivideUnchecked heads keyValueHeads) in
1160        (let unrestricted planeWords = (naturalMultiply saHeadWidth seq) in
1161        (let unrestricted planeBytes = (naturalMultiply planeWords 4) in
1162        (saS2R saTid (constructor SM86SpecialRegister SM86ThreadIdX)
1163        (saS2R 2 (constructor SM86SpecialRegister SM86CooperativeThreadArrayIdX)
1164        (saS2R 3 (constructor SM86SpecialRegister SM86CooperativeThreadArrayIdY)
1165        (saMovImmAfter 7 0 saWait5
1166        (saImad 4 2 saThreads saTid
1167        (saImad 5 3 planeWords 4
1168        (saImad 6 3 (naturalMultiply groups planeWords) 4
1169        (saMovImm 11 4
1170        (saWide 12 6 11 (saArgument 0)
1171        (saWide 14 5 11 (saArgument 1)
1172        (saMovImm 9 0
1173        (saFor groups (lambda unrestricted group : Nat .
1174          (lambda unrestricted rest : (family SM86Program) .
1175            (saLoad 8 12 (naturalMultiply group planeBytes) saSB0 saWaitNone
1176              (saFadd 9 9 8 saWait0 rest))))
1177        (saOp (saStore 14 9 0 saWaitNone)
1178        (saOp (constructor SM86InstructionBody SM86Exit (saAfter saWaitAll)) sm86ProgramEmpty))))))))))))))))))))
1179
1180def streamingAttentionReduceGroupedGradientSM86 =
1181  (lambda unrestricted seq : Nat . (lambda unrestricted heads : Nat .
1182    (lambda unrestricted keyValueHeads : Nat .
1183      (lambda erased admitted : (equal Nat (streamingAttentionGroupedSM86Admitted seq heads keyValueHeads) 1) .
1184        (streamingAttentionReduceGroupedGradientUncheckedSM86 seq heads keyValueHeads)))))
1185
1186def streamingAttentionReduceGroupedGradientSM86Registers : Nat = 24
1187def streamingAttentionReduceGroupedGradientSM86Threads : Nat = saThreads
1188def streamingAttentionReduceGroupedGradientSM86ConstantBytes : Nat = (saArgument 2)
1189def streamingAttentionReduceSourceArgument : Nat = 0
1190def streamingAttentionReduceDestinationArgument : Nat = 1
1191def streamingAttentionReduceGroupedGradientSM86Blocks =
1192  (lambda unrestricted seq : Nat .
1193    (naturalDivideUnchecked (naturalMultiply saHeadWidth seq) saThreads))
1194
1195-- These dimensions are part of the realization's contract, not facts for a
1196-- pairing to restate. The admission above guards the divisor and geometry.
1197def streamingAttentionGroupedSM86QueryBlocks =
1198  (lambda unrestricted seq : Nat . (naturalDivideUnchecked seq saTile))
1199def streamingAttentionGroupedSM86KeyBlocks =
1200  (lambda unrestricted seq : Nat . (naturalDivideUnchecked seq skSubTile))
1201def streamingAttentionGroupedSM86Groups =
1202  (lambda unrestricted heads : Nat . (lambda unrestricted keyValueHeads : Nat .
1203    (naturalDivideUnchecked heads keyValueHeads)))

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.