Source/Packages

Learning.Checked.MomentCorrection

packages/learning/src/Learning/Checked/MomentCorrection.alpha

72 lines5 declarations3.1 KiBSHA-256 61091b557e77

Complete file · line 65

MomentCorrection.alpha

Definition view
1module Learning.Checked.MomentCorrection
2
3import Float32Exact
4import Std.Natural
5
6-- For positive update t, the two correction factors used by a decoupled
7-- moment update are lr * sqrt(1 - beta2^t) / (1 - beta1^t) and
8-- eps * sqrt(1 - beta2^t). Keep the square root interval exact until both
9-- bounds round to the same F32 word; a host or GPU must never substitute a
10-- host-language floating-point approximation for these launch scalars.
11def momentNaturalPower =
12  (lambda unrestricted base : Nat .
13    (lambda unrestricted exponent : Nat .
14      (nat-eliminate
15        (lambda unrestricted current : Nat . Nat)
16        1
17        (lambda unrestricted predecessor : Nat .
18          (lambda unrestricted result : Nat .
19            (naturalMultiply base result)))
20        exponent)))
21
22def momentPowerGap =
23  (lambda unrestricted numerator : Nat .
24    (lambda unrestricted denominator : Nat .
25      (lambda unrestricted step : Nat .
26        (naturalSaturatingSubtract
27          (momentNaturalPower denominator step)
28          (momentNaturalPower numerator step)))))
29
30def momentScaledRootWord =
31  (lambda unrestricted scaleNumerator : Nat .
32    (lambda unrestricted scaleDenominator : Nat .
33      (lambda unrestricted beta2Numerator : Nat .
34        (lambda unrestricted beta2Denominator : Nat .
35          (lambda unrestricted step : Nat .
36            (let unrestricted gap =
37              (momentPowerGap beta2Numerator beta2Denominator step) in
38            (let unrestricted denominator =
39              (momentNaturalPower beta2Denominator step) in
40            (let unrestricted root =
41              (float32ExactRootLow gap denominator) in
42              (float32ExactBetween
43                (naturalMultiply scaleNumerator root)
44                (naturalMultiply scaleDenominator
45                  (naturalMultiply denominator float32RootScale))
46                (naturalMultiply scaleNumerator (succ root))
47                (naturalMultiply scaleDenominator
48                  (naturalMultiply denominator float32RootScale)))))))))))
49
50def momentCorrectedStepSizeWord =
51  (lambda unrestricted rateNumerator : Nat .
52    (lambda unrestricted rateDenominator : Nat .
53      (lambda unrestricted beta1Numerator : Nat .
54        (lambda unrestricted beta1Denominator : Nat .
55          (lambda unrestricted beta2Numerator : Nat .
56            (lambda unrestricted beta2Denominator : Nat .
57              (lambda unrestricted step : Nat .
58                (momentScaledRootWord
59                  (naturalMultiply rateNumerator
60                    (momentNaturalPower beta1Denominator step))
61                  (naturalMultiply rateDenominator
62                    (momentPowerGap beta1Numerator beta1Denominator step))
63                  beta2Numerator beta2Denominator step))))))))
64
65def momentScaledEpsilonWord =
66  (lambda unrestricted epsilonNumerator : Nat .
67    (lambda unrestricted epsilonDenominator : Nat .
68      (lambda unrestricted beta2Numerator : Nat .
69        (lambda unrestricted beta2Denominator : Nat .
70          (lambda unrestricted step : Nat .
71            (momentScaledRootWord epsilonNumerator epsilonDenominator
72              beta2Numerator beta2Denominator step))))))

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.