Source/Packages

Realization.Nvidia.SM86.TiledProductSM86

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

712 lines127 declarations45.7 KiBSHA-256 3cb01e59fc37

Complete file

TiledProductSM86.alpha

Definition view
1module Realization.Nvidia.SM86.TiledProductSM86
2
3import Accelerator.SM86.Control
4import Accelerator.SM86.Immediate
5import Accelerator.SM86.Instruction
6import Accelerator.SM86.NumericSemantics
7import Accelerator.SM86.Operands
8import Accelerator.SM86.Program
9import Accelerator.SM86.Types
10import Realization.Nvidia.SM86.StreamingAttentionSM86
11import Std.Natural
12import Std.Foundation
13
14-- A matrix product that reuses its operands (PRD-PERFORMANCE P4):
15-- C (M x N, binary32, row-major) = A (M x K) B (K x N), half operands, a
16-- block per BM x BN tile of C.  Each operand is given in either layout:
17--
18--   A as rows    (M x K, row-major: the activations of a forward product)
19--   A as columns (K x M, row-major: A^T stored, e.g. dY^T of a weight
20--                 gradient, read from dY as it is)
21--   B as rows    (N x K, row-major: C = A B^T, a weight used forward)
22--   B as columns (K x N, row-major: a weight used backward, an activation of
23--                 a weight gradient)
24--
25-- so no operand is transposed into another plane first.  The block walks K
26-- in steps of 32: each step's A and B tiles go global -> shared memory
27-- asynchronously (cp.async: LDGSTS, a group per step, stages - 1 steps ahead
28-- through the tile's stages; one barrier a step), and every warp takes its fragments with
29-- LDSM -- .T for an operand given as columns, whose shared tile is stored the
30-- other way round -- and multiplies them on the tensor cores (HMMA.16816,
31-- binary32 accumulation, fixed-latency: no scoreboard), the next k-half's
32-- fragment loads and the next stage's copies between its products.  Shared
33-- rows are swizzled by 16-byte chunks so the eight rows an LDSM reads fall
34-- in distinct banks.  The k-steps run in a loop of one step a stage, so the
35-- program's length does not grow with K.
36--
37-- The tile is a parameter (TiledProductTile): warps along M and N, each
38-- warp's 16-row m-tiles and 8-column n-tiles, the stages.  A step's operand
39-- tile is extent x 64 bytes, extent x 4 16-byte chunks, which the block's
40-- threads copy evenly: 1, 2 or 4 chunks a thread an operand
41-- (tiledProductSM86ChunksFit), so a tile may be square (128 x 128, 64 x 64)
42-- or rectangular (128 x 64 with eight warps: two chunks of A a thread, one
43-- of B) and as large as the registers allow (256 x 128, each warp 64 x 64).
44-- M and N are multiples of the tile, K of 32 stages.  Parameters (constant
45-- bank 0 from the SM86 parameter base): C, A, B pointers.
46
47family TiledProductLayout : Type 0
48constructor TiledProductRows
49constructor TiledProductColumns
50end-family
51
52family TiledProductTile : Type 0
53-- warps along M, warps along N, 16-row m-tiles and 8-column n-tiles per warp
54constructor TiledProductTileValue
55field unrestricted tiledProductWarpsM : Nat
56field unrestricted tiledProductWarpsN : Nat
57field unrestricted tiledProductMTiles : Nat
58field unrestricted tiledProductNTiles : Nat
59-- the shared-memory stages the copies run ahead through (a step's tiles
60-- are copied stages - 1 steps before its products)
61field unrestricted tiledProductStages : Nat
62end-family
63
64def tiledProductTile = (lambda unrestricted warpsM : Nat . (lambda unrestricted warpsN : Nat . (lambda unrestricted mTiles : Nat .
65  (lambda unrestricted nTiles : Nat . (lambda unrestricted stages : Nat .
66    (constructor TiledProductTile TiledProductTileValue warpsM warpsN mTiles nTiles stages))))))
67def tiledProductLarge : (family TiledProductTile) = (tiledProductTile 2 4 4 4 4)
68def tiledProductSmall : (family TiledProductTile) = (tiledProductTile 2 2 2 4 4)
69
70def tpWarpsM = (lambda unrestricted t : (family TiledProductTile) .
71  (eliminate TiledProductTile (lambda unrestricted c : (family TiledProductTile) . Nat) t (branch TiledProductTileValue a b c d e . a)))
72def tpWarpsN = (lambda unrestricted t : (family TiledProductTile) .
73  (eliminate TiledProductTile (lambda unrestricted c : (family TiledProductTile) . Nat) t (branch TiledProductTileValue a b c d e . b)))
74def tpMTiles = (lambda unrestricted t : (family TiledProductTile) .
75  (eliminate TiledProductTile (lambda unrestricted c : (family TiledProductTile) . Nat) t (branch TiledProductTileValue a b c d e . c)))
76def tpNTiles = (lambda unrestricted t : (family TiledProductTile) .
77  (eliminate TiledProductTile (lambda unrestricted c : (family TiledProductTile) . Nat) t (branch TiledProductTileValue a b c d e . d)))
78def tpStagesOf = (lambda unrestricted t : (family TiledProductTile) .
79  (eliminate TiledProductTile (lambda unrestricted c : (family TiledProductTile) . Nat) t (branch TiledProductTileValue a b c d e . e)))
80def tpThreads = (lambda unrestricted t : (family TiledProductTile) . (naturalMultiply 32 (naturalMultiply (tpWarpsM t) (tpWarpsN t))))
81def tpBlockM = (lambda unrestricted t : (family TiledProductTile) . (naturalMultiply (tpWarpsM t) (naturalMultiply 16 (tpMTiles t))))
82def tpBlockN = (lambda unrestricted t : (family TiledProductTile) . (naturalMultiply (tpWarpsN t) (naturalMultiply 8 (tpNTiles t))))
83def tpIsColumns = (lambda unrestricted l : (family TiledProductLayout) .
84  (eliminate TiledProductLayout (lambda unrestricted c : (family TiledProductLayout) . Nat) l
85    (branch TiledProductRows . 0) (branch TiledProductColumns . 1)))
86
87def tpStep : Nat = 32
88-- the base-2 logarithm of a power of two (at most 2^16)
89def tpLog2 = (lambda unrestricted n : Nat .
90  (app
91    (nat-eliminate (lambda unrestricted fuel : Nat . (pi unrestricted m : Nat . Nat))
92      (lambda unrestricted m : Nat . 0)
93      (lambda unrestricted p : Nat . (lambda unrestricted rest : (pi unrestricted m : Nat . Nat) .
94        (lambda unrestricted m : Nat . (naturalSelect (naturalLessOrEqual m 1) 0 (succ (rest (naturalDivideUnchecked m 2)))))))
95      16)
96    n))
97-- a shared tile of `extent` rows of 32 halves (64 bytes, 4 chunks), or of
98-- 32 rows of `extent` halves (extent / 8 chunks): either way extent x 64
99-- bytes; two operands a stage, the tile's stages (a step's tiles are copied
100-- stages - 1 steps before its products)
101def tpAheadOf = (lambda unrestricted t : (family TiledProductTile) . (naturalSaturatingSubtract (tpStagesOf t) 1))
102def tpTileBytes = (lambda unrestricted extent : Nat . (naturalMultiply extent 64))
103def tpSharedBytes = (lambda unrestricted t : (family TiledProductTile) .
104  (naturalMultiply (tpStagesOf t) (naturalAdd (tpTileBytes (tpBlockM t)) (tpTileBytes (tpBlockN t)))))
105-- what an SM has for a block's shared memory on the GB10 and the RTX 3090
106-- (100 KiB), less the two kilobytes the sm_121 realization reserves (the
107-- lowering bases shared memory at 0x400; QMD V05 declares one more)
108def tpSharedCapacity : Nat = (naturalSaturatingSubtract (naturalMultiply 100 1024) 2048)
109-- 1 when a tile's stages fit
110def tiledProductSM86SharedFits = (lambda unrestricted t : (family TiledProductTile) .
111  (naturalLessOrEqual (tpSharedBytes t) tpSharedCapacity))
112def tpBufferA = (lambda unrestricted t : (family TiledProductTile) . (lambda unrestricted buffer : Nat .
113  (naturalMultiply buffer (naturalAdd (tpTileBytes (tpBlockM t)) (tpTileBytes (tpBlockN t))))))
114def tpBufferB = (lambda unrestricted t : (family TiledProductTile) . (lambda unrestricted buffer : Nat .
115  (naturalAdd (tpBufferA t buffer) (tpTileBytes (tpBlockM t)))))
116
117-- the chunks a thread copies of an operand tile of `extent` rows (rows
118-- layout) or columns (columns layout) a step: extent x 4 over the threads
119def tpChunks = (lambda unrestricted t : (family TiledProductTile) . (lambda unrestricted extent : Nat .
120  (naturalDivideUnchecked (naturalMultiply extent 4) (tpThreads t))))
121def tpChunksA = (lambda unrestricted t : (family TiledProductTile) . (tpChunks t (tpBlockM t)))
122def tpChunksB = (lambda unrestricted t : (family TiledProductTile) . (tpChunks t (tpBlockN t)))
123-- the registers below hold up to four chunks an operand
124def tpChunkRoom : Nat = 4
125-- 1 when every thread copies the same whole number of chunks of each
126-- operand, at most the room: 1, 2 or 4 (a power of two, so the copies'
127-- rows and chunks split by shifts)
128def tpChunkCountFits = (lambda unrestricted t : (family TiledProductTile) . (lambda unrestricted extent : Nat .
129  (let unrestricted chunks = (tpChunks t extent) in
130  (naturalAnd (naturalIsZero (naturalModuloUnchecked (naturalMultiply extent 4) (tpThreads t)))
131  (naturalAnd (naturalLess 0 chunks)
132  (naturalAnd (naturalLessOrEqual chunks tpChunkRoom)
133    (naturalIsZero (naturalModuloUnchecked tpChunkRoom chunks))))))))
134def tiledProductSM86ChunksFit = (lambda unrestricted t : (family TiledProductTile) .
135  (naturalAnd (tpChunkCountFits t (tpBlockM t)) (tpChunkCountFits t (tpBlockN t))))
136
137-- ---- registers ----
138def tpTid : Nat = 0
139def tpLane : Nat = 1
140def tpWarp : Nat = 2
141def tpWarpM : Nat = 3
142def tpWarpN : Nat = 4
143def tpOne : Nat = 5
144def tpTileM : Nat = 6
145def tpTileN : Nat = 7
146def tpScratch : Nat = 8
147def tpScratch2 : Nat = 9
148-- the chunks a thread moves of each operand (up to four): global pointers
149-- (pairs)
150def tpAPointer = (lambda unrestricted c : Nat . (naturalAdd 10 (naturalMultiply 2 c)))
151def tpBPointer = (lambda unrestricted c : Nat . (naturalAdd 18 (naturalMultiply 2 c)))
152-- their shared byte addresses (swizzled; the stage's offset an immediate)
153def tpAStore = (lambda unrestricted c : Nat . (naturalAdd 26 c))
154def tpBStore = (lambda unrestricted c : Nat . (naturalAdd 30 c))
155-- LDSM byte addresses: A's per k-half (rows) or per m-tile (columns), B's
156-- per k-half (rows) or per n-tile pair (columns); up to four each
157def tpAMatrix = (lambda unrestricted i : Nat . (naturalAdd 34 i))
158def tpBMatrix = (lambda unrestricted i : Nat . (naturalAdd 38 i))
159def tpC : Nat = 42
160-- the accumulators: m-tile mt, n-tile nt, four binary32
161def tpAccBase : Nat = 48
162def tpAcc = (lambda unrestricted t : (family TiledProductTile) . (lambda unrestricted mt : Nat . (lambda unrestricted nt : Nat .
163  (naturalAdd tpAccBase (naturalMultiply 4 (naturalAdd (naturalMultiply mt (tpNTiles t)) nt))))))
164def tpAccCount = (lambda unrestricted t : (family TiledProductTile) . (naturalMultiply 4 (naturalMultiply (tpMTiles t) (tpNTiles t))))
165-- A's fragments: m-tile mt, k-half kk (four registers); B's: n-tile nt,
166-- k-half kk (two)
167def tpAFragment = (lambda unrestricted t : (family TiledProductTile) . (lambda unrestricted mt : Nat . (lambda unrestricted kk : Nat .
168  (naturalAdd (naturalAdd tpAccBase (tpAccCount t)) (naturalMultiply 4 (naturalAdd (naturalMultiply kk (tpMTiles t)) mt))))))
169def tpBFragment = (lambda unrestricted t : (family TiledProductTile) . (lambda unrestricted nt : Nat . (lambda unrestricted kk : Nat .
170  (naturalAdd (naturalAdd (naturalAdd tpAccBase (tpAccCount t)) (naturalMultiply 8 (tpMTiles t)))
171    (naturalMultiply 2 (naturalAdd (naturalMultiply kk (tpNTiles t)) nt))))))
172-- each chunk's byte offset from its operand's pointer (A's, then B's: a
173-- register a chunk, so the largest tiles fit sm86RegisterLimit),
174-- after the fragments: each copy re-forms the pointers from them for the
175-- next step
176def tpOffsetBase = (lambda unrestricted t : (family TiledProductTile) .
177  (naturalAdd (naturalAdd tpAccBase (tpAccCount t)) (naturalAdd (naturalMultiply 8 (tpMTiles t)) (naturalMultiply 4 (tpNTiles t)))))
178def tpAOffset = (lambda unrestricted t : (family TiledProductTile) . (lambda unrestricted c : Nat . (naturalAdd (tpOffsetBase t) c)))
179def tpBOffset = (lambda unrestricted t : (family TiledProductTile) . (lambda unrestricted c : Nat . (naturalAdd (tpOffsetBase t) (naturalAdd (tpChunksA t) c))))
180-- the iterations left
181def tpCounter : Nat = 44
182def tpRegisters = (lambda unrestricted t : (family TiledProductTile) . (naturalAdd (tpOffsetBase t) (naturalAdd (tpChunksA t) (tpChunksB t))))
183
184def tpLdsm = (lambda unrestricted d : Nat . (lambda unrestricted address : Nat . (lambda unrestricted offset : Nat .
185  (lambda unrestricted transpose : Nat . (lambda unrestricted wait : Nat .
186    (saOp (constructor SM86InstructionBody SM86LoadSharedMatrix (saR d) (saR address) (saU offset)
187      (constructor SM86SharedMatrixCount SM86SharedMatrix4)
188      (eliminate StdBool (lambda unrestricted c : (family StdBool) . (family SM86SharedMatrixTranspose)) (stdBoolFromNatural transpose)
189        (branch StdTrue . (constructor SM86SharedMatrixTranspose SM86SharedMatrixTransposed))
190        (branch StdFalse . (constructor SM86SharedMatrixTranspose SM86SharedMatrixNotTransposed)))
191      (saCtl saSB1 wait))))))))
192-- LDGSTS of a 16-byte chunk (cp.async.cg): its registers read under SB3,
193-- which re-forming the pointer waits on
194def tpCopy = (lambda unrestricted shared : Nat . (lambda unrestricted offset : Nat . (lambda unrestricted source : Nat .
195  (saOp (constructor SM86InstructionBody SM86LoadGlobalToShared (saR shared) (saU offset) (saR source) (saU 0)
196    (constructor SM86Control SM86ControlValue (byte 15) (constructor SM86YieldMode SM86Continue)
197      saNone saSB3 (byte 0) (byte 0)))))))
198-- LDGDEPBAR: the step's copies a group, counted on SB0
199def tpCommit = (saOp (constructor SM86InstructionBody SM86CommitAsyncGroup (saCtl saSB0 saWaitNone)))
200-- DEPBAR.LE SB0, n: at most n groups in flight
201def tpWaitGroups = (lambda unrestricted n : Nat .
202  (saOp (constructor SM86InstructionBody SM86WaitAsyncGroups (nat-to-byte n) saPlain)))
203def tpBarrier = (lambda unrestricted wait : Nat . (saOp (constructor SM86InstructionBody SM86BarrierSynchronize (saAfter wait))))
204
205-- ---- moving one k-step's tiles ----
206-- chunk c (of the two a thread moves) of an operand tile: its row and
207-- 16-byte chunk in the global layout.  Rows layout: `extent` rows of 64
208-- bytes (4 chunks); columns layout: 32 rows of `extent` halves.
209-- The global offset (bytes) of a step's chunk past the operand's pointer
210-- for the chunk at step 0: rows move 64 bytes along K a step, columns 32
211-- rows (32 stride bytes)
212def tpStepOffset = (lambda unrestricted layout : (family TiledProductLayout) . (lambda unrestricted stride : Nat . (lambda unrestricted step : Nat .
213  (naturalSelect (tpIsColumns layout)
214    (naturalMultiply step (naturalMultiply tpStep (naturalMultiply 2 stride)))
215    (naturalMultiply step 64)))))
216
217-- d (pair) = a * b + c[0][offset], after `wait`
218def tpWideAfter = (lambda unrestricted d : Nat . (lambda unrestricted a : Nat . (lambda unrestricted b : Nat .
219  (lambda unrestricted offset : Nat . (lambda unrestricted wait : Nat .
220    (saOp (constructor SM86InstructionBody SM86IntegerMultiplyAddWideConstant (saR d) (saR a) (saR b) (byte 0) (saU offset)
221      (saAfter wait))))))))
222
223-- `yes tail` when `condition` is nonzero, else `no tail`
224def tpChooseWith = (lambda unrestricted condition : Nat .
225  (lambda unrestricted yes : (pi unrestricted tail : (family SM86Program) . (family SM86Program)) .
226  (lambda unrestricted no : (pi unrestricted tail : (family SM86Program) . (family SM86Program)) .
227  (lambda unrestricted tail : (family SM86Program) .
228    (nat-eliminate (lambda unrestricted n : Nat . (family SM86Program)) (no tail)
229      (lambda unrestricted p : Nat . (lambda unrestricted ignored : (family SM86Program) . (yes tail))) condition)))))
230
231-- the next step's tiles copied into stage `stage` (A's chunks, then B's,
232-- each at its swizzled shared word), then each chunk's offset a step on and
233-- its pointer re-formed from it, A's and B's chunk c in turn (the first
234-- after the copies have read the pointers), then the group closed
235def tpCopyStep = (lambda unrestricted t : (family TiledProductTile) .
236  (lambda unrestricted layoutA : (family TiledProductLayout) . (lambda unrestricted layoutB : (family TiledProductLayout) .
237  (lambda unrestricted strideA : Nat . (lambda unrestricted strideB : Nat . (lambda unrestricted stage : Nat .
238  (lambda unrestricted tail : (family SM86Program) .
239    (saFor (tpChunksA t) (lambda unrestricted c : Nat . (tpCopy (tpAStore c) (tpBufferA t stage) (tpAPointer c)))
240    (saFor (tpChunksB t) (lambda unrestricted c : Nat . (tpCopy (tpBStore c) (tpBufferB t stage) (tpBPointer c)))
241    (saFor (naturalSelect (naturalLess (tpChunksA t) (tpChunksB t)) (tpChunksB t) (tpChunksA t)) (lambda unrestricted c : Nat . (lambda unrestricted rest : (family SM86Program) .
242      (tpChooseWith (naturalLess c (tpChunksA t))
243        (lambda unrestricted next : (family SM86Program) .
244          (saAddImm (tpAOffset t c) (tpAOffset t c) (tpStepOffset layoutA strideA 1)
245          (tpWideAfter (tpAPointer c) (tpAOffset t c) tpOne (saArgument 1) (naturalSelect c saWaitNone saWait3) next)))
246        (lambda unrestricted next : (family SM86Program) . next)
247      (tpChooseWith (naturalLess c (tpChunksB t))
248        (lambda unrestricted next : (family SM86Program) .
249          (saAddImm (tpBOffset t c) (tpBOffset t c) (tpStepOffset layoutB strideB 1)
250          (tpWideAfter (tpBPointer c) (tpBOffset t c) tpOne (saArgument 2) saWaitNone next)))
251        (lambda unrestricted next : (family SM86Program) . next)
252        rest))))
253    (tpCommit tail)))))))))))
254-- ---- one k-step's products ----
255-- A's fragments for k-half kk: rows layout -- m-tile mt at +mt 16 rows
256-- (1024 bytes) from the k-half's address; columns layout -- m-tile mt's own
257-- address, the k-half at +16 rows (16 x extent x 2 bytes)
258def tpLoadA = (lambda unrestricted t : (family TiledProductTile) . (lambda unrestricted layout : (family TiledProductLayout) .
259  (lambda unrestricted buffer : Nat . (lambda unrestricted kk : Nat . (lambda unrestricted tail : (family SM86Program) .
260    (saFor (tpMTiles t) (lambda unrestricted mt : Nat .
261        (tpLdsm (tpAFragment t mt kk)
262          (naturalSelect (tpIsColumns layout) (tpAMatrix mt) (tpAMatrix kk))
263          (naturalAdd (tpBufferA t buffer)
264            (naturalSelect (tpIsColumns layout)
265              (naturalMultiply kk (naturalMultiply 16 (naturalMultiply (tpBlockM t) 2)))
266              (naturalMultiply mt 1024)))
267          (tpIsColumns layout)
268          saWaitNone))
269      tail))))))
270
271-- B's fragments for k-half kk, two n-tiles an LDSM: rows layout -- the
272-- pair at +16 rows from the k-half's address; columns layout -- the pair's
273-- own address, the k-half at +16 rows
274def tpLoadB = (lambda unrestricted t : (family TiledProductTile) . (lambda unrestricted layout : (family TiledProductLayout) .
275  (lambda unrestricted buffer : Nat . (lambda unrestricted kk : Nat . (lambda unrestricted tail : (family SM86Program) .
276    (saFor (naturalDivideUnchecked (tpNTiles t) 2) (lambda unrestricted pair : Nat .
277        (tpLdsm (tpBFragment t (naturalMultiply 2 pair) kk)
278          (naturalSelect (tpIsColumns layout) (tpBMatrix pair) (tpBMatrix kk))
279          (naturalAdd (tpBufferB t buffer)
280            (naturalSelect (tpIsColumns layout)
281              (naturalMultiply kk (naturalMultiply 16 (naturalMultiply (tpBlockN t) 2)))
282              (naturalMultiply pair 1024)))
283          (tpIsColumns layout)
284          saWaitNone))
285      tail))))))
286
287-- k-half kk's products, every m-tile x n-tile, an m-tile's n-tiles in a
288-- run; the first waits for the fragments.  HMMA is fixed-latency
289-- (Accelerator.SM86.Operands.sm86TensorLatencyCycles): no scoreboard, the
290-- accumulators' next product sixteen products later.  A run's products
291-- share their A fragment, so each but the run's last keeps it in the
292-- operand-reuse cache for the next (a run is never split: tpInterleave
293-- takes it whole)
294def tpHmma = (lambda unrestricted d : Nat . (lambda unrestricted a : Nat . (lambda unrestricted b : Nat .
295  (lambda unrestricted wait : Nat . (lambda unrestricted reuse : Nat .
296    (saOp (constructor SM86InstructionBody SM86TensorCoreHalfMatrixMultiplyAccumulate16x8x16Float32
297      (saR d) (saR a) (saR b) (saR d)
298      (constructor SM86Control SM86ControlValue (byte 15) (constructor SM86YieldMode SM86Continue)
299        saNone saNone (nat-to-byte wait) (nat-to-byte reuse)))))))))
300def tpProducts = (lambda unrestricted t : (family TiledProductTile) . (lambda unrestricted kk : Nat .
301  (lambda unrestricted tail : (family SM86Program) .
302    (saFor (naturalMultiply (tpMTiles t) (tpNTiles t)) (lambda unrestricted i : Nat .
303      (let unrestricted mt = (naturalDivideUnchecked i (tpNTiles t)) in (let unrestricted nt = (naturalModuloUnchecked i (tpNTiles t)) in
304        (tpHmma (tpAcc t mt nt) (tpAFragment t mt kk) (tpBFragment t nt kk)
305          (naturalSelect i saWaitNone saWait1)
306          (naturalSelect (naturalLess (succ nt) (tpNTiles t)) 1 0)))))
307    tail))))
308
309-- the first n instructions of p; p without them
310def tpTake = (lambda unrestricted n : Nat .
311  (nat-eliminate (lambda unrestricted c : Nat . (pi unrestricted p : (family SM86Program) . (family SM86Program)))
312    (lambda unrestricted p : (family SM86Program) . sm86ProgramEmpty)
313    (lambda unrestricted predecessor : Nat . (lambda unrestricted induction : (pi unrestricted p : (family SM86Program) . (family SM86Program)) .
314      (lambda unrestricted p : (family SM86Program) .
315        (eliminate SM86Program (lambda unrestricted c : (family SM86Program) . (family SM86Program)) p
316          (branch SM86ProgramEnd . sm86ProgramEmpty)
317          (branch SM86ProgramNext head rest ignored . (constructor SM86Program SM86ProgramNext head (induction rest)))))))
318    n))
319def tpDrop = (lambda unrestricted n : Nat .
320  (nat-eliminate (lambda unrestricted c : Nat . (pi unrestricted p : (family SM86Program) . (family SM86Program)))
321    (lambda unrestricted p : (family SM86Program) . p)
322    (lambda unrestricted predecessor : Nat . (lambda unrestricted induction : (pi unrestricted p : (family SM86Program) . (family SM86Program)) .
323      (lambda unrestricted p : (family SM86Program) .
324        (eliminate SM86Program (lambda unrestricted c : (family SM86Program) . (family SM86Program)) p
325          (branch SM86ProgramEnd . sm86ProgramEmpty)
326          (branch SM86ProgramNext head rest ignored . (induction rest))))))
327    n))
328
329-- `groups` times: n instructions of p, then m of q (as many as are left);
330-- then whatever remains of either
331def tpInterleave = (lambda unrestricted groups : Nat . (lambda unrestricted n : Nat . (lambda unrestricted m : Nat .
332  (nat-eliminate
333    (lambda unrestricted c : Nat . (pi unrestricted p : (family SM86Program) . (pi unrestricted q : (family SM86Program) . (family SM86Program))))
334    (lambda unrestricted p : (family SM86Program) . (lambda unrestricted q : (family SM86Program) . (sm86ProgramAppend p q)))
335    (lambda unrestricted predecessor : Nat .
336      (lambda unrestricted induction : (pi unrestricted p : (family SM86Program) . (pi unrestricted q : (family SM86Program) . (family SM86Program))) .
337        (lambda unrestricted p : (family SM86Program) . (lambda unrestricted q : (family SM86Program) .
338          (sm86ProgramAppend (tpTake n p)
339            (sm86ProgramAppend (tpTake m q) (induction (tpDrop n p) (tpDrop m q))))))))
340    groups))))
341
342-- k-half kk's products, q spread over the gaps after their runs
343def tpProductsWith = (lambda unrestricted t : (family TiledProductTile) . (lambda unrestricted kk : Nat .
344  (lambda unrestricted q : (family SM86Program) .
345    (let unrestricted groups = (tpMTiles t) in
346    (tpInterleave groups (tpNTiles t)
347      (naturalDivideUnchecked (naturalAdd (sm86ProgramCount q) (naturalSaturatingSubtract groups 1)) groups)
348      (tpProducts t kk sm86ProgramEmpty) q)))))
349
350-- a k-step's products from buffer `buffer`, the work between them hidden
351-- under the tensor pipe's issue gaps: the first k-half's fragments; its
352-- products with the second k-half's fragment loads between their runs; the
353-- second k-half's products with `between` between theirs
354def tpStepProducts = (lambda unrestricted t : (family TiledProductTile) .
355  (lambda unrestricted layoutA : (family TiledProductLayout) . (lambda unrestricted layoutB : (family TiledProductLayout) .
356  (lambda unrestricted buffer : Nat . (lambda unrestricted between : (family SM86Program) . (lambda unrestricted tail : (family SM86Program) .
357    (tpLoadA t layoutA buffer 0 (tpLoadB t layoutB buffer 0
358    (sm86ProgramAppend
359      (tpProductsWith t 0 (tpLoadA t layoutA buffer 1 (tpLoadB t layoutB buffer 1 sm86ProgramEmpty)))
360      (sm86ProgramAppend (tpProductsWith t 1 between) tail))))))))))
361
362-- ---- addresses ----
363def tpZero : Nat = 45
364-- sixteen temporaries while the addresses are formed: the accumulators',
365-- which the prologue zeroes after the addresses (every tile has at least
366-- sixteen: an m-tile by two n-tiles, four each)
367def tpT = (lambda unrestricted n : Nat . (naturalAdd tpAccBase n))
368-- d = a ^ b (LOP3 0x3c)
369def tpXor = (lambda unrestricted d : Nat . (lambda unrestricted a : Nat . (lambda unrestricted b : Nat .
370  (saOp (constructor SM86InstructionBody SM86LogicThreeInputTruthTable (saR d) (saR a) (saR b) (byte 60) saPlain)))))
371-- d = a + b
372def tpAdd = (lambda unrestricted d : Nat . (lambda unrestricted a : Nat . (lambda unrestricted b : Nat .
373  (saOp (constructor SM86InstructionBody SM86IntegerAddThreeRegister (saR d) (saR a) (saR b) saPlain)))))
374-- d = a & immediate (through temporary 15)
375def tpAndImm = (lambda unrestricted d : Nat . (lambda unrestricted a : Nat . (lambda unrestricted immediate : Nat .
376  (lambda unrestricted tail : (family SM86Program) .
377    (saMovImm (tpT 15) immediate (saAnd d a (tpT 15) tail))))))
378
379-- An operand's chunks: `corner` holds the tile's corner in bytes (rows
380-- layout: its first row's offset; columns layout: its first column's), the
381-- operand's pointer is `argument`.  Chunk tid + c threads is row
382-- chunk / perRow, chunk chunk % perRow of the step's tile (rows layout:
383-- `extent` rows of 4 chunks; columns layout: 32 rows of extent / 8).
384-- Global: pointer + corner + row stride 2 + chunk 16.  Shared bytes (LDGSTS
385-- has no .X4): row (perRow 16) + 16 (chunk ^ swizzle), swizzle = (row / 2) % 4 for rows,
386-- row % 8 for columns.
387def tpChunkAddresses = (lambda unrestricted t : (family TiledProductTile) . (lambda unrestricted layout : (family TiledProductLayout) .
388  (lambda unrestricted extent : Nat . (lambda unrestricted stride : Nat . (lambda unrestricted corner : Nat .
389  (lambda unrestricted argument : Nat . (lambda unrestricted pointer : (pi unrestricted c : Nat . Nat) .
390  (lambda unrestricted offset : (pi unrestricted c : Nat . Nat) .
391  (lambda unrestricted store : (pi unrestricted c : Nat . Nat) . (lambda unrestricted tail : (family SM86Program) .
392    (let unrestricted columns = (tpIsColumns layout) in
393    (let unrestricted perRow = (naturalSelect columns (naturalDivideUnchecked extent 8) 4) in
394    (saFor (tpChunks t extent) (lambda unrestricted c : Nat . (lambda unrestricted rest : (family SM86Program) .
395      -- T0: the chunk; T1: its row; T2: its chunk within the row
396      (saAddImm (tpT 0) tpTid (naturalMultiply c (tpThreads t))
397      (saShr (tpT 1) (tpT 0) (tpLog2 perRow)
398      (tpAndImm (tpT 2) (tpT 0) (naturalSaturatingSubtract perRow 1)
399      -- global: corner + row stride 2 + chunk 16, kept as the chunk's
400      -- offset (the wide multiply-add's destination is not its source)
401      (saImad (offset c) (tpT 1) (naturalMultiply 2 stride) corner
402      (saImad (offset c) (tpT 2) 16 (offset c)
403      (saWide (pointer c) (offset c) tpOne argument
404      -- the swizzle: (row / 2) % 4 or row % 8, into T3
405      (saShr (tpT 3) (tpT 1) (naturalSelect columns 0 1)
406      (tpAndImm (tpT 3) (tpT 3) (naturalSelect columns 7 3)
407      (tpXor (tpT 3) (tpT 3) (tpT 2)
408      -- the shared byte: row perRow 16 + 16 swizzled chunk
409      (saImad (tpT 3) (tpT 3) 16 tpZero
410      (saImad (store c) (tpT 1) (naturalMultiply perRow 16) (tpT 3)
411        rest)))))))))))))
412      tail)))))))))))))
413
414-- The warp's LDSM addresses (bytes within a buffer's tile).  q = lane / 8
415-- selects the matrix of an LDSM.x4, r = lane % 8 its row; one bit of q
416-- picks the fragment's row half, the other its k half.  `origin` holds the
417-- warp's first row (rows layout) or column (columns layout) of the tile.
418-- Rows layout: an address per k-half kk -- row origin + 8 rowBit + r, chunk
419-- 2 kk + chunkBit, swizzled with (row / 2) % 4, at row 64 + 16 chunk.
420-- Columns layout: an address per m-tile (A) or pair of n-tiles (B) i --
421-- row 8 rowBit + r, chunk origin / 8 + 2 i + chunkBit, swizzled with r, at
422-- row (extent 2) + 16 chunk.  rowBit is q % 2 for A as rows and B as
423-- columns, q / 2 for B as rows and A as columns (the order LDSM's four
424-- matrices fill the fragments in).
425def tpMatrixAddresses = (lambda unrestricted layout : (family TiledProductLayout) . (lambda unrestricted isB : Nat .
426  (lambda unrestricted extent : Nat . (lambda unrestricted origin : Nat . (lambda unrestricted count : Nat .
427  (lambda unrestricted address : (pi unrestricted i : Nat . Nat) . (lambda unrestricted tail : (family SM86Program) .
428    (let unrestricted columns = (tpIsColumns layout) in
429    (let unrestricted rowIsHigh = (naturalSelect (naturalEqual columns isB) 0 1) in
430    -- T4: q; T5: rowBit; T6: chunkBit; T7: r
431    (saShr (tpT 4) tpLane 3
432    (tpAndImm (tpT 7) tpLane 7
433    (saShr (tpT 5) (tpT 4) (naturalSelect rowIsHigh 1 0)
434    (tpAndImm (tpT 5) (tpT 5) 1
435    (saShr (tpT 6) (tpT 4) (naturalSelect rowIsHigh 0 1)
436    (tpAndImm (tpT 6) (tpT 6) 1
437    -- T8: the row (rows: origin + 8 rowBit + r; columns: 8 rowBit + r)
438    (saImad (tpT 8) (tpT 5) 8 (tpT 7)
439    (nat-eliminate (lambda unrestricted n : Nat . (family SM86Program))
440      -- rows layout
441      (tpAdd (tpT 8) (tpT 8) origin
442      (saShr (tpT 9) (tpT 8) 1
443      (tpAndImm (tpT 9) (tpT 9) 3
444      (saFor count (lambda unrestricted i : Nat . (lambda unrestricted rest : (family SM86Program) .
445        (saAddImm (tpT 10) (tpT 6) (naturalMultiply 2 i)
446        (tpXor (tpT 10) (tpT 10) (tpT 9)
447        (saImad (tpT 10) (tpT 10) 16 tpZero
448        (saImad (address i) (tpT 8) 64 (tpT 10)
449          rest))))))
450        tail))))
451      -- columns layout
452      (lambda unrestricted p : Nat . (lambda unrestricted ignored : (family SM86Program) .
453        (saShr (tpT 9) origin 3
454        (saFor count (lambda unrestricted i : Nat . (lambda unrestricted rest : (family SM86Program) .
455          (tpAdd (tpT 10) (tpT 9) (tpT 6)
456          (saAddImm (tpT 10) (tpT 10) (naturalMultiply 2 i)
457          (tpXor (tpT 10) (tpT 10) (tpT 7)
458          (saImad (tpT 10) (tpT 10) 16 tpZero
459          (saImad (address i) (tpT 8) (naturalMultiply extent 2) (tpT 10)
460            rest)))))))
461          tail))))
462      columns)))))))))))))))))
463
464-- ---- the program ----
465-- C (M x N) = A B for the layouts; strides are the stored rows' lengths in
466-- elements (A: K for rows, M for columns; B: K for rows, N for columns;
467-- C: N).  Grid: x over N / BN, y over M / BM.  K is a multiple of 64 (an
468-- even number of k-steps).
469--
470-- The k-steps run in a loop of two (one per shared buffer), so the program's
471-- length does not grow with K:
472--
473--   prologue   addresses; steps 0 and 1 loaded into staging sets 0 and 1
474--              (SB0, SB2); set 0 into buffer 0; the iteration count,
475--              K / 64 - 1
476--   body       (the pointers at step 2i) step 2i+2 into set 0, buffer 0's
477--              products, set 1 (step 2i+1) into buffer 1; step 2i+3 into
478--              set 1, buffer 1's products, set 0 (step 2i+2) into buffer
479--              0; the pointers two steps on
480--   control    one iteration fewer; back to the body while any are left
481--   tail       buffer 0's products, set 1 into buffer 1; buffer 1's
482--              products; C
483--
484-- A piece is finished (e.g. its stalls compacted and drained) on its own;
485-- the loop's control keeps the largest stalls, so the path round the back
486-- edge leaves every fixed-latency result ready.  The checked program is the
487-- same finished pieces with the loop unrolled twice and no branch: it covers
488-- both ways into the body (from the prologue and from itself), and every
489-- iteration leaves the same things in flight (each staging set's loads
490-- under its scoreboard, the shared stores' reads under SB3; the products
491-- are fixed-latency), so two iterations reach the state every later one
492-- starts from.
493
494-- one k-step on stage `stage`: its group retired (at most stages - 2
495-- younger in flight), the barrier (every thread's copies of the stage visible, every
496-- warp done with the stage the copies below overwrite), then the products
497-- with `copies` -- the step stages - 1 on into the stage the step before
498-- used, or an empty group -- between the second k-half's
499def tpStageStep = (lambda unrestricted t : (family TiledProductTile) .
500  (lambda unrestricted layoutA : (family TiledProductLayout) . (lambda unrestricted layoutB : (family TiledProductLayout) .
501  (lambda unrestricted stage : Nat . (lambda unrestricted copies : (family SM86Program) . (lambda unrestricted tail : (family SM86Program) .
502    (tpWaitGroups (naturalSaturatingSubtract (tpAheadOf t) 1)
503    (tpBarrier saWaitNone
504    (tpStepProducts t layoutA layoutB stage copies tail)))))))))
505-- the copies of the step stages - 1 on, into the stage before `stage`
506def tpCopiesFor = (lambda unrestricted t : (family TiledProductTile) .
507  (lambda unrestricted layoutA : (family TiledProductLayout) . (lambda unrestricted layoutB : (family TiledProductLayout) .
508  (lambda unrestricted strideA : Nat . (lambda unrestricted strideB : Nat . (lambda unrestricted stage : Nat .
509    (tpCopyStep t layoutA layoutB strideA strideB (naturalModuloUnchecked (naturalAdd stage (tpAheadOf t)) (tpStagesOf t)) sm86ProgramEmpty)))))))
510def tpEmptyGroup : (family SM86Program) = (tpCommit sm86ProgramEmpty)
511-- `yes` when `condition` is nonzero, else `no`
512def tpProgramIf = (lambda unrestricted condition : Nat .
513  (lambda unrestricted yes : (family SM86Program) . (lambda unrestricted no : (family SM86Program) .
514    (nat-eliminate (lambda unrestricted n : Nat . (family SM86Program)) no
515      (lambda unrestricted p : Nat . (lambda unrestricted ignored : (family SM86Program) . yes)) condition))))
516
517-- the loop's iterations: every `stages` steps but the last `stages`
518def tpIterations = (lambda unrestricted t : (family TiledProductTile) . (lambda unrestricted k : Nat .
519  (naturalSaturatingSubtract (naturalDivideUnchecked k (naturalMultiply (tpStagesOf t) tpStep)) 1)))
520def tpPrologue = (lambda unrestricted t : (family TiledProductTile) .
521  (lambda unrestricted layoutA : (family TiledProductLayout) . (lambda unrestricted layoutB : (family TiledProductLayout) .
522  (lambda unrestricted k : Nat . (lambda unrestricted strideA : Nat . (lambda unrestricted strideB : Nat . (lambda unrestricted strideC : Nat .
523    (let unrestricted blockM = (tpBlockM t) in (let unrestricted blockN = (tpBlockN t) in
524    (saS2R tpTid (constructor SM86SpecialRegister SM86ThreadIdX)
525    (saS2R tpTileN (constructor SM86SpecialRegister SM86CooperativeThreadArrayIdX)
526    (saS2R tpTileM (constructor SM86SpecialRegister SM86CooperativeThreadArrayIdY)
527    (saMovImm tpOne 1
528    (saMovImmAfter tpZero 0 saWait5
529    (tpAndImm tpLane tpTid 31
530    (saShr tpWarp tpTid 5
531    (saShr tpWarpM tpWarp (tpLog2 (tpWarpsN t))
532    (tpAndImm tpWarpN tpWarp (naturalSaturatingSubtract (tpWarpsN t) 1)
533    -- the corners: A's rows or columns, B's, in bytes
534    (saImad tpScratch tpTileM (naturalMultiply blockM (naturalSelect (tpIsColumns layoutA) 2 (naturalMultiply 2 strideA))) tpZero
535    (tpChunkAddresses t layoutA blockM strideA tpScratch (saArgument 1) tpAPointer (tpAOffset t) tpAStore
536    (saImad tpScratch tpTileN (naturalMultiply blockN (naturalSelect (tpIsColumns layoutB) 2 (naturalMultiply 2 strideB))) tpZero
537    (tpChunkAddresses t layoutB blockN strideB tpScratch (saArgument 2) tpBPointer (tpBOffset t) tpBStore
538    -- the warp's fragments' addresses
539    (saImad tpScratch tpWarpM (naturalMultiply 16 (tpMTiles t)) tpZero
540    (tpMatrixAddresses layoutA 0 blockM tpScratch (naturalSelect (tpIsColumns layoutA) (tpMTiles t) 2) tpAMatrix
541    (saImad tpScratch tpWarpN (naturalMultiply 8 (tpNTiles t)) tpZero
542    (tpMatrixAddresses layoutB 1 blockN tpScratch (naturalSelect (tpIsColumns layoutB) (naturalDivideUnchecked (tpNTiles t) 2) 2) tpBMatrix
543    -- C: row tile M + 16 MT warpM + lane / 4, column tile N + 8 NT warpN +
544    -- 2 (lane % 4), four bytes each
545    (saShr tpScratch tpLane 2
546    (saImad tpScratch tpWarpM (naturalMultiply 16 (tpMTiles t)) tpScratch
547    (saImad tpScratch tpTileM blockM tpScratch
548    (tpAndImm tpScratch2 tpLane 3
549    (saImad tpScratch2 tpScratch2 2 tpZero
550    (saImad tpScratch2 tpWarpN (naturalMultiply 8 (tpNTiles t)) tpScratch2
551    (saImad tpScratch2 tpTileN blockN tpScratch2
552    (saImad tpScratch tpScratch strideC tpScratch2
553    (saImad tpScratch tpScratch 4 tpZero
554    (saWide tpC tpScratch tpOne (saArgument 0)
555    (saFor (tpAccCount t) (lambda unrestricted i : Nat . (saMovImm (naturalAdd tpAccBase i) 0))
556    -- steps 0 .. stages - 2 copied into their stages, a group each
557    (saFor (tpAheadOf t) (lambda unrestricted stage : Nat .
558      (tpCopyStep t layoutA layoutB strideA strideB stage))
559    (saMovImm tpCounter (tpIterations t k)
560      sm86ProgramEmpty)))))))))))))))))))))))))))))))))))))))
561
562-- a step a stage, each copying the step stages - 1 on
563def tpBody = (lambda unrestricted t : (family TiledProductTile) .
564  (lambda unrestricted layoutA : (family TiledProductLayout) . (lambda unrestricted layoutB : (family TiledProductLayout) .
565  (lambda unrestricted strideA : Nat . (lambda unrestricted strideB : Nat .
566    (saFor (tpStagesOf t) (lambda unrestricted stage : Nat .
567      (tpStageStep t layoutA layoutB stage (tpCopiesFor t layoutA layoutB strideA strideB stage)))
568      sm86ProgramEmpty))))))
569
570-- one iteration fewer and P1 = any left, with the largest stalls; the
571-- branch back over `body` instructions and these (its displacement counts
572-- from its successor, in 16-byte instructions)
573def tpControl = (lambda unrestricted tail : (family SM86Program) .
574  (saAddImm tpCounter tpCounter 4294967295
575  (saGreater saP1 tpCounter 0 tail)))
576def tpControlCount : Nat = 2
577def tpBranch = (lambda unrestricted body : Nat .
578  (saWhen saP1 (constructor SM86InstructionBody SM86Branch
579    (saU (naturalSaturatingSubtract 4294967296 (naturalMultiply 16 (naturalAdd body (naturalAdd tpControlCount 1)))))
580    (sm86Unsigned32 (byte 255) (byte 255) (byte 131) (byte 3))
581    sm86SafeControl)
582    sm86ProgramEmpty))
583
584-- C's pair i (of 2 MT NT a thread): its accumulators and its byte offset
585-- from the thread's first
586def tpPairAccumulator = (lambda unrestricted t : (family TiledProductTile) . (lambda unrestricted i : Nat .
587  (let unrestricted j = (naturalDivideUnchecked i 2) in
588    (naturalAdd (tpAcc t (naturalDivideUnchecked j (tpNTiles t)) (naturalModuloUnchecked j (tpNTiles t))) (naturalMultiply 2 (naturalModuloUnchecked i 2))))))
589def tpPairOffset = (lambda unrestricted t : (family TiledProductTile) . (lambda unrestricted strideC : Nat . (lambda unrestricted i : Nat .
590  (let unrestricted h = (naturalModuloUnchecked i 2) in (let unrestricted j = (naturalDivideUnchecked i 2) in
591    (naturalAdd (naturalMultiply (naturalAdd (naturalMultiply 16 (naturalDivideUnchecked j (tpNTiles t))) (naturalMultiply 8 h)) (naturalMultiply strideC 4))
592      (naturalMultiply (naturalModuloUnchecked j (tpNTiles t)) 32)))))))
593def tpPairs = (lambda unrestricted t : (family TiledProductTile) . (naturalMultiply 2 (naturalMultiply (tpMTiles t) (tpNTiles t))))
594
595-- C = the accumulators
596def tpStoreC = (lambda unrestricted t : (family TiledProductTile) . (lambda unrestricted strideC : Nat . (lambda unrestricted tail : (family SM86Program) .
597  (saFor (tpPairs t) (lambda unrestricted i : Nat .
598    (saOp (sbStore64 tpC (tpPairAccumulator t i) (tpPairOffset t strideC i) saWaitNone)))
599    tail))))
600
601-- C = the accumulators + R (R laid out as C, its pointer the fourth
602-- parameter): the fragment registers are free once the last products have
603-- issued; the first pair holds R's pointer (from the prologue's C offset,
604-- tpScratch), the rest eight pairs of R at a time, loaded under SB2 (SB0
605-- counts the copies' groups), added
606-- into the accumulators and stored
607def tpResidualBatch : Nat = 8
608def tpResidualC = (lambda unrestricted t : (family TiledProductTile) . (lambda unrestricted strideC : Nat . (lambda unrestricted tail : (family SM86Program) .
609  (let unrestricted pointer = (tpAFragment t 0 0) in
610  (let unrestricted value = (lambda unrestricted k : Nat . (naturalAdd pointer (naturalAdd 2 k))) in
611  (saWide pointer tpScratch tpOne (saArgument 3)
612  (saFor (naturalDivideUnchecked (tpPairs t) tpResidualBatch) (lambda unrestricted b : Nat . (lambda unrestricted rest : (family SM86Program) .
613    (let unrestricted pairAt = (lambda unrestricted k : Nat . (naturalAdd (naturalMultiply b tpResidualBatch) k)) in
614    (saFor (naturalMultiply 2 tpResidualBatch) (lambda unrestricted w : Nat .
615      (saLoad (value w) pointer (naturalAdd (tpPairOffset t strideC (pairAt (naturalDivideUnchecked w 2))) (naturalMultiply 4 (naturalModuloUnchecked w 2)))
616        saSB2 saWaitNone))
617    (saFor (naturalMultiply 2 tpResidualBatch) (lambda unrestricted w : Nat .
618      (let unrestricted accumulator = (naturalAdd (tpPairAccumulator t (pairAt (naturalDivideUnchecked w 2))) (naturalModuloUnchecked w 2)) in
619        (saFadd accumulator accumulator (value w) (naturalSelect w saWaitNone saWait2))))
620    (saFor tpResidualBatch (lambda unrestricted k : Nat .
621      (saOp (sbStore64 tpC (tpPairAccumulator t (pairAt k)) (tpPairOffset t strideC (pairAt k)) saWaitNone)))
622      rest))))))
623    tail)))))))
624
625def tpTail = (lambda unrestricted residual : Nat . (lambda unrestricted t : (family TiledProductTile) .
626  (lambda unrestricted layoutA : (family TiledProductLayout) . (lambda unrestricted layoutB : (family TiledProductLayout) .
627  (lambda unrestricted strideA : Nat . (lambda unrestricted strideB : Nat . (lambda unrestricted strideC : Nat .
628    -- the last `stages` steps: the first copies the last step, the rest close
629    -- empty groups (the count each step waits for stays stages - 2)
630    (saFor (tpStagesOf t) (lambda unrestricted stage : Nat .
631      (tpStageStep t layoutA layoutB stage
632        (tpProgramIf stage tpEmptyGroup (tpCopiesFor t layoutA layoutB strideA strideB stage))))
633    -- C, once the products are done and every group retired
634    (tpWaitGroups 0
635    (tpChooseWith residual (tpResidualC t strideC) (tpStoreC t strideC)
636    (saOp (constructor SM86InstructionBody SM86Exit (saAfter saWaitAll)) sm86ProgramEmpty)))))))))))
637
638-- `yes` when `condition` is nonzero, else `no` (only the chosen one built)
639def tpChoose = (lambda unrestricted condition : Nat . (lambda unrestricted yes : (family SM86Program) . (lambda unrestricted no : (family SM86Program) .
640  (nat-eliminate (lambda unrestricted n : Nat . (family SM86Program)) no
641    (lambda unrestricted p : Nat . (lambda unrestricted ignored : (family SM86Program) . yes)) condition))))
642
643-- the program, each piece finished by `finish`: `looped` 1 gives the loop,
644-- 0 the checked unrolling; `residual` 1 adds R to C
645def tpAssemble = (lambda unrestricted finish : (pi unrestricted piece : (family SM86Program) . (family SM86Program)) .
646  (lambda unrestricted looped : Nat . (lambda unrestricted residual : Nat .
647  (lambda unrestricted t : (family TiledProductTile) .
648  (lambda unrestricted layoutA : (family TiledProductLayout) . (lambda unrestricted layoutB : (family TiledProductLayout) .
649  (lambda unrestricted k : Nat . (lambda unrestricted strideA : Nat . (lambda unrestricted strideB : Nat . (lambda unrestricted strideC : Nat .
650    (let unrestricted iterations = (tpIterations t k) in
651    (let unrestricted body = (finish (tpBody t layoutA layoutB strideA strideB)) in
652    (let unrestricted once = (sm86ProgramAppend body (tpControl sm86ProgramEmpty)) in
653    (let unrestricted loop =
654      (tpChoose looped
655        (sm86ProgramAppend once (tpBranch (sm86ProgramCount body)))
656        (sm86ProgramAppend once (tpChoose (naturalLess 1 iterations) once sm86ProgramEmpty))) in
657    (sm86ProgramAppend (finish (tpPrologue t layoutA layoutB k strideA strideB strideC))
658      (sm86ProgramAppend (tpChoose iterations loop sm86ProgramEmpty)
659        (finish (tpTail residual t layoutA layoutB strideA strideB strideC))))))))))))))))))
660
661def tiledProductSM86 = (lambda unrestricted finish : (pi unrestricted piece : (family SM86Program) . (family SM86Program)) .
662  (tpAssemble finish 1 0))
663-- what the schedule checks read: the loop unrolled twice, no branch
664def tiledProductSM86Checked = (lambda unrestricted finish : (pi unrestricted piece : (family SM86Program) . (family SM86Program)) .
665  (tpAssemble finish 0 0))
666-- C = A B + R (a residual added in the epilogue: C and R distinct planes,
667-- R laid out as C, its pointer the fourth parameter)
668def tiledProductSM86Residual = (lambda unrestricted finish : (pi unrestricted piece : (family SM86Program) . (family SM86Program)) .
669  (tpAssemble finish 1 1))
670def tiledProductSM86ResidualChecked = (lambda unrestricted finish : (pi unrestricted piece : (family SM86Program) . (family SM86Program)) .
671  (tpAssemble finish 0 1))
672
673def tiledProductSM86Registers = tpRegisters
674def tiledProductSM86Threads = tpThreads
675def tiledProductSM86SharedBytes = tpSharedBytes
676def tiledProductSM86BlockM = tpBlockM
677def tiledProductSM86BlockN = tpBlockN
678def tiledProductSM86Stages = tpStagesOf
679def tiledProductSM86ChunksA = tpChunksA
680def tiledProductSM86ChunksB = tpChunksB
681def tiledProductSM86ConstantBytes : Nat = (saArgument 3)
682def tiledProductSM86ResidualConstantBytes : Nat = (saArgument 4)
683def tiledProductSM86OutputArgument : Nat = 0
684def tiledProductSM86LeftArgument : Nat = 1
685def tiledProductSM86RightArgument : Nat = 2
686def tiledProductSM86ResidualArgument : Nat = 3
687
688-- Admission belongs to the product, since the loop's full tiles and staging
689-- contract are the same for every model. The caller supplies the target's
690-- usable shared-memory ceiling; this module does not choose a card or a tile.
691-- The grid may be large (for a vocabulary projection), but each addressable
692-- matrix extent must fit the kernel's 32-bit element arithmetic.
693def tiledProductSM86ShapeAdmitted =
694  (lambda unrestricted tile : (family TiledProductTile) .
695    (lambda unrestricted m : Nat . (lambda unrestricted n : Nat .
696      (lambda unrestricted k : Nat . (lambda unrestricted sharedLimit : Nat .
697        (naturalAnd (naturalNonzero m)
698        (naturalAnd (naturalNonzero n)
699        (naturalAnd (naturalNonzero k)
700        (naturalAnd (naturalIsZero (naturalModuloUnchecked m (tpBlockM tile)))
701        (naturalAnd (naturalIsZero (naturalModuloUnchecked n (tpBlockN tile)))
702        (naturalAnd (naturalIsZero
703          (naturalModuloUnchecked k (naturalMultiply tpStep (tpStagesOf tile))))
704        (naturalAnd (tiledProductSM86ChunksFit tile)
705        (naturalAnd (sm86RegisterSpanAdmitted (tiledProductSM86Registers tile))
706        (naturalAnd (tiledProductSM86SharedFits tile)
707        (naturalAnd (naturalLessOrEqual (tpSharedBytes tile) sharedLimit)
708        (naturalAnd
709          (naturalLess (naturalMultiply 4 (naturalMultiply m n)) (naturalPowerOfTwo 32))
710        (naturalAnd
711          (naturalLess (naturalMultiply 2 (naturalMultiply m k)) (naturalPowerOfTwo 32))
712          (naturalLess (naturalMultiply 2 (naturalMultiply n k)) (naturalPowerOfTwo 32)))))))))))))))))))

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.