module Learning.Checked.MomentCorrection import Float32Exact import Std.Natural -- For positive update t, the two correction factors used by a decoupled -- moment update are lr * sqrt(1 - beta2^t) / (1 - beta1^t) and -- eps * sqrt(1 - beta2^t). Keep the square root interval exact until both -- bounds round to the same F32 word; a host or GPU must never substitute a -- host-language floating-point approximation for these launch scalars. def momentNaturalPower = (lambda unrestricted base : Nat . (lambda unrestricted exponent : Nat . (nat-eliminate (lambda unrestricted current : Nat . Nat) 1 (lambda unrestricted predecessor : Nat . (lambda unrestricted result : Nat . (naturalMultiply base result))) exponent))) def momentPowerGap = (lambda unrestricted numerator : Nat . (lambda unrestricted denominator : Nat . (lambda unrestricted step : Nat . (naturalSaturatingSubtract (momentNaturalPower denominator step) (momentNaturalPower numerator step))))) def momentScaledRootWord = (lambda unrestricted scaleNumerator : Nat . (lambda unrestricted scaleDenominator : Nat . (lambda unrestricted beta2Numerator : Nat . (lambda unrestricted beta2Denominator : Nat . (lambda unrestricted step : Nat . (let unrestricted gap = (momentPowerGap beta2Numerator beta2Denominator step) in (let unrestricted denominator = (momentNaturalPower beta2Denominator step) in (let unrestricted root = (float32ExactRootLow gap denominator) in (float32ExactBetween (naturalMultiply scaleNumerator root) (naturalMultiply scaleDenominator (naturalMultiply denominator float32RootScale)) (naturalMultiply scaleNumerator (succ root)) (naturalMultiply scaleDenominator (naturalMultiply denominator float32RootScale))))))))))) def momentCorrectedStepSizeWord = (lambda unrestricted rateNumerator : Nat . (lambda unrestricted rateDenominator : Nat . (lambda unrestricted beta1Numerator : Nat . (lambda unrestricted beta1Denominator : Nat . (lambda unrestricted beta2Numerator : Nat . (lambda unrestricted beta2Denominator : Nat . (lambda unrestricted step : Nat . (momentScaledRootWord (naturalMultiply rateNumerator (momentNaturalPower beta1Denominator step)) (naturalMultiply rateDenominator (momentPowerGap beta1Numerator beta1Denominator step)) beta2Numerator beta2Denominator step)))))))) def momentScaledEpsilonWord = (lambda unrestricted epsilonNumerator : Nat . (lambda unrestricted epsilonDenominator : Nat . (lambda unrestricted beta2Numerator : Nat . (lambda unrestricted beta2Denominator : Nat . (lambda unrestricted step : Nat . (momentScaledRootWord epsilonNumerator epsilonDenominator beta2Numerator beta2Denominator step))))))