Source/Packages

Platform.Linux.Nvidia.PlanHostAdamW

packages/hardware/platforms/linux-nvidia/src/Platform/Linux/Nvidia/PlanHostAdamW.alpha

336 lines52 declarations19.0 KiBSHA-256 1034d4a38685

Complete file · line 272

PlanHostAdamW.alpha

Definition view
1module Platform.Linux.Nvidia.PlanHostAdamW
2
3import Checkpoint.Envelope
4import Data.Bytes
5import Data.Float32Bits
6import Model.Config
7import Model.Parameter
8import Model.Word64
9import Platform.Linux.Nvidia.PlanHost
10import Platform.Linux.Nvidia.PlanHostRequest
11import Runtime.NativePhysicalProgram
12import Std.Float
13import Std.List
14import Std.Natural
15
16-- AdamW's step-dependent scalars, supplied by the host at run time.  The
17-- bias-corrected step size lr sqrt(1 - beta2^t) / (1 - beta1^t) and the
18-- scaled epsilon eps sqrt(1 - beta2^t) depend on the update's number t,
19-- which counts on across invocations: the checkpoint's header carries the
20-- updates taken (Checkpoint.Envelope), so an invocation's k-th update is
21-- t = that count + k + 1.  After the recipes are expanded, and before
22-- anything is submitted, the host computes both for each of its updates
23-- and writes them into the parameter word of every AdamW launch of that
24-- update (the realization reports where it placed each launch's parameter
25-- block: NvidiaWholeProgramParameterAddresses).  A fresh start reads a
26-- zero count, so it takes t = 1, 2, ...
27--
28-- The arithmetic is the learner's: binary64 throughout, rounded once to
29-- binary32 at the end, or binary32 at every step (the binary64 operation
30-- rounded to binary32, which is exactly the binary32 operation).  The
31-- powers are running products, taken first over the updates already done.
32
33family NvidiaPlanHostArithmetic : Type 0
34constructor NvidiaPlanHostBinary64
35constructor NvidiaPlanHostBinary32
36end-family
37
38-- a hyperparameter: a ratio of naturals (the nearest value of the
39-- arithmetic's precision), or a binary32's bits
40family NvidiaPlanHostConstant : Type 0
41constructor NvidiaPlanHostRatio
42field unrestricted nvidiaPlanHostRatioNumerator : Nat
43field unrestricted nvidiaPlanHostRatioDenominator : Nat
44constructor NvidiaPlanHostBinary32Bits
45field unrestricted nvidiaPlanHostBinary32BitsValue : Nat
46end-family
47
48family NvidiaPlanHostAdamW : Type 0
49constructor NvidiaPlanHostAdamWValue
50field unrestricted nvidiaPlanHostAdamWArithmetic : (family NvidiaPlanHostArithmetic)
51field unrestricted nvidiaPlanHostAdamWLearningRate : (family NvidiaPlanHostConstant)
52field unrestricted nvidiaPlanHostAdamWBeta1 : (family NvidiaPlanHostConstant)
53field unrestricted nvidiaPlanHostAdamWBeta2 : (family NvidiaPlanHostConstant)
54field unrestricted nvidiaPlanHostAdamWEpsilon : (family NvidiaPlanHostConstant)
55end-family
56
57-- a binary32's bits, for NvidiaPlanHostBinary32Bits
58def nvidiaPlanHostBinary32Of =
59  (lambda unrestricted value : F32 . (nvidiaPlanHostLittleNatural 4 (dataFloat32EncodeLE value)))
60
61-- ---- where the words go ----
62-- entry `ordinal` of a realization's per-launch table of 8-byte words: a
63-- launch's parameter block's device address
64-- (NvidiaWholeProgramParameterAddresses), or a word in it
65-- (NvidiaWholeProgramParameterWords)
66def nvidiaPlanHostParameterBlock =
67  (lambda unrestricted addresses : Bytes .
68    (lambda unrestricted ordinal : Nat .
69      (nvidiaPlanHostLittleNatural 8 (dataBytesDropValidated (naturalMultiply 8 ordinal) addresses))))
70
71-- the entries at `ordinals` (ascending) of such a table, in one walk
72-- (the table is dropped through once, not once per entry)
73def nvidiaPlanHostTableEntries =
74  (lambda unrestricted table : Bytes .
75    (lambda unrestricted ordinals : (family StdList Nat) .
76      (app
77        (app
78          (eliminate StdList
79            (lambda unrestricted current : (family StdList Nat) .
80              (pi unrestricted at : Nat . (pi unrestricted rest : Bytes . (family StdList Nat))))
81            ordinals
82            (branch StdListEmpty .
83              (lambda unrestricted at : Nat . (lambda unrestricted rest : Bytes . (constructor StdList StdListEmpty Nat))))
84            (branch StdListCons ordinal tail induction .
85              (lambda unrestricted at : Nat .
86                (lambda unrestricted rest : Bytes .
87                  (let unrestricted here = (dataBytesDropValidated (naturalMultiply 8 (naturalSaturatingSubtract ordinal at)) rest) in
88                  (constructor StdList StdListCons Nat (nvidiaPlanHostLittleNatural 8 here) (induction ordinal here)))))))
89          0)
90        table)))
91
92-- every update's launches, in order (the updates' launches ascend), and
93-- the stretch of such a list (or of its entries) that is update k's
94def adAllLaunches =
95  (lambda unrestricted updates : Nat .
96    (lambda unrestricted launches : (pi unrestricted k : Nat . (family StdList Nat)) .
97      (nat-eliminate
98        (lambda unrestricted current : Nat . (family StdList Nat))
99        (constructor StdList StdListEmpty Nat)
100        (lambda unrestricted k : Nat .
101          (lambda unrestricted earlier : (family StdList Nat) . (stdListAppend Nat earlier (launches k))))
102        updates)))
103
104def adUpdateStretch =
105  (lambda unrestricted launches : (pi unrestricted k : Nat . (family StdList Nat)) .
106    (lambda unrestricted all : (family StdList Nat) .
107      (lambda unrestricted k : Nat .
108        (stdListTake Nat (stdListLength Nat (launches k))
109          (stdListDrop Nat (stdListLength Nat (adAllLaunches k launches)) all)))))
110
111-- whether a buffer maps a device address
112def nvidiaPlanHostBufferMaps =
113  (lambda unrestricted buffer : (family NvidiaPlanHostBuffer) .
114    (lambda unrestricted device : Nat .
115      (naturalAnd (naturalLessOrEqual (phBufferGPU buffer) device)
116        (naturalLess device (naturalAdd (phBufferGPU buffer) (phBufferExtent buffer))))))
117
118def adMappedHost =
119  (lambda unrestricted buffer : (family NvidiaPlanHostBuffer) .
120    (lambda unrestricted device : Nat .
121      (naturalSelect (nvidiaPlanHostBufferMaps buffer device)
122        (naturalAdd (phBufferHost buffer) (naturalSaturatingSubtract device (phBufferGPU buffer)))
123        0)))
124
125-- a device address's host address in a host-mapped data buffer (else 0)
126def adMappedDataHost =
127  (lambda unrestricted buffers : (family NvidiaPlanHostBuffers) .
128    (lambda unrestricted device : Nat .
129      (eliminate NvidiaPlanHostBuffers (lambda unrestricted current : (family NvidiaPlanHostBuffers) . Nat) buffers
130        (branch NvidiaPlanHostBuffersEnd . 0)
131        (branch NvidiaPlanHostBuffersNext head rest induction .
132          (naturalAdd (naturalSelect (naturalNonzero (phBufferHost head)) (adMappedHost head device) 0) induction)))))
133
134-- Parameter blocks live in the program buffer or past the QMD table -- in
135-- the QMD buffer, or in a mapped data buffer after it when the table and
136-- its spill outgrow it (Coppelius's QMD overflow, under the GB10's 512-byte
137-- QMD records): a device address's host address there (else 0).
138def nvidiaPlanHostDeviceToHost =
139  (lambda unrestricted layout : (family NvidiaPlanHostLayout) .
140    (lambda unrestricted device : Nat .
141      (eliminate NvidiaPlanHostLayout (lambda unrestricted current : (family NvidiaPlanHostLayout) . Nat) layout
142        (branch NvidiaPlanHostLayoutValue gpfifo push sem program qmd userd data abi lifecycle expectedName errorHost .
143          (naturalAdd (adMappedHost program device)
144            (naturalAdd (adMappedHost qmd device) (adMappedDataHost data device)))))))
145
146-- update k's sites, from its launches (ordinals): the host address of the
147-- word at `offset` in each launch's parameter block
148def nvidiaPlanHostAdamWHostSites =
149  (lambda unrestricted layout : (family NvidiaPlanHostLayout) .
150    (lambda unrestricted addresses : Bytes .
151      (lambda unrestricted offset : Nat .
152        (lambda unrestricted updates : Nat .
153          (lambda unrestricted launches : (pi unrestricted k : Nat . (family StdList Nat)) .
154            (let unrestricted sites =
155              (stdListMap Nat Nat
156                (lambda unrestricted block : Nat . (nvidiaPlanHostDeviceToHost layout (naturalAdd block offset)))
157                (nvidiaPlanHostTableEntries addresses (adAllLaunches updates launches))) in
158            (adUpdateStretch launches sites)))))))
159
160-- What admits the schedule (1): every update has launches (ascending); each one's
161-- site is mapped; and each holds, as realized, the word `baked k` the plan
162-- put there for its update (the first invocation's; `words` is the
163-- realization's NvidiaWholeProgramParameterWords at the site's offset) --
164-- so a site is the word the launch's AdamW reads, not a neighbour.
165def nvidiaPlanHostAdamWSitesAdmitted =
166  (lambda unrestricted layout : (family NvidiaPlanHostLayout) .
167    (lambda unrestricted addresses : Bytes .
168      (lambda unrestricted words : Bytes .
169        (lambda unrestricted offset : Nat .
170          (lambda unrestricted updates : Nat .
171            (lambda unrestricted launches : (pi unrestricted k : Nat . (family StdList Nat)) .
172              (lambda unrestricted baked : (pi unrestricted k : Nat . Nat) .
173                (let unrestricted all = (adAllLaunches updates launches) in
174                (let unrestricted blocks = (nvidiaPlanHostTableEntries addresses all) in
175                (let unrestricted realized = (nvidiaPlanHostTableEntries words all) in
176                (nat-eliminate
177                  (lambda unrestricted current : Nat . Nat)
178                  1
179                  (lambda unrestricted k : Nat .
180                    (lambda unrestricted induction : Nat .
181                      (naturalAnd induction
182                        (naturalAnd (naturalNonzero (stdListLength Nat (launches k)))
183                          (naturalAnd
184                            (stdListFold Nat Nat
185                              (lambda unrestricted block : Nat .
186                                (lambda unrestricted rest : Nat .
187                                  (naturalAnd (naturalNonzero (nvidiaPlanHostDeviceToHost layout (naturalAdd block offset))) rest)))
188                              1
189                              (adUpdateStretch launches blocks k))
190                            (stdListFold Nat Nat
191                              (lambda unrestricted word : Nat .
192                                (lambda unrestricted rest : Nat . (naturalAnd (naturalEqual word (baked k)) rest)))
193                              1
194                              (adUpdateStretch launches realized k)))))))
195                  updates)))))))))))
196
197-- ---- the computation ----
198-- the scratch words (PlanHostRequest.prLearnerScratch)
199def adWord = (lambda unrestricted index : Nat . (naturalAdd prLearnerScratch (naturalMultiply 8 index)))
200def adOne : Nat = (adWord 0)
201def adBeta1 : Nat = (adWord 1)
202def adBeta2 : Nat = (adWord 2)
203def adRate : Nat = (adWord 3)
204def adEpsilon : Nat = (adWord 4)
205def adPower1 : Nat = (adWord 5)
206def adPower2 : Nat = (adWord 6)
207def adGap1 : Nat = (adWord 7)
208def adRoot : Nat = (adWord 8)
209def adStep : Nat = (adWord 9)
210def adScaled : Nat = (adWord 10)
211def adDenominator : Nat = (adWord 11)
212-- the parameter word: the step size's binary32, then the epsilon's (the
213-- second store's high half lands in the word after)
214def adOut : Nat = (adWord 12)
215
216def adOp =
217  (lambda unrestricted kind : (family NativePhysicalFloat64Operation) .
218    (lambda unrestricted destination : Nat .
219      (lambda unrestricted left : (family NativePhysicalOperand) .
220        (lambda unrestricted right : (family NativePhysicalOperand) .
221          (lambda unrestricted tail : (family NativePhysicalCommands) .
222            (phNext (constructor NativePhysicalOperation NativePhysicalFloat64 kind (phState destination) left right) tail))))))
223
224
225-- the arithmetic's rounding of a word just computed
226def adRound =
227  (lambda unrestricted arithmetic : (family NvidiaPlanHostArithmetic) .
228    (lambda unrestricted at : Nat .
229      (lambda unrestricted tail : (family NativePhysicalCommands) .
230        (eliminate NvidiaPlanHostArithmetic (lambda unrestricted current : (family NvidiaPlanHostArithmetic) . (family NativePhysicalCommands)) arithmetic
231          (branch NvidiaPlanHostBinary64 . tail)
232          (branch NvidiaPlanHostBinary32 .
233            (adOp (constructor NativePhysicalFloat64Operation NativePhysicalFloat64ToBinary32) at (phLoad at) (phImm 0)
234              (adOp (constructor NativePhysicalFloat64Operation NativePhysicalFloat64FromBinary32) at (phLoad at) (phImm 0) tail)))))))
235
236def adArith =
237  (lambda unrestricted arithmetic : (family NvidiaPlanHostArithmetic) .
238    (lambda unrestricted kind : (family NativePhysicalFloat64Operation) .
239      (lambda unrestricted destination : Nat .
240        (lambda unrestricted left : Nat .
241          (lambda unrestricted right : Nat .
242            (lambda unrestricted tail : (family NativePhysicalCommands) .
243              (adOp kind destination (phLoad left) (phLoad right) (adRound arithmetic destination tail))))))))
244
245def adConstant =
246  (lambda unrestricted arithmetic : (family NvidiaPlanHostArithmetic) .
247    (lambda unrestricted constant : (family NvidiaPlanHostConstant) .
248      (lambda unrestricted destination : Nat .
249        (lambda unrestricted tail : (family NativePhysicalCommands) .
250          (eliminate NvidiaPlanHostConstant (lambda unrestricted current : (family NvidiaPlanHostConstant) . (family NativePhysicalCommands)) constant
251            (branch NvidiaPlanHostRatio numerator denominator .
252              (adOp (constructor NativePhysicalFloat64Operation NativePhysicalFloat64FromNatural) destination (phImm numerator) (phImm 0)
253                (adOp (constructor NativePhysicalFloat64Operation NativePhysicalFloat64FromNatural) adDenominator (phImm denominator) (phImm 0)
254                  (adArith arithmetic (constructor NativePhysicalFloat64Operation NativePhysicalFloat64Divide) destination destination adDenominator tail))))
255            (branch NvidiaPlanHostBinary32Bits bits .
256              (phNext (constructor NativePhysicalOperation NativePhysicalStoreWord64 (phState destination) (phImm bits))
257                (adOp (constructor NativePhysicalFloat64Operation NativePhysicalFloat64FromBinary32) destination (phLoad destination) (phImm 0) tail))))))))
258
259def adMultiply = (constructor NativePhysicalFloat64Operation NativePhysicalFloat64Multiply)
260def adSubtract = (constructor NativePhysicalFloat64Operation NativePhysicalFloat64Subtract)
261def adDivide = (constructor NativePhysicalFloat64Operation NativePhysicalFloat64Divide)
262
263-- one more update's powers
264def adAdvance =
265  (lambda unrestricted arithmetic : (family NvidiaPlanHostArithmetic) .
266    (lambda unrestricted tail : (family NativePhysicalCommands) .
267      (adArith arithmetic adMultiply adPower1 adPower1 adBeta1
268        (adArith arithmetic adMultiply adPower2 adPower2 adBeta2 tail))))
269
270-- update k: its powers, its two scalars, their parameter word into every
271-- site of the update
272def adUpdate =
273  (lambda unrestricted arithmetic : (family NvidiaPlanHostArithmetic) .
274    (lambda unrestricted sites : (family StdList Nat) .
275      (lambda unrestricted tail : (family NativePhysicalCommands) .
276        (adAdvance arithmetic
277          (adArith arithmetic adSubtract adRoot adOne adPower2
278            (adOp (constructor NativePhysicalFloat64Operation NativePhysicalFloat64SquareRoot) adRoot (phImm 0) (phLoad adRoot)
279              (adRound arithmetic adRoot
280                (adArith arithmetic adSubtract adGap1 adOne adPower1
281                  (adArith arithmetic adMultiply adStep adRate adRoot
282                    (adArith arithmetic adDivide adStep adStep adGap1
283                      (adArith arithmetic adMultiply adScaled adEpsilon adRoot
284                        (adOp (constructor NativePhysicalFloat64Operation NativePhysicalFloat64ToBinary32) adOut (phLoad adStep) (phImm 0)
285                          (adOp (constructor NativePhysicalFloat64Operation NativePhysicalFloat64ToBinary32) (naturalAdd adOut 4) (phLoad adScaled) (phImm 0)
286                            (stdListFold Nat (family NativePhysicalCommands)
287                              (lambda unrestricted site : Nat .
288                                (lambda unrestricted rest : (family NativePhysicalCommands) .
289                                  (phNext (constructor NativePhysicalOperation NativePhysicalStoreWord64 (phImm site) (phLoad adOut)) rest)))
290                              tail
291                              sites))))))))))))))
292
293-- The commands: the constants, the powers after the updates the
294-- checkpoint's header counts, then each of the invocation's `updates` in
295-- order, update k's word written to `sites k` (host addresses).
296def nvidiaPlanHostAdamWCommands =
297  (lambda unrestricted adamW : (family NvidiaPlanHostAdamW) .
298    (lambda unrestricted updates : Nat .
299      (lambda unrestricted sites : (pi unrestricted k : Nat . (family StdList Nat)) .
300        (lambda unrestricted tail : (family NativePhysicalCommands) .
301          (eliminate NvidiaPlanHostAdamW (lambda unrestricted current : (family NvidiaPlanHostAdamW) . (family NativePhysicalCommands)) adamW
302            (branch NvidiaPlanHostAdamWValue arithmetic rate beta1 beta2 epsilon .
303              (adOp (constructor NativePhysicalFloat64Operation NativePhysicalFloat64FromNatural) adOne (phImm 1) (phImm 0)
304                (adConstant arithmetic beta1 adBeta1
305                  (adConstant arithmetic beta2 adBeta2
306                    (adConstant arithmetic rate adRate
307                      (adConstant arithmetic epsilon adEpsilon
308                        (phNext (constructor NativePhysicalOperation NativePhysicalStoreWord64 (phState adPower1) (phLoad adOne))
309                          (phNext (constructor NativePhysicalOperation NativePhysicalStoreWord64 (phState adPower2) (phLoad adOne))
310                            (phNext (constructor NativePhysicalOperation NativePhysicalRepeatBeginCounted
311                                      (phLoad (naturalAdd prCheckpointHeaderIn checkpointEnvelopeUpdatesAt)))
312                              (adAdvance arithmetic
313                                (phNext (constructor NativePhysicalOperation NativePhysicalRepeatEnd)
314                                  (app
315                                    (nat-eliminate
316                                      (lambda unrestricted current : Nat .
317                                        (pi unrestricted after : (family NativePhysicalCommands) . (family NativePhysicalCommands)))
318                                      (lambda unrestricted after : (family NativePhysicalCommands) . after)
319                                      (lambda unrestricted k : Nat .
320                                        (lambda unrestricted induction : (pi unrestricted after : (family NativePhysicalCommands) . (family NativePhysicalCommands)) .
321                                          (lambda unrestricted after : (family NativePhysicalCommands) .
322                                            (induction (adUpdate arithmetic (sites k) after)))))
323                                      updates)
324                                    tail)))))))))))))))))
325
326-- One more update's commands, for a run loop's body
327-- (PlanHostRequest.NvidiaPlanHostStepCommands): the powers advanced, the
328-- two scalars computed, their word written into every one of `sites`.  The
329-- constants and the powers the checkpoint counts come first, from
330-- `nvidiaPlanHostAdamWCommands` with no updates of its own.
331def nvidiaPlanHostAdamWStepCommands =
332  (lambda unrestricted adamW : (family NvidiaPlanHostAdamW) .
333    (lambda unrestricted sites : (family StdList Nat) .
334      (eliminate NvidiaPlanHostAdamW (lambda unrestricted current : (family NvidiaPlanHostAdamW) . (family NativePhysicalCommands)) adamW
335        (branch NvidiaPlanHostAdamWValue arithmetic rate beta1 beta2 epsilon .
336          (adUpdate arithmetic sites (constructor NativePhysicalCommands NativePhysicalCommandsEnd))))))

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.