Source/Packages

Realization.Nvidia.SM86.GatedGELUSM86

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

170 lines18 declarations7.4 KiBSHA-256 fe6a7ec05912

Complete file · line 38

GatedGELUSM86.alpha

Definition view
1module Realization.Nvidia.SM86.GatedGELUSM86
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-- A gated feed-forward's activation in one pass each way,
13-- the same arithmetic as the separate launches it replaces
14-- (Realization.Nvidia.SM86.GELUSM86's forward and backward, the elementwise
15-- product, the binary32 -> half cast), instruction for instruction, so the
16-- results are bit-identical:
17--
18--   forward    gelu = GELU(gate)                      (binary32, kept for the
19--                                                      backward)
20--              gated = half(up gelu)                  (the down product's
21--                                                      operand)
22--   backward   upGradient = half(dGated gelu)
23--              gateGradient = half(GELU'(gate) (dGated up))
24--
25-- with GELU(x) = x ((tanh(s (x^3 c + x)) + 1) h) and GELU'(x) g =
26-- g ((1 + t) h + x (1 - t^2) (x^2 3c + 1) s h), t = tanh(s (x^3 c + x)), the
27-- scalars s = sqrt(2/pi), c = 0.044715, h = 1/2, 1 and 3c as GELUSM86 reads
28-- them.  The half packing is the cast's: F2FP with element 2i + 1 high,
29-- which the realization's half format makes binary16 or bfloat16.  Two
30-- elements a thread, 256 threads a block, the element count a multiple of
31-- 512.  Parameters (constant bank 0 from the SM86 parameter base): the
32-- pointers (forward: gelu, gated, up, gate; backward: upGradient,
33-- gateGradient, dGated, gelu, up, gate), then after six pointer slots the
34-- scalars s, c, h, 1, 3c.
35
36def ggThreads : Nat = 256
37def ggElementsPerThread : Nat = 2
38def ggScalar = (lambda unrestricted j : Nat . (naturalAdd (saArgument 6) (naturalMultiply 4 j)))
39
40-- thread index T = block 256 + thread; T 8 and T 4: its binary32 pair's and
41-- half pair's byte offsets.  R0 thread, R1 block, R2 T, R3 one, R19 zero,
42-- R20 T 8, R21 T 4
43def ggPrologue = (lambda unrestricted tail : (family SM86Program) .
44  (saS2R 0 (constructor SM86SpecialRegister SM86ThreadIdX)
45  (saS2R 1 (constructor SM86SpecialRegister SM86CooperativeThreadArrayIdX)
46  (saMovImm 19 0
47  (saMovImmAfter 3 1 saWait5
48  (saImad 2 1 ggThreads 0
49  (saImad 20 2 8 19
50  (saImad 21 2 4 19 tail))))))))
51
52-- the pair at `pointer` (a register pair) into d, d + 1, under SB0
53def ggLoadPair = (lambda unrestricted d : Nat . (lambda unrestricted pointer : Nat . (lambda unrestricted tail : (family SM86Program) .
54  (saLoad d pointer 0 saSB0 saWaitNone (saLoad (succ d) pointer 4 saSB0 saWaitNone tail)))))
55
56def ggFfma = (lambda unrestricted d : Nat . (lambda unrestricted a : Nat . (lambda unrestricted b : Nat . (lambda unrestricted c : Nat .
57  (saFfma d a b c)))))
58def ggTanh = (lambda unrestricted d : Nat . (lambda unrestricted a : Nat .
59  (saOp (constructor SM86InstructionBody SM86MultiFunctionUnitApproximation (saR d) (saR a)
60    (constructor SM86MultiFunction SM86HyperbolicTangent) (saCtl saSB1 saWaitNone)))))
61def ggPack = (lambda unrestricted d : Nat . (lambda unrestricted high : Nat . (lambda unrestricted low : Nat .
62  (saPack d high low))))
63def ggStore = (lambda unrestricted pointer : Nat . (lambda unrestricted value : Nat . (lambda unrestricted tail : (family SM86Program) .
64  (saOp (saStore pointer value 0 saWaitNone) tail))))
65
66-- ---- forward ----
67-- R4:R5 up, R6:R7 gate, R8:R9 gelu, R10:R11 gated; R14, R15 up; R16, R17
68-- gate; R22 s, R23 c, R24 h, R25 1; element e's temporaries from 26 + 4e;
69-- R38, R39 gelu; R40, R41 up gelu; R42 their half pair
70def ggForwardElement = (lambda unrestricted e : Nat . (lambda unrestricted tail : (family SM86Program) .
71  (let unrestricted x = (naturalAdd 16 e) in (let unrestricted t = (naturalAdd 26 (naturalMultiply 4 e)) in
72  (saFmul t x x (naturalSelect e saWaitNone saWait0)
73  (saFmul (naturalAdd t 1) t x saWaitNone
74  (ggFfma (naturalAdd t 2) (naturalAdd t 1) 23 x
75  (saFmul (naturalAdd t 2) (naturalAdd t 2) 22 saWaitNone
76  (ggTanh (naturalAdd t 3) (naturalAdd t 2)
77  (saFadd (naturalAdd t 3) (naturalAdd t 3) 25 saWait1
78  (saFmul (naturalAdd t 3) (naturalAdd t 3) 24 saWaitNone
79  (saFmul (naturalAdd 38 e) x (naturalAdd t 3) saWaitNone
80  (saFmul (naturalAdd 40 e) (naturalAdd 14 e) (naturalAdd 38 e) saWaitNone tail)))))))))))))
81
82def gatedGELUForwardSM86 : (family SM86Program) =
83  (ggPrologue
84  (saWide 4 20 3 (saArgument 2)
85  (saWide 6 20 3 (saArgument 3)
86  (saWide 8 20 3 (saArgument 0)
87  (saWide 10 21 3 (saArgument 1)
88  (ggLoadPair 14 4
89  (ggLoadPair 16 6
90  (saMovConst 22 (ggScalar 0)
91  (saMovConst 23 (ggScalar 1)
92  (saMovConst 24 (ggScalar 2)
93  (saMovConst 25 (ggScalar 3)
94  (ggForwardElement 0
95  (ggForwardElement 1
96  (saOp (sbStore64 8 38 0 saWaitNone)
97  (ggPack 42 41 40
98  (ggStore 10 42
99  (saOp (constructor SM86InstructionBody SM86Exit (saAfter saWaitAll)) sm86ProgramEmpty)))))))))))))))))
100
101def gatedGELUForwardSM86Registers : Nat = 43
102
103-- ---- backward ----
104-- R4:R5 dGated, R6:R7 gelu, R8:R9 up, R10:R11 gate, R12:R13 upGradient,
105-- R14:R15 gateGradient; R22, R23 dGated; R24, R25 gelu; R26, R27 up; R28,
106-- R29 gate; R30 s, R31 c, R32 h, R33 1, R34 3c; element e's from 36 + 16e;
107-- R70, R71 the half pairs
108def ggBackwardElement = (lambda unrestricted e : Nat . (lambda unrestricted tail : (family SM86Program) .
109  (let unrestricted b = (naturalAdd 36 (naturalMultiply 16 e)) in
110  (let unrestricted x = (naturalAdd 28 e) in
111  (let unrestricted r = (lambda unrestricted i : Nat . (naturalAdd b i)) in
112  -- upGradient = dGated gelu; g = dGated up
113  (saFmul (r 0) (naturalAdd 22 e) (naturalAdd 24 e) (naturalSelect e saWaitNone saWait0)
114  (saFmul (r 1) (naturalAdd 22 e) (naturalAdd 26 e) saWaitNone
115  -- t = tanh(s (x^3 c + x))
116  (saFmul (r 2) x x saWaitNone
117  (saFmul (r 3) (r 2) x saWaitNone
118  (ggFfma (r 4) (r 3) 31 x
119  (saFmul (r 4) (r 4) 30 saWaitNone
120  (ggTanh (r 5) (r 4)
121  -- 1 - t^2
122  (saFmul (r 6) (r 5) (r 5) saWait1
123  (saFneg (r 6) (r 6)
124  (saFadd (r 6) 33 (r 6) saWaitNone
125  -- (x^2 3c + 1) s
126  (ggFfma (r 7) (r 2) 34 33
127  (saFmul (r 7) (r 7) 30 saWaitNone
128  -- (1 + t) h
129  (saFadd (r 8) 33 (r 5) saWaitNone
130  (saFmul (r 8) (r 8) 32 saWaitNone
131  -- x (1 - t^2) (x^2 3c + 1) s h
132  (saFmul (r 9) x (r 6) saWaitNone
133  (saFmul (r 9) (r 9) (r 7) saWaitNone
134  (saFmul (r 9) (r 9) 32 saWaitNone
135  (saFadd (r 10) (r 8) (r 9) saWaitNone
136  (saFmul (r 11) (r 1) (r 10) saWaitNone tail))))))))))))))))))))))))
137
138def gatedGELUBackwardSM86 : (family SM86Program) =
139  (ggPrologue
140  (saWide 4 20 3 (saArgument 2)
141  (saWide 6 20 3 (saArgument 3)
142  (saWide 8 20 3 (saArgument 4)
143  (saWide 10 20 3 (saArgument 5)
144  (saWide 12 21 3 (saArgument 0)
145  (saWide 14 21 3 (saArgument 1)
146  (ggLoadPair 22 4
147  (ggLoadPair 24 6
148  (ggLoadPair 26 8
149  (ggLoadPair 28 10
150  (saMovConst 30 (ggScalar 0)
151  (saMovConst 31 (ggScalar 1)
152  (saMovConst 32 (ggScalar 2)
153  (saMovConst 33 (ggScalar 3)
154  (saMovConst 34 (ggScalar 4)
155  (ggBackwardElement 0
156  (ggBackwardElement 1
157  (ggPack 70 52 36
158  (ggPack 71 63 47
159  (ggStore 12 70
160  (ggStore 14 71
161  (saOp (constructor SM86InstructionBody SM86Exit (saAfter saWaitAll)) sm86ProgramEmpty)))))))))))))))))))))))
162
163def gatedGELUBackwardSM86Registers : Nat = 72
164
165def gatedGELUSM86Threads : Nat = ggThreads
166-- the grid for `elements` (a multiple of 512)
167def gatedGELUSM86Blocks = (lambda unrestricted elements : Nat .
168  (naturalDivideUnchecked elements (naturalMultiply ggThreads ggElementsPerThread)))
169-- through the five scalars, in whole 8-byte parameter words
170def gatedGELUSM86ConstantBytes : Nat = (ggScalar 6)

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.