Source/Packages

Realization.Nvidia.SM86.StreamingAttentionSM86

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

1,129 lines205 declarations66.1 KiBSHA-256 aefcbc3b00ca

Complete file · line 1061

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
181def saMaskBase : Nat = 16
182def saScratch : Nat = 17
183def saScale : Nat = 18
184def saTileIndex : Nat = 19
185def saHead : Nat = 192
186def saRow : Nat = 193
187-- Query backward occupies R194..R243, so the grouped K/V selector lives
188-- above that entire allocation rather than aliasing its transpose pointer.
189def saKeyValueHead : Nat = 244
190-- the A fragments of Q: k-slice kt's four registers
191def saQ = (lambda unrestricted kt : Nat . (naturalAdd 20 (naturalMultiply kt 4)))
192def saQReg = (lambda unrestricted kt : Nat . (lambda unrestricted k : Nat . (naturalAdd (saQ kt) k)))
193-- the score accumulators: n-tile nt's four binary32 registers
194def saS = (lambda unrestricted nt : Nat . (naturalAdd 36 (naturalMultiply nt 4)))
195def saO = (lambda unrestricted nt : Nat . (naturalAdd 68 (naturalMultiply nt 4)))
196-- K's (then V^T's) B fragments: slice k, n-tile nt, half h
197def saB = (lambda unrestricted k : Nat . (lambda unrestricted nt : Nat .
198  (naturalAdd 100 (naturalMultiply (naturalAdd (naturalMultiply k 8) nt) 2))))
199-- P's A fragments: k-slice kk
200def saPA = (lambda unrestricted kk : Nat . (naturalAdd 164 (naturalMultiply kk 4)))
201def saM = (lambda unrestricted h : Nat . (naturalAdd 180 h))
202def saL = (lambda unrestricted h : Nat . (naturalAdd 182 h))
203def saTmp = (lambda unrestricted n : Nat . (naturalAdd 184 n))
204def saRegisters : Nat = 248
205
206-- the element (row half h, column b) of an accumulator tile: c0 c1 the
207-- first row, c2 c3 the row 8 below
208def saElem = (lambda unrestricted tile : Nat . (lambda unrestricted h : Nat . (lambda unrestricted b : Nat .
209  (naturalAdd tile (naturalAdd (naturalMultiply 2 h) b)))))
210
211-- ---- one key tile ----
212-- the tile's K fragments (all 32 slices x tiles x halves), S zeroed
213def saLoadKeys = (lambda unrestricted tail : (family SM86Program) .
214  (saFor 32 (lambda unrestricted i : Nat .
215    (let unrestricted k = (naturalDivideUnchecked i 8) in (let unrestricted nt = (naturalModuloUnchecked i 8) in
216      (lambda unrestricted rest : (family SM86Program) .
217        (saLoad (saB k nt) saKeyPointer (naturalAdd (naturalMultiply nt 1024) (naturalMultiply k 32)) saSB0 (naturalSelect (naturalIsZero i) saWait3 saWaitNone)
218        (saLoad (succ (saB k nt)) saKeyPointer (naturalAdd (naturalMultiply nt 1024) (naturalAdd (naturalMultiply k 32) 16)) saSB0 saWaitNone
219          rest))))))
220  (saFor 32 (lambda unrestricted i : Nat . (saMovImm (naturalAdd 36 i) 0))
221    tail)))
222
223-- S = Q K^T: slice by slice, the eight n-tiles of a slice in flight together
224def saScores = (lambda unrestricted tail : (family SM86Program) .
225  (saFor 4 (lambda unrestricted k : Nat .
226    (saFor 8 (lambda unrestricted nt : Nat .
227      (saHmma (saS nt) (saQ k) (saB k nt) saSB2
228        (naturalSelect (naturalIsZero nt) (naturalSelect (naturalIsZero k) (naturalAdd saWait0 saWait1) saWait2) saWaitNone)))))
229    tail))
230
231-- the tile's V^T fragments into the registers K's had (the products that
232-- read them have finished: the first load waits for them)
233def saLoadValues = (lambda unrestricted seq : Nat . (lambda unrestricted tail : (family SM86Program) .
234  (saFor 32 (lambda unrestricted i : Nat .
235    (let unrestricted k = (naturalDivideUnchecked i 8) in (let unrestricted nt = (naturalModuloUnchecked i 8) in
236      (let unrestricted rowOffset = (naturalMultiply nt (naturalMultiply 8 (naturalMultiply seq saHalfBytes))) in
237      (lambda unrestricted rest : (family SM86Program) .
238        (saLoad (saB k nt) saValuePointer (naturalAdd rowOffset (naturalMultiply k 32)) saSB0 (naturalSelect (naturalIsZero i) saWait2 saWaitNone)
239        (saLoad (succ (saB k nt)) saValuePointer (naturalAdd rowOffset (naturalAdd (naturalMultiply k 32) 16)) saSB0 saWaitNone
240          rest)))))))
241    tail)))
242
243-- the diagonal tile's mask: element (row r + 8h, column 8 nt + 2 (lane % 4) + b)
244-- is masked when its column exceeds its row: when 8 nt + b - 8 h exceeds
245-- 16 warp + lane / 4 - 2 (lane % 4).  saMask holds that plus 64, plus 4096
246-- for every tile after this one, so no element of an earlier tile is masked
247def saMaskTile = (lambda unrestricted tail : (family SM86Program) .
248  (saFor 32 (lambda unrestricted i : Nat .
249    (let unrestricted nt = (naturalDivideUnchecked i 4) in (let unrestricted h = (naturalDivideUnchecked (naturalModuloUnchecked i 4) 2) in
250      (let unrestricted b = (naturalModuloUnchecked i 2) in
251      (let unrestricted threshold = (naturalSaturatingSubtract (naturalAdd (naturalAdd (naturalMultiply 8 nt) b) 63) (naturalMultiply 8 h)) in
252      (lambda unrestricted rest : (family SM86Program) .
253        (saGreater saP0 saMask threshold
254          (saUnless saP0 (constructor SM86InstructionBody SM86MoveImmediate (saR (saElem (saS nt) h b)) (saU saMinusInfinity) saPlain)
255            rest))))))))
256    tail))
257
258-- one row half h: its sixteen scores, scaled to base 2 (the products have
259-- finished: the first multiply waits for them)
260def saScaleRow = (lambda unrestricted h : Nat . (lambda unrestricted tail : (family SM86Program) .
261  (saFor 16 (lambda unrestricted i : Nat .
262    (let unrestricted e = (saElem (saS (naturalDivideUnchecked i 2)) h (naturalModuloUnchecked i 2)) in
263      (saFmul e e saScale saWaitNone)))
264    tail)))
265
266-- the row's maximum over the quad's 64 columns, in tmp h
267def saRowMax = (lambda unrestricted h : Nat . (lambda unrestricted tail : (family SM86Program) .
268  (let unrestricted into = (saTmp h) in
269  (saFmax into (saElem (saS 0) h 0) (saElem (saS 0) h 1) saWaitNone
270    (saFor 14 (lambda unrestricted j : Nat .
271      (let unrestricted i = (naturalAdd j 2) in
272        (saFmax into into (saElem (saS (naturalDivideUnchecked i 2)) h (naturalModuloUnchecked i 2)) saWaitNone)))
273      (saShfl (saTmp 2) into 1 saWaitNone
274      (saFmax into into (saTmp 2) saWait5
275      (saShfl (saTmp 2) into 2 saWaitNone
276      (saFmax into into (saTmp 2) saWait5
277        tail)))))))))
278
279-- m' = max(m, the row's maximum) in tmp h; alpha = 2^(m - m') in tmp 4 + h;
280-- m = m'; -m' in tmp 6 + h
281def saRescaleFactor = (lambda unrestricted h : Nat . (lambda unrestricted tail : (family SM86Program) .
282  (saFmax (saTmp h) (saTmp h) (saM h) saWaitNone
283  (saFneg (saTmp (naturalAdd 6 h)) (saTmp h)
284  (saFadd (saTmp (naturalAdd 4 h)) (saM h) (saTmp (naturalAdd 6 h)) saWaitNone
285  (saEx2 (saTmp (naturalAdd 4 h)) (saTmp (naturalAdd 4 h)) saWaitNone
286  (saOp (constructor SM86InstructionBody SM86IntegerAddThreeImmediate (saR (saM h)) (saR (saTmp h)) (saU 0) saPlain)
287    tail)))))))
288
289-- P = 2^(t - m') in place of t; the row's sum of P into l, rescaled
290def saExponentials = (lambda unrestricted h : Nat . (lambda unrestricted tail : (family SM86Program) .
291  (saFor 16 (lambda unrestricted i : Nat .
292    (let unrestricted e = (saElem (saS (naturalDivideUnchecked i 2)) h (naturalModuloUnchecked i 2)) in
293      (lambda unrestricted rest : (family SM86Program) .
294        (saFadd e e (saTmp (naturalAdd 6 h)) saWaitNone
295          (saEx2 e e saWaitNone rest)))))
296    -- l = l alpha + (the sum of the sixteen P), once the exponentials are in
297    (saFmul (saL h) (saL h) (saTmp (naturalAdd 4 h)) saWait4
298    (saFor 16 (lambda unrestricted i : Nat .
299      (saFadd (saL h) (saL h) (saElem (saS (naturalDivideUnchecked i 2)) h (naturalModuloUnchecked i 2)) saWaitNone))
300      tail)))))
301
302-- O's row half h rescaled by alpha (the previous tile's products have
303-- finished: the first multiply waits for them)
304def saRescaleOutput = (lambda unrestricted h : Nat . (lambda unrestricted tail : (family SM86Program) .
305  (saFor 16 (lambda unrestricted i : Nat .
306    (let unrestricted e = (saElem (saO (naturalDivideUnchecked i 2)) h (naturalModuloUnchecked i 2)) in
307      (saFmul e e (saTmp (naturalAdd 4 h)) (naturalSelect (naturalIsZero i) saWait3 saWaitNone))))
308    tail)))
309
310-- P's A fragments for P V: k-slice kk is n-tiles 2 kk and 2 kk + 1
311def saPackP = (lambda unrestricted tail : (family SM86Program) .
312  (saFor 4 (lambda unrestricted kk : Nat .
313    (let unrestricted left = (saS (naturalMultiply 2 kk)) in (let unrestricted right = (saS (succ (naturalMultiply 2 kk))) in
314      (lambda unrestricted rest : (family SM86Program) .
315        (saPack (saPA kk) (succ left) left
316        (saPack (succ (saPA kk)) (naturalAdd left 3) (naturalAdd left 2)
317        (saPack (naturalAdd (saPA kk) 2) (succ right) right
318        (saPack (naturalAdd (saPA kk) 3) (naturalAdd right 3) (naturalAdd right 2)
319          rest))))))))
320    tail))
321
322-- O += P V: the first product waits for V^T's loads
323def saValues = (lambda unrestricted tail : (family SM86Program) .
324  (saFor 4 (lambda unrestricted k : Nat .
325    (saFor 8 (lambda unrestricted nt : Nat .
326      (saHmma (saO nt) (saPA k) (saB k nt) saSB3
327        (naturalSelect (naturalIsZero nt) (naturalSelect (naturalIsZero k) saWait0 saWait3) saWaitNone)))))
328    tail))
329
330def saTileBody = (lambda unrestricted seq : Nat . (lambda unrestricted tail : (family SM86Program) .
331  (saWide saKeyPointer saKeyOffset saOne (saArgument 1)
332  (saWide saValuePointer saValueOffset saOne (saArgument 2)
333  (saImad saMask saCount 4096 saMaskBase
334  (saLoadKeys
335  (saScores
336  (saLoadValues seq
337  (saMaskTile
338  (saScaleRow 0 (saScaleRow 1
339  (saRowMax 0 (saRowMax 1
340  (saRescaleFactor 0 (saRescaleFactor 1
341  (saExponentials 0 (saExponentials 1
342  (saRescaleOutput 0 (saRescaleOutput 1
343  (saPackP
344  (saValues
345    tail)))))))))))))))))))))
346
347-- ---- the loop's tail: the next key tile, one fewer left ----
348def saLoopTail = (lambda unrestricted seq : Nat . (lambda unrestricted bodyInstructions : Nat .
349  (saAddImm saKeyOffset saKeyOffset (naturalMultiply saTile (naturalMultiply saHeadWidth saHalfBytes))
350  (saAddImm saValueOffset saValueOffset (naturalMultiply saTile saHalfBytes)
351  (saAddImm saCount saCount 4294967295
352  (saGreater saP1 saCount 0
353  (saWhen saP1
354    (constructor SM86InstructionBody SM86Branch
355      (saU (naturalSaturatingSubtract 4294967296 (naturalMultiply 16 (naturalAdd bodyInstructions 5))))
356      (sm86Unsigned32 (byte 255) (byte 255) (byte 131) (byte 3))
357      saPlain)
358    sm86ProgramEmpty)))))))
359
360-- ---- the epilogue: O / l to the merged plane, m + log2 l ----
361def saEpilogue = (lambda unrestricted seq : Nat . (lambda unrestricted heads : Nat . (lambda unrestricted tail : (family SM86Program) .
362  (let unrestricted rowBytes = (naturalMultiply heads (naturalMultiply saHeadWidth saHalfBytes)) in
363  (saFor 2 (lambda unrestricted h : Nat . (lambda unrestricted rest : (family SM86Program) .
364    (saShfl (saTmp 2) (saL h) 1 saWait3
365    (saFadd (saL h) (saL h) (saTmp 2) saWait5
366    (saShfl (saTmp 2) (saL h) 2 saWaitNone
367    (saFadd (saL h) (saL h) (saTmp 2) saWait5
368    (saMufu (saTmp (naturalAdd 4 h)) (saL h) (constructor SM86MultiFunction SM86Reciprocal) saWaitNone
369    (saMufu (saTmp (naturalAdd 6 h)) (saL h) (constructor SM86MultiFunction SM86LogarithmBase2) saWaitNone
370    (saFor 16 (lambda unrestricted i : Nat .
371      (let unrestricted e = (saElem (saO (naturalDivideUnchecked i 2)) h (naturalModuloUnchecked i 2)) in
372        (saFmul e e (saTmp (naturalAdd 4 h)) (naturalSelect (naturalIsZero i) saWait4 saWaitNone))))
373    (saFadd (saM h) (saM h) (saTmp (naturalAdd 6 h)) saWaitNone
374      rest))))))))))
375  -- the output: this thread's rows r and r + 8, columns 8 nt + 2 (lane % 4)
376  (saWide saPointer saScratch saOne (saArgument 3)
377  (saFor 16 (lambda unrestricted i : Nat .
378    (let unrestricted nt = (naturalDivideUnchecked i 2) in (let unrestricted h = (naturalModuloUnchecked i 2) in
379      (lambda unrestricted rest : (family SM86Program) .
380        (saPack (saTmp 2) (succ (saElem (saO nt) h 0)) (saElem (saO nt) h 0)
381        (saOp (saStore saPointer (saTmp 2) (naturalAdd (naturalMultiply nt 16) (naturalMultiply h (naturalMultiply 8 rowBytes))) saWaitNone)
382          rest))))))
383  -- the log-sum-exp, by the quad's first lane
384  (saWide saPointer saRow saOne (saArgument 4)
385  (saGreater saP0 saQuadLane 0
386  (saUnless saP0 (saStore saPointer (saM 0) 0 saWaitNone)
387  (saUnless saP0 (saStore saPointer (saM 1) 32 saWaitNone)
388    tail)))))))))))
389
390-- ---- the program ----
391-- seq a multiple of 64, heads the number of 64-wide heads
392-- Grid Y selects the compact K/V head and grid Z selects one of its query
393-- heads.  This avoids a device integer divide and keeps K/V planes compact.
394-- The equal-head case uses one group and a unit Z dimension.
395def streamingAttentionGroupedSM86Admitted =
396  (lambda unrestricted seq : Nat .
397    (lambda unrestricted heads : Nat .
398      (lambda unrestricted keyValueHeads : Nat .
399        (naturalAnd (naturalNonzero keyValueHeads)
400          (naturalAnd (naturalNonzero heads)
401            (naturalAnd (naturalNonzero seq)
402              (naturalAnd (naturalEqual (naturalModuloUnchecked seq saTile) 0)
403                (naturalAnd
404                  (naturalEqual
405                    (naturalModuloUnchecked heads
406                      (naturalSelect (naturalNonzero keyValueHeads) keyValueHeads 1)) 0)
407                  (naturalAnd (naturalLessOrEqual heads 65535)
408                    (naturalLessOrEqual
409                      (naturalMultiply heads (naturalMultiply saHeadWidth (naturalMultiply seq 4)))
410                      4294967295))))))))))
411
412def streamingAttentionForwardGroupedUncheckedSM86 = (lambda unrestricted seq : Nat . (lambda unrestricted heads : Nat . (lambda unrestricted keyValueHeads : Nat .
413  (let unrestricted groups = (naturalDivideUnchecked heads keyValueHeads) in
414  (let unrestricted headBytes = (naturalMultiply seq (naturalMultiply saHeadWidth saHalfBytes)) in
415  (let unrestricted rowBytes = (naturalMultiply heads (naturalMultiply saHeadWidth saHalfBytes)) in
416  (let unrestricted body = (saTileBody seq sm86ProgramEmpty) in
417  (saS2R saTid (constructor SM86SpecialRegister SM86ThreadIdX)
418  (saS2R saTileIndex (constructor SM86SpecialRegister SM86CooperativeThreadArrayIdX)
419  (saS2R saKeyValueHead (constructor SM86SpecialRegister SM86CooperativeThreadArrayIdY)
420  (saS2R saHead (constructor SM86SpecialRegister SM86CooperativeThreadArrayIdZ)
421  (saMovImm saOne 1
422  -- the special registers are in once this waits
423  (saMovImmAfter saScratch 31 saWait5
424  (saImad saHead saKeyValueHead groups saHead
425  (saAnd saLane saTid saScratch
426  (saShr saWarp saTid 5
427  (saShr saQuad saLane 2
428  (saMovImm saScratch 3
429  (saAnd saQuadLane saLane saScratch
430  -- the thread's query row: 64 tile + 16 warp + lane / 4
431  (saImad saRow saWarp 16 saQuad
432  (saImad saRow saTileIndex saTile saRow
433  -- Q: head h's plane, the row, the quad lane's pair of columns
434  (saMovImm saScratch 0
435  (saImad saScratch saRow (naturalMultiply saHeadWidth saHalfBytes) saScratch
436  (saImad saScratch saQuadLane 4 saScratch
437  (saImad saScratch saHead headBytes saScratch
438  (saWide saPointer saScratch saOne (saArgument 0)
439  (saFor 16 (lambda unrestricted i : Nat .
440    (let unrestricted kt = (naturalDivideUnchecked i 4) in (let unrestricted j = (naturalModuloUnchecked i 4) in
441      (saLoad (saQReg kt j) saPointer
442        (naturalAdd (naturalMultiply kt 32)
443          (naturalAdd (naturalMultiply (naturalModuloUnchecked j 2) (naturalMultiply 8 (naturalMultiply saHeadWidth saHalfBytes)))
444            (naturalMultiply (naturalDivideUnchecked j 2) 16)))
445        saSB1 saWaitNone))))
446  -- the output's offset in the merged plane (kept in saScratch): the row,
447  -- head h's 64 columns, the quad lane's pair
448  (saMovImm saScratch 0
449  (saImad saScratch saRow rowBytes saScratch
450  (saImad saScratch saHead (naturalMultiply saHeadWidth saHalfBytes) saScratch
451  (saImad saScratch saQuadLane 4 saScratch
452  -- the log-sum-exp's (kept in saRow): 4 (h seq + row)
453  (saImad saRow saHead seq saRow
454  (saMovImm saKeyOffset 0
455  (saImad saRow saRow 4 saKeyOffset
456  -- K from key 0: the quad's row lane / 4, its pair of columns
457  (saImad saKeyOffset saQuad (naturalMultiply saHeadWidth saHalfBytes) saKeyOffset
458  (saImad saKeyOffset saQuadLane 4 saKeyOffset
459  (saImad saKeyOffset saKeyValueHead headBytes saKeyOffset
460  -- V^T from key 0: row lane / 4 of head h's 64 x seq plane
461  (saMovImm saValueOffset 0
462  (saImad saValueOffset saQuad (naturalMultiply seq saHalfBytes) saValueOffset
463  (saImad saValueOffset saQuadLane 4 saValueOffset
464  (saImad saValueOffset saKeyValueHead headBytes saValueOffset
465  -- the tiles up to the diagonal one; the mask's base, 16 warp + lane / 4
466  -- - 2 (lane % 4) + 64 - 4096
467  (saAddImm saCount saTileIndex 1
468  (saImad saMaskBase saWarp 16 saQuad
469  (saImad saMaskBase saQuadLane 4294967294 saMaskBase
470  (saAddImm saMaskBase saMaskBase (naturalSaturatingSubtract 4294967296 (naturalSaturatingSubtract 4096 64))
471  (saMovConst saScale (saArgument 5)
472  (saMovImm (saM 0) saMinusInfinity (saMovImm (saM 1) saMinusInfinity
473  (saMovImm (saL 0) 0 (saMovImm (saL 1) 0
474  (saFor 32 (lambda unrestricted i : Nat . (saMovImm (naturalAdd 68 i) 0))
475  (sm86ProgramAppend body
476    (sm86ProgramAppend (saLoopTail seq (sm86ProgramCount body))
477      (saEpilogue seq heads
478        (saOp (constructor SM86InstructionBody SM86Exit (saAfter saWaitAll)) sm86ProgramEmpty)))))))))))))))))))))))))))))))))))))))))))))))))))))))
479
480def streamingAttentionForwardGroupedSM86 =
481  (lambda unrestricted seq : Nat . (lambda unrestricted heads : Nat .
482    (lambda unrestricted keyValueHeads : Nat .
483      (lambda erased admitted : (equal Nat (streamingAttentionGroupedSM86Admitted seq heads keyValueHeads) 1) .
484        (streamingAttentionForwardGroupedUncheckedSM86 seq heads keyValueHeads)))))
485def streamingAttentionForwardSM86 = (lambda unrestricted seq : Nat . (lambda unrestricted heads : Nat .
486  (streamingAttentionForwardGroupedUncheckedSM86 seq heads heads)))
487
488def streamingAttentionForwardSM86Registers : Nat = saRegisters
489def streamingAttentionForwardSM86Threads : Nat = saThreads
490def streamingAttentionForwardSM86ConstantBytes : Nat = saConstantBytes
491def streamingAttentionForwardQueryArgument : Nat = 0
492def streamingAttentionForwardKeyArgument : Nat = 1
493def streamingAttentionForwardValueArgument : Nat = 2
494def streamingAttentionForwardOutputArgument : Nat = 3
495def streamingAttentionForwardLogSumArgument : Nat = 4
496def streamingAttentionForwardScaleArgument : Nat = 5
497
498-- ===========================================================================
499-- The backward pass, deterministic and in two launches (no atomics).  With
500-- P = 2^(S c - L) recomputed from the forward's base-2 log-sum-exp L, the
501-- probability gradient dP = dO V^T and the row term D = rowsum(dO * O)
502-- (= rowsum(P * dP)), the score gradient is dS = P * (dP - D); then
503-- dQ = s dS K, dK = s dS^T Q and dV = P^T dO, s the score scale.
504--
505-- The query launch (streamingAttentionQuerySM86): a block per 64 queries of
506-- a query head, grid (T / 64, K/V heads, groups). It forms D for its rows from dO's
507-- fragments and the forward's output O (and writes it for the key launch),
508-- then walks the key tiles up to the diagonal as the forward does, keeping
509-- dQ.  Operands: Q, K, V, dO ([head][T][64], half), K^T ([head][64][T],
510-- half), O ([T][heads x 64], half, the forward's), L ([head][T]); results
511-- dQ ([head][T][64], binary32) and D ([head][T], binary32).  Parameters:
512-- Q, K, V, K^T, dO, O, L, dQ, D, then c and s.
513
514def sbStore64 = (lambda unrestricted address : Nat . (lambda unrestricted value : Nat . (lambda unrestricted offset : Nat .
515  (lambda unrestricted wait : Nat .
516    (constructor SM86InstructionBody SM86StoreGlobal64 (saR address) (saR value) (saU offset) (saAfter wait))))))
517def sbWiden = (lambda unrestricted d : Nat . (lambda unrestricted a : Nat . (lambda unrestricted high : Nat . (lambda unrestricted wait : Nat .
518  (saOp (constructor SM86InstructionBody SM86HalfToFloat (saR d) (saR a)
519    (eliminate StdBool (lambda unrestricted current : (family StdBool) . (family SM86HalfSelector)) (stdBoolFromNatural high)
520      (branch StdTrue . (constructor SM86HalfSelector SM86HighHalf))
521      (branch StdFalse . (constructor SM86HalfSelector SM86LowHalf)))
522    (saAfter wait)))))))
523def sbNegate = (lambda unrestricted d : Nat . (lambda unrestricted a : Nat . (lambda unrestricted wait : Nat .
524  (saOp (constructor SM86InstructionBody SM86FloatNegate (saR d) (saR a) (saAfter wait))))))
525
526-- registers beyond the forward's
527def sqKTOffset : Nat = 194
528def sqDOut = (lambda unrestricted kt : Nat . (naturalAdd 196 (naturalMultiply kt 4)))
529def sqDP = (lambda unrestricted nt : Nat . (naturalAdd 212 (naturalMultiply nt 4)))
530-- the row terms: -L (184 + h) and -D (186 + h) for the rows r and r + 8
531def sqNegL = (lambda unrestricted h : Nat . (naturalAdd 184 h))
532def sqNegD = (lambda unrestricted h : Nat . (naturalAdd 186 h))
533def sqTmp = (lambda unrestricted n : Nat . (naturalAdd 188 n))
534def sqRegisters : Nat = 248
535-- the dQ accumulators are the forward's O's; dS's A fragments P's
536
537-- K^T's rows as the B fragments of dQ += dS K: slice kk (keys 16 kk ..),
538-- n-tile nt (columns 8 nt ..) of head h's 64 x seq plane
539def sqLoadKeysT = (lambda unrestricted seq : Nat . (lambda unrestricted tail : (family SM86Program) .
540  (saFor 32 (lambda unrestricted i : Nat .
541    (let unrestricted k = (naturalDivideUnchecked i 8) in (let unrestricted nt = (naturalModuloUnchecked i 8) in
542      (let unrestricted rowOffset = (naturalMultiply nt (naturalMultiply 8 (naturalMultiply seq saHalfBytes))) in
543      (lambda unrestricted rest : (family SM86Program) .
544        (saLoad (saB k nt) saValuePointer (naturalAdd rowOffset (naturalMultiply k 32)) saSB0 (naturalSelect (naturalIsZero i) saWait2 saWaitNone)
545        (saLoad (succ (saB k nt)) saValuePointer (naturalAdd rowOffset (naturalAdd (naturalMultiply k 32) 16)) saSB0 saWaitNone
546          rest)))))))
547    tail)))
548
549-- a tile's rows (K's or V's) as B fragments, the first load waiting on `wait`
550def sqLoadRows = (lambda unrestricted pointer : Nat . (lambda unrestricted wait : Nat . (lambda unrestricted tail : (family SM86Program) .
551  (saFor 32 (lambda unrestricted i : Nat .
552    (let unrestricted k = (naturalDivideUnchecked i 8) in (let unrestricted nt = (naturalModuloUnchecked i 8) in
553      (lambda unrestricted rest : (family SM86Program) .
554        (saLoad (saB k nt) pointer (naturalAdd (naturalMultiply nt 1024) (naturalMultiply k 32)) saSB0 (naturalSelect (naturalIsZero i) wait saWaitNone)
555        (saLoad (succ (saB k nt)) pointer (naturalAdd (naturalMultiply nt 1024) (naturalAdd (naturalMultiply k 32) 16)) saSB0 saWaitNone
556          rest))))))
557    tail))))
558
559-- accumulators = A (4-register fragments from `a`) x the B fragments, slice
560-- by slice; the first product waits on `first`, each slice after on SB2
561def sqProducts = (lambda unrestricted accumulator : (pi unrestricted nt : Nat . Nat) . (lambda unrestricted a : Nat .
562  (lambda unrestricted first : Nat . (lambda unrestricted tail : (family SM86Program) .
563    (saFor 4 (lambda unrestricted k : Nat .
564      (saFor 8 (lambda unrestricted nt : Nat .
565        (saHmma (accumulator nt) (naturalAdd a (naturalMultiply k 4)) (saB k nt) saSB2
566          (naturalSelect (naturalIsZero nt) (naturalSelect (naturalIsZero k) first saWait2) saWaitNone)))))
567      tail)))))
568
569-- P = 2^(t - L) and dS = P (dP - D), into S's registers (the products have
570-- finished: the first operation waits for them)
571def sqScoreGradient = (lambda unrestricted tail : (family SM86Program) .
572  (saFor 32 (lambda unrestricted i : Nat .
573    (let unrestricted nt = (naturalDivideUnchecked i 4) in (let unrestricted h = (naturalDivideUnchecked (naturalModuloUnchecked i 4) 2) in
574      (let unrestricted b = (naturalModuloUnchecked i 2) in
575      (let unrestricted e = (saElem (saS nt) h b) in
576      (lambda unrestricted rest : (family SM86Program) .
577        (saFfma e e saScale (sqNegL h)
578          (saEx2 e e saWaitNone rest))))))))
579  (saFor 32 (lambda unrestricted i : Nat .
580    (let unrestricted nt = (naturalDivideUnchecked i 4) in (let unrestricted h = (naturalDivideUnchecked (naturalModuloUnchecked i 4) 2) in
581      (let unrestricted b = (naturalModuloUnchecked i 2) in
582      (let unrestricted e = (saElem (saS nt) h b) in (let unrestricted g = (saElem (sqDP nt) h b) in
583      (lambda unrestricted rest : (family SM86Program) .
584        (saFadd g g (sqNegD h) saWaitNone
585          (saFmul e e g (naturalSelect (naturalIsZero i) saWait4 saWaitNone) rest)))))))))
586    tail)))
587
588def sqTileBody = (lambda unrestricted seq : Nat . (lambda unrestricted tail : (family SM86Program) .
589  (saImad saMask saCount 4096 saMaskBase
590  -- S = Q K^T
591  (saWide saKeyPointer saKeyOffset saOne (saArgument 1)
592  (sqLoadRows saKeyPointer saWait3
593  (saFor 32 (lambda unrestricted i : Nat . (saMovImm (naturalAdd 36 i) 0))
594  (saFor 32 (lambda unrestricted i : Nat . (saMovImm (naturalAdd 212 i) 0))
595  (sqProducts saS (saQ 0) (naturalAdd saWait0 saWait1)
596  -- dP = dO V^T (V's rows at K's offset)
597  (saWide saKeyPointer saKeyOffset saOne (saArgument 2)
598  (sqLoadRows saKeyPointer saWait2
599  (sqProducts sqDP (sqDOut 0) saWait0
600  -- K^T's fragments for dQ, once dP's products are done
601  (saWide saValuePointer sqKTOffset saOne (saArgument 3)
602  (sqLoadKeysT seq
603  (saMaskTile
604  (sqScoreGradient
605  (saPackP
606  (saValues
607    tail)))))))))))))))))
608
609def sqLoopTail = (lambda unrestricted bodyInstructions : Nat .
610  (saAddImm saKeyOffset saKeyOffset (naturalMultiply saTile (naturalMultiply saHeadWidth saHalfBytes))
611  (saAddImm sqKTOffset sqKTOffset (naturalMultiply saTile saHalfBytes)
612  (saAddImm saCount saCount 4294967295
613  (saGreater saP1 saCount 0
614  (saWhen saP1
615    (constructor SM86InstructionBody SM86Branch
616      (saU (naturalSaturatingSubtract 4294967296 (naturalMultiply 16 (naturalAdd bodyInstructions 5))))
617      (sm86Unsigned32 (byte 255) (byte 255) (byte 131) (byte 3))
618      saPlain)
619    sm86ProgramEmpty))))))
620
621-- D for the thread's rows: its sixteen dO values of row r (and of r + 8)
622-- against O's, summed over the quad; -D and -L kept, D written by the
623-- quad's first lane
624def sqRowTerms = (lambda unrestricted seq : Nat . (lambda unrestricted heads : Nat . (lambda unrestricted tail : (family SM86Program) .
625  (let unrestricted rowBytes = (naturalMultiply heads (naturalMultiply saHeadWidth saHalfBytes)) in
626  -- O in dO's fragment pattern, into the B registers (free before the loop)
627  (saWide saPointer saScratch saOne (saArgument 5)
628  (saFor 16 (lambda unrestricted i : Nat .
629    (let unrestricted kt = (naturalDivideUnchecked i 4) in (let unrestricted j = (naturalModuloUnchecked i 4) in
630      (saLoad (naturalAdd 100 i) saPointer
631        (naturalAdd (naturalMultiply kt 32)
632          (naturalAdd (naturalMultiply (naturalModuloUnchecked j 2) (naturalMultiply 8 rowBytes))
633            (naturalMultiply (naturalDivideUnchecked j 2) 16)))
634        saSB0 saWaitNone))))
635  (saMovImm (sqTmp 0) 0 (saMovImm (sqTmp 1) 0
636  -- register j of slice kt is row r (j even) or r + 8 (j odd); both halves
637  (saFor 32 (lambda unrestricted i : Nat .
638    (let unrestricted r = (naturalDivideUnchecked i 2) in (let unrestricted high = (naturalModuloUnchecked i 2) in
639      (let unrestricted h = (naturalModuloUnchecked (naturalModuloUnchecked r 4) 2) in
640      (lambda unrestricted rest : (family SM86Program) .
641        (sbWiden (sqTmp 2) (naturalAdd 100 r) high (naturalSelect (naturalIsZero i) (naturalAdd saWait0 saWait1) saWaitNone)
642        (sbWiden (sqTmp 3) (naturalAdd (sqDOut 0) r) high saWaitNone
643        (saFfma (sqTmp h) (sqTmp 2) (sqTmp 3) (sqTmp h)
644          rest))))))))
645  (saFor 2 (lambda unrestricted h : Nat . (lambda unrestricted rest : (family SM86Program) .
646    (saShfl (sqTmp 2) (sqTmp h) 1 saWaitNone
647    (saFadd (sqTmp h) (sqTmp h) (sqTmp 2) saWait5
648    (saShfl (sqTmp 2) (sqTmp h) 2 saWaitNone
649    (saFadd (sqTmp h) (sqTmp h) (sqTmp 2) saWait5
650    (saFneg (sqNegD h) (sqTmp h)
651      rest)))))))
652  -- D and L at 4 (h seq + row): D by the quad's first lane; -L
653  (saWide saPointer saRow saOne (saArgument 8)
654  (saGreater saP0 saQuadLane 0
655  (saUnless saP0 (saStore saPointer (sqTmp 0) 0 saWaitNone)
656  (saUnless saP0 (saStore saPointer (sqTmp 1) 32 saWaitNone)
657  (saWide saPointer saRow saOne (saArgument 6)
658  (saLoad (sqTmp 2) saPointer 0 saSB1 saWaitNone
659  (saLoad (sqTmp 3) saPointer 32 saSB1 saWaitNone
660  (sbNegate (sqNegL 0) (sqTmp 2) saWait1
661  (sbNegate (sqNegL 1) (sqTmp 3) saWaitNone
662    tail)))))))))))))))))))
663
664-- dQ = s (the sum): the thread's rows r and r + 8, pairs of binary32 columns
665def sqEpilogue = (lambda unrestricted tail : (family SM86Program) .
666  (saMovConst (sqTmp 4) (saArgument 10)
667  (saFor 32 (lambda unrestricted i : Nat .
668    (saFmul (naturalAdd 68 i) (naturalAdd 68 i) (sqTmp 4) (naturalSelect (naturalIsZero i) saWait3 saWaitNone)))
669  (saWide saPointer saScratch saOne (saArgument 7)
670  (saFor 16 (lambda unrestricted i : Nat .
671    (let unrestricted nt = (naturalDivideUnchecked i 2) in (let unrestricted h = (naturalModuloUnchecked i 2) in
672      (saOp (sbStore64 saPointer (saElem (saO nt) h 0)
673        (naturalAdd (naturalMultiply nt 32) (naturalMultiply h (naturalMultiply 8 (naturalMultiply saHeadWidth 4))))
674        saWaitNone)))))
675    tail)))))
676
677def streamingAttentionQueryGroupedUncheckedSM86 = (lambda unrestricted seq : Nat . (lambda unrestricted heads : Nat . (lambda unrestricted keyValueHeads : Nat .
678  (let unrestricted groups = (naturalDivideUnchecked heads keyValueHeads) in
679  (let unrestricted headBytes = (naturalMultiply seq (naturalMultiply saHeadWidth saHalfBytes)) in
680  (let unrestricted rowBytes = (naturalMultiply heads (naturalMultiply saHeadWidth saHalfBytes)) in
681  (let unrestricted body = (sqTileBody seq sm86ProgramEmpty) in
682  (saS2R saTid (constructor SM86SpecialRegister SM86ThreadIdX)
683  (saS2R saTileIndex (constructor SM86SpecialRegister SM86CooperativeThreadArrayIdX)
684  (saS2R saKeyValueHead (constructor SM86SpecialRegister SM86CooperativeThreadArrayIdY)
685  (saS2R saHead (constructor SM86SpecialRegister SM86CooperativeThreadArrayIdZ)
686  (saMovImm saOne 1
687  (saMovImmAfter saScratch 31 saWait5
688  (saImad saHead saKeyValueHead groups saHead
689  (saAnd saLane saTid saScratch
690  (saShr saWarp saTid 5
691  (saShr saQuad saLane 2
692  (saMovImm saScratch 3
693  (saAnd saQuadLane saLane saScratch
694  (saImad saRow saWarp 16 saQuad
695  (saImad saRow saTileIndex saTile saRow
696  -- Q's and dO's A fragments: head h's plane, the row, the quad lane's pair
697  (saMovImm saScratch 0
698  (saImad saScratch saRow (naturalMultiply saHeadWidth saHalfBytes) saScratch
699  (saImad saScratch saQuadLane 4 saScratch
700  (saImad saScratch saHead headBytes saScratch
701  (saWide saPointer saScratch saOne (saArgument 0)
702  (saFor 16 (lambda unrestricted i : Nat .
703    (let unrestricted kt = (naturalDivideUnchecked i 4) in (let unrestricted j = (naturalModuloUnchecked i 4) in
704      (saLoad (saQReg kt j) saPointer
705        (naturalAdd (naturalMultiply kt 32)
706          (naturalAdd (naturalMultiply (naturalModuloUnchecked j 2) (naturalMultiply 8 (naturalMultiply saHeadWidth saHalfBytes)))
707            (naturalMultiply (naturalDivideUnchecked j 2) 16)))
708        saSB1 saWaitNone))))
709  (saWide saPointer saScratch saOne (saArgument 4)
710  (saFor 16 (lambda unrestricted i : Nat .
711    (let unrestricted kt = (naturalDivideUnchecked i 4) in (let unrestricted j = (naturalModuloUnchecked i 4) in
712      (saLoad (naturalAdd (sqDOut kt) j) saPointer
713        (naturalAdd (naturalMultiply kt 32)
714          (naturalAdd (naturalMultiply (naturalModuloUnchecked j 2) (naturalMultiply 8 (naturalMultiply saHeadWidth saHalfBytes)))
715            (naturalMultiply (naturalDivideUnchecked j 2) 16)))
716        saSB1 saWaitNone))))
717  -- dQ's offset (binary32 rows), kept in saKeyPointer's partner until the
718  -- epilogue: saScratch becomes O's (merged plane) for the row terms
719  (saMovImm saKeyOffset 0
720  (saImad saKeyOffset saRow (naturalMultiply saHeadWidth 4) saKeyOffset
721  (saImad saKeyOffset saQuadLane 8 saKeyOffset
722  (saImad saKeyOffset saHead (naturalMultiply seq (naturalMultiply saHeadWidth 4)) saKeyOffset
723  (saMovImm saScratch 0
724  (saImad saScratch saRow rowBytes saScratch
725  (saImad saScratch saHead (naturalMultiply saHeadWidth saHalfBytes) saScratch
726  (saImad saScratch saQuadLane 4 saScratch
727  -- the row terms' offset: 4 (h seq + row)
728  (saImad saRow saHead seq saRow
729  (saMovImm saValueOffset 0
730  (saImad saRow saRow 4 saValueOffset
731  (sqRowTerms seq heads
732  -- dQ's offset into saScratch; K's and V's from key 0, K^T's
733  (saAddImm saScratch saKeyOffset 0
734  (saMovImm saKeyOffset 0
735  (saImad saKeyOffset saQuad (naturalMultiply saHeadWidth saHalfBytes) saKeyOffset
736  (saImad saKeyOffset saQuadLane 4 saKeyOffset
737  (saImad saKeyOffset saKeyValueHead headBytes saKeyOffset
738  (saMovImm sqKTOffset 0
739  (saImad sqKTOffset saQuad (naturalMultiply seq saHalfBytes) sqKTOffset
740  (saImad sqKTOffset saQuadLane 4 sqKTOffset
741  (saImad sqKTOffset saKeyValueHead headBytes sqKTOffset
742  (saAddImm saCount saTileIndex 1
743  (saImad saMaskBase saWarp 16 saQuad
744  (saImad saMaskBase saQuadLane 4294967294 saMaskBase
745  (saAddImm saMaskBase saMaskBase (naturalSaturatingSubtract 4294967296 (naturalSaturatingSubtract 4096 64))
746  (saMovConst saScale (saArgument 9)
747  (saFor 32 (lambda unrestricted i : Nat . (saMovImm (naturalAdd 68 i) 0))
748  (sm86ProgramAppend body
749    (sm86ProgramAppend (sqLoopTail (sm86ProgramCount body))
750      (sqEpilogue
751        (saOp (constructor SM86InstructionBody SM86Exit (saAfter saWaitAll)) sm86ProgramEmpty))))))))))))))))))))))))))))))))))))))))))))))))))))))))))))
752
753def streamingAttentionQueryGroupedSM86 =
754  (lambda unrestricted seq : Nat . (lambda unrestricted heads : Nat .
755    (lambda unrestricted keyValueHeads : Nat .
756      (lambda erased admitted : (equal Nat (streamingAttentionGroupedSM86Admitted seq heads keyValueHeads) 1) .
757        (streamingAttentionQueryGroupedUncheckedSM86 seq heads keyValueHeads)))))
758def streamingAttentionQuerySM86 = (lambda unrestricted seq : Nat . (lambda unrestricted heads : Nat .
759  (streamingAttentionQueryGroupedUncheckedSM86 seq heads heads)))
760
761def streamingAttentionQuerySM86Registers : Nat = sqRegisters
762def streamingAttentionQuerySM86Threads : Nat = saThreads
763def streamingAttentionQuerySM86ConstantBytes : Nat = (saArgument 11)
764def streamingAttentionQueryQueryArgument : Nat = 0
765def streamingAttentionQueryKeyArgument : Nat = 1
766def streamingAttentionQueryValueArgument : Nat = 2
767def streamingAttentionQueryTransposedKeyArgument : Nat = 3
768def streamingAttentionQueryOutputGradientArgument : Nat = 4
769def streamingAttentionQueryAttentionOutputArgument : Nat = 5
770def streamingAttentionQueryLogSumArgument : Nat = 6
771def streamingAttentionQueryQueryGradientArgument : Nat = 7
772def streamingAttentionQueryRowTermArgument : Nat = 8
773def streamingAttentionQueryBaseTwoScaleArgument : Nat = 9
774def streamingAttentionQueryNaturalScaleArgument : Nat = 10
775
776-- The key launch (streamingAttentionKeySM86): a block of two warps per 32
777-- keys of a K/V head, grid (T / 32, K/V heads, groups). A warp holds its 16 keys' K and V
778-- rows as A fragments and walks the 32-query sub-tiles from the last one
779-- down to its diagonal (the last, masked): S^T = K Q^T, dP^T = V dO^T,
780-- P^T = 2^(S^T c - L), dS^T = P^T (dP^T - D) with L and D the queries',
781-- dV += P^T dO and dK += dS^T Q.  Operands: Q, K, V, dO ([head][T][64],
782-- half), Q^T, dO^T ([head][64][T], half), L and D ([head][T]); results dK
783-- ([head][T][64]) and dV^T ([head][64][T]), binary32.  Parameters: Q, K,
784-- V, dO, Q^T, dO^T, L, D, dK, dV^T, then c and s.
785
786def skK = (lambda unrestricted kt : Nat . (naturalAdd 20 (naturalMultiply kt 4)))
787def skV = (lambda unrestricted kt : Nat . (naturalAdd 36 (naturalMultiply kt 4)))
788def skS = (lambda unrestricted nt : Nat . (naturalAdd 52 (naturalMultiply nt 4)))
789def skDP = (lambda unrestricted nt : Nat . (naturalAdd 68 (naturalMultiply nt 4)))
790def skDV = (lambda unrestricted nt : Nat . (naturalAdd 84 (naturalMultiply nt 4)))
791def skDK = (lambda unrestricted nt : Nat . (naturalAdd 116 (naturalMultiply nt 4)))
792-- the B fragments of slice k, n-tile nt among `tiles` n-tiles
793def skB = (lambda unrestricted tiles : Nat . (lambda unrestricted k : Nat . (lambda unrestricted nt : Nat .
794  (naturalAdd 148 (naturalMultiply (naturalAdd (naturalMultiply k tiles) nt) 2)))))
795def skPA = (lambda unrestricted kk : Nat . (naturalAdd 180 (naturalMultiply kk 4)))
796def skSA = (lambda unrestricted kk : Nat . (naturalAdd 188 (naturalMultiply kk 4)))
797-- -L and -D of the thread's columns: n-tile nt, column b
798def skNegL = (lambda unrestricted nt : Nat . (lambda unrestricted b : Nat . (naturalAdd 196 (naturalAdd (naturalMultiply nt 2) b))))
799def skNegD = (lambda unrestricted nt : Nat . (lambda unrestricted b : Nat . (naturalAdd 204 (naturalAdd (naturalMultiply nt 2) b))))
800def skTmp = (lambda unrestricted n : Nat . (naturalAdd 212 n))
801def skHead : Nat = 220
802def skRow : Nat = 221
803def skLogOffset : Nat = 222
804def skKeyValueHead : Nat = 223
805def skRegisters : Nat = 232
806def skSubTile : Nat = 32
807def skThreads : Nat = 64
808
809-- rows ([T][64] half) as the B fragments of a product over the 64 columns:
810-- slices k 0..3, n-tiles nt 0..3 (the sub-tile's 32 rows)
811def skLoadRows = (lambda unrestricted pointer : Nat . (lambda unrestricted wait : Nat . (lambda unrestricted tail : (family SM86Program) .
812  (saFor 16 (lambda unrestricted i : Nat .
813    (let unrestricted k = (naturalDivideUnchecked i 4) in (let unrestricted nt = (naturalModuloUnchecked i 4) in
814      (lambda unrestricted rest : (family SM86Program) .
815        (saLoad (skB 4 k nt) pointer (naturalAdd (naturalMultiply nt 1024) (naturalMultiply k 32)) saSB0 (naturalSelect (naturalIsZero i) wait saWaitNone)
816        (saLoad (succ (skB 4 k nt)) pointer (naturalAdd (naturalMultiply nt 1024) (naturalAdd (naturalMultiply k 32) 16)) saSB0 saWaitNone
817          rest))))))
818    tail))))
819
820-- transposed rows ([64][T] half) as the B fragments of a product over the
821-- sub-tile's 32 positions: slices kk 0..1, n-tiles nt 0..7 (the 64 columns)
822def skLoadColumns = (lambda unrestricted seq : Nat . (lambda unrestricted pointer : Nat . (lambda unrestricted wait : Nat .
823  (lambda unrestricted tail : (family SM86Program) .
824    (saFor 16 (lambda unrestricted i : Nat .
825      (let unrestricted k = (naturalDivideUnchecked i 8) in (let unrestricted nt = (naturalModuloUnchecked i 8) in
826        (let unrestricted rowOffset = (naturalMultiply nt (naturalMultiply 8 (naturalMultiply seq saHalfBytes))) in
827        (lambda unrestricted rest : (family SM86Program) .
828          (saLoad (skB 8 k nt) pointer (naturalAdd rowOffset (naturalMultiply k 32)) saSB0 (naturalSelect (naturalIsZero i) wait saWaitNone)
829          (saLoad (succ (skB 8 k nt)) pointer (naturalAdd rowOffset (naturalAdd (naturalMultiply k 32) 16)) saSB0 saWaitNone
830            rest)))))))
831      tail)))))
832
833def skProducts = (lambda unrestricted accumulator : (pi unrestricted nt : Nat . Nat) . (lambda unrestricted a : (pi unrestricted k : Nat . Nat) .
834  (lambda unrestricted slices : Nat . (lambda unrestricted tiles : Nat . (lambda unrestricted barrier : (family SM86Barrier) .
835  (lambda unrestricted wait : Nat . (lambda unrestricted first : Nat . (lambda unrestricted tail : (family SM86Program) .
836    (saFor slices (lambda unrestricted k : Nat .
837      (saFor tiles (lambda unrestricted nt : Nat .
838        (saHmma (accumulator nt) (a k) (skB tiles k nt) barrier
839          (naturalSelect (naturalIsZero nt) (naturalSelect (naturalIsZero k) first wait) saWaitNone)))))
840      tail)))))))))
841
842-- the diagonal sub-tile's mask: element (key r + 8h, query 8 nt + 2 (lane % 4)
843-- + b) is kept when the query is not before the key: saMask holds
844-- 2 (lane % 4) - 16 warp - lane / 4 + 128, plus 4096 for every sub-tile
845-- before the diagonal one
846def skMaskTile = (lambda unrestricted tail : (family SM86Program) .
847  (saFor 16 (lambda unrestricted i : Nat .
848    (let unrestricted nt = (naturalDivideUnchecked i 4) in (let unrestricted h = (naturalDivideUnchecked (naturalModuloUnchecked i 4) 2) in
849      (let unrestricted b = (naturalModuloUnchecked i 2) in
850      (let unrestricted threshold = (naturalSaturatingSubtract (naturalAdd 127 (naturalMultiply 8 h)) (naturalAdd (naturalMultiply 8 nt) b)) in
851      (lambda unrestricted rest : (family SM86Program) .
852        (saGreater saP0 saMask threshold
853          (saUnless saP0 (constructor SM86InstructionBody SM86MoveImmediate (saR (saElem (skS nt) h b)) (saU saMinusInfinity) saPlain)
854            rest))))))))
855    tail))
856
857-- P^T = 2^(t - L), dS^T = P^T (dP^T - D), in S^T's registers
858def skScoreGradient = (lambda unrestricted tail : (family SM86Program) .
859  (saFor 16 (lambda unrestricted i : Nat .
860    (let unrestricted nt = (naturalDivideUnchecked i 4) in (let unrestricted h = (naturalDivideUnchecked (naturalModuloUnchecked i 4) 2) in
861      (let unrestricted b = (naturalModuloUnchecked i 2) in
862      (let unrestricted e = (saElem (skS nt) h b) in
863      (lambda unrestricted rest : (family SM86Program) .
864        (saFfma e e saScale (skNegL nt b)
865          (saEx2 e e saWaitNone rest))))))))
866  -- P^T's fragments, before dS^T replaces it
867  (saFor 2 (lambda unrestricted kk : Nat .
868    (let unrestricted left = (skS (naturalMultiply 2 kk)) in (let unrestricted right = (skS (succ (naturalMultiply 2 kk))) in
869      (lambda unrestricted rest : (family SM86Program) .
870        (saOp (constructor SM86InstructionBody SM86FloatPairToPackedHalfPair (saR (skPA kk)) (saR (succ left)) (saR left)
871          (saAfter (naturalSelect (naturalIsZero kk) saWait4 saWaitNone)))
872        (saPack (succ (skPA kk)) (naturalAdd left 3) (naturalAdd left 2)
873        (saPack (naturalAdd (skPA kk) 2) (succ right) right
874        (saPack (naturalAdd (skPA kk) 3) (naturalAdd right 3) (naturalAdd right 2)
875          rest))))))))
876  (saFor 16 (lambda unrestricted i : Nat .
877    (let unrestricted nt = (naturalDivideUnchecked i 4) in (let unrestricted h = (naturalDivideUnchecked (naturalModuloUnchecked i 4) 2) in
878      (let unrestricted b = (naturalModuloUnchecked i 2) in
879      (let unrestricted e = (saElem (skS nt) h b) in (let unrestricted g = (saElem (skDP nt) h b) in
880      (lambda unrestricted rest : (family SM86Program) .
881        (saFadd g g (skNegD nt b) saWaitNone
882          (saFmul e e g saWaitNone rest)))))))))
883  (saFor 2 (lambda unrestricted kk : Nat .
884    (let unrestricted left = (skS (naturalMultiply 2 kk)) in (let unrestricted right = (skS (succ (naturalMultiply 2 kk))) in
885      (lambda unrestricted rest : (family SM86Program) .
886        (saPack (skSA kk) (succ left) left
887        (saPack (succ (skSA kk)) (naturalAdd left 3) (naturalAdd left 2)
888        (saPack (naturalAdd (skSA kk) 2) (succ right) right
889        (saPack (naturalAdd (skSA kk) 3) (naturalAdd right 3) (naturalAdd right 2)
890          rest))))))))
891    tail)))))
892
893def skTileBody = (lambda unrestricted seq : Nat . (lambda unrestricted tail : (family SM86Program) .
894  (saImad saMask saCount 4096 saMaskBase
895  -- the sub-tile's L and D (their negatives)
896  (saWide saPointer skLogOffset saOne (saArgument 6)
897  (saFor 8 (lambda unrestricted i : Nat .
898    (saLoad (skTmp i) saPointer (naturalAdd (naturalMultiply (naturalDivideUnchecked i 2) 32) (naturalMultiply (naturalModuloUnchecked i 2) 4))
899      saSB1 (naturalSelect (naturalIsZero i) saWait3 saWaitNone)))
900  (saFor 8 (lambda unrestricted i : Nat .
901    (sbNegate (skNegL (naturalDivideUnchecked i 2) (naturalModuloUnchecked i 2)) (skTmp i) (naturalSelect (naturalIsZero i) saWait1 saWaitNone)))
902  (saWide saPointer skLogOffset saOne (saArgument 7)
903  (saFor 8 (lambda unrestricted i : Nat .
904    (saLoad (skTmp i) saPointer (naturalAdd (naturalMultiply (naturalDivideUnchecked i 2) 32) (naturalMultiply (naturalModuloUnchecked i 2) 4))
905      saSB1 saWaitNone))
906  (saFor 8 (lambda unrestricted i : Nat .
907    (sbNegate (skNegD (naturalDivideUnchecked i 2) (naturalModuloUnchecked i 2)) (skTmp i) (naturalSelect (naturalIsZero i) saWait1 saWaitNone)))
908  -- S^T = K Q^T
909  (saWide saKeyPointer saKeyOffset saOne (saArgument 0)
910  (skLoadRows saKeyPointer saWait3
911  (saFor 16 (lambda unrestricted i : Nat . (saMovImm (naturalAdd 52 i) 0))
912  (saFor 16 (lambda unrestricted i : Nat . (saMovImm (naturalAdd 68 i) 0))
913  (skProducts skS skK 4 4 saSB2 saWait2 (naturalAdd saWait0 saWait1)
914  -- dP^T = V dO^T
915  (saWide saKeyPointer saKeyOffset saOne (saArgument 3)
916  (skLoadRows saKeyPointer saWait2
917  (skProducts skDP skV 4 4 saSB2 saWait2 saWait0
918  -- dO^T's fragments for dV, once dP^T's products are done
919  (saWide saValuePointer saValueOffset saOne (saArgument 5)
920  (skLoadColumns seq saValuePointer saWait2
921  (skMaskTile
922  (skScoreGradient
923  (skProducts skDV skPA 2 8 saSB3 saWait3 saWait0
924  -- Q^T's fragments for dK, once dV's products have read dO^T's
925  (saWide saValuePointer saValueOffset saOne (saArgument 4)
926  (skLoadColumns seq saValuePointer saWait3
927  (skProducts skDK skSA 2 8 saSB3 saWait3 saWait0
928    tail)))))))))))))))))))))))))
929
930def skLoopTail = (lambda unrestricted bodyInstructions : Nat .
931  (saAddImm saKeyOffset saKeyOffset (naturalSaturatingSubtract 4294967296 (naturalMultiply skSubTile (naturalMultiply saHeadWidth saHalfBytes)))
932  (saAddImm saValueOffset saValueOffset (naturalSaturatingSubtract 4294967296 (naturalMultiply skSubTile saHalfBytes))
933  (saAddImm skLogOffset skLogOffset (naturalSaturatingSubtract 4294967296 (naturalMultiply skSubTile 4))
934  (saAddImm saCount saCount 4294967295
935  (saGreater saP1 saCount 0
936  (saWhen saP1
937    (constructor SM86InstructionBody SM86Branch
938      (saU (naturalSaturatingSubtract 4294967296 (naturalMultiply 16 (naturalAdd bodyInstructions 6))))
939      (sm86Unsigned32 (byte 255) (byte 255) (byte 131) (byte 3))
940      saPlain)
941    sm86ProgramEmpty)))))))
942
943-- dK = s (the sum), rows; dV^T, columns: element (key r + 8h, column
944-- 8 nt + 2 (lane % 4) + b) at (8 nt + b) seq + 8 h words past the thread's
945def skEpilogue = (lambda unrestricted seq : Nat . (lambda unrestricted tail : (family SM86Program) .
946  (saMovConst (skTmp 0) (saArgument 11)
947  (saFor 32 (lambda unrestricted i : Nat .
948    (saFmul (naturalAdd 116 i) (naturalAdd 116 i) (skTmp 0) (naturalSelect (naturalIsZero i) saWait3 saWaitNone)))
949  (saWide saPointer saScratch saOne (saArgument 8)
950  (saFor 16 (lambda unrestricted i : Nat .
951    (let unrestricted nt = (naturalDivideUnchecked i 2) in (let unrestricted h = (naturalModuloUnchecked i 2) in
952      (saOp (sbStore64 saPointer (saElem (skDK nt) h 0)
953        (naturalAdd (naturalMultiply nt 32) (naturalMultiply h (naturalMultiply 8 (naturalMultiply saHeadWidth 4))))
954        saWaitNone)))))
955  -- dV^T's offset: head h's 64 x seq plane, the quad lane's first column,
956  -- the key
957  (saMovImm (skTmp 1) 0
958  (saImad (skTmp 1) saQuadLane (naturalMultiply 8 seq) (skTmp 1)
959  (saImad (skTmp 1) skHead (naturalMultiply saHeadWidth (naturalMultiply seq 4)) (skTmp 1)
960  (saImad (skTmp 1) skRow 4 (skTmp 1)
961  (saWide saPointer (skTmp 1) saOne (saArgument 9)
962  (saFor 32 (lambda unrestricted i : Nat .
963    (let unrestricted nt = (naturalDivideUnchecked i 4) in (let unrestricted h = (naturalDivideUnchecked (naturalModuloUnchecked i 4) 2) in
964      (let unrestricted b = (naturalModuloUnchecked i 2) in
965      (saOp (saStore saPointer (saElem (skDV nt) h b)
966        (naturalAdd (naturalMultiply (naturalAdd (naturalMultiply 8 nt) b) (naturalMultiply seq 4)) (naturalMultiply h 32))
967        saWaitNone))))))
968    tail))))))))))))
969
970-- Each query head writes a separate dK/dV plane.  A later reduction sums
971-- those planes into the compact K/V gradient; concurrent CTAs never race on
972-- one compact gradient address.
973def streamingAttentionKeyGroupedUncheckedSM86 = (lambda unrestricted seq : Nat . (lambda unrestricted heads : Nat . (lambda unrestricted keyValueHeads : Nat .
974  (let unrestricted groups = (naturalDivideUnchecked heads keyValueHeads) in
975  (let unrestricted headBytes = (naturalMultiply seq (naturalMultiply saHeadWidth saHalfBytes)) in
976  (let unrestricted subTiles = (naturalDivideUnchecked seq skSubTile) in
977  (let unrestricted body = (skTileBody seq sm86ProgramEmpty) in
978  (saS2R saTid (constructor SM86SpecialRegister SM86ThreadIdX)
979  (saS2R saTileIndex (constructor SM86SpecialRegister SM86CooperativeThreadArrayIdX)
980  (saS2R skKeyValueHead (constructor SM86SpecialRegister SM86CooperativeThreadArrayIdY)
981  (saS2R skHead (constructor SM86SpecialRegister SM86CooperativeThreadArrayIdZ)
982  (saMovImm saOne 1
983  (saMovImmAfter saScratch 31 saWait5
984  (saImad skHead skKeyValueHead groups skHead
985  (saAnd saLane saTid saScratch
986  (saShr saWarp saTid 5
987  (saShr saQuad saLane 2
988  (saMovImm saScratch 3
989  (saAnd saQuadLane saLane saScratch
990  -- the thread's key row: 32 tile + 16 warp + lane / 4
991  (saImad skRow saWarp 16 saQuad
992  (saImad skRow saTileIndex skSubTile skRow
993  -- K's and V's A fragments
994  (saMovImm saScratch 0
995  (saImad saScratch skRow (naturalMultiply saHeadWidth saHalfBytes) saScratch
996  (saImad saScratch saQuadLane 4 saScratch
997  (saImad saScratch skKeyValueHead headBytes saScratch
998  (saWide saPointer saScratch saOne (saArgument 1)
999  (saFor 16 (lambda unrestricted i : Nat .
1000    (let unrestricted kt = (naturalDivideUnchecked i 4) in (let unrestricted j = (naturalModuloUnchecked i 4) in
1001      (saLoad (naturalAdd (skK kt) j) saPointer
1002        (naturalAdd (naturalMultiply kt 32)
1003          (naturalAdd (naturalMultiply (naturalModuloUnchecked j 2) (naturalMultiply 8 (naturalMultiply saHeadWidth saHalfBytes)))
1004            (naturalMultiply (naturalDivideUnchecked j 2) 16)))
1005        saSB1 saWaitNone))))
1006  (saWide saPointer saScratch saOne (saArgument 2)
1007  (saFor 16 (lambda unrestricted i : Nat .
1008    (let unrestricted kt = (naturalDivideUnchecked i 4) in (let unrestricted j = (naturalModuloUnchecked i 4) in
1009      (saLoad (naturalAdd (skV kt) j) saPointer
1010        (naturalAdd (naturalMultiply kt 32)
1011          (naturalAdd (naturalMultiply (naturalModuloUnchecked j 2) (naturalMultiply 8 (naturalMultiply saHeadWidth saHalfBytes)))
1012            (naturalMultiply (naturalDivideUnchecked j 2) 16)))
1013        saSB1 saWaitNone))))
1014  -- dK's offset (binary32 rows) in saScratch
1015  (saMovImm saScratch 0
1016  (saImad saScratch skRow (naturalMultiply saHeadWidth 4) saScratch
1017  (saImad saScratch saQuadLane 8 saScratch
1018  (saImad saScratch skHead (naturalMultiply seq (naturalMultiply saHeadWidth 4)) saScratch
1019  -- the queries start at the last sub-tile: Q's and dO's rows (row lane / 4,
1020  -- the quad lane's pair), Q^T's and dO^T's columns (row lane / 4), L's and D's
1021  (saMovImm saKeyOffset 0
1022  (saImad saKeyOffset saQuad (naturalMultiply saHeadWidth saHalfBytes) saKeyOffset
1023  (saImad saKeyOffset saQuadLane 4 saKeyOffset
1024  (saImad saKeyOffset skHead headBytes saKeyOffset
1025  (saAddImm saKeyOffset saKeyOffset (naturalMultiply (naturalSaturatingSubtract subTiles 1) (naturalMultiply skSubTile (naturalMultiply saHeadWidth saHalfBytes)))
1026  (saMovImm saValueOffset 0
1027  (saImad saValueOffset saQuad (naturalMultiply seq saHalfBytes) saValueOffset
1028  (saImad saValueOffset saQuadLane 4 saValueOffset
1029  (saImad saValueOffset skHead headBytes saValueOffset
1030  (saAddImm saValueOffset saValueOffset (naturalMultiply (naturalSaturatingSubtract subTiles 1) (naturalMultiply skSubTile saHalfBytes))
1031  (saMovImm skLogOffset 0
1032  (saImad skLogOffset saQuadLane 8 skLogOffset
1033  (saImad skLogOffset skHead (naturalMultiply seq 4) skLogOffset
1034  (saAddImm skLogOffset skLogOffset (naturalMultiply (naturalSaturatingSubtract subTiles 1) (naturalMultiply skSubTile 4))
1035  -- the sub-tiles from the last down to the diagonal; the mask's base,
1036  -- 2 (lane % 4) - 16 warp - lane / 4 + 128 - 4096
1037  (saMovImm saCount subTiles
1038  (saImad saCount saTileIndex 4294967295 saCount
1039  (saMovImm saMaskBase 0
1040  (saImad saMaskBase saQuadLane 2 saMaskBase
1041  (saImad saMaskBase saWarp 4294967280 saMaskBase
1042  (saImad saMaskBase saQuad 4294967295 saMaskBase
1043  (saAddImm saMaskBase saMaskBase (naturalSaturatingSubtract 4294967296 (naturalSaturatingSubtract 4096 128))
1044  (saMovConst saScale (saArgument 10)
1045  (saFor 64 (lambda unrestricted i : Nat . (saMovImm (naturalAdd 84 i) 0))
1046  (sm86ProgramAppend body
1047    (sm86ProgramAppend (skLoopTail (sm86ProgramCount body))
1048      (skEpilogue seq
1049        (saOp (constructor SM86InstructionBody SM86Exit (saAfter saWaitAll)) sm86ProgramEmpty))))))))))))))))))))))))))))))))))))))))))))))))))))))))))))
1050
1051def streamingAttentionKeyGroupedSM86 =
1052  (lambda unrestricted seq : Nat . (lambda unrestricted heads : Nat .
1053    (lambda unrestricted keyValueHeads : Nat .
1054      (lambda erased admitted : (equal Nat (streamingAttentionGroupedSM86Admitted seq heads keyValueHeads) 1) .
1055        (streamingAttentionKeyGroupedUncheckedSM86 seq heads keyValueHeads)))))
1056def streamingAttentionKeySM86 = (lambda unrestricted seq : Nat . (lambda unrestricted heads : Nat .
1057  (streamingAttentionKeyGroupedUncheckedSM86 seq heads heads)))
1058
1059def streamingAttentionKeySM86Registers : Nat = skRegisters
1060def streamingAttentionKeySM86Threads : Nat = skThreads
1061def streamingAttentionKeySM86ConstantBytes : Nat = (saArgument 12)
1062def streamingAttentionKeyQueryArgument : Nat = 0
1063def streamingAttentionKeyKeyArgument : Nat = 1
1064def streamingAttentionKeyValueArgument : Nat = 2
1065def streamingAttentionKeyOutputGradientArgument : Nat = 3
1066def streamingAttentionKeyTransposedQueryArgument : Nat = 4
1067def streamingAttentionKeyTransposedOutputGradientArgument : Nat = 5
1068def streamingAttentionKeyLogSumArgument : Nat = 6
1069def streamingAttentionKeyRowTermArgument : Nat = 7
1070def streamingAttentionKeyKeyGradientArgument : Nat = 8
1071def streamingAttentionKeyValueGradientArgument : Nat = 9
1072def streamingAttentionKeyBaseTwoScaleArgument : Nat = 10
1073def streamingAttentionKeyNaturalScaleArgument : Nat = 11
1074
1075-- Sum the query-head dK or dV planes of one K/V group into its compact
1076-- gradient plane. Both backward outputs have 64*T binary32 elements per
1077-- head, though dV is transposed within that plane. A CTA owns 128 elements
1078-- of one K/V head; no atomics or cross-CTA ordering are needed. The caller
1079-- supplies disjoint source/destination regions and launches grid
1080-- (64*T/128, keyValueHeads, 1) for T divisible by 64.
1081def streamingAttentionReduceGroupedGradientUncheckedSM86 =
1082  (lambda unrestricted seq : Nat .
1083    (lambda unrestricted heads : Nat .
1084      (lambda unrestricted keyValueHeads : Nat .
1085        (let unrestricted groups = (naturalDivideUnchecked heads keyValueHeads) in
1086        (let unrestricted planeWords = (naturalMultiply saHeadWidth seq) in
1087        (let unrestricted planeBytes = (naturalMultiply planeWords 4) in
1088        (saS2R saTid (constructor SM86SpecialRegister SM86ThreadIdX)
1089        (saS2R 2 (constructor SM86SpecialRegister SM86CooperativeThreadArrayIdX)
1090        (saS2R 3 (constructor SM86SpecialRegister SM86CooperativeThreadArrayIdY)
1091        (saMovImmAfter 7 0 saWait5
1092        (saImad 4 2 saThreads saTid
1093        (saImad 5 3 planeWords 4
1094        (saImad 6 3 (naturalMultiply groups planeWords) 4
1095        (saMovImm 11 4
1096        (saWide 12 6 11 (saArgument 0)
1097        (saWide 14 5 11 (saArgument 1)
1098        (saMovImm 9 0
1099        (saFor groups (lambda unrestricted group : Nat .
1100          (lambda unrestricted rest : (family SM86Program) .
1101            (saLoad 8 12 (naturalMultiply group planeBytes) saSB0 saWaitNone
1102              (saFadd 9 9 8 saWait0 rest))))
1103        (saOp (saStore 14 9 0 saWaitNone)
1104        (saOp (constructor SM86InstructionBody SM86Exit (saAfter saWaitAll)) sm86ProgramEmpty))))))))))))))))))))
1105
1106def streamingAttentionReduceGroupedGradientSM86 =
1107  (lambda unrestricted seq : Nat . (lambda unrestricted heads : Nat .
1108    (lambda unrestricted keyValueHeads : Nat .
1109      (lambda erased admitted : (equal Nat (streamingAttentionGroupedSM86Admitted seq heads keyValueHeads) 1) .
1110        (streamingAttentionReduceGroupedGradientUncheckedSM86 seq heads keyValueHeads)))))
1111
1112def streamingAttentionReduceGroupedGradientSM86Registers : Nat = 24
1113def streamingAttentionReduceGroupedGradientSM86Threads : Nat = saThreads
1114def streamingAttentionReduceGroupedGradientSM86ConstantBytes : Nat = (saArgument 2)
1115def streamingAttentionReduceSourceArgument : Nat = 0
1116def streamingAttentionReduceDestinationArgument : Nat = 1
1117def streamingAttentionReduceGroupedGradientSM86Blocks =
1118  (lambda unrestricted seq : Nat .
1119    (naturalDivideUnchecked (naturalMultiply saHeadWidth seq) saThreads))
1120
1121-- These dimensions are part of the realization's contract, not facts for a
1122-- pairing to restate. The admission above guards the divisor and geometry.
1123def streamingAttentionGroupedSM86QueryBlocks =
1124  (lambda unrestricted seq : Nat . (naturalDivideUnchecked seq saTile))
1125def streamingAttentionGroupedSM86KeyBlocks =
1126  (lambda unrestricted seq : Nat . (naturalDivideUnchecked seq skSubTile))
1127def streamingAttentionGroupedSM86Groups =
1128  (lambda unrestricted heads : Nat . (lambda unrestricted keyValueHeads : Nat .
1129    (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.