module Float32Exact import Float32Model import FloatLiteralSpec import Std.Natural -- binary32 constants from their exact values. Each is the binary32 nearest -- the value, computed by these functions with unbounded naturals -- the -- executable's naturals are 64-bit words and these exact values are not -- -- so a plan takes a constant as `(compile-time (derivation))`: the compiler -- evaluates it and hands the program the word. pi and ln 2 are bounded to -- 40 places, enough that both ends of their interval round alike. def float32NotRounded : Nat = 0x7fc0dead def float32ExactRational = (lambda unrestricted numerator : Nat . (lambda unrestricted denominator : Nat . (modelRoundBits 0 numerator denominator))) -- a value known to lie in [lo, hi]: its binary32 when both ends round to it def float32ExactBetween = (lambda unrestricted loN : Nat . (lambda unrestricted loD : Nat . (lambda unrestricted hiN : Nat . (lambda unrestricted hiD : Nat . (let unrestricted lo = (float32ExactRational loN loD) in (let unrestricted hi = (float32ExactRational hiN hiD) in (naturalSelect (naturalEqual lo hi) lo float32NotRounded))))))) -- sqrt(n / d) in [floor(sqrt(n d S^2)) / (d S), that + 1 / (d S)] def float32RootScale : Nat = 1267650600228229401496703205376 def float32ExactRootLow = (lambda unrestricted n : Nat . (lambda unrestricted d : Nat . (modelIntegerRoot (naturalMultiply (naturalMultiply n d) (naturalMultiply float32RootScale float32RootScale))))) def float32Places : Nat = 10000000000000000000000000000000000000000 def float32Ln2Digits : Nat = 6931471805599453094172321214581765680755 def float32PiDigits : Nat = 31415926535897932384626433832795028841971 def float32ExactInverseRoot = (lambda unrestricted n : Nat . (let unrestricted low = (float32ExactRootLow 1 n) in (float32ExactBetween low (naturalMultiply n float32RootScale) (succ low) (naturalMultiply n float32RootScale)))) -- The EX2 attention path needs score/sqrt(n) in base-two units. Bound both -- sqrt(n) and ln(2) from their rational intervals, then require that the -- bounds round to the same F32 word. This avoids rounding 1/sqrt(n) before -- multiplying by log2(e), which can choose a different final word. def float32ExactInverseRootNaturalLogTwo = (lambda unrestricted n : Nat . (let unrestricted rootLow = (float32ExactRootLow n 1) in (float32ExactBetween (naturalMultiply float32RootScale float32Places) (naturalMultiply (succ rootLow) (succ float32Ln2Digits)) (naturalMultiply float32RootScale float32Places) (naturalMultiply rootLow float32Ln2Digits)))) def float32Ln2 : Nat = (compile-time (float32ExactBetween float32Ln2Digits float32Places (succ float32Ln2Digits) float32Places)) def float32Log2E : Nat = (compile-time (float32ExactBetween float32Places (succ float32Ln2Digits) float32Places float32Ln2Digits)) -- floor((n / d)^(1 / k) S): the largest r with r^k d <= n S^k, by bisection. -- The root is below S (n / d + 1), so its bit length is at most the sum of -- theirs -- the bound and the fuel; the target n S^k itself can run past -- specBitLength's reach (S = 2^96, k = 32 is 3,073 bits) def float32ExactRootKLow = (lambda unrestricted k : Nat . (lambda unrestricted n : Nat . (lambda unrestricted d : Nat . (lambda unrestricted scale : Nat . (let unrestricted target = (nat-multiply n (specPower scale k)) in (let unrestricted bits = (nat-add (specBitLength scale) (specBitLength (nat-add (nat-divide n d) 1))) in (app (nat-eliminate (lambda unrestricted current : Nat . (pi unrestricted low : Nat . (pi unrestricted high : Nat . Nat))) (lambda unrestricted low : Nat . (lambda unrestricted high : Nat . low)) (lambda unrestricted predecessor : Nat . (lambda unrestricted induction : (pi unrestricted low : Nat . (pi unrestricted high : Nat . Nat)) . (lambda unrestricted low : Nat . (lambda unrestricted high : Nat . (specSelect (nat-less-than (nat-add low 1) high) (specSelect (nat-less-than target (nat-multiply (specPower (nat-divide (nat-add low high) 2) k) d)) (induction low (nat-divide (nat-add low high) 2)) (induction (nat-divide (nat-add low high) 2) high)) low))))) (nat-add bits 2)) zero (specPow2 bits)))))))) -- a value in [lo, hi] / d less the value of `word` (a positive binary32 -- below 2^24, exponent field at most 150): the binary32 nearest the -- difference, its sign set when the word is the larger; zero when the -- interval reaches the word's value (the difference is below its width def float32ExactRemainder = (lambda unrestricted lo : Nat . (lambda unrestricted hi : Nat . (lambda unrestricted d : Nat . (lambda unrestricted word : Nat . (let unrestricted shift = (specPow2 (nat-subtract 150 (modelExponent word))) in (let unrestricted value = (nat-multiply (nat-add (modelFraction word) modelPow2Twenty3) d) in (let unrestricted low = (nat-multiply lo shift) in (let unrestricted high = (nat-multiply hi shift) in (let unrestricted scale = (nat-multiply d shift) in (specSelect (nat-less-than value low) (float32ExactBetween (nat-subtract low value) scale (nat-subtract high value) scale) (specSelect (nat-less-than high value) (nat-add modelPow2Thirty1 (float32ExactBetween (nat-subtract value high) scale (nat-subtract value low) scale)) 0))))))))))) -- 2 pi in two words (the second the nearest to what the first misses), and -- 1 / (2 pi): pi in [P, P + 1] / 10^40 def float32TwoPiHigh : Nat = (compile-time (float32ExactBetween (nat-multiply 2 float32PiDigits) float32Places (nat-multiply 2 (succ float32PiDigits)) float32Places)) def float32TwoPiLow : Nat = (compile-time (float32ExactRemainder (nat-multiply 2 float32PiDigits) (nat-multiply 2 (succ float32PiDigits)) float32Places float32TwoPiHigh)) def float32InverseTwoPi : Nat = (compile-time (float32ExactBetween float32Places (nat-multiply 2 (succ float32PiDigits)) float32Places (nat-multiply 2 float32PiDigits)))