module Realization.Nvidia.SM86.IndexedRowScatterOrderedSM86 import Accelerator.SM86.Control import Accelerator.SM86.Immediate import Accelerator.SM86.Instruction import Accelerator.SM86.Types import Realization.Nvidia.SM86.IndexedRowScatterSM86 import Realization.Nvidia.SM86.LinearStepLoopSM86 import Realization.Nvidia.SM86.LinearStepSM86 import Realization.Nvidia.SM86.StallCompaction import Std.Natural -- The indexed-row scatter with a DETERMINISTIC reduction: destination row -- ids[j] += source row j for every row j, as Realization.Nvidia.SM86. -- IndexedRowScatterSM86 computes it with RED.E.ADD.F32 -- whose additions -- into a row that several source rows share land in whatever order the -- blocks run -- but in one order, fixed by the input: each destination row -- belongs to the block of the FIRST source row that names it, which adds -- every source row naming it, in ascending order, to the row's old words. -- The other blocks exit. So each destination word is written by one thread -- once, no two blocks touch the same row, and the result is -- dest[v][c] = (((dest[v][c] + src[j0][c]) + src[j1][c]) + ...) -- for the rows j0 < j1 < ... that name v -- the same bits every run -- (Learning.Checked.IndexedRowScatter; SM86.IndexedRowScatterCheck decides -- it on the block model, block by block in either order). -- -- Same ABI as the atomic form (destination, source and index pointers at -- c[0x160], c[0x168], c[0x170]); one block per source row, one thread per -- component. -- -- R0 = t (the block) R1 = c (the thread) R5 = 4 R12 = j = 0 -- R11 = v = ids[t] R25 = t + 1 -- again: u = ids[j]; exit when u = v and j != t; j += 1; again while j != t + 1 -- j = t; R20 = dest[v][c] -- again: u = ids[j]; when u = v: R20 += src[j][c]; j += 1; again while j <= rows - 1 -- dest[v][c] = R20; exit -- -- Equality is a nonzero exclusive or (ISETP.GT 0, unsigned); every loop's -- backward branch is its own length (LinearStepLoopSM86.llLoop). The scan -- takes rows 0 .. t, never none, so no branch jumps forward: only the -- backward form's encoding is proven on the card (a forward jump with the -- backward form's sign bits is a fault the thread model refuses). def irsT : Nat = 0 def irsC : Nat = 1 def irsRow : Nat = 3 def irsFour : Nat = 5 def irsIdAddress : Nat = 6 def irsIdScan : Nat = 8 def irsV : Nat = 11 def irsJ : Nat = 12 def irsU : Nat = 14 def irsDifference : Nat = 15 def irsSourceAddress : Nat = 16 def irsSource : Nat = 18 def irsAccumulator : Nat = 20 def irsDestinationAddress : Nat = 22 def irsPastT : Nat = 25 def irsEarlierFlag : Nat = 26 def irsP1 : (family SM86Predicate) = (constructor SM86Predicate SM86Predicate1) def irsP2 : (family SM86Predicate) = (constructor SM86Predicate SM86Predicate2) def irsP3 : (family SM86Predicate) = (constructor SM86Predicate SM86Predicate3) -- a ^ b (LOP3.LUT 0x3c: the third input unused) def irsExclusiveOr = (lambda unrestricted destination : Nat . (lambda unrestricted left : Nat . (lambda unrestricted right : Nat . (lambda unrestricted control : (family SM86Control) . (constructor SM86InstructionBody SM86LogicThreeInputTruthTable (lsR destination) (lsR left) (lsR right) (byte 60) control))))) -- the pair at `destination` = index x 4 + the address in parameter word `parameter` def irsAddressOf = (lambda unrestricted destination : Nat . (lambda unrestricted index : Nat . (lambda unrestricted parameter : Nat . (constructor SM86InstructionBody SM86IntegerMultiplyAddWideConstant (lsR destination) (lsR index) (lsR irsFour) (byte 0) (lsU parameter) lsControl)))) -- destination = row x width + c def irsElement = (lambda unrestricted destination : Nat . (lambda unrestricted row : Nat . (lambda unrestricted width : Nat . (constructor SM86InstructionBody SM86IntegerMultiplyAddImmediate (lsR destination) (lsR row) (lsU width) (lsR irsC) lsControl)))) def irsWhen = (lambda unrestricted predicate : (family SM86Predicate) . (lambda unrestricted body : (family SM86InstructionBody) . (lambda unrestricted tail : (family SM86Program) . (constructor SM86Program SM86ProgramNext (sm86PredicatedInstruction predicate body) tail)))) def irsUnless = (lambda unrestricted predicate : (family SM86Predicate) . (lambda unrestricted body : (family SM86InstructionBody) . (lambda unrestricted tail : (family SM86Program) . (constructor SM86Program SM86ProgramNext (sm86NegatedPredicatedInstruction predicate body) tail)))) def irsBranch = (lambda unrestricted offset : Nat . (constructor SM86InstructionBody SM86Branch (lsU offset) (sm86Unsigned32 (byte 255) (byte 255) (byte 131) (byte 3)) sm86BranchControl)) -- ids[j] into R14, and P2 = (ids[j] != v) def irsCompare = (lambda unrestricted tail : (family SM86Program) . (lsNext (irsAddressOf irsIdScan irsJ indexedRowScatterSM86ClassIdPointerNatural) (lsNext (lsLoad irsU irsIdScan 0) (lsNext (irsExclusiveOr irsDifference irsU irsV lsControlWaitSB1) (lsNext (llGreater irsP2 irsDifference 0) tail))))) -- the scan of rows 0 .. t: exit at a row before t that names v (the flag is -- j ^ t, nonzero, when the row is equal and earlier) def irsEarlierBody = (lambda unrestricted tail : (family SM86Program) . (irsCompare (lsNext (llMove irsEarlierFlag 0) (irsUnless irsP2 (irsExclusiveOr irsEarlierFlag irsJ irsT lsControl) (lsNext (llGreater irsP3 irsEarlierFlag 0) (irsWhen irsP3 (constructor SM86InstructionBody SM86Exit lsControl) (lsNext (llAddImmediate irsJ 1) (lsNext (irsExclusiveOr irsDifference irsJ irsPastT lsControl) (lsNext (llGreater irsP1 irsDifference 0) tail))))))))) -- again while j != t + 1: @P1 BRA back over the body and the branch def irsEarlier = (lambda unrestricted tail : (family SM86Program) . (irsEarlierBody (irsWhen irsP1 (irsBranch (naturalSaturatingSubtract 4294967296 (naturalMultiply 16 (succ (llLength irsEarlierBody))))) tail))) -- the accumulation over rows t .. rows - 1 def irsAccumulate = (lambda unrestricted width : Nat . (lambda unrestricted rows : Nat . (llLoop irsP1 (lambda unrestricted tail : (family SM86Program) . (irsCompare (irsUnless irsP2 (irsElement irsRow irsJ width) (irsUnless irsP2 (irsAddressOf irsSourceAddress irsRow indexedRowScatterSM86SourcePointerNatural) (irsUnless irsP2 (lsLoad irsSource irsSourceAddress 0) (irsUnless irsP2 (constructor SM86InstructionBody SM86FloatAdd (lsR irsAccumulator) (lsR irsAccumulator) (lsR irsSource) lsControlWaitSB1) (lsNext (llAddImmediate irsJ 1) (lsNext (llGreater irsP1 irsJ (naturalSaturatingSubtract rows 1)) tail))))))))))) -- The program for `width` components and `rows` source rows. def indexedRowScatterOrderedSM86ProgramRaw = (lambda unrestricted width : Nat . (lambda unrestricted rows : Nat . (lsNext (constructor SM86InstructionBody SM86SpecialToRegister (lsR irsT) (constructor SM86SpecialRegister SM86CooperativeThreadArrayIdX) lsControlSetSB0) (lsNext (constructor SM86InstructionBody SM86SpecialToRegister (lsR irsC) (constructor SM86SpecialRegister SM86ThreadIdX) lsControlSetSB0) (lsNext (llMove irsFour 4) (lsNext (llMove irsJ 0) (lsNext (constructor SM86InstructionBody SM86IntegerMultiplyAddWideConstant (lsR irsIdAddress) (lsR irsT) (lsR irsFour) (byte 0) (lsU indexedRowScatterSM86ClassIdPointerNatural) lsControlWaitSB0) (lsNext (lsLoad irsV irsIdAddress 0) (lsNext (constructor SM86InstructionBody SM86IntegerAddThreeImmediate (lsR irsPastT) (lsR irsT) (lsU 1) lsControl) (irsEarlier (lsNext (constructor SM86InstructionBody SM86IntegerAddThreeImmediate (lsR irsJ) (lsR irsT) (lsU 0) lsControl) (lsNext (irsElement irsRow irsV width) (lsNext (irsAddressOf irsDestinationAddress irsRow indexedRowScatterSM86DestinationPointerNatural) (lsNext (lsLoad irsAccumulator irsDestinationAddress 0) (irsAccumulate width rows (lsNext (constructor SM86InstructionBody SM86StoreGlobal (lsR irsDestinationAddress) (lsR irsAccumulator) (lsU 0) lsControlWaitSB1) (lsNext (constructor SM86InstructionBody SM86Exit lsControl) (constructor SM86Program SM86ProgramEnd)))))))))))))))))) -- The scan reuses address registers across global loads. An LDG write -- barrier retires its result but does not guarantee that its address read -- has finished; serialize those late reads before reusing an address. -- This only changes controls, so the loop's relative branch lengths stay -- valid. See the read-barrier ownership check in SM86.Scoreboard. def indexedRowScatterOrderedSM86Program = (lambda unrestricted width : Nat . (lambda unrestricted rows : Nat . (sm86GuardLateReads (indexedRowScatterOrderedSM86ProgramRaw width rows))))