module Platform.Linux.Nvidia.PlanHostAdamW import Checkpoint.Envelope import Data.Bytes import Data.Float32Bits import Model.Config import Model.Parameter import Model.Word64 import Platform.Linux.Nvidia.PlanHost import Platform.Linux.Nvidia.PlanHostRequest import Runtime.NativePhysicalProgram import Std.Float import Std.List import Std.Natural -- AdamW's step-dependent scalars, supplied by the host at run time. The -- bias-corrected step size lr sqrt(1 - beta2^t) / (1 - beta1^t) and the -- scaled epsilon eps sqrt(1 - beta2^t) depend on the update's number t, -- which counts on across invocations: the checkpoint's header carries the -- updates taken (Checkpoint.Envelope), so an invocation's k-th update is -- t = that count + k + 1. After the recipes are expanded, and before -- anything is submitted, the host computes both for each of its updates -- and writes them into the parameter word of every AdamW launch of that -- update (the realization reports where it placed each launch's parameter -- block: NvidiaWholeProgramParameterAddresses). A fresh start reads a -- zero count, so it takes t = 1, 2, ... -- -- The arithmetic is the learner's: binary64 throughout, rounded once to -- binary32 at the end, or binary32 at every step (the binary64 operation -- rounded to binary32, which is exactly the binary32 operation). The -- powers are running products, taken first over the updates already done. family NvidiaPlanHostArithmetic : Type 0 constructor NvidiaPlanHostBinary64 constructor NvidiaPlanHostBinary32 end-family -- a hyperparameter: a ratio of naturals (the nearest value of the -- arithmetic's precision), or a binary32's bits family NvidiaPlanHostConstant : Type 0 constructor NvidiaPlanHostRatio field unrestricted nvidiaPlanHostRatioNumerator : Nat field unrestricted nvidiaPlanHostRatioDenominator : Nat constructor NvidiaPlanHostBinary32Bits field unrestricted nvidiaPlanHostBinary32BitsValue : Nat end-family family NvidiaPlanHostAdamW : Type 0 constructor NvidiaPlanHostAdamWValue field unrestricted nvidiaPlanHostAdamWArithmetic : (family NvidiaPlanHostArithmetic) field unrestricted nvidiaPlanHostAdamWLearningRate : (family NvidiaPlanHostConstant) field unrestricted nvidiaPlanHostAdamWBeta1 : (family NvidiaPlanHostConstant) field unrestricted nvidiaPlanHostAdamWBeta2 : (family NvidiaPlanHostConstant) field unrestricted nvidiaPlanHostAdamWEpsilon : (family NvidiaPlanHostConstant) end-family -- a binary32's bits, for NvidiaPlanHostBinary32Bits def nvidiaPlanHostBinary32Of = (lambda unrestricted value : F32 . (nvidiaPlanHostLittleNatural 4 (dataFloat32EncodeLE value))) -- ---- where the words go ---- -- entry `ordinal` of a realization's per-launch table of 8-byte words: a -- launch's parameter block's device address -- (NvidiaWholeProgramParameterAddresses), or a word in it -- (NvidiaWholeProgramParameterWords) def nvidiaPlanHostParameterBlock = (lambda unrestricted addresses : Bytes . (lambda unrestricted ordinal : Nat . (nvidiaPlanHostLittleNatural 8 (dataBytesDropValidated (naturalMultiply 8 ordinal) addresses)))) -- the entries at `ordinals` (ascending) of such a table, in one walk -- (the table is dropped through once, not once per entry) def nvidiaPlanHostTableEntries = (lambda unrestricted table : Bytes . (lambda unrestricted ordinals : (family StdList Nat) . (app (app (eliminate StdList (lambda unrestricted current : (family StdList Nat) . (pi unrestricted at : Nat . (pi unrestricted rest : Bytes . (family StdList Nat)))) ordinals (branch StdListEmpty . (lambda unrestricted at : Nat . (lambda unrestricted rest : Bytes . (constructor StdList StdListEmpty Nat)))) (branch StdListCons ordinal tail induction . (lambda unrestricted at : Nat . (lambda unrestricted rest : Bytes . (let unrestricted here = (dataBytesDropValidated (naturalMultiply 8 (naturalSaturatingSubtract ordinal at)) rest) in (constructor StdList StdListCons Nat (nvidiaPlanHostLittleNatural 8 here) (induction ordinal here))))))) 0) table))) -- every update's launches, in order (the updates' launches ascend), and -- the stretch of such a list (or of its entries) that is update k's def adAllLaunches = (lambda unrestricted updates : Nat . (lambda unrestricted launches : (pi unrestricted k : Nat . (family StdList Nat)) . (nat-eliminate (lambda unrestricted current : Nat . (family StdList Nat)) (constructor StdList StdListEmpty Nat) (lambda unrestricted k : Nat . (lambda unrestricted earlier : (family StdList Nat) . (stdListAppend Nat earlier (launches k)))) updates))) def adUpdateStretch = (lambda unrestricted launches : (pi unrestricted k : Nat . (family StdList Nat)) . (lambda unrestricted all : (family StdList Nat) . (lambda unrestricted k : Nat . (stdListTake Nat (stdListLength Nat (launches k)) (stdListDrop Nat (stdListLength Nat (adAllLaunches k launches)) all))))) -- whether a buffer maps a device address def nvidiaPlanHostBufferMaps = (lambda unrestricted buffer : (family NvidiaPlanHostBuffer) . (lambda unrestricted device : Nat . (naturalAnd (naturalLessOrEqual (phBufferGPU buffer) device) (naturalLess device (naturalAdd (phBufferGPU buffer) (phBufferExtent buffer)))))) def adMappedHost = (lambda unrestricted buffer : (family NvidiaPlanHostBuffer) . (lambda unrestricted device : Nat . (naturalSelect (nvidiaPlanHostBufferMaps buffer device) (naturalAdd (phBufferHost buffer) (naturalSaturatingSubtract device (phBufferGPU buffer))) 0))) -- a device address's host address in a host-mapped data buffer (else 0) def adMappedDataHost = (lambda unrestricted buffers : (family NvidiaPlanHostBuffers) . (lambda unrestricted device : Nat . (eliminate NvidiaPlanHostBuffers (lambda unrestricted current : (family NvidiaPlanHostBuffers) . Nat) buffers (branch NvidiaPlanHostBuffersEnd . 0) (branch NvidiaPlanHostBuffersNext head rest induction . (naturalAdd (naturalSelect (naturalNonzero (phBufferHost head)) (adMappedHost head device) 0) induction))))) -- Parameter blocks live in the program buffer or past the QMD table -- in -- the QMD buffer, or in a mapped data buffer after it when the table and -- its spill outgrow it (Coppelius's QMD overflow, under the GB10's 512-byte -- QMD records): a device address's host address there (else 0). def nvidiaPlanHostDeviceToHost = (lambda unrestricted layout : (family NvidiaPlanHostLayout) . (lambda unrestricted device : Nat . (eliminate NvidiaPlanHostLayout (lambda unrestricted current : (family NvidiaPlanHostLayout) . Nat) layout (branch NvidiaPlanHostLayoutValue gpfifo push sem program qmd userd data abi lifecycle expectedName errorHost . (naturalAdd (adMappedHost program device) (naturalAdd (adMappedHost qmd device) (adMappedDataHost data device))))))) -- update k's sites, from its launches (ordinals): the host address of the -- word at `offset` in each launch's parameter block def nvidiaPlanHostAdamWHostSites = (lambda unrestricted layout : (family NvidiaPlanHostLayout) . (lambda unrestricted addresses : Bytes . (lambda unrestricted offset : Nat . (lambda unrestricted updates : Nat . (lambda unrestricted launches : (pi unrestricted k : Nat . (family StdList Nat)) . (let unrestricted sites = (stdListMap Nat Nat (lambda unrestricted block : Nat . (nvidiaPlanHostDeviceToHost layout (naturalAdd block offset))) (nvidiaPlanHostTableEntries addresses (adAllLaunches updates launches))) in (adUpdateStretch launches sites))))))) -- What admits the schedule (1): every update has launches (ascending); each one's -- site is mapped; and each holds, as realized, the word `baked k` the plan -- put there for its update (the first invocation's; `words` is the -- realization's NvidiaWholeProgramParameterWords at the site's offset) -- -- so a site is the word the launch's AdamW reads, not a neighbour. def nvidiaPlanHostAdamWSitesAdmitted = (lambda unrestricted layout : (family NvidiaPlanHostLayout) . (lambda unrestricted addresses : Bytes . (lambda unrestricted words : Bytes . (lambda unrestricted offset : Nat . (lambda unrestricted updates : Nat . (lambda unrestricted launches : (pi unrestricted k : Nat . (family StdList Nat)) . (lambda unrestricted baked : (pi unrestricted k : Nat . Nat) . (let unrestricted all = (adAllLaunches updates launches) in (let unrestricted blocks = (nvidiaPlanHostTableEntries addresses all) in (let unrestricted realized = (nvidiaPlanHostTableEntries words all) in (nat-eliminate (lambda unrestricted current : Nat . Nat) 1 (lambda unrestricted k : Nat . (lambda unrestricted induction : Nat . (naturalAnd induction (naturalAnd (naturalNonzero (stdListLength Nat (launches k))) (naturalAnd (stdListFold Nat Nat (lambda unrestricted block : Nat . (lambda unrestricted rest : Nat . (naturalAnd (naturalNonzero (nvidiaPlanHostDeviceToHost layout (naturalAdd block offset))) rest))) 1 (adUpdateStretch launches blocks k)) (stdListFold Nat Nat (lambda unrestricted word : Nat . (lambda unrestricted rest : Nat . (naturalAnd (naturalEqual word (baked k)) rest))) 1 (adUpdateStretch launches realized k))))))) updates))))))))))) -- ---- the computation ---- -- the scratch words (PlanHostRequest.prLearnerScratch) def adWord = (lambda unrestricted index : Nat . (naturalAdd prLearnerScratch (naturalMultiply 8 index))) def adOne : Nat = (adWord 0) def adBeta1 : Nat = (adWord 1) def adBeta2 : Nat = (adWord 2) def adRate : Nat = (adWord 3) def adEpsilon : Nat = (adWord 4) def adPower1 : Nat = (adWord 5) def adPower2 : Nat = (adWord 6) def adGap1 : Nat = (adWord 7) def adRoot : Nat = (adWord 8) def adStep : Nat = (adWord 9) def adScaled : Nat = (adWord 10) def adDenominator : Nat = (adWord 11) -- the parameter word: the step size's binary32, then the epsilon's (the -- second store's high half lands in the word after) def adOut : Nat = (adWord 12) def adOp = (lambda unrestricted kind : (family NativePhysicalFloat64Operation) . (lambda unrestricted destination : Nat . (lambda unrestricted left : (family NativePhysicalOperand) . (lambda unrestricted right : (family NativePhysicalOperand) . (lambda unrestricted tail : (family NativePhysicalCommands) . (phNext (constructor NativePhysicalOperation NativePhysicalFloat64 kind (phState destination) left right) tail)))))) -- the arithmetic's rounding of a word just computed def adRound = (lambda unrestricted arithmetic : (family NvidiaPlanHostArithmetic) . (lambda unrestricted at : Nat . (lambda unrestricted tail : (family NativePhysicalCommands) . (eliminate NvidiaPlanHostArithmetic (lambda unrestricted current : (family NvidiaPlanHostArithmetic) . (family NativePhysicalCommands)) arithmetic (branch NvidiaPlanHostBinary64 . tail) (branch NvidiaPlanHostBinary32 . (adOp (constructor NativePhysicalFloat64Operation NativePhysicalFloat64ToBinary32) at (phLoad at) (phImm 0) (adOp (constructor NativePhysicalFloat64Operation NativePhysicalFloat64FromBinary32) at (phLoad at) (phImm 0) tail))))))) def adArith = (lambda unrestricted arithmetic : (family NvidiaPlanHostArithmetic) . (lambda unrestricted kind : (family NativePhysicalFloat64Operation) . (lambda unrestricted destination : Nat . (lambda unrestricted left : Nat . (lambda unrestricted right : Nat . (lambda unrestricted tail : (family NativePhysicalCommands) . (adOp kind destination (phLoad left) (phLoad right) (adRound arithmetic destination tail)))))))) def adConstant = (lambda unrestricted arithmetic : (family NvidiaPlanHostArithmetic) . (lambda unrestricted constant : (family NvidiaPlanHostConstant) . (lambda unrestricted destination : Nat . (lambda unrestricted tail : (family NativePhysicalCommands) . (eliminate NvidiaPlanHostConstant (lambda unrestricted current : (family NvidiaPlanHostConstant) . (family NativePhysicalCommands)) constant (branch NvidiaPlanHostRatio numerator denominator . (adOp (constructor NativePhysicalFloat64Operation NativePhysicalFloat64FromNatural) destination (phImm numerator) (phImm 0) (adOp (constructor NativePhysicalFloat64Operation NativePhysicalFloat64FromNatural) adDenominator (phImm denominator) (phImm 0) (adArith arithmetic (constructor NativePhysicalFloat64Operation NativePhysicalFloat64Divide) destination destination adDenominator tail)))) (branch NvidiaPlanHostBinary32Bits bits . (phNext (constructor NativePhysicalOperation NativePhysicalStoreWord64 (phState destination) (phImm bits)) (adOp (constructor NativePhysicalFloat64Operation NativePhysicalFloat64FromBinary32) destination (phLoad destination) (phImm 0) tail)))))))) def adMultiply = (constructor NativePhysicalFloat64Operation NativePhysicalFloat64Multiply) def adSubtract = (constructor NativePhysicalFloat64Operation NativePhysicalFloat64Subtract) def adDivide = (constructor NativePhysicalFloat64Operation NativePhysicalFloat64Divide) -- one more update's powers def adAdvance = (lambda unrestricted arithmetic : (family NvidiaPlanHostArithmetic) . (lambda unrestricted tail : (family NativePhysicalCommands) . (adArith arithmetic adMultiply adPower1 adPower1 adBeta1 (adArith arithmetic adMultiply adPower2 adPower2 adBeta2 tail)))) -- update k: its powers, its two scalars, their parameter word into every -- site of the update def adUpdate = (lambda unrestricted arithmetic : (family NvidiaPlanHostArithmetic) . (lambda unrestricted sites : (family StdList Nat) . (lambda unrestricted tail : (family NativePhysicalCommands) . (adAdvance arithmetic (adArith arithmetic adSubtract adRoot adOne adPower2 (adOp (constructor NativePhysicalFloat64Operation NativePhysicalFloat64SquareRoot) adRoot (phImm 0) (phLoad adRoot) (adRound arithmetic adRoot (adArith arithmetic adSubtract adGap1 adOne adPower1 (adArith arithmetic adMultiply adStep adRate adRoot (adArith arithmetic adDivide adStep adStep adGap1 (adArith arithmetic adMultiply adScaled adEpsilon adRoot (adOp (constructor NativePhysicalFloat64Operation NativePhysicalFloat64ToBinary32) adOut (phLoad adStep) (phImm 0) (adOp (constructor NativePhysicalFloat64Operation NativePhysicalFloat64ToBinary32) (naturalAdd adOut 4) (phLoad adScaled) (phImm 0) (stdListFold Nat (family NativePhysicalCommands) (lambda unrestricted site : Nat . (lambda unrestricted rest : (family NativePhysicalCommands) . (phNext (constructor NativePhysicalOperation NativePhysicalStoreWord64 (phImm site) (phLoad adOut)) rest))) tail sites)))))))))))))) -- The commands: the constants, the powers after the updates the -- checkpoint's header counts, then each of the invocation's `updates` in -- order, update k's word written to `sites k` (host addresses). def nvidiaPlanHostAdamWCommands = (lambda unrestricted adamW : (family NvidiaPlanHostAdamW) . (lambda unrestricted updates : Nat . (lambda unrestricted sites : (pi unrestricted k : Nat . (family StdList Nat)) . (lambda unrestricted tail : (family NativePhysicalCommands) . (eliminate NvidiaPlanHostAdamW (lambda unrestricted current : (family NvidiaPlanHostAdamW) . (family NativePhysicalCommands)) adamW (branch NvidiaPlanHostAdamWValue arithmetic rate beta1 beta2 epsilon . (adOp (constructor NativePhysicalFloat64Operation NativePhysicalFloat64FromNatural) adOne (phImm 1) (phImm 0) (adConstant arithmetic beta1 adBeta1 (adConstant arithmetic beta2 adBeta2 (adConstant arithmetic rate adRate (adConstant arithmetic epsilon adEpsilon (phNext (constructor NativePhysicalOperation NativePhysicalStoreWord64 (phState adPower1) (phLoad adOne)) (phNext (constructor NativePhysicalOperation NativePhysicalStoreWord64 (phState adPower2) (phLoad adOne)) (phNext (constructor NativePhysicalOperation NativePhysicalRepeatBeginCounted (phLoad (naturalAdd prCheckpointHeaderIn checkpointEnvelopeUpdatesAt))) (adAdvance arithmetic (phNext (constructor NativePhysicalOperation NativePhysicalRepeatEnd) (app (nat-eliminate (lambda unrestricted current : Nat . (pi unrestricted after : (family NativePhysicalCommands) . (family NativePhysicalCommands))) (lambda unrestricted after : (family NativePhysicalCommands) . after) (lambda unrestricted k : Nat . (lambda unrestricted induction : (pi unrestricted after : (family NativePhysicalCommands) . (family NativePhysicalCommands)) . (lambda unrestricted after : (family NativePhysicalCommands) . (induction (adUpdate arithmetic (sites k) after))))) updates) tail))))))))))))))))) -- One more update's commands, for a run loop's body -- (PlanHostRequest.NvidiaPlanHostStepCommands): the powers advanced, the -- two scalars computed, their word written into every one of `sites`. The -- constants and the powers the checkpoint counts come first, from -- `nvidiaPlanHostAdamWCommands` with no updates of its own. def nvidiaPlanHostAdamWStepCommands = (lambda unrestricted adamW : (family NvidiaPlanHostAdamW) . (lambda unrestricted sites : (family StdList Nat) . (eliminate NvidiaPlanHostAdamW (lambda unrestricted current : (family NvidiaPlanHostAdamW) . (family NativePhysicalCommands)) adamW (branch NvidiaPlanHostAdamWValue arithmetic rate beta1 beta2 epsilon . (adUpdate arithmetic sites (constructor NativePhysicalCommands NativePhysicalCommandsEnd))))))