Source/Packages

Realization.Nvidia.SM86.LinearStepLoopSM86

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

178 lines33 declarations8.0 KiBSHA-256 059e5870f8c0

Complete file · line 39

LinearStepLoopSM86.alpha

Definition view
1module Realization.Nvidia.SM86.LinearStepLoopSM86
2
3import Accelerator.SM86.Control
4import Accelerator.SM86.Immediate
5import Accelerator.SM86.Instruction
6import Accelerator.SM86.Types
7import Realization.Nvidia.SM86.LinearStepSM86
8import Std.Natural
9
10-- The linear step with LOOPS: the same computation as
11-- Realization.Nvidia.SM86.LinearStepSM86 -- the same operations on the same
12-- operands in the same order, so the record is the specification's word for
13-- word -- in a program of constant length and constant register demand.
14-- Where the unrolled program keeps W, x, t, dW and W' in registers (848 of
15-- them for 16 x 16; the register file admits 38 x 1 .. 1 x 57), this one
16-- walks the outputs in an outer loop and the inputs in two inner loops,
17-- loading each element when it is used and storing each result when it is
18-- made, with 64-bit addresses from IMAD.WIDE (index x 4 + the parameter
19-- word's address).  The loops close with ISETP.GT (an unsigned compare --
20-- Checked.IntegerCompareProbe) and a backward BRA.
21--
22--   prologue (tid, the lane guard, the parameter words; LinearStepSM86)
23--   R11 = -eta   R12 = 0.5   R13 = loss = +0   R35 = 4   R14 = i = 0   R16 = e = 0
24--   outer:  R28 = y = +0   R15 = j = 0
25--     first:   W_e, x_j loaded; y = fma(W_e, x_j, y); e += 1; j += 1; again while j <= k-1
26--     y_i stored; t_i loaded; d = y + (-t_i); loss = fma(d, d, loss); e -= k; j = 0
27--     second:  x_j, W_e loaded; dW = d x x_j; W' = fma(dW, -eta, W_e); both stored;
28--              e += 1; j += 1; again while j <= k-1
29--     i += 1; again while i <= m-1
30--   loss = 0.5 x loss, stored at the record's word m (i = m by then); exit
31--
32-- Each backward branch's offset is its loop's length in bytes, negated,
33-- counted from the loop's own block (llLoop).
34
35def llR11 : Nat = 11
36def llHalfRegister : Nat = 12
37def llLoss : Nat = 13
38def llI : Nat = 14
39def llJ : Nat = 15
40def llE : Nat = 16
41def llWAddress : Nat = 18
42def llXAddress : Nat = 20
43def llYAddress : Nat = 22
44def llTAddress : Nat = 24
45def llOutAddress : Nat = 26
46def llY : Nat = 28
47def llW : Nat = 29
48def llX : Nat = 30
49def llT : Nat = 31
50def llD : Nat = 32
51def llDW : Nat = 33
52def llWp : Nat = 34
53def llFour : Nat = 35
54
55def llP1 : (family SM86Predicate) = (constructor SM86Predicate SM86Predicate1)
56def llP2 : (family SM86Predicate) = (constructor SM86Predicate SM86Predicate2)
57
58def llMove =
59  (lambda unrestricted destination : Nat . (lambda unrestricted value : Nat .
60    (constructor SM86InstructionBody SM86MoveImmediate (lsR destination) (lsU value) lsControl)))
61
62def llAddImmediate =
63  (lambda unrestricted destination : Nat . (lambda unrestricted value : Nat .
64    (constructor SM86InstructionBody SM86IntegerAddThreeImmediate (lsR destination) (lsR destination) (lsU value) lsControl)))
65
66-- the pair at `destination` = index x 4 + the address in parameter word `parameter`
67def llAddressOf =
68  (lambda unrestricted destination : Nat . (lambda unrestricted index : Nat . (lambda unrestricted parameter : Nat .
69    (constructor SM86InstructionBody SM86IntegerMultiplyAddWideConstant (lsR destination) (lsR index) (lsR llFour) (byte 0) (lsU parameter) lsControl))))
70
71def llGreater =
72  (lambda unrestricted predicate : (family SM86Predicate) . (lambda unrestricted source : Nat . (lambda unrestricted bound : Nat .
73    (constructor SM86InstructionBody SM86PredicateGreaterThanImmediate predicate (lsR source) (lsU bound) lsControl))))
74
75-- @!P BRA back `instructions` (counted from the branch's successor)
76def llBackUnless =
77  (lambda unrestricted predicate : (family SM86Predicate) . (lambda unrestricted instructions : Nat . (lambda unrestricted tail : (family SM86Program) .
78    (constructor SM86Program SM86ProgramNext
79      (sm86NegatedPredicatedInstruction predicate
80        (constructor SM86InstructionBody SM86Branch
81          (lsU (naturalSaturatingSubtract 4294967296 (naturalMultiply 16 instructions)))
82          (sm86Unsigned32 (byte 255) (byte 255) (byte 131) (byte 3))
83          sm86BranchControl))
84      tail))))
85
86-- the number of instructions a block prepends
87def llLength =
88  (lambda unrestricted block : (pi unrestricted tail : (family SM86Program) . (family SM86Program)) .
89    (eliminate SM86Program (lambda unrestricted current : (family SM86Program) . Nat)
90      (block (constructor SM86Program SM86ProgramEnd))
91      (branch SM86ProgramEnd . zero)
92      (branch SM86ProgramNext head rest induction . (succ induction))))
93
94-- A loop: the block, then @!P BRA back to its first instruction -- the
95-- distance is the block's own length and the branch, so no count is written
96-- by hand.
97def llLoop =
98  (lambda unrestricted predicate : (family SM86Predicate) .
99    (lambda unrestricted block : (pi unrestricted tail : (family SM86Program) . (family SM86Program)) .
100      (lambda unrestricted tail : (family SM86Program) .
101        (block (llBackUnless predicate (succ (llLength block)) tail)))))
102
103-- the first inner loop: W_e, x_j loaded; y = fma(W_e, x_j, y); e, j advanced
104def llFirstLoop =
105  (lambda unrestricted k : Nat .
106    (llLoop llP1 (lambda unrestricted tail : (family SM86Program) .
107    (lsNext (llAddressOf llWAddress llE 352)
108    (lsNext (lsLoad llW llWAddress 0)
109    (lsNext (llAddressOf llXAddress llJ 360)
110    (lsNext (lsLoad llX llXAddress 0)
111    (lsNext (lsFusedWith lsControlWaitSB1 llY llW llX llY)
112    (lsNext (llAddImmediate llE 1)
113    (lsNext (llAddImmediate llJ 1)
114    (lsNext (llGreater llP1 llJ (naturalSaturatingSubtract k 1))
115      tail)))))))))))
116
117-- the second inner loop: x_j, W_e loaded; dW = d x x_j; W' = fma(dW, -eta, W_e);
118-- both stored; e, j advanced
119def llSecondLoop =
120  (lambda unrestricted m : Nat . (lambda unrestricted k : Nat .
121    (llLoop llP1 (lambda unrestricted tail : (family SM86Program) .
122    (lsNext (llAddressOf llXAddress llJ 360)
123    (lsNext (lsLoad llX llXAddress 0)
124    (lsNext (llAddressOf llWAddress llE 352)
125    (lsNext (lsLoad llW llWAddress 0)
126    (lsNext (constructor SM86InstructionBody SM86FloatMultiply (lsR llDW) (lsR llD) (lsR llX) lsControlWaitSB1)
127    (lsNext (lsFused llWp llDW llR11 llW)
128    (lsNext (llAddressOf llOutAddress llE 376)
129    (lsNext (lsStore llOutAddress llDW (lsGradientOffset m k 0 0))
130    (lsNext (lsStore llOutAddress llWp (lsUpdatedOffset m k 0 0))
131    (lsNext (llAddImmediate llE 1)
132    (lsNext (llAddImmediate llJ 1)
133    (lsNext (llGreater llP1 llJ (naturalSaturatingSubtract k 1))
134      tail))))))))))))))))
135
136-- the outer loop: y = +0, j = 0, the first inner loop, y_i stored, d and the
137-- loss, e back to the row's first element, j = 0, the second inner loop, i
138-- advanced
139def llOuterLoop =
140  (lambda unrestricted m : Nat . (lambda unrestricted k : Nat .
141    (llLoop llP2 (lambda unrestricted tail : (family SM86Program) .
142    (lsNext (lsZero llY)
143    (lsNext (llMove llJ 0)
144    (llFirstLoop k
145      (lsNext (llAddressOf llYAddress llI 376)
146      (lsNext (lsStore llYAddress llY 0)
147      (lsNext (llAddressOf llTAddress llI 368)
148      (lsNext (lsLoad llT llTAddress 0)
149      (lsNext (constructor SM86InstructionBody SM86FloatNegate (lsR llD) (lsR llT) lsControlWaitSB1)
150      (lsNext (lsAdd llD llY llD)
151      (lsNext (lsFused llLoss llD llD llLoss)
152      (lsNext (llAddImmediate llE (naturalSaturatingSubtract 4294967296 k))
153      (lsNext (llMove llJ 0)
154      (llSecondLoop m k
155        (lsNext (llAddImmediate llI 1)
156        (lsNext (llGreater llP2 llI (naturalSaturatingSubtract m 1))
157          tail)))))))))))))))))))
158
159def linearStepLoopSM86ProgramFor =
160  (lambda unrestricted m : Nat . (lambda unrestricted k : Nat .
161    (lsPrologue k m
162      (lsNext (lsNegate llR11 10)
163      (lsNext (llMove llHalfRegister 1056964608)
164      (lsNext (lsZero llLoss)
165      (lsNext (llMove llFour 4)
166      (lsNext (llMove llI 0)
167      (lsNext (llMove llE 0)
168      (llOuterLoop m k
169        (lsNext (lsMultiply llLoss llHalfRegister llLoss)
170        (lsNext (llAddressOf llYAddress llI 376)
171        (lsNext (lsStore llYAddress llLoss 0)
172          lsExit)))))))))))))
173
174-- instructions run: 18 before the outer loop; per output 2 + 9k + 9 + 13k + 3;
175-- 4 after; with a margin
176def linearStepLoopSM86FuelFor =
177  (lambda unrestricted m : Nat . (lambda unrestricted k : Nat .
178    (naturalAdd 32 (naturalMultiply m (naturalAdd 16 (naturalMultiply 22 k))))))

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.