module Realization.Nvidia.SM86.AdamWHalfSM86 import Accelerator.SM86.Control import Accelerator.SM86.Immediate import Accelerator.SM86.Instruction import Accelerator.SM86.NumericSemantics import Accelerator.SM86.Program import Accelerator.SM86.Types import Realization.Nvidia.SM86.StreamingAttentionSM86 import Std.Natural -- AdamW over the whole flat bank with the parameters' half copies written -- as it goes: the half cast that followed the -- update, fused into it. Each element's update is -- Realization.Nvidia.SM86.AdamWSM86's (adamWF32MathProgram), instruction -- for instruction, and each half pair is the cast's (F2FP, element 2i + 1 -- high, in the realization's half format), so the parameters, the moments -- and the halves are bit-identical to the update followed by the cast: -- -- g = g one; m = m beta1 + g (1 - beta1); v = v beta2 + g^2 (1 - beta2) -- p = p decay + -(m stepSize (1 / (sqrt v + epsilon))) -- half[i] = half(p[i]) -- -- The half bank is parallel to the parameter bank: element i's half at byte -- 2 i of it. A thread takes the quad 4T .. 4T + 3 (T = block 256 + thread): -- each bank's four elements in one 128-bit load and store, the four halves -- in one 64-bit store -- the bank is memory-bound, and wide accesses keep -- more of it in flight a thread -- then T + stride, ... `iterations` times; -- quads at or past `quads` are untouched. The gradient is read, never -- written (every gradient the step reads was written whole by the backward -- before it). -- -- Parameters (constant bank 0 from the SM86 parameter base): the pointers -- P, G, M, V (arguments 0-3, 16-byte aligned); the scalars one, beta1, -- 1 - beta1, beta2, 1 - beta2, decay, stepSize, epsilon (words from argument -- 4, where AdamWSM86's ABI puts them, so the host writes the step's scalars -- at the same offsets); the half bank's pointer (argument 9, 8-byte -- aligned). -- -- Registers: R0 thread, R1 block, R2 T, R3 sixteen, R4 eight; R6:R7, R8:R9, -- R10:R11, R12:R13 the quad's P, G, M, V; R18..R25 the scalars; R29 the -- iterations left; R30:R31 the halves' address; R32..R35, R36..R39, -- R40..R43, R44..R47 the quad's p, g, m, v; element e's temporaries -- R48 + 3e .. R50 + 3e; R60, R61 the half pairs. Loads under SB0..SB3 (P, -- G, M, V), the square roots and reciprocals under SB4; the stores' reads -- under SB5, which the next iteration's first load waits for. def ahThreads : Nat = 256 def adamWHalfSM86ParameterArgument : Nat = 0 def adamWHalfSM86GradientArgument : Nat = 1 def adamWHalfSM86FirstMomentArgument : Nat = 2 def adamWHalfSM86SecondMomentArgument : Nat = 3 def adamWHalfSM86ScalarPairArgument = (lambda unrestricted pair : Nat . (naturalAdd 4 pair)) def adamWHalfSM86HalfArgument : Nat = 9 def adamWHalfSM86ArgumentCount : Nat = 10 def ahScalar = (lambda unrestricted j : Nat . (naturalAdd (saArgument (adamWHalfSM86ScalarPairArgument 0)) (naturalMultiply 4 j))) def ahHalfPointer : Nat = (saArgument adamWHalfSM86HalfArgument) -- element e's p, g, m, v (bank b = 0..3) and temporaries (i = 0..2) def ahValue = (lambda unrestricted e : Nat . (lambda unrestricted b : Nat . (naturalAdd (naturalAdd 32 (naturalMultiply 4 b)) e))) def ahTemporary = (lambda unrestricted e : Nat . (lambda unrestricted i : Nat . (naturalAdd (naturalAdd 48 (naturalMultiply 3 e)) i))) def ahHalfPairs : Nat = 60 def ahCounter : Nat = 29 def ahIterations = (lambda unrestricted quads : Nat . (lambda unrestricted stride : Nat . (naturalDivideUnchecked (naturalAdd quads (naturalSaturatingSubtract stride 1)) stride))) -- ---- the prologue: T, the constants, the scalars, the count ---- def ahPrologue = (lambda unrestricted iterations : Nat . (lambda unrestricted tail : (family SM86Program) . (saS2R 0 (constructor SM86SpecialRegister SM86ThreadIdX) (saS2R 1 (constructor SM86SpecialRegister SM86CooperativeThreadArrayIdX) (saMovImm 3 16 (saMovImm 4 8 (saMovImm ahCounter iterations (saMovConst 18 (ahScalar 0) (saMovConst 19 (ahScalar 1) (saMovConst 20 (ahScalar 2) (saMovConst 21 (ahScalar 3) (saMovConst 22 (ahScalar 4) (saMovConst 23 (ahScalar 5) (saMovConst 24 (ahScalar 6) (saMovConst 25 (ahScalar 7) (saMovImmAfter 5 0 saWait5 (saImad 2 1 ahThreads 0 tail))))))))))))))))) -- ---- the body: one quad ---- def ahGuarded = (lambda unrestricted body : (family SM86InstructionBody) . (lambda unrestricted tail : (family SM86Program) . (saUnless saP0 body tail))) -- a store whose register read SB5 counts def ahStoreControl : (family SM86Control) = (constructor SM86Control SM86ControlValue (byte 15) (constructor SM86YieldMode SM86Continue) saNone saSB5 (byte 0) (byte 0)) -- bounds, the four pointers and the halves' def ahAddresses = (lambda unrestricted quads : Nat . (lambda unrestricted tail : (family SM86Program) . (saGreater saP0 2 (naturalPredecessor quads) (saWide 6 2 3 (saArgument adamWHalfSM86ParameterArgument) (saWide 8 2 3 (saArgument adamWHalfSM86GradientArgument) (saWide 10 2 3 (saArgument adamWHalfSM86FirstMomentArgument) (saWide 12 2 3 (saArgument adamWHalfSM86SecondMomentArgument) (saWide 30 2 4 ahHalfPointer tail)))))))) -- bank b's quad (pointer 6 + 2b) into R32 + 4b, under SB b; the first waits -- for the last iteration's stores' reads def ahLoadBank = (lambda unrestricted b : Nat . (lambda unrestricted barrier : (family SM86Barrier) . (lambda unrestricted wait : Nat . (lambda unrestricted tail : (family SM86Program) . (ahGuarded (constructor SM86InstructionBody SM86LoadGlobalWide (saR (ahValue 0 b)) (saR (naturalAdd 6 (naturalMultiply 2 b))) (saU 0) (saCtl barrier wait)) tail))))) def ahLoads = (lambda unrestricted tail : (family SM86Program) . (ahLoadBank 0 saSB0 saWait5 (ahLoadBank 1 saSB1 saWaitNone (ahLoadBank 2 saSB2 saWaitNone (ahLoadBank 3 saSB3 saWaitNone tail))))) -- element e's update: adamWF32MathProgram's instructions on its registers def ahUpdate = (lambda unrestricted e : Nat . (lambda unrestricted tail : (family SM86Program) . (let unrestricted p = (ahValue e 0) in (let unrestricted g = (ahValue e 1) in (let unrestricted m = (ahValue e 2) in (let unrestricted v = (ahValue e 3) in (let unrestricted t0 = (ahTemporary e 0) in (let unrestricted t1 = (ahTemporary e 1) in (let unrestricted t2 = (ahTemporary e 2) in (saFmul g g 18 (naturalSelect e saWaitNone (naturalAdd saWait0 (naturalAdd saWait1 (naturalAdd saWait2 saWait3)))) (saFmul t0 g 20 saWaitNone (saFfma m m 19 t0 (saFmul t2 g g saWaitNone (saFmul t2 t2 22 saWaitNone (saFfma v v 21 t2 (saMufu t1 v (constructor SM86MultiFunction SM86SquareRoot) saWaitNone (saFadd t1 t1 25 saWait4 (saMufu t1 t1 (constructor SM86MultiFunction SM86Reciprocal) saWaitNone (saFmul t0 m 24 saWaitNone (saFmul t0 t0 t1 saWait4 (saFneg t0 t0 (saFmul p p 23 saWaitNone (saFadd p p t0 saWaitNone tail))))))))))))))))))))))) -- the quad's P, M, V (128 bits each) and its two half pairs (64 bits) def ahStores = (lambda unrestricted tail : (family SM86Program) . (saPack ahHalfPairs (ahValue 1 0) (ahValue 0 0) (saPack (succ ahHalfPairs) (ahValue 3 0) (ahValue 2 0) (ahGuarded (constructor SM86InstructionBody SM86StoreGlobalWide (saR 6) (saR (ahValue 0 0)) (saU 0) ahStoreControl) (ahGuarded (constructor SM86InstructionBody SM86StoreGlobalWide (saR 10) (saR (ahValue 0 2)) (saU 0) ahStoreControl) (ahGuarded (constructor SM86InstructionBody SM86StoreGlobalWide (saR 12) (saR (ahValue 0 3)) (saU 0) ahStoreControl) (ahGuarded (constructor SM86InstructionBody SM86StoreGlobal64 (saR 30) (saR ahHalfPairs) (saU 0) ahStoreControl) tail))))))) def ahBody = (lambda unrestricted quads : Nat . (ahAddresses quads (ahLoads (ahUpdate 0 (ahUpdate 1 (ahUpdate 2 (ahUpdate 3 (ahStores sm86ProgramEmpty)))))))) -- ---- the loop's control: the next quad, one iteration fewer ---- def ahControl = (lambda unrestricted stride : Nat . (lambda unrestricted tail : (family SM86Program) . (saAddImm 2 2 stride (saAddImm ahCounter ahCounter 4294967295 (saGreater saP1 ahCounter 0 tail))))) def ahControlCount : Nat = 3 -- back to the body's first instruction (the displacement counts the body, -- the control and the branch, from the branch's successor) def ahBranch = (lambda unrestricted body : Nat . (saWhen saP1 (constructor SM86InstructionBody SM86Branch (saU (naturalSaturatingSubtract 4294967296 (naturalMultiply 16 (naturalAdd body (naturalAdd ahControlCount 1))))) (sm86Unsigned32 (byte 255) (byte 255) (byte 131) (byte 3)) sm86SafeControl) sm86ProgramEmpty)) def ahExit : (family SM86Program) = (saOp (constructor SM86InstructionBody SM86Exit (saAfter saWaitAll)) sm86ProgramEmpty) -- `finish` is applied to the prologue and the body on their own (e.g. the -- stalls compacted and drained); the control keeps the largest stalls. -- looped 1: the loop with its back edge; looped 0: what the schedule checks -- read, the body twice and no branch (it covers both ways into the body, -- and every iteration leaves the same things in flight: the stores' reads -- under SB5, nothing else) def ahAssemble = (lambda unrestricted finish : (pi unrestricted piece : (family SM86Program) . (family SM86Program)) . (lambda unrestricted looped : Nat . (lambda unrestricted quads : Nat . (lambda unrestricted stride : Nat . (let unrestricted body = (finish (ahBody quads)) in (let unrestricted once = (sm86ProgramAppend body (ahControl stride sm86ProgramEmpty)) in (sm86ProgramAppend (finish (ahPrologue (ahIterations quads stride) sm86ProgramEmpty)) (sm86ProgramAppend once (nat-eliminate (lambda unrestricted n : Nat . (family SM86Program)) (sm86ProgramAppend once ahExit) (lambda unrestricted predecessor : Nat . (lambda unrestricted induction : (family SM86Program) . (sm86ProgramAppend (ahBranch (sm86ProgramCount body)) ahExit))) looped))))))))) -- the image for `quads` element quads (the bank's elements over four) over -- a grid whose threads number `stride` def adamWHalfSM86 = (lambda unrestricted finish : (pi unrestricted piece : (family SM86Program) . (family SM86Program)) . (ahAssemble finish 1)) def adamWHalfSM86Checked = (lambda unrestricted finish : (pi unrestricted piece : (family SM86Program) . (family SM86Program)) . (ahAssemble finish 0)) def adamWHalfSM86Threads : Nat = ahThreads def adamWHalfSM86Registers : Nat = 64 -- the elements a thread takes an iteration def adamWHalfSM86Elements : Nat = 4 def adamWHalfSM86Iterations = ahIterations -- through the half pointer, in whole 8-byte parameter words def adamWHalfSM86ConstantBytes : Nat = (saArgument adamWHalfSM86ArgumentCount) -- where the step-dependent scalars sit in the parameter block (the host -- writes each update's there) def adamWHalfSM86StepSizeOffset : Nat = (ahScalar 6) def adamWHalfSM86EpsilonOffset : Nat = (ahScalar 7)