module Realization.Nvidia.SM86.GatedGELUSM86 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 -- A gated feed-forward's activation in one pass each way, -- the same arithmetic as the separate launches it replaces -- (Realization.Nvidia.SM86.GELUSM86's forward and backward, the elementwise -- product, the binary32 -> half cast), instruction for instruction, so the -- results are bit-identical: -- -- forward gelu = GELU(gate) (binary32, kept for the -- backward) -- gated = half(up gelu) (the down product's -- operand) -- backward upGradient = half(dGated gelu) -- gateGradient = half(GELU'(gate) (dGated up)) -- -- with GELU(x) = x ((tanh(s (x^3 c + x)) + 1) h) and GELU'(x) g = -- g ((1 + t) h + x (1 - t^2) (x^2 3c + 1) s h), t = tanh(s (x^3 c + x)), the -- scalars s = sqrt(2/pi), c = 0.044715, h = 1/2, 1 and 3c as GELUSM86 reads -- them. The half packing is the cast's: F2FP with element 2i + 1 high, -- which the realization's half format makes binary16 or bfloat16. Two -- elements a thread, 256 threads a block, the element count a multiple of -- 512. Parameters (constant bank 0 from the SM86 parameter base): the -- pointers (forward: gelu, gated, up, gate; backward: upGradient, -- gateGradient, dGated, gelu, up, gate), then after six pointer slots the -- scalars s, c, h, 1, 3c. def ggThreads : Nat = 256 def ggElementsPerThread : Nat = 2 def ggScalar = (lambda unrestricted j : Nat . (naturalAdd (saArgument 6) (naturalMultiply 4 j))) -- thread index T = block 256 + thread; T 8 and T 4: its binary32 pair's and -- half pair's byte offsets. R0 thread, R1 block, R2 T, R3 one, R19 zero, -- R20 T 8, R21 T 4 def ggPrologue = (lambda unrestricted tail : (family SM86Program) . (saS2R 0 (constructor SM86SpecialRegister SM86ThreadIdX) (saS2R 1 (constructor SM86SpecialRegister SM86CooperativeThreadArrayIdX) (saMovImm 19 0 (saMovImmAfter 3 1 saWait5 (saImad 2 1 ggThreads 0 (saImad 20 2 8 19 (saImad 21 2 4 19 tail)))))))) -- the pair at `pointer` (a register pair) into d, d + 1, under SB0 def ggLoadPair = (lambda unrestricted d : Nat . (lambda unrestricted pointer : Nat . (lambda unrestricted tail : (family SM86Program) . (saLoad d pointer 0 saSB0 saWaitNone (saLoad (succ d) pointer 4 saSB0 saWaitNone tail))))) def ggFfma = (lambda unrestricted d : Nat . (lambda unrestricted a : Nat . (lambda unrestricted b : Nat . (lambda unrestricted c : Nat . (saFfma d a b c))))) def ggTanh = (lambda unrestricted d : Nat . (lambda unrestricted a : Nat . (saOp (constructor SM86InstructionBody SM86MultiFunctionUnitApproximation (saR d) (saR a) (constructor SM86MultiFunction SM86HyperbolicTangent) (saCtl saSB1 saWaitNone))))) def ggPack = (lambda unrestricted d : Nat . (lambda unrestricted high : Nat . (lambda unrestricted low : Nat . (saPack d high low)))) def ggStore = (lambda unrestricted pointer : Nat . (lambda unrestricted value : Nat . (lambda unrestricted tail : (family SM86Program) . (saOp (saStore pointer value 0 saWaitNone) tail)))) -- ---- forward ---- -- R4:R5 up, R6:R7 gate, R8:R9 gelu, R10:R11 gated; R14, R15 up; R16, R17 -- gate; R22 s, R23 c, R24 h, R25 1; element e's temporaries from 26 + 4e; -- R38, R39 gelu; R40, R41 up gelu; R42 their half pair def ggForwardElement = (lambda unrestricted e : Nat . (lambda unrestricted tail : (family SM86Program) . (let unrestricted x = (naturalAdd 16 e) in (let unrestricted t = (naturalAdd 26 (naturalMultiply 4 e)) in (saFmul t x x (naturalSelect e saWaitNone saWait0) (saFmul (naturalAdd t 1) t x saWaitNone (ggFfma (naturalAdd t 2) (naturalAdd t 1) 23 x (saFmul (naturalAdd t 2) (naturalAdd t 2) 22 saWaitNone (ggTanh (naturalAdd t 3) (naturalAdd t 2) (saFadd (naturalAdd t 3) (naturalAdd t 3) 25 saWait1 (saFmul (naturalAdd t 3) (naturalAdd t 3) 24 saWaitNone (saFmul (naturalAdd 38 e) x (naturalAdd t 3) saWaitNone (saFmul (naturalAdd 40 e) (naturalAdd 14 e) (naturalAdd 38 e) saWaitNone tail))))))))))))) def gatedGELUForwardSM86 : (family SM86Program) = (ggPrologue (saWide 4 20 3 (saArgument 2) (saWide 6 20 3 (saArgument 3) (saWide 8 20 3 (saArgument 0) (saWide 10 21 3 (saArgument 1) (ggLoadPair 14 4 (ggLoadPair 16 6 (saMovConst 22 (ggScalar 0) (saMovConst 23 (ggScalar 1) (saMovConst 24 (ggScalar 2) (saMovConst 25 (ggScalar 3) (ggForwardElement 0 (ggForwardElement 1 (saOp (sbStore64 8 38 0 saWaitNone) (ggPack 42 41 40 (ggStore 10 42 (saOp (constructor SM86InstructionBody SM86Exit (saAfter saWaitAll)) sm86ProgramEmpty))))))))))))))))) def gatedGELUForwardSM86Registers : Nat = 43 -- ---- backward ---- -- R4:R5 dGated, R6:R7 gelu, R8:R9 up, R10:R11 gate, R12:R13 upGradient, -- R14:R15 gateGradient; R22, R23 dGated; R24, R25 gelu; R26, R27 up; R28, -- R29 gate; R30 s, R31 c, R32 h, R33 1, R34 3c; element e's from 36 + 16e; -- R70, R71 the half pairs def ggBackwardElement = (lambda unrestricted e : Nat . (lambda unrestricted tail : (family SM86Program) . (let unrestricted b = (naturalAdd 36 (naturalMultiply 16 e)) in (let unrestricted x = (naturalAdd 28 e) in (let unrestricted r = (lambda unrestricted i : Nat . (naturalAdd b i)) in -- upGradient = dGated gelu; g = dGated up (saFmul (r 0) (naturalAdd 22 e) (naturalAdd 24 e) (naturalSelect e saWaitNone saWait0) (saFmul (r 1) (naturalAdd 22 e) (naturalAdd 26 e) saWaitNone -- t = tanh(s (x^3 c + x)) (saFmul (r 2) x x saWaitNone (saFmul (r 3) (r 2) x saWaitNone (ggFfma (r 4) (r 3) 31 x (saFmul (r 4) (r 4) 30 saWaitNone (ggTanh (r 5) (r 4) -- 1 - t^2 (saFmul (r 6) (r 5) (r 5) saWait1 (saFneg (r 6) (r 6) (saFadd (r 6) 33 (r 6) saWaitNone -- (x^2 3c + 1) s (ggFfma (r 7) (r 2) 34 33 (saFmul (r 7) (r 7) 30 saWaitNone -- (1 + t) h (saFadd (r 8) 33 (r 5) saWaitNone (saFmul (r 8) (r 8) 32 saWaitNone -- x (1 - t^2) (x^2 3c + 1) s h (saFmul (r 9) x (r 6) saWaitNone (saFmul (r 9) (r 9) (r 7) saWaitNone (saFmul (r 9) (r 9) 32 saWaitNone (saFadd (r 10) (r 8) (r 9) saWaitNone (saFmul (r 11) (r 1) (r 10) saWaitNone tail)))))))))))))))))))))))) def gatedGELUBackwardSM86 : (family SM86Program) = (ggPrologue (saWide 4 20 3 (saArgument 2) (saWide 6 20 3 (saArgument 3) (saWide 8 20 3 (saArgument 4) (saWide 10 20 3 (saArgument 5) (saWide 12 21 3 (saArgument 0) (saWide 14 21 3 (saArgument 1) (ggLoadPair 22 4 (ggLoadPair 24 6 (ggLoadPair 26 8 (ggLoadPair 28 10 (saMovConst 30 (ggScalar 0) (saMovConst 31 (ggScalar 1) (saMovConst 32 (ggScalar 2) (saMovConst 33 (ggScalar 3) (saMovConst 34 (ggScalar 4) (ggBackwardElement 0 (ggBackwardElement 1 (ggPack 70 52 36 (ggPack 71 63 47 (ggStore 12 70 (ggStore 14 71 (saOp (constructor SM86InstructionBody SM86Exit (saAfter saWaitAll)) sm86ProgramEmpty))))))))))))))))))))))) def gatedGELUBackwardSM86Registers : Nat = 72 def gatedGELUSM86Threads : Nat = ggThreads -- the grid for `elements` (a multiple of 512) def gatedGELUSM86Blocks = (lambda unrestricted elements : Nat . (naturalDivideUnchecked elements (naturalMultiply ggThreads ggElementsPerThread))) -- through the five scalars, in whole 8-byte parameter words def gatedGELUSM86ConstantBytes : Nat = (ggScalar 6)