Source/Packages

Realization.Nvidia.SM86.IndexedRowScatterOrderedSM86

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

162 lines31 declarations8.5 KiBSHA-256 2c5e7fb343d9

Complete file · line 54

IndexedRowScatterOrderedSM86.alpha

Definition view
1module Realization.Nvidia.SM86.IndexedRowScatterOrderedSM86
2
3import Accelerator.SM86.Control
4import Accelerator.SM86.Immediate
5import Accelerator.SM86.Instruction
6import Accelerator.SM86.Types
7import Realization.Nvidia.SM86.IndexedRowScatterSM86
8import Realization.Nvidia.SM86.LinearStepLoopSM86
9import Realization.Nvidia.SM86.LinearStepSM86
10import Realization.Nvidia.SM86.StallCompaction
11import Std.Natural
12
13-- The indexed-row scatter with a DETERMINISTIC reduction: destination row
14-- ids[j] += source row j for every row j, as Realization.Nvidia.SM86.
15-- IndexedRowScatterSM86 computes it with RED.E.ADD.F32 -- whose additions
16-- into a row that several source rows share land in whatever order the
17-- blocks run -- but in one order, fixed by the input: each destination row
18-- belongs to the block of the FIRST source row that names it, which adds
19-- every source row naming it, in ascending order, to the row's old words.
20-- The other blocks exit.  So each destination word is written by one thread
21-- once, no two blocks touch the same row, and the result is
22--   dest[v][c] = (((dest[v][c] + src[j0][c]) + src[j1][c]) + ...)
23-- for the rows j0 < j1 < ... that name v -- the same bits every run
24-- (Learning.Checked.IndexedRowScatter; SM86.IndexedRowScatterCheck decides
25-- it on the block model, block by block in either order).
26--
27-- Same ABI as the atomic form (destination, source and index pointers at
28-- c[0x160], c[0x168], c[0x170]); one block per source row, one thread per
29-- component.
30--
31--   R0 = t (the block)   R1 = c (the thread)   R5 = 4   R12 = j = 0
32--   R11 = v = ids[t]   R25 = t + 1
33--   again:  u = ids[j]; exit when u = v and j != t; j += 1; again while j != t + 1
34--   j = t;  R20 = dest[v][c]
35--   again:  u = ids[j]; when u = v: R20 += src[j][c];  j += 1; again while j <= rows - 1
36--   dest[v][c] = R20; exit
37--
38-- Equality is a nonzero exclusive or (ISETP.GT 0, unsigned); every loop's
39-- backward branch is its own length (LinearStepLoopSM86.llLoop).  The scan
40-- takes rows 0 .. t, never none, so no branch jumps forward: only the
41-- backward form's encoding is proven on the card (a forward jump with the
42-- backward form's sign bits is a fault the thread model refuses).
43
44def irsT : Nat = 0
45def irsC : Nat = 1
46def irsRow : Nat = 3
47def irsFour : Nat = 5
48def irsIdAddress : Nat = 6
49def irsIdScan : Nat = 8
50def irsV : Nat = 11
51def irsJ : Nat = 12
52def irsU : Nat = 14
53def irsDifference : Nat = 15
54def irsSourceAddress : Nat = 16
55def irsSource : Nat = 18
56def irsAccumulator : Nat = 20
57def irsDestinationAddress : Nat = 22
58def irsPastT : Nat = 25
59def irsEarlierFlag : Nat = 26
60
61def irsP1 : (family SM86Predicate) = (constructor SM86Predicate SM86Predicate1)
62def irsP2 : (family SM86Predicate) = (constructor SM86Predicate SM86Predicate2)
63def irsP3 : (family SM86Predicate) = (constructor SM86Predicate SM86Predicate3)
64
65-- a ^ b (LOP3.LUT 0x3c: the third input unused)
66def irsExclusiveOr =
67  (lambda unrestricted destination : Nat . (lambda unrestricted left : Nat . (lambda unrestricted right : Nat . (lambda unrestricted control : (family SM86Control) .
68    (constructor SM86InstructionBody SM86LogicThreeInputTruthTable (lsR destination) (lsR left) (lsR right) (byte 60) control)))))
69
70-- the pair at `destination` = index x 4 + the address in parameter word `parameter`
71def irsAddressOf =
72  (lambda unrestricted destination : Nat . (lambda unrestricted index : Nat . (lambda unrestricted parameter : Nat .
73    (constructor SM86InstructionBody SM86IntegerMultiplyAddWideConstant (lsR destination) (lsR index) (lsR irsFour) (byte 0) (lsU parameter) lsControl))))
74
75-- destination = row x width + c
76def irsElement =
77  (lambda unrestricted destination : Nat . (lambda unrestricted row : Nat . (lambda unrestricted width : Nat .
78    (constructor SM86InstructionBody SM86IntegerMultiplyAddImmediate (lsR destination) (lsR row) (lsU width) (lsR irsC) lsControl))))
79
80def irsWhen =
81  (lambda unrestricted predicate : (family SM86Predicate) . (lambda unrestricted body : (family SM86InstructionBody) . (lambda unrestricted tail : (family SM86Program) .
82    (constructor SM86Program SM86ProgramNext (sm86PredicatedInstruction predicate body) tail))))
83
84def irsUnless =
85  (lambda unrestricted predicate : (family SM86Predicate) . (lambda unrestricted body : (family SM86InstructionBody) . (lambda unrestricted tail : (family SM86Program) .
86    (constructor SM86Program SM86ProgramNext (sm86NegatedPredicatedInstruction predicate body) tail))))
87
88def irsBranch =
89  (lambda unrestricted offset : Nat .
90    (constructor SM86InstructionBody SM86Branch (lsU offset) (sm86Unsigned32 (byte 255) (byte 255) (byte 131) (byte 3)) sm86BranchControl))
91
92-- ids[j] into R14, and P2 = (ids[j] != v)
93def irsCompare =
94  (lambda unrestricted tail : (family SM86Program) .
95    (lsNext (irsAddressOf irsIdScan irsJ indexedRowScatterSM86ClassIdPointerNatural)
96    (lsNext (lsLoad irsU irsIdScan 0)
97    (lsNext (irsExclusiveOr irsDifference irsU irsV lsControlWaitSB1)
98    (lsNext (llGreater irsP2 irsDifference 0)
99      tail)))))
100
101-- the scan of rows 0 .. t: exit at a row before t that names v (the flag is
102-- j ^ t, nonzero, when the row is equal and earlier)
103def irsEarlierBody =
104  (lambda unrestricted tail : (family SM86Program) .
105    (irsCompare
106    (lsNext (llMove irsEarlierFlag 0)
107    (irsUnless irsP2 (irsExclusiveOr irsEarlierFlag irsJ irsT lsControl)
108    (lsNext (llGreater irsP3 irsEarlierFlag 0)
109    (irsWhen irsP3 (constructor SM86InstructionBody SM86Exit lsControl)
110    (lsNext (llAddImmediate irsJ 1)
111    (lsNext (irsExclusiveOr irsDifference irsJ irsPastT lsControl)
112    (lsNext (llGreater irsP1 irsDifference 0)
113      tail)))))))))
114
115-- again while j != t + 1: @P1 BRA back over the body and the branch
116def irsEarlier =
117  (lambda unrestricted tail : (family SM86Program) .
118    (irsEarlierBody
119      (irsWhen irsP1 (irsBranch (naturalSaturatingSubtract 4294967296 (naturalMultiply 16 (succ (llLength irsEarlierBody)))))
120        tail)))
121
122-- the accumulation over rows t .. rows - 1
123def irsAccumulate =
124  (lambda unrestricted width : Nat . (lambda unrestricted rows : Nat .
125    (llLoop irsP1 (lambda unrestricted tail : (family SM86Program) .
126      (irsCompare
127      (irsUnless irsP2 (irsElement irsRow irsJ width)
128      (irsUnless irsP2 (irsAddressOf irsSourceAddress irsRow indexedRowScatterSM86SourcePointerNatural)
129      (irsUnless irsP2 (lsLoad irsSource irsSourceAddress 0)
130      (irsUnless irsP2 (constructor SM86InstructionBody SM86FloatAdd (lsR irsAccumulator) (lsR irsAccumulator) (lsR irsSource) lsControlWaitSB1)
131      (lsNext (llAddImmediate irsJ 1)
132      (lsNext (llGreater irsP1 irsJ (naturalSaturatingSubtract rows 1))
133        tail)))))))))))
134
135-- The program for `width` components and `rows` source rows.
136def indexedRowScatterOrderedSM86ProgramRaw =
137  (lambda unrestricted width : Nat . (lambda unrestricted rows : Nat .
138    (lsNext (constructor SM86InstructionBody SM86SpecialToRegister (lsR irsT) (constructor SM86SpecialRegister SM86CooperativeThreadArrayIdX) lsControlSetSB0)
139    (lsNext (constructor SM86InstructionBody SM86SpecialToRegister (lsR irsC) (constructor SM86SpecialRegister SM86ThreadIdX) lsControlSetSB0)
140    (lsNext (llMove irsFour 4)
141    (lsNext (llMove irsJ 0)
142    (lsNext (constructor SM86InstructionBody SM86IntegerMultiplyAddWideConstant (lsR irsIdAddress) (lsR irsT) (lsR irsFour) (byte 0) (lsU indexedRowScatterSM86ClassIdPointerNatural) lsControlWaitSB0)
143    (lsNext (lsLoad irsV irsIdAddress 0)
144    (lsNext (constructor SM86InstructionBody SM86IntegerAddThreeImmediate (lsR irsPastT) (lsR irsT) (lsU 1) lsControl)
145    (irsEarlier
146    (lsNext (constructor SM86InstructionBody SM86IntegerAddThreeImmediate (lsR irsJ) (lsR irsT) (lsU 0) lsControl)
147    (lsNext (irsElement irsRow irsV width)
148    (lsNext (irsAddressOf irsDestinationAddress irsRow indexedRowScatterSM86DestinationPointerNatural)
149    (lsNext (lsLoad irsAccumulator irsDestinationAddress 0)
150    (irsAccumulate width rows
151    (lsNext (constructor SM86InstructionBody SM86StoreGlobal (lsR irsDestinationAddress) (lsR irsAccumulator) (lsU 0) lsControlWaitSB1)
152    (lsNext (constructor SM86InstructionBody SM86Exit lsControl)
153      (constructor SM86Program SM86ProgramEnd))))))))))))))))))
154
155-- The scan reuses address registers across global loads. An LDG write
156-- barrier retires its result but does not guarantee that its address read
157-- has finished; serialize those late reads before reusing an address.
158-- This only changes controls, so the loop's relative branch lengths stay
159-- valid. See the read-barrier ownership check in SM86.Scoreboard.
160def indexedRowScatterOrderedSM86Program =
161  (lambda unrestricted width : Nat . (lambda unrestricted rows : Nat .
162    (sm86GuardLateReads (indexedRowScatterOrderedSM86ProgramRaw width rows))))

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.