Source/Packages

Realization.Nvidia.SM86.AdamWHalfSM86

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

205 lines38 declarations10.6 KiBSHA-256 557f99a5eb69

Complete file · line 52

AdamWHalfSM86.alpha

Definition view
1module Realization.Nvidia.SM86.AdamWHalfSM86
2
3import Accelerator.SM86.Control
4import Accelerator.SM86.Immediate
5import Accelerator.SM86.Instruction
6import Accelerator.SM86.NumericSemantics
7import Accelerator.SM86.Program
8import Accelerator.SM86.Types
9import Realization.Nvidia.SM86.StreamingAttentionSM86
10import Std.Natural
11
12-- AdamW over the whole flat bank with the parameters' half copies written
13-- as it goes: the half cast that followed the
14-- update, fused into it.  Each element's update is
15-- Realization.Nvidia.SM86.AdamWSM86's (adamWF32MathProgram), instruction
16-- for instruction, and each half pair is the cast's (F2FP, element 2i + 1
17-- high, in the realization's half format), so the parameters, the moments
18-- and the halves are bit-identical to the update followed by the cast:
19--
20--   g = g one;  m = m beta1 + g (1 - beta1);  v = v beta2 + g^2 (1 - beta2)
21--   p = p decay + -(m stepSize (1 / (sqrt v + epsilon)))
22--   half[i] = half(p[i])
23--
24-- The half bank is parallel to the parameter bank: element i's half at byte
25-- 2 i of it.  A thread takes the quad 4T .. 4T + 3 (T = block 256 + thread):
26-- each bank's four elements in one 128-bit load and store, the four halves
27-- in one 64-bit store -- the bank is memory-bound, and wide accesses keep
28-- more of it in flight a thread -- then T + stride, ... `iterations` times;
29-- quads at or past `quads` are untouched.  The gradient is read, never
30-- written (every gradient the step reads was written whole by the backward
31-- before it).
32--
33-- Parameters (constant bank 0 from the SM86 parameter base): the pointers
34-- P, G, M, V (arguments 0-3, 16-byte aligned); the scalars one, beta1,
35-- 1 - beta1, beta2, 1 - beta2, decay, stepSize, epsilon (words from argument
36-- 4, where AdamWSM86's ABI puts them, so the host writes the step's scalars
37-- at the same offsets); the half bank's pointer (argument 9, 8-byte
38-- aligned).
39--
40-- Registers: R0 thread, R1 block, R2 T, R3 sixteen, R4 eight; R6:R7, R8:R9,
41-- R10:R11, R12:R13 the quad's P, G, M, V; R18..R25 the scalars; R29 the
42-- iterations left; R30:R31 the halves' address; R32..R35, R36..R39,
43-- R40..R43, R44..R47 the quad's p, g, m, v; element e's temporaries
44-- R48 + 3e .. R50 + 3e; R60, R61 the half pairs.  Loads under SB0..SB3 (P,
45-- G, M, V), the square roots and reciprocals under SB4; the stores' reads
46-- under SB5, which the next iteration's first load waits for.
47
48def ahThreads : Nat = 256
49def adamWHalfSM86ParameterArgument : Nat = 0
50def adamWHalfSM86GradientArgument : Nat = 1
51def adamWHalfSM86FirstMomentArgument : Nat = 2
52def adamWHalfSM86SecondMomentArgument : Nat = 3
53def adamWHalfSM86ScalarPairArgument =
54  (lambda unrestricted pair : Nat . (naturalAdd 4 pair))
55def adamWHalfSM86HalfArgument : Nat = 9
56def adamWHalfSM86ArgumentCount : Nat = 10
57def ahScalar = (lambda unrestricted j : Nat .
58  (naturalAdd (saArgument (adamWHalfSM86ScalarPairArgument 0))
59    (naturalMultiply 4 j)))
60def ahHalfPointer : Nat = (saArgument adamWHalfSM86HalfArgument)
61
62-- element e's p, g, m, v (bank b = 0..3) and temporaries (i = 0..2)
63def ahValue = (lambda unrestricted e : Nat . (lambda unrestricted b : Nat . (naturalAdd (naturalAdd 32 (naturalMultiply 4 b)) e)))
64def ahTemporary = (lambda unrestricted e : Nat . (lambda unrestricted i : Nat . (naturalAdd (naturalAdd 48 (naturalMultiply 3 e)) i)))
65def ahHalfPairs : Nat = 60
66def ahCounter : Nat = 29
67
68def ahIterations = (lambda unrestricted quads : Nat . (lambda unrestricted stride : Nat .
69  (naturalDivideUnchecked (naturalAdd quads (naturalSaturatingSubtract stride 1)) stride)))
70
71-- ---- the prologue: T, the constants, the scalars, the count ----
72def ahPrologue = (lambda unrestricted iterations : Nat . (lambda unrestricted tail : (family SM86Program) .
73  (saS2R 0 (constructor SM86SpecialRegister SM86ThreadIdX)
74  (saS2R 1 (constructor SM86SpecialRegister SM86CooperativeThreadArrayIdX)
75  (saMovImm 3 16
76  (saMovImm 4 8
77  (saMovImm ahCounter iterations
78  (saMovConst 18 (ahScalar 0)
79  (saMovConst 19 (ahScalar 1)
80  (saMovConst 20 (ahScalar 2)
81  (saMovConst 21 (ahScalar 3)
82  (saMovConst 22 (ahScalar 4)
83  (saMovConst 23 (ahScalar 5)
84  (saMovConst 24 (ahScalar 6)
85  (saMovConst 25 (ahScalar 7)
86  (saMovImmAfter 5 0 saWait5
87  (saImad 2 1 ahThreads 0 tail)))))))))))))))))
88
89-- ---- the body: one quad ----
90def ahGuarded = (lambda unrestricted body : (family SM86InstructionBody) . (lambda unrestricted tail : (family SM86Program) .
91  (saUnless saP0 body tail)))
92-- a store whose register read SB5 counts
93def ahStoreControl : (family SM86Control) =
94  (constructor SM86Control SM86ControlValue (byte 15) (constructor SM86YieldMode SM86Continue)
95    saNone saSB5 (byte 0) (byte 0))
96
97-- bounds, the four pointers and the halves'
98def ahAddresses = (lambda unrestricted quads : Nat . (lambda unrestricted tail : (family SM86Program) .
99  (saGreater saP0 2 (naturalPredecessor quads)
100  (saWide 6 2 3 (saArgument adamWHalfSM86ParameterArgument)
101  (saWide 8 2 3 (saArgument adamWHalfSM86GradientArgument)
102  (saWide 10 2 3 (saArgument adamWHalfSM86FirstMomentArgument)
103  (saWide 12 2 3 (saArgument adamWHalfSM86SecondMomentArgument)
104  (saWide 30 2 4 ahHalfPointer tail))))))))
105
106-- bank b's quad (pointer 6 + 2b) into R32 + 4b, under SB b; the first waits
107-- for the last iteration's stores' reads
108def ahLoadBank = (lambda unrestricted b : Nat . (lambda unrestricted barrier : (family SM86Barrier) .
109  (lambda unrestricted wait : Nat . (lambda unrestricted tail : (family SM86Program) .
110    (ahGuarded (constructor SM86InstructionBody SM86LoadGlobalWide (saR (ahValue 0 b)) (saR (naturalAdd 6 (naturalMultiply 2 b))) (saU 0)
111      (saCtl barrier wait)) tail)))))
112def ahLoads = (lambda unrestricted tail : (family SM86Program) .
113  (ahLoadBank 0 saSB0 saWait5
114  (ahLoadBank 1 saSB1 saWaitNone
115  (ahLoadBank 2 saSB2 saWaitNone
116  (ahLoadBank 3 saSB3 saWaitNone tail)))))
117
118-- element e's update: adamWF32MathProgram's instructions on its registers
119def ahUpdate = (lambda unrestricted e : Nat . (lambda unrestricted tail : (family SM86Program) .
120  (let unrestricted p = (ahValue e 0) in (let unrestricted g = (ahValue e 1) in
121  (let unrestricted m = (ahValue e 2) in (let unrestricted v = (ahValue e 3) in
122  (let unrestricted t0 = (ahTemporary e 0) in (let unrestricted t1 = (ahTemporary e 1) in
123  (let unrestricted t2 = (ahTemporary e 2) in
124  (saFmul g g 18 (naturalSelect e saWaitNone (naturalAdd saWait0 (naturalAdd saWait1 (naturalAdd saWait2 saWait3))))
125  (saFmul t0 g 20 saWaitNone
126  (saFfma m m 19 t0
127  (saFmul t2 g g saWaitNone
128  (saFmul t2 t2 22 saWaitNone
129  (saFfma v v 21 t2
130  (saMufu t1 v (constructor SM86MultiFunction SM86SquareRoot) saWaitNone
131  (saFadd t1 t1 25 saWait4
132  (saMufu t1 t1 (constructor SM86MultiFunction SM86Reciprocal) saWaitNone
133  (saFmul t0 m 24 saWaitNone
134  (saFmul t0 t0 t1 saWait4
135  (saFneg t0 t0
136  (saFmul p p 23 saWaitNone
137  (saFadd p p t0 saWaitNone tail)))))))))))))))))))))))
138
139-- the quad's P, M, V (128 bits each) and its two half pairs (64 bits)
140def ahStores = (lambda unrestricted tail : (family SM86Program) .
141  (saPack ahHalfPairs (ahValue 1 0) (ahValue 0 0)
142  (saPack (succ ahHalfPairs) (ahValue 3 0) (ahValue 2 0)
143  (ahGuarded (constructor SM86InstructionBody SM86StoreGlobalWide (saR 6) (saR (ahValue 0 0)) (saU 0) ahStoreControl)
144  (ahGuarded (constructor SM86InstructionBody SM86StoreGlobalWide (saR 10) (saR (ahValue 0 2)) (saU 0) ahStoreControl)
145  (ahGuarded (constructor SM86InstructionBody SM86StoreGlobalWide (saR 12) (saR (ahValue 0 3)) (saU 0) ahStoreControl)
146  (ahGuarded (constructor SM86InstructionBody SM86StoreGlobal64 (saR 30) (saR ahHalfPairs) (saU 0) ahStoreControl) tail)))))))
147
148def ahBody = (lambda unrestricted quads : Nat .
149  (ahAddresses quads (ahLoads (ahUpdate 0 (ahUpdate 1 (ahUpdate 2 (ahUpdate 3 (ahStores sm86ProgramEmpty))))))))
150
151-- ---- the loop's control: the next quad, one iteration fewer ----
152def ahControl = (lambda unrestricted stride : Nat . (lambda unrestricted tail : (family SM86Program) .
153  (saAddImm 2 2 stride
154  (saAddImm ahCounter ahCounter 4294967295
155  (saGreater saP1 ahCounter 0 tail)))))
156def ahControlCount : Nat = 3
157-- back to the body's first instruction (the displacement counts the body,
158-- the control and the branch, from the branch's successor)
159def ahBranch = (lambda unrestricted body : Nat .
160  (saWhen saP1 (constructor SM86InstructionBody SM86Branch
161    (saU (naturalSaturatingSubtract 4294967296 (naturalMultiply 16 (naturalAdd body (naturalAdd ahControlCount 1)))))
162    (sm86Unsigned32 (byte 255) (byte 255) (byte 131) (byte 3))
163    sm86SafeControl)
164    sm86ProgramEmpty))
165
166def ahExit : (family SM86Program) =
167  (saOp (constructor SM86InstructionBody SM86Exit (saAfter saWaitAll)) sm86ProgramEmpty)
168
169-- `finish` is applied to the prologue and the body on their own (e.g. the
170-- stalls compacted and drained); the control keeps the largest stalls.
171-- looped 1: the loop with its back edge; looped 0: what the schedule checks
172-- read, the body twice and no branch (it covers both ways into the body,
173-- and every iteration leaves the same things in flight: the stores' reads
174-- under SB5, nothing else)
175def ahAssemble = (lambda unrestricted finish : (pi unrestricted piece : (family SM86Program) . (family SM86Program)) .
176  (lambda unrestricted looped : Nat .
177  (lambda unrestricted quads : Nat . (lambda unrestricted stride : Nat .
178    (let unrestricted body = (finish (ahBody quads)) in
179    (let unrestricted once = (sm86ProgramAppend body (ahControl stride sm86ProgramEmpty)) in
180    (sm86ProgramAppend (finish (ahPrologue (ahIterations quads stride) sm86ProgramEmpty))
181      (sm86ProgramAppend once
182        (nat-eliminate (lambda unrestricted n : Nat . (family SM86Program))
183          (sm86ProgramAppend once ahExit)
184          (lambda unrestricted predecessor : Nat . (lambda unrestricted induction : (family SM86Program) .
185            (sm86ProgramAppend (ahBranch (sm86ProgramCount body)) ahExit)))
186          looped)))))))))
187
188-- the image for `quads` element quads (the bank's elements over four) over
189-- a grid whose threads number `stride`
190def adamWHalfSM86 = (lambda unrestricted finish : (pi unrestricted piece : (family SM86Program) . (family SM86Program)) .
191  (ahAssemble finish 1))
192def adamWHalfSM86Checked = (lambda unrestricted finish : (pi unrestricted piece : (family SM86Program) . (family SM86Program)) .
193  (ahAssemble finish 0))
194
195def adamWHalfSM86Threads : Nat = ahThreads
196def adamWHalfSM86Registers : Nat = 64
197-- the elements a thread takes an iteration
198def adamWHalfSM86Elements : Nat = 4
199def adamWHalfSM86Iterations = ahIterations
200-- through the half pointer, in whole 8-byte parameter words
201def adamWHalfSM86ConstantBytes : Nat = (saArgument adamWHalfSM86ArgumentCount)
202-- where the step-dependent scalars sit in the parameter block (the host
203-- writes each update's there)
204def adamWHalfSM86StepSizeOffset : Nat = (ahScalar 6)
205def adamWHalfSM86EpsilonOffset : Nat = (ahScalar 7)

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.