module Realization.Nvidia.SM86.TiledProductSM86 import Accelerator.SM86.Control import Accelerator.SM86.Immediate import Accelerator.SM86.Instruction import Accelerator.SM86.NumericSemantics import Accelerator.SM86.Operands import Accelerator.SM86.Program import Accelerator.SM86.Types import Realization.Nvidia.SM86.StreamingAttentionSM86 import Std.Natural import Std.Foundation -- A matrix product that reuses its operands (PRD-PERFORMANCE P4): -- C (M x N, binary32, row-major) = A (M x K) B (K x N), half operands, a -- block per BM x BN tile of C. Each operand is given in either layout: -- -- A as rows (M x K, row-major: the activations of a forward product) -- A as columns (K x M, row-major: A^T stored, e.g. dY^T of a weight -- gradient, read from dY as it is) -- B as rows (N x K, row-major: C = A B^T, a weight used forward) -- B as columns (K x N, row-major: a weight used backward, an activation of -- a weight gradient) -- -- so no operand is transposed into another plane first. The block walks K -- in steps of 32: each step's A and B tiles go global -> shared memory -- asynchronously (cp.async: LDGSTS, a group per step, stages - 1 steps ahead -- through the tile's stages; one barrier a step), and every warp takes its fragments with -- LDSM -- .T for an operand given as columns, whose shared tile is stored the -- other way round -- and multiplies them on the tensor cores (HMMA.16816, -- binary32 accumulation, fixed-latency: no scoreboard), the next k-half's -- fragment loads and the next stage's copies between its products. Shared -- rows are swizzled by 16-byte chunks so the eight rows an LDSM reads fall -- in distinct banks. The k-steps run in a loop of one step a stage, so the -- program's length does not grow with K. -- -- The tile is a parameter (TiledProductTile): warps along M and N, each -- warp's 16-row m-tiles and 8-column n-tiles, the stages. A step's operand -- tile is extent x 64 bytes, extent x 4 16-byte chunks, which the block's -- threads copy evenly: 1, 2 or 4 chunks a thread an operand -- (tiledProductSM86ChunksFit), so a tile may be square (128 x 128, 64 x 64) -- or rectangular (128 x 64 with eight warps: two chunks of A a thread, one -- of B) and as large as the registers allow (256 x 128, each warp 64 x 64). -- M and N are multiples of the tile, K of 32 stages. Parameters (constant -- bank 0 from the SM86 parameter base): C, A, B pointers. family TiledProductLayout : Type 0 constructor TiledProductRows constructor TiledProductColumns end-family family TiledProductTile : Type 0 -- warps along M, warps along N, 16-row m-tiles and 8-column n-tiles per warp constructor TiledProductTileValue field unrestricted tiledProductWarpsM : Nat field unrestricted tiledProductWarpsN : Nat field unrestricted tiledProductMTiles : Nat field unrestricted tiledProductNTiles : Nat -- the shared-memory stages the copies run ahead through (a step's tiles -- are copied stages - 1 steps before its products) field unrestricted tiledProductStages : Nat end-family def tiledProductTile = (lambda unrestricted warpsM : Nat . (lambda unrestricted warpsN : Nat . (lambda unrestricted mTiles : Nat . (lambda unrestricted nTiles : Nat . (lambda unrestricted stages : Nat . (constructor TiledProductTile TiledProductTileValue warpsM warpsN mTiles nTiles stages)))))) def tiledProductLarge : (family TiledProductTile) = (tiledProductTile 2 4 4 4 4) def tiledProductSmall : (family TiledProductTile) = (tiledProductTile 2 2 2 4 4) def tpWarpsM = (lambda unrestricted t : (family TiledProductTile) . (eliminate TiledProductTile (lambda unrestricted c : (family TiledProductTile) . Nat) t (branch TiledProductTileValue a b c d e . a))) def tpWarpsN = (lambda unrestricted t : (family TiledProductTile) . (eliminate TiledProductTile (lambda unrestricted c : (family TiledProductTile) . Nat) t (branch TiledProductTileValue a b c d e . b))) def tpMTiles = (lambda unrestricted t : (family TiledProductTile) . (eliminate TiledProductTile (lambda unrestricted c : (family TiledProductTile) . Nat) t (branch TiledProductTileValue a b c d e . c))) def tpNTiles = (lambda unrestricted t : (family TiledProductTile) . (eliminate TiledProductTile (lambda unrestricted c : (family TiledProductTile) . Nat) t (branch TiledProductTileValue a b c d e . d))) def tpStagesOf = (lambda unrestricted t : (family TiledProductTile) . (eliminate TiledProductTile (lambda unrestricted c : (family TiledProductTile) . Nat) t (branch TiledProductTileValue a b c d e . e))) def tpThreads = (lambda unrestricted t : (family TiledProductTile) . (naturalMultiply 32 (naturalMultiply (tpWarpsM t) (tpWarpsN t)))) def tpBlockM = (lambda unrestricted t : (family TiledProductTile) . (naturalMultiply (tpWarpsM t) (naturalMultiply 16 (tpMTiles t)))) def tpBlockN = (lambda unrestricted t : (family TiledProductTile) . (naturalMultiply (tpWarpsN t) (naturalMultiply 8 (tpNTiles t)))) def tpIsColumns = (lambda unrestricted l : (family TiledProductLayout) . (eliminate TiledProductLayout (lambda unrestricted c : (family TiledProductLayout) . Nat) l (branch TiledProductRows . 0) (branch TiledProductColumns . 1))) def tpStep : Nat = 32 -- the base-2 logarithm of a power of two (at most 2^16) def tpLog2 = (lambda unrestricted n : Nat . (app (nat-eliminate (lambda unrestricted fuel : Nat . (pi unrestricted m : Nat . Nat)) (lambda unrestricted m : Nat . 0) (lambda unrestricted p : Nat . (lambda unrestricted rest : (pi unrestricted m : Nat . Nat) . (lambda unrestricted m : Nat . (naturalSelect (naturalLessOrEqual m 1) 0 (succ (rest (naturalDivideUnchecked m 2))))))) 16) n)) -- a shared tile of `extent` rows of 32 halves (64 bytes, 4 chunks), or of -- 32 rows of `extent` halves (extent / 8 chunks): either way extent x 64 -- bytes; two operands a stage, the tile's stages (a step's tiles are copied -- stages - 1 steps before its products) def tpAheadOf = (lambda unrestricted t : (family TiledProductTile) . (naturalSaturatingSubtract (tpStagesOf t) 1)) def tpTileBytes = (lambda unrestricted extent : Nat . (naturalMultiply extent 64)) def tpSharedBytes = (lambda unrestricted t : (family TiledProductTile) . (naturalMultiply (tpStagesOf t) (naturalAdd (tpTileBytes (tpBlockM t)) (tpTileBytes (tpBlockN t))))) -- what an SM has for a block's shared memory on the GB10 and the RTX 3090 -- (100 KiB), less the two kilobytes the sm_121 realization reserves (the -- lowering bases shared memory at 0x400; QMD V05 declares one more) def tpSharedCapacity : Nat = (naturalSaturatingSubtract (naturalMultiply 100 1024) 2048) -- 1 when a tile's stages fit def tiledProductSM86SharedFits = (lambda unrestricted t : (family TiledProductTile) . (naturalLessOrEqual (tpSharedBytes t) tpSharedCapacity)) def tpBufferA = (lambda unrestricted t : (family TiledProductTile) . (lambda unrestricted buffer : Nat . (naturalMultiply buffer (naturalAdd (tpTileBytes (tpBlockM t)) (tpTileBytes (tpBlockN t)))))) def tpBufferB = (lambda unrestricted t : (family TiledProductTile) . (lambda unrestricted buffer : Nat . (naturalAdd (tpBufferA t buffer) (tpTileBytes (tpBlockM t))))) -- the chunks a thread copies of an operand tile of `extent` rows (rows -- layout) or columns (columns layout) a step: extent x 4 over the threads def tpChunks = (lambda unrestricted t : (family TiledProductTile) . (lambda unrestricted extent : Nat . (naturalDivideUnchecked (naturalMultiply extent 4) (tpThreads t)))) def tpChunksA = (lambda unrestricted t : (family TiledProductTile) . (tpChunks t (tpBlockM t))) def tpChunksB = (lambda unrestricted t : (family TiledProductTile) . (tpChunks t (tpBlockN t))) -- the registers below hold up to four chunks an operand def tpChunkRoom : Nat = 4 -- 1 when every thread copies the same whole number of chunks of each -- operand, at most the room: 1, 2 or 4 (a power of two, so the copies' -- rows and chunks split by shifts) def tpChunkCountFits = (lambda unrestricted t : (family TiledProductTile) . (lambda unrestricted extent : Nat . (let unrestricted chunks = (tpChunks t extent) in (naturalAnd (naturalIsZero (naturalModuloUnchecked (naturalMultiply extent 4) (tpThreads t))) (naturalAnd (naturalLess 0 chunks) (naturalAnd (naturalLessOrEqual chunks tpChunkRoom) (naturalIsZero (naturalModuloUnchecked tpChunkRoom chunks)))))))) def tiledProductSM86ChunksFit = (lambda unrestricted t : (family TiledProductTile) . (naturalAnd (tpChunkCountFits t (tpBlockM t)) (tpChunkCountFits t (tpBlockN t)))) -- ---- registers ---- def tpTid : Nat = 0 def tpLane : Nat = 1 def tpWarp : Nat = 2 def tpWarpM : Nat = 3 def tpWarpN : Nat = 4 def tpOne : Nat = 5 def tpTileM : Nat = 6 def tpTileN : Nat = 7 def tpScratch : Nat = 8 def tpScratch2 : Nat = 9 -- the chunks a thread moves of each operand (up to four): global pointers -- (pairs) def tpAPointer = (lambda unrestricted c : Nat . (naturalAdd 10 (naturalMultiply 2 c))) def tpBPointer = (lambda unrestricted c : Nat . (naturalAdd 18 (naturalMultiply 2 c))) -- their shared byte addresses (swizzled; the stage's offset an immediate) def tpAStore = (lambda unrestricted c : Nat . (naturalAdd 26 c)) def tpBStore = (lambda unrestricted c : Nat . (naturalAdd 30 c)) -- LDSM byte addresses: A's per k-half (rows) or per m-tile (columns), B's -- per k-half (rows) or per n-tile pair (columns); up to four each def tpAMatrix = (lambda unrestricted i : Nat . (naturalAdd 34 i)) def tpBMatrix = (lambda unrestricted i : Nat . (naturalAdd 38 i)) def tpC : Nat = 42 -- the accumulators: m-tile mt, n-tile nt, four binary32 def tpAccBase : Nat = 48 def tpAcc = (lambda unrestricted t : (family TiledProductTile) . (lambda unrestricted mt : Nat . (lambda unrestricted nt : Nat . (naturalAdd tpAccBase (naturalMultiply 4 (naturalAdd (naturalMultiply mt (tpNTiles t)) nt)))))) def tpAccCount = (lambda unrestricted t : (family TiledProductTile) . (naturalMultiply 4 (naturalMultiply (tpMTiles t) (tpNTiles t)))) -- A's fragments: m-tile mt, k-half kk (four registers); B's: n-tile nt, -- k-half kk (two) def tpAFragment = (lambda unrestricted t : (family TiledProductTile) . (lambda unrestricted mt : Nat . (lambda unrestricted kk : Nat . (naturalAdd (naturalAdd tpAccBase (tpAccCount t)) (naturalMultiply 4 (naturalAdd (naturalMultiply kk (tpMTiles t)) mt)))))) def tpBFragment = (lambda unrestricted t : (family TiledProductTile) . (lambda unrestricted nt : Nat . (lambda unrestricted kk : Nat . (naturalAdd (naturalAdd (naturalAdd tpAccBase (tpAccCount t)) (naturalMultiply 8 (tpMTiles t))) (naturalMultiply 2 (naturalAdd (naturalMultiply kk (tpNTiles t)) nt)))))) -- each chunk's byte offset from its operand's pointer (A's, then B's: a -- register a chunk, so the largest tiles fit sm86RegisterLimit), -- after the fragments: each copy re-forms the pointers from them for the -- next step def tpOffsetBase = (lambda unrestricted t : (family TiledProductTile) . (naturalAdd (naturalAdd tpAccBase (tpAccCount t)) (naturalAdd (naturalMultiply 8 (tpMTiles t)) (naturalMultiply 4 (tpNTiles t))))) def tpAOffset = (lambda unrestricted t : (family TiledProductTile) . (lambda unrestricted c : Nat . (naturalAdd (tpOffsetBase t) c))) def tpBOffset = (lambda unrestricted t : (family TiledProductTile) . (lambda unrestricted c : Nat . (naturalAdd (tpOffsetBase t) (naturalAdd (tpChunksA t) c)))) -- the iterations left def tpCounter : Nat = 44 def tpRegisters = (lambda unrestricted t : (family TiledProductTile) . (naturalAdd (tpOffsetBase t) (naturalAdd (tpChunksA t) (tpChunksB t)))) def tpLdsm = (lambda unrestricted d : Nat . (lambda unrestricted address : Nat . (lambda unrestricted offset : Nat . (lambda unrestricted transpose : Nat . (lambda unrestricted wait : Nat . (saOp (constructor SM86InstructionBody SM86LoadSharedMatrix (saR d) (saR address) (saU offset) (constructor SM86SharedMatrixCount SM86SharedMatrix4) (eliminate StdBool (lambda unrestricted c : (family StdBool) . (family SM86SharedMatrixTranspose)) (stdBoolFromNatural transpose) (branch StdTrue . (constructor SM86SharedMatrixTranspose SM86SharedMatrixTransposed)) (branch StdFalse . (constructor SM86SharedMatrixTranspose SM86SharedMatrixNotTransposed))) (saCtl saSB1 wait)))))))) -- LDGSTS of a 16-byte chunk (cp.async.cg): its registers read under SB3, -- which re-forming the pointer waits on def tpCopy = (lambda unrestricted shared : Nat . (lambda unrestricted offset : Nat . (lambda unrestricted source : Nat . (saOp (constructor SM86InstructionBody SM86LoadGlobalToShared (saR shared) (saU offset) (saR source) (saU 0) (constructor SM86Control SM86ControlValue (byte 15) (constructor SM86YieldMode SM86Continue) saNone saSB3 (byte 0) (byte 0))))))) -- LDGDEPBAR: the step's copies a group, counted on SB0 def tpCommit = (saOp (constructor SM86InstructionBody SM86CommitAsyncGroup (saCtl saSB0 saWaitNone))) -- DEPBAR.LE SB0, n: at most n groups in flight def tpWaitGroups = (lambda unrestricted n : Nat . (saOp (constructor SM86InstructionBody SM86WaitAsyncGroups (nat-to-byte n) saPlain))) def tpBarrier = (lambda unrestricted wait : Nat . (saOp (constructor SM86InstructionBody SM86BarrierSynchronize (saAfter wait)))) -- ---- moving one k-step's tiles ---- -- chunk c (of the two a thread moves) of an operand tile: its row and -- 16-byte chunk in the global layout. Rows layout: `extent` rows of 64 -- bytes (4 chunks); columns layout: 32 rows of `extent` halves. -- The global offset (bytes) of a step's chunk past the operand's pointer -- for the chunk at step 0: rows move 64 bytes along K a step, columns 32 -- rows (32 stride bytes) def tpStepOffset = (lambda unrestricted layout : (family TiledProductLayout) . (lambda unrestricted stride : Nat . (lambda unrestricted step : Nat . (naturalSelect (tpIsColumns layout) (naturalMultiply step (naturalMultiply tpStep (naturalMultiply 2 stride))) (naturalMultiply step 64))))) -- d (pair) = a * b + c[0][offset], after `wait` def tpWideAfter = (lambda unrestricted d : Nat . (lambda unrestricted a : Nat . (lambda unrestricted b : Nat . (lambda unrestricted offset : Nat . (lambda unrestricted wait : Nat . (saOp (constructor SM86InstructionBody SM86IntegerMultiplyAddWideConstant (saR d) (saR a) (saR b) (byte 0) (saU offset) (saAfter wait)))))))) -- `yes tail` when `condition` is nonzero, else `no tail` def tpChooseWith = (lambda unrestricted condition : Nat . (lambda unrestricted yes : (pi unrestricted tail : (family SM86Program) . (family SM86Program)) . (lambda unrestricted no : (pi unrestricted tail : (family SM86Program) . (family SM86Program)) . (lambda unrestricted tail : (family SM86Program) . (nat-eliminate (lambda unrestricted n : Nat . (family SM86Program)) (no tail) (lambda unrestricted p : Nat . (lambda unrestricted ignored : (family SM86Program) . (yes tail))) condition))))) -- the next step's tiles copied into stage `stage` (A's chunks, then B's, -- each at its swizzled shared word), then each chunk's offset a step on and -- its pointer re-formed from it, A's and B's chunk c in turn (the first -- after the copies have read the pointers), then the group closed def tpCopyStep = (lambda unrestricted t : (family TiledProductTile) . (lambda unrestricted layoutA : (family TiledProductLayout) . (lambda unrestricted layoutB : (family TiledProductLayout) . (lambda unrestricted strideA : Nat . (lambda unrestricted strideB : Nat . (lambda unrestricted stage : Nat . (lambda unrestricted tail : (family SM86Program) . (saFor (tpChunksA t) (lambda unrestricted c : Nat . (tpCopy (tpAStore c) (tpBufferA t stage) (tpAPointer c))) (saFor (tpChunksB t) (lambda unrestricted c : Nat . (tpCopy (tpBStore c) (tpBufferB t stage) (tpBPointer c))) (saFor (naturalSelect (naturalLess (tpChunksA t) (tpChunksB t)) (tpChunksB t) (tpChunksA t)) (lambda unrestricted c : Nat . (lambda unrestricted rest : (family SM86Program) . (tpChooseWith (naturalLess c (tpChunksA t)) (lambda unrestricted next : (family SM86Program) . (saAddImm (tpAOffset t c) (tpAOffset t c) (tpStepOffset layoutA strideA 1) (tpWideAfter (tpAPointer c) (tpAOffset t c) tpOne (saArgument 1) (naturalSelect c saWaitNone saWait3) next))) (lambda unrestricted next : (family SM86Program) . next) (tpChooseWith (naturalLess c (tpChunksB t)) (lambda unrestricted next : (family SM86Program) . (saAddImm (tpBOffset t c) (tpBOffset t c) (tpStepOffset layoutB strideB 1) (tpWideAfter (tpBPointer c) (tpBOffset t c) tpOne (saArgument 2) saWaitNone next))) (lambda unrestricted next : (family SM86Program) . next) rest)))) (tpCommit tail))))))))))) -- ---- one k-step's products ---- -- A's fragments for k-half kk: rows layout -- m-tile mt at +mt 16 rows -- (1024 bytes) from the k-half's address; columns layout -- m-tile mt's own -- address, the k-half at +16 rows (16 x extent x 2 bytes) def tpLoadA = (lambda unrestricted t : (family TiledProductTile) . (lambda unrestricted layout : (family TiledProductLayout) . (lambda unrestricted buffer : Nat . (lambda unrestricted kk : Nat . (lambda unrestricted tail : (family SM86Program) . (saFor (tpMTiles t) (lambda unrestricted mt : Nat . (tpLdsm (tpAFragment t mt kk) (naturalSelect (tpIsColumns layout) (tpAMatrix mt) (tpAMatrix kk)) (naturalAdd (tpBufferA t buffer) (naturalSelect (tpIsColumns layout) (naturalMultiply kk (naturalMultiply 16 (naturalMultiply (tpBlockM t) 2))) (naturalMultiply mt 1024))) (tpIsColumns layout) saWaitNone)) tail)))))) -- B's fragments for k-half kk, two n-tiles an LDSM: rows layout -- the -- pair at +16 rows from the k-half's address; columns layout -- the pair's -- own address, the k-half at +16 rows def tpLoadB = (lambda unrestricted t : (family TiledProductTile) . (lambda unrestricted layout : (family TiledProductLayout) . (lambda unrestricted buffer : Nat . (lambda unrestricted kk : Nat . (lambda unrestricted tail : (family SM86Program) . (saFor (naturalDivideUnchecked (tpNTiles t) 2) (lambda unrestricted pair : Nat . (tpLdsm (tpBFragment t (naturalMultiply 2 pair) kk) (naturalSelect (tpIsColumns layout) (tpBMatrix pair) (tpBMatrix kk)) (naturalAdd (tpBufferB t buffer) (naturalSelect (tpIsColumns layout) (naturalMultiply kk (naturalMultiply 16 (naturalMultiply (tpBlockN t) 2))) (naturalMultiply pair 1024))) (tpIsColumns layout) saWaitNone)) tail)))))) -- k-half kk's products, every m-tile x n-tile, an m-tile's n-tiles in a -- run; the first waits for the fragments. HMMA is fixed-latency -- (Accelerator.SM86.Operands.sm86TensorLatencyCycles): no scoreboard, the -- accumulators' next product sixteen products later. A run's products -- share their A fragment, so each but the run's last keeps it in the -- operand-reuse cache for the next (a run is never split: tpInterleave -- takes it whole) def tpHmma = (lambda unrestricted d : Nat . (lambda unrestricted a : Nat . (lambda unrestricted b : Nat . (lambda unrestricted wait : Nat . (lambda unrestricted reuse : Nat . (saOp (constructor SM86InstructionBody SM86TensorCoreHalfMatrixMultiplyAccumulate16x8x16Float32 (saR d) (saR a) (saR b) (saR d) (constructor SM86Control SM86ControlValue (byte 15) (constructor SM86YieldMode SM86Continue) saNone saNone (nat-to-byte wait) (nat-to-byte reuse))))))))) def tpProducts = (lambda unrestricted t : (family TiledProductTile) . (lambda unrestricted kk : Nat . (lambda unrestricted tail : (family SM86Program) . (saFor (naturalMultiply (tpMTiles t) (tpNTiles t)) (lambda unrestricted i : Nat . (let unrestricted mt = (naturalDivideUnchecked i (tpNTiles t)) in (let unrestricted nt = (naturalModuloUnchecked i (tpNTiles t)) in (tpHmma (tpAcc t mt nt) (tpAFragment t mt kk) (tpBFragment t nt kk) (naturalSelect i saWaitNone saWait1) (naturalSelect (naturalLess (succ nt) (tpNTiles t)) 1 0))))) tail)))) -- the first n instructions of p; p without them def tpTake = (lambda unrestricted n : Nat . (nat-eliminate (lambda unrestricted c : Nat . (pi unrestricted p : (family SM86Program) . (family SM86Program))) (lambda unrestricted p : (family SM86Program) . sm86ProgramEmpty) (lambda unrestricted predecessor : Nat . (lambda unrestricted induction : (pi unrestricted p : (family SM86Program) . (family SM86Program)) . (lambda unrestricted p : (family SM86Program) . (eliminate SM86Program (lambda unrestricted c : (family SM86Program) . (family SM86Program)) p (branch SM86ProgramEnd . sm86ProgramEmpty) (branch SM86ProgramNext head rest ignored . (constructor SM86Program SM86ProgramNext head (induction rest))))))) n)) def tpDrop = (lambda unrestricted n : Nat . (nat-eliminate (lambda unrestricted c : Nat . (pi unrestricted p : (family SM86Program) . (family SM86Program))) (lambda unrestricted p : (family SM86Program) . p) (lambda unrestricted predecessor : Nat . (lambda unrestricted induction : (pi unrestricted p : (family SM86Program) . (family SM86Program)) . (lambda unrestricted p : (family SM86Program) . (eliminate SM86Program (lambda unrestricted c : (family SM86Program) . (family SM86Program)) p (branch SM86ProgramEnd . sm86ProgramEmpty) (branch SM86ProgramNext head rest ignored . (induction rest)))))) n)) -- `groups` times: n instructions of p, then m of q (as many as are left); -- then whatever remains of either def tpInterleave = (lambda unrestricted groups : Nat . (lambda unrestricted n : Nat . (lambda unrestricted m : Nat . (nat-eliminate (lambda unrestricted c : Nat . (pi unrestricted p : (family SM86Program) . (pi unrestricted q : (family SM86Program) . (family SM86Program)))) (lambda unrestricted p : (family SM86Program) . (lambda unrestricted q : (family SM86Program) . (sm86ProgramAppend p q))) (lambda unrestricted predecessor : Nat . (lambda unrestricted induction : (pi unrestricted p : (family SM86Program) . (pi unrestricted q : (family SM86Program) . (family SM86Program))) . (lambda unrestricted p : (family SM86Program) . (lambda unrestricted q : (family SM86Program) . (sm86ProgramAppend (tpTake n p) (sm86ProgramAppend (tpTake m q) (induction (tpDrop n p) (tpDrop m q)))))))) groups)))) -- k-half kk's products, q spread over the gaps after their runs def tpProductsWith = (lambda unrestricted t : (family TiledProductTile) . (lambda unrestricted kk : Nat . (lambda unrestricted q : (family SM86Program) . (let unrestricted groups = (tpMTiles t) in (tpInterleave groups (tpNTiles t) (naturalDivideUnchecked (naturalAdd (sm86ProgramCount q) (naturalSaturatingSubtract groups 1)) groups) (tpProducts t kk sm86ProgramEmpty) q))))) -- a k-step's products from buffer `buffer`, the work between them hidden -- under the tensor pipe's issue gaps: the first k-half's fragments; its -- products with the second k-half's fragment loads between their runs; the -- second k-half's products with `between` between theirs def tpStepProducts = (lambda unrestricted t : (family TiledProductTile) . (lambda unrestricted layoutA : (family TiledProductLayout) . (lambda unrestricted layoutB : (family TiledProductLayout) . (lambda unrestricted buffer : Nat . (lambda unrestricted between : (family SM86Program) . (lambda unrestricted tail : (family SM86Program) . (tpLoadA t layoutA buffer 0 (tpLoadB t layoutB buffer 0 (sm86ProgramAppend (tpProductsWith t 0 (tpLoadA t layoutA buffer 1 (tpLoadB t layoutB buffer 1 sm86ProgramEmpty))) (sm86ProgramAppend (tpProductsWith t 1 between) tail)))))))))) -- ---- addresses ---- def tpZero : Nat = 45 -- sixteen temporaries while the addresses are formed: the accumulators', -- which the prologue zeroes after the addresses (every tile has at least -- sixteen: an m-tile by two n-tiles, four each) def tpT = (lambda unrestricted n : Nat . (naturalAdd tpAccBase n)) -- d = a ^ b (LOP3 0x3c) def tpXor = (lambda unrestricted d : Nat . (lambda unrestricted a : Nat . (lambda unrestricted b : Nat . (saOp (constructor SM86InstructionBody SM86LogicThreeInputTruthTable (saR d) (saR a) (saR b) (byte 60) saPlain))))) -- d = a + b def tpAdd = (lambda unrestricted d : Nat . (lambda unrestricted a : Nat . (lambda unrestricted b : Nat . (saOp (constructor SM86InstructionBody SM86IntegerAddThreeRegister (saR d) (saR a) (saR b) saPlain))))) -- d = a & immediate (through temporary 15) def tpAndImm = (lambda unrestricted d : Nat . (lambda unrestricted a : Nat . (lambda unrestricted immediate : Nat . (lambda unrestricted tail : (family SM86Program) . (saMovImm (tpT 15) immediate (saAnd d a (tpT 15) tail)))))) -- An operand's chunks: `corner` holds the tile's corner in bytes (rows -- layout: its first row's offset; columns layout: its first column's), the -- operand's pointer is `argument`. Chunk tid + c threads is row -- chunk / perRow, chunk chunk % perRow of the step's tile (rows layout: -- `extent` rows of 4 chunks; columns layout: 32 rows of extent / 8). -- Global: pointer + corner + row stride 2 + chunk 16. Shared bytes (LDGSTS -- has no .X4): row (perRow 16) + 16 (chunk ^ swizzle), swizzle = (row / 2) % 4 for rows, -- row % 8 for columns. def tpChunkAddresses = (lambda unrestricted t : (family TiledProductTile) . (lambda unrestricted layout : (family TiledProductLayout) . (lambda unrestricted extent : Nat . (lambda unrestricted stride : Nat . (lambda unrestricted corner : Nat . (lambda unrestricted argument : Nat . (lambda unrestricted pointer : (pi unrestricted c : Nat . Nat) . (lambda unrestricted offset : (pi unrestricted c : Nat . Nat) . (lambda unrestricted store : (pi unrestricted c : Nat . Nat) . (lambda unrestricted tail : (family SM86Program) . (let unrestricted columns = (tpIsColumns layout) in (let unrestricted perRow = (naturalSelect columns (naturalDivideUnchecked extent 8) 4) in (saFor (tpChunks t extent) (lambda unrestricted c : Nat . (lambda unrestricted rest : (family SM86Program) . -- T0: the chunk; T1: its row; T2: its chunk within the row (saAddImm (tpT 0) tpTid (naturalMultiply c (tpThreads t)) (saShr (tpT 1) (tpT 0) (tpLog2 perRow) (tpAndImm (tpT 2) (tpT 0) (naturalSaturatingSubtract perRow 1) -- global: corner + row stride 2 + chunk 16, kept as the chunk's -- offset (the wide multiply-add's destination is not its source) (saImad (offset c) (tpT 1) (naturalMultiply 2 stride) corner (saImad (offset c) (tpT 2) 16 (offset c) (saWide (pointer c) (offset c) tpOne argument -- the swizzle: (row / 2) % 4 or row % 8, into T3 (saShr (tpT 3) (tpT 1) (naturalSelect columns 0 1) (tpAndImm (tpT 3) (tpT 3) (naturalSelect columns 7 3) (tpXor (tpT 3) (tpT 3) (tpT 2) -- the shared byte: row perRow 16 + 16 swizzled chunk (saImad (tpT 3) (tpT 3) 16 tpZero (saImad (store c) (tpT 1) (naturalMultiply perRow 16) (tpT 3) rest))))))))))))) tail))))))))))))) -- The warp's LDSM addresses (bytes within a buffer's tile). q = lane / 8 -- selects the matrix of an LDSM.x4, r = lane % 8 its row; one bit of q -- picks the fragment's row half, the other its k half. `origin` holds the -- warp's first row (rows layout) or column (columns layout) of the tile. -- Rows layout: an address per k-half kk -- row origin + 8 rowBit + r, chunk -- 2 kk + chunkBit, swizzled with (row / 2) % 4, at row 64 + 16 chunk. -- Columns layout: an address per m-tile (A) or pair of n-tiles (B) i -- -- row 8 rowBit + r, chunk origin / 8 + 2 i + chunkBit, swizzled with r, at -- row (extent 2) + 16 chunk. rowBit is q % 2 for A as rows and B as -- columns, q / 2 for B as rows and A as columns (the order LDSM's four -- matrices fill the fragments in). def tpMatrixAddresses = (lambda unrestricted layout : (family TiledProductLayout) . (lambda unrestricted isB : Nat . (lambda unrestricted extent : Nat . (lambda unrestricted origin : Nat . (lambda unrestricted count : Nat . (lambda unrestricted address : (pi unrestricted i : Nat . Nat) . (lambda unrestricted tail : (family SM86Program) . (let unrestricted columns = (tpIsColumns layout) in (let unrestricted rowIsHigh = (naturalSelect (naturalEqual columns isB) 0 1) in -- T4: q; T5: rowBit; T6: chunkBit; T7: r (saShr (tpT 4) tpLane 3 (tpAndImm (tpT 7) tpLane 7 (saShr (tpT 5) (tpT 4) (naturalSelect rowIsHigh 1 0) (tpAndImm (tpT 5) (tpT 5) 1 (saShr (tpT 6) (tpT 4) (naturalSelect rowIsHigh 0 1) (tpAndImm (tpT 6) (tpT 6) 1 -- T8: the row (rows: origin + 8 rowBit + r; columns: 8 rowBit + r) (saImad (tpT 8) (tpT 5) 8 (tpT 7) (nat-eliminate (lambda unrestricted n : Nat . (family SM86Program)) -- rows layout (tpAdd (tpT 8) (tpT 8) origin (saShr (tpT 9) (tpT 8) 1 (tpAndImm (tpT 9) (tpT 9) 3 (saFor count (lambda unrestricted i : Nat . (lambda unrestricted rest : (family SM86Program) . (saAddImm (tpT 10) (tpT 6) (naturalMultiply 2 i) (tpXor (tpT 10) (tpT 10) (tpT 9) (saImad (tpT 10) (tpT 10) 16 tpZero (saImad (address i) (tpT 8) 64 (tpT 10) rest)))))) tail)))) -- columns layout (lambda unrestricted p : Nat . (lambda unrestricted ignored : (family SM86Program) . (saShr (tpT 9) origin 3 (saFor count (lambda unrestricted i : Nat . (lambda unrestricted rest : (family SM86Program) . (tpAdd (tpT 10) (tpT 9) (tpT 6) (saAddImm (tpT 10) (tpT 10) (naturalMultiply 2 i) (tpXor (tpT 10) (tpT 10) (tpT 7) (saImad (tpT 10) (tpT 10) 16 tpZero (saImad (address i) (tpT 8) (naturalMultiply extent 2) (tpT 10) rest))))))) tail)))) columns))))))))))))))))) -- ---- the program ---- -- C (M x N) = A B for the layouts; strides are the stored rows' lengths in -- elements (A: K for rows, M for columns; B: K for rows, N for columns; -- C: N). Grid: x over N / BN, y over M / BM. K is a multiple of 64 (an -- even number of k-steps). -- -- The k-steps run in a loop of two (one per shared buffer), so the program's -- length does not grow with K: -- -- prologue addresses; steps 0 and 1 loaded into staging sets 0 and 1 -- (SB0, SB2); set 0 into buffer 0; the iteration count, -- K / 64 - 1 -- body (the pointers at step 2i) step 2i+2 into set 0, buffer 0's -- products, set 1 (step 2i+1) into buffer 1; step 2i+3 into -- set 1, buffer 1's products, set 0 (step 2i+2) into buffer -- 0; the pointers two steps on -- control one iteration fewer; back to the body while any are left -- tail buffer 0's products, set 1 into buffer 1; buffer 1's -- products; C -- -- A piece is finished (e.g. its stalls compacted and drained) on its own; -- the loop's control keeps the largest stalls, so the path round the back -- edge leaves every fixed-latency result ready. The checked program is the -- same finished pieces with the loop unrolled twice and no branch: it covers -- both ways into the body (from the prologue and from itself), and every -- iteration leaves the same things in flight (each staging set's loads -- under its scoreboard, the shared stores' reads under SB3; the products -- are fixed-latency), so two iterations reach the state every later one -- starts from. -- one k-step on stage `stage`: its group retired (at most stages - 2 -- younger in flight), the barrier (every thread's copies of the stage visible, every -- warp done with the stage the copies below overwrite), then the products -- with `copies` -- the step stages - 1 on into the stage the step before -- used, or an empty group -- between the second k-half's def tpStageStep = (lambda unrestricted t : (family TiledProductTile) . (lambda unrestricted layoutA : (family TiledProductLayout) . (lambda unrestricted layoutB : (family TiledProductLayout) . (lambda unrestricted stage : Nat . (lambda unrestricted copies : (family SM86Program) . (lambda unrestricted tail : (family SM86Program) . (tpWaitGroups (naturalSaturatingSubtract (tpAheadOf t) 1) (tpBarrier saWaitNone (tpStepProducts t layoutA layoutB stage copies tail))))))))) -- the copies of the step stages - 1 on, into the stage before `stage` def tpCopiesFor = (lambda unrestricted t : (family TiledProductTile) . (lambda unrestricted layoutA : (family TiledProductLayout) . (lambda unrestricted layoutB : (family TiledProductLayout) . (lambda unrestricted strideA : Nat . (lambda unrestricted strideB : Nat . (lambda unrestricted stage : Nat . (tpCopyStep t layoutA layoutB strideA strideB (naturalModuloUnchecked (naturalAdd stage (tpAheadOf t)) (tpStagesOf t)) sm86ProgramEmpty))))))) def tpEmptyGroup : (family SM86Program) = (tpCommit sm86ProgramEmpty) -- `yes` when `condition` is nonzero, else `no` def tpProgramIf = (lambda unrestricted condition : Nat . (lambda unrestricted yes : (family SM86Program) . (lambda unrestricted no : (family SM86Program) . (nat-eliminate (lambda unrestricted n : Nat . (family SM86Program)) no (lambda unrestricted p : Nat . (lambda unrestricted ignored : (family SM86Program) . yes)) condition)))) -- the loop's iterations: every `stages` steps but the last `stages` def tpIterations = (lambda unrestricted t : (family TiledProductTile) . (lambda unrestricted k : Nat . (naturalSaturatingSubtract (naturalDivideUnchecked k (naturalMultiply (tpStagesOf t) tpStep)) 1))) def tpPrologue = (lambda unrestricted t : (family TiledProductTile) . (lambda unrestricted layoutA : (family TiledProductLayout) . (lambda unrestricted layoutB : (family TiledProductLayout) . (lambda unrestricted k : Nat . (lambda unrestricted strideA : Nat . (lambda unrestricted strideB : Nat . (lambda unrestricted strideC : Nat . (let unrestricted blockM = (tpBlockM t) in (let unrestricted blockN = (tpBlockN t) in (saS2R tpTid (constructor SM86SpecialRegister SM86ThreadIdX) (saS2R tpTileN (constructor SM86SpecialRegister SM86CooperativeThreadArrayIdX) (saS2R tpTileM (constructor SM86SpecialRegister SM86CooperativeThreadArrayIdY) (saMovImm tpOne 1 (saMovImmAfter tpZero 0 saWait5 (tpAndImm tpLane tpTid 31 (saShr tpWarp tpTid 5 (saShr tpWarpM tpWarp (tpLog2 (tpWarpsN t)) (tpAndImm tpWarpN tpWarp (naturalSaturatingSubtract (tpWarpsN t) 1) -- the corners: A's rows or columns, B's, in bytes (saImad tpScratch tpTileM (naturalMultiply blockM (naturalSelect (tpIsColumns layoutA) 2 (naturalMultiply 2 strideA))) tpZero (tpChunkAddresses t layoutA blockM strideA tpScratch (saArgument 1) tpAPointer (tpAOffset t) tpAStore (saImad tpScratch tpTileN (naturalMultiply blockN (naturalSelect (tpIsColumns layoutB) 2 (naturalMultiply 2 strideB))) tpZero (tpChunkAddresses t layoutB blockN strideB tpScratch (saArgument 2) tpBPointer (tpBOffset t) tpBStore -- the warp's fragments' addresses (saImad tpScratch tpWarpM (naturalMultiply 16 (tpMTiles t)) tpZero (tpMatrixAddresses layoutA 0 blockM tpScratch (naturalSelect (tpIsColumns layoutA) (tpMTiles t) 2) tpAMatrix (saImad tpScratch tpWarpN (naturalMultiply 8 (tpNTiles t)) tpZero (tpMatrixAddresses layoutB 1 blockN tpScratch (naturalSelect (tpIsColumns layoutB) (naturalDivideUnchecked (tpNTiles t) 2) 2) tpBMatrix -- C: row tile M + 16 MT warpM + lane / 4, column tile N + 8 NT warpN + -- 2 (lane % 4), four bytes each (saShr tpScratch tpLane 2 (saImad tpScratch tpWarpM (naturalMultiply 16 (tpMTiles t)) tpScratch (saImad tpScratch tpTileM blockM tpScratch (tpAndImm tpScratch2 tpLane 3 (saImad tpScratch2 tpScratch2 2 tpZero (saImad tpScratch2 tpWarpN (naturalMultiply 8 (tpNTiles t)) tpScratch2 (saImad tpScratch2 tpTileN blockN tpScratch2 (saImad tpScratch tpScratch strideC tpScratch2 (saImad tpScratch tpScratch 4 tpZero (saWide tpC tpScratch tpOne (saArgument 0) (saFor (tpAccCount t) (lambda unrestricted i : Nat . (saMovImm (naturalAdd tpAccBase i) 0)) -- steps 0 .. stages - 2 copied into their stages, a group each (saFor (tpAheadOf t) (lambda unrestricted stage : Nat . (tpCopyStep t layoutA layoutB strideA strideB stage)) (saMovImm tpCounter (tpIterations t k) sm86ProgramEmpty))))))))))))))))))))))))))))))))))))))) -- a step a stage, each copying the step stages - 1 on def tpBody = (lambda unrestricted t : (family TiledProductTile) . (lambda unrestricted layoutA : (family TiledProductLayout) . (lambda unrestricted layoutB : (family TiledProductLayout) . (lambda unrestricted strideA : Nat . (lambda unrestricted strideB : Nat . (saFor (tpStagesOf t) (lambda unrestricted stage : Nat . (tpStageStep t layoutA layoutB stage (tpCopiesFor t layoutA layoutB strideA strideB stage))) sm86ProgramEmpty)))))) -- one iteration fewer and P1 = any left, with the largest stalls; the -- branch back over `body` instructions and these (its displacement counts -- from its successor, in 16-byte instructions) def tpControl = (lambda unrestricted tail : (family SM86Program) . (saAddImm tpCounter tpCounter 4294967295 (saGreater saP1 tpCounter 0 tail))) def tpControlCount : Nat = 2 def tpBranch = (lambda unrestricted body : Nat . (saWhen saP1 (constructor SM86InstructionBody SM86Branch (saU (naturalSaturatingSubtract 4294967296 (naturalMultiply 16 (naturalAdd body (naturalAdd tpControlCount 1))))) (sm86Unsigned32 (byte 255) (byte 255) (byte 131) (byte 3)) sm86SafeControl) sm86ProgramEmpty)) -- C's pair i (of 2 MT NT a thread): its accumulators and its byte offset -- from the thread's first def tpPairAccumulator = (lambda unrestricted t : (family TiledProductTile) . (lambda unrestricted i : Nat . (let unrestricted j = (naturalDivideUnchecked i 2) in (naturalAdd (tpAcc t (naturalDivideUnchecked j (tpNTiles t)) (naturalModuloUnchecked j (tpNTiles t))) (naturalMultiply 2 (naturalModuloUnchecked i 2)))))) def tpPairOffset = (lambda unrestricted t : (family TiledProductTile) . (lambda unrestricted strideC : Nat . (lambda unrestricted i : Nat . (let unrestricted h = (naturalModuloUnchecked i 2) in (let unrestricted j = (naturalDivideUnchecked i 2) in (naturalAdd (naturalMultiply (naturalAdd (naturalMultiply 16 (naturalDivideUnchecked j (tpNTiles t))) (naturalMultiply 8 h)) (naturalMultiply strideC 4)) (naturalMultiply (naturalModuloUnchecked j (tpNTiles t)) 32))))))) def tpPairs = (lambda unrestricted t : (family TiledProductTile) . (naturalMultiply 2 (naturalMultiply (tpMTiles t) (tpNTiles t)))) -- C = the accumulators def tpStoreC = (lambda unrestricted t : (family TiledProductTile) . (lambda unrestricted strideC : Nat . (lambda unrestricted tail : (family SM86Program) . (saFor (tpPairs t) (lambda unrestricted i : Nat . (saOp (sbStore64 tpC (tpPairAccumulator t i) (tpPairOffset t strideC i) saWaitNone))) tail)))) -- C = the accumulators + R (R laid out as C, its pointer the fourth -- parameter): the fragment registers are free once the last products have -- issued; the first pair holds R's pointer (from the prologue's C offset, -- tpScratch), the rest eight pairs of R at a time, loaded under SB2 (SB0 -- counts the copies' groups), added -- into the accumulators and stored def tpResidualBatch : Nat = 8 def tpResidualC = (lambda unrestricted t : (family TiledProductTile) . (lambda unrestricted strideC : Nat . (lambda unrestricted tail : (family SM86Program) . (let unrestricted pointer = (tpAFragment t 0 0) in (let unrestricted value = (lambda unrestricted k : Nat . (naturalAdd pointer (naturalAdd 2 k))) in (saWide pointer tpScratch tpOne (saArgument 3) (saFor (naturalDivideUnchecked (tpPairs t) tpResidualBatch) (lambda unrestricted b : Nat . (lambda unrestricted rest : (family SM86Program) . (let unrestricted pairAt = (lambda unrestricted k : Nat . (naturalAdd (naturalMultiply b tpResidualBatch) k)) in (saFor (naturalMultiply 2 tpResidualBatch) (lambda unrestricted w : Nat . (saLoad (value w) pointer (naturalAdd (tpPairOffset t strideC (pairAt (naturalDivideUnchecked w 2))) (naturalMultiply 4 (naturalModuloUnchecked w 2))) saSB2 saWaitNone)) (saFor (naturalMultiply 2 tpResidualBatch) (lambda unrestricted w : Nat . (let unrestricted accumulator = (naturalAdd (tpPairAccumulator t (pairAt (naturalDivideUnchecked w 2))) (naturalModuloUnchecked w 2)) in (saFadd accumulator accumulator (value w) (naturalSelect w saWaitNone saWait2)))) (saFor tpResidualBatch (lambda unrestricted k : Nat . (saOp (sbStore64 tpC (tpPairAccumulator t (pairAt k)) (tpPairOffset t strideC (pairAt k)) saWaitNone))) rest)))))) tail))))))) def tpTail = (lambda unrestricted residual : Nat . (lambda unrestricted t : (family TiledProductTile) . (lambda unrestricted layoutA : (family TiledProductLayout) . (lambda unrestricted layoutB : (family TiledProductLayout) . (lambda unrestricted strideA : Nat . (lambda unrestricted strideB : Nat . (lambda unrestricted strideC : Nat . -- the last `stages` steps: the first copies the last step, the rest close -- empty groups (the count each step waits for stays stages - 2) (saFor (tpStagesOf t) (lambda unrestricted stage : Nat . (tpStageStep t layoutA layoutB stage (tpProgramIf stage tpEmptyGroup (tpCopiesFor t layoutA layoutB strideA strideB stage)))) -- C, once the products are done and every group retired (tpWaitGroups 0 (tpChooseWith residual (tpResidualC t strideC) (tpStoreC t strideC) (saOp (constructor SM86InstructionBody SM86Exit (saAfter saWaitAll)) sm86ProgramEmpty))))))))))) -- `yes` when `condition` is nonzero, else `no` (only the chosen one built) def tpChoose = (lambda unrestricted condition : Nat . (lambda unrestricted yes : (family SM86Program) . (lambda unrestricted no : (family SM86Program) . (nat-eliminate (lambda unrestricted n : Nat . (family SM86Program)) no (lambda unrestricted p : Nat . (lambda unrestricted ignored : (family SM86Program) . yes)) condition)))) -- the program, each piece finished by `finish`: `looped` 1 gives the loop, -- 0 the checked unrolling; `residual` 1 adds R to C def tpAssemble = (lambda unrestricted finish : (pi unrestricted piece : (family SM86Program) . (family SM86Program)) . (lambda unrestricted looped : Nat . (lambda unrestricted residual : Nat . (lambda unrestricted t : (family TiledProductTile) . (lambda unrestricted layoutA : (family TiledProductLayout) . (lambda unrestricted layoutB : (family TiledProductLayout) . (lambda unrestricted k : Nat . (lambda unrestricted strideA : Nat . (lambda unrestricted strideB : Nat . (lambda unrestricted strideC : Nat . (let unrestricted iterations = (tpIterations t k) in (let unrestricted body = (finish (tpBody t layoutA layoutB strideA strideB)) in (let unrestricted once = (sm86ProgramAppend body (tpControl sm86ProgramEmpty)) in (let unrestricted loop = (tpChoose looped (sm86ProgramAppend once (tpBranch (sm86ProgramCount body))) (sm86ProgramAppend once (tpChoose (naturalLess 1 iterations) once sm86ProgramEmpty))) in (sm86ProgramAppend (finish (tpPrologue t layoutA layoutB k strideA strideB strideC)) (sm86ProgramAppend (tpChoose iterations loop sm86ProgramEmpty) (finish (tpTail residual t layoutA layoutB strideA strideB strideC)))))))))))))))))) def tiledProductSM86 = (lambda unrestricted finish : (pi unrestricted piece : (family SM86Program) . (family SM86Program)) . (tpAssemble finish 1 0)) -- what the schedule checks read: the loop unrolled twice, no branch def tiledProductSM86Checked = (lambda unrestricted finish : (pi unrestricted piece : (family SM86Program) . (family SM86Program)) . (tpAssemble finish 0 0)) -- C = A B + R (a residual added in the epilogue: C and R distinct planes, -- R laid out as C, its pointer the fourth parameter) def tiledProductSM86Residual = (lambda unrestricted finish : (pi unrestricted piece : (family SM86Program) . (family SM86Program)) . (tpAssemble finish 1 1)) def tiledProductSM86ResidualChecked = (lambda unrestricted finish : (pi unrestricted piece : (family SM86Program) . (family SM86Program)) . (tpAssemble finish 0 1)) def tiledProductSM86Registers = tpRegisters def tiledProductSM86Threads = tpThreads def tiledProductSM86SharedBytes = tpSharedBytes def tiledProductSM86BlockM = tpBlockM def tiledProductSM86BlockN = tpBlockN def tiledProductSM86Stages = tpStagesOf def tiledProductSM86ChunksA = tpChunksA def tiledProductSM86ChunksB = tpChunksB def tiledProductSM86ConstantBytes : Nat = (saArgument 3) def tiledProductSM86ResidualConstantBytes : Nat = (saArgument 4) def tiledProductSM86OutputArgument : Nat = 0 def tiledProductSM86LeftArgument : Nat = 1 def tiledProductSM86RightArgument : Nat = 2 def tiledProductSM86ResidualArgument : Nat = 3 -- Admission belongs to the product, since the loop's full tiles and staging -- contract are the same for every model. The caller supplies the target's -- usable shared-memory ceiling; this module does not choose a card or a tile. -- The grid may be large (for a vocabulary projection), but each addressable -- matrix extent must fit the kernel's 32-bit element arithmetic. def tiledProductSM86ShapeAdmitted = (lambda unrestricted tile : (family TiledProductTile) . (lambda unrestricted m : Nat . (lambda unrestricted n : Nat . (lambda unrestricted k : Nat . (lambda unrestricted sharedLimit : Nat . (naturalAnd (naturalNonzero m) (naturalAnd (naturalNonzero n) (naturalAnd (naturalNonzero k) (naturalAnd (naturalIsZero (naturalModuloUnchecked m (tpBlockM tile))) (naturalAnd (naturalIsZero (naturalModuloUnchecked n (tpBlockN tile))) (naturalAnd (naturalIsZero (naturalModuloUnchecked k (naturalMultiply tpStep (tpStagesOf tile)))) (naturalAnd (tiledProductSM86ChunksFit tile) (naturalAnd (sm86RegisterSpanAdmitted (tiledProductSM86Registers tile)) (naturalAnd (tiledProductSM86SharedFits tile) (naturalAnd (naturalLessOrEqual (tpSharedBytes tile) sharedLimit) (naturalAnd (naturalLess (naturalMultiply 4 (naturalMultiply m n)) (naturalPowerOfTwo 32)) (naturalAnd (naturalLess (naturalMultiply 2 (naturalMultiply m k)) (naturalPowerOfTwo 32)) (naturalLess (naturalMultiply 2 (naturalMultiply n k)) (naturalPowerOfTwo 32)))))))))))))))))))