oracle.py 7.4 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229
  1. """Independent binary64 references. No host libm and no calls to dmath.
  2. Integer/Fraction arithmetic handles exact operations. Decimal handles roots,
  3. exponentials and logs; Machin's formula and convergent series handle angles.
  4. The generator requires identical rounded results at two different precisions.
  5. """
  6. from decimal import Decimal as D, getcontext, ROUND_HALF_EVEN
  7. from fractions import Fraction
  8. from functools import lru_cache
  9. import struct
  10. SIGN = 1 << 63
  11. INF = 0x7FF0000000000000
  12. NAN = 0x7FF8000000000000
  13. MASK = SIGN - 1
  14. def bits(x):
  15. return struct.unpack('>Q', struct.pack('>d', x))[0]
  16. def value(u):
  17. return struct.unpack('>d', struct.pack('>Q', u))[0]
  18. def rounded(x):
  19. if isinstance(x, D) and x.is_nan():
  20. return NAN
  21. try:
  22. return bits(float(x))
  23. except OverflowError:
  24. return INF | (SIGN if x < 0 else 0)
  25. def atan_series(x):
  26. term = total = x
  27. square = -x * x
  28. n = 1
  29. while True:
  30. term *= square
  31. after = total + term / (2 * n + 1)
  32. if after == total:
  33. return total
  34. total = after
  35. n += 1
  36. @lru_cache(None)
  37. def pi(precision):
  38. assert precision == getcontext().prec
  39. return 16 * atan_series(D(1) / 5) - 4 * atan_series(D(1) / 239)
  40. def atan(x):
  41. if x.is_zero():
  42. return x
  43. if x.is_signed():
  44. return -atan(-x)
  45. if x > 1:
  46. return pi(getcontext().prec) / 2 - atan(1 / x)
  47. scale = 1
  48. while x > D('0.03125'):
  49. x = x / (1 + (1 + x * x).sqrt())
  50. scale *= 2
  51. return scale * atan_series(x)
  52. def atan2(y, x):
  53. p = pi(getcontext().prec)
  54. if y.is_nan() or x.is_nan():
  55. return D('NaN')
  56. if y.is_zero():
  57. return p.copy_sign(y) if x.is_signed() else y
  58. if x.is_zero():
  59. return (p / 2).copy_sign(y)
  60. if y.is_infinite():
  61. angle = p / 4 if x.is_infinite() else p / 2
  62. if x.is_infinite() and x.is_signed():
  63. angle = 3 * p / 4
  64. return angle.copy_sign(y)
  65. if x.is_infinite():
  66. return (p if x.is_signed() else D(0)).copy_sign(y)
  67. angle = atan(abs(y / x))
  68. if x.is_signed():
  69. angle = p - angle
  70. return angle.copy_sign(y)
  71. def sincos(x):
  72. if x.is_zero():
  73. return x, D(1)
  74. half_pi = pi(getcontext().prec) / 2
  75. quadrant = (x / half_pi).to_integral_value(rounding=ROUND_HALF_EVEN)
  76. r = x - quadrant * half_pi
  77. sin_term = sine = r
  78. cos_term = cosine = D(1)
  79. square = -r * r
  80. n = 1
  81. while True:
  82. sin_term = sin_term * square / ((2 * n) * (2 * n + 1))
  83. cos_term = cos_term * square / ((2 * n - 1) * (2 * n))
  84. s, c = sine + sin_term, cosine + cos_term
  85. if s == sine and c == cosine:
  86. break
  87. sine, cosine = s, c
  88. n += 1
  89. return [(sine, cosine), (cosine, -sine), (-sine, -cosine),
  90. (-cosine, sine)][int(quadrant) % 4]
  91. def exponential(x):
  92. # These bounds are outside the finite/nonzero binary64 result range.
  93. if x > 1500:
  94. return D('Infinity')
  95. if x < -1500:
  96. return D(0)
  97. return x.exp()
  98. def reference(name, a, b=0):
  99. """Return output words (two for modf/sincos, one otherwise)."""
  100. x, y = value(a), value(b)
  101. ax, ay = a & MASK, b & MASK
  102. nx, ny = ax > INF, ay > INF
  103. if name == 'isnan':
  104. return (int(nx),)
  105. if name == 'isinf':
  106. return (int(ax == INF),)
  107. if name == 'isfinite':
  108. return (int(ax < INF),)
  109. if name == 'isnormal':
  110. return (int(0x0010000000000000 <= ax < INF),)
  111. if name == 'fabs':
  112. return (ax,)
  113. if name == 'copysign':
  114. return (ax | (b & SIGN),)
  115. if name in ('fmin', 'fmax'):
  116. if nx:
  117. return (NAN if ny else b,)
  118. if ny:
  119. return (a,)
  120. if ax == 0 and ay == 0:
  121. return ((a | b) if name == 'fmin' else (a & b),)
  122. return (a if (x < y if name == 'fmin' else x > y) else b,)
  123. if name == 'pow':
  124. if y == 0 or x == 1:
  125. return (bits(1.0),)
  126. if nx or ny:
  127. return (NAN,)
  128. if ay == INF:
  129. if abs(x) == 1:
  130. return (bits(1.0),)
  131. return (INF if (abs(x) > 1) == (y > 0) else 0,)
  132. odd = y.is_integer() and abs(y) < 2**53 and int(y) % 2 != 0
  133. sign = SIGN if (a & SIGN) and odd else 0
  134. if ax == 0:
  135. return (sign | (INF if y < 0 else 0),)
  136. if ax == INF:
  137. return (sign | (INF if y > 0 else 0),)
  138. if x < 0 and not y.is_integer():
  139. return (NAN,)
  140. if y.is_integer() and abs(y) <= 4097:
  141. return (rounded(Fraction(x) ** int(y)),)
  142. z = exponential(D.from_float(y) * D.from_float(abs(x)).ln())
  143. return (rounded(z) | sign,)
  144. if nx or (name in ('atan2', 'fmod', 'log_base') and ny):
  145. return (NAN, NAN) if name in ('modf', 'sincos') else (NAN,)
  146. if name in ('ceil', 'floor', 'trunc', 'modf'):
  147. if ax == INF or ax == 0:
  148. return (a & SIGN, a) if name == 'modf' else (a,)
  149. exact = Fraction(x)
  150. integer = int(exact)
  151. if name == 'ceil':
  152. integer = -((-exact.numerator) // exact.denominator)
  153. if name == 'floor':
  154. integer = exact.numerator // exact.denominator
  155. integral = bits(float(integer)) if integer else a & SIGN
  156. if name == 'modf':
  157. fraction = exact - integer
  158. return (rounded(fraction) if fraction else a & SIGN, integral)
  159. return (integral,)
  160. if name == 'fmod':
  161. if ay == 0 or ax == INF:
  162. return (NAN,)
  163. if ay == INF or ax == 0:
  164. return (a,)
  165. q = int(Fraction(x) / Fraction(y))
  166. remainder = Fraction(x) - q * Fraction(y)
  167. return (rounded(remainder) if remainder else a & SIGN,)
  168. dx, dy = D.from_float(x), D.from_float(y)
  169. if name == 'sqrt':
  170. return (rounded(dx.sqrt()) if x >= 0 else NAN,)
  171. if name == 'cbrt':
  172. if ax == 0 or ax == INF:
  173. return (a,)
  174. return (rounded(exponential(abs(dx).ln() / 3)) | (a & SIGN),)
  175. if name in ('exp', 'exp2', 'exp10'):
  176. # Resolve exact halfway/subnormal cases using rationals, rather than
  177. # allowing the last Decimal rounding error to decide a binary64 tie.
  178. if name == 'exp2' and ax < INF and x.is_integer() and abs(x) <= 1100:
  179. return (rounded(Fraction(2) ** int(x)),)
  180. if name == 'exp10' and ax < INF and x.is_integer() and abs(x) <= 500:
  181. return (rounded(Fraction(10) ** int(x)),)
  182. factor = {'exp': D(1), 'exp2': D(2).ln(), 'exp10': D(10).ln()}[name]
  183. return (rounded(exponential(dx * factor)),)
  184. if name in ('log', 'log2', 'log10', 'log_base'):
  185. logarithm = dx.ln()
  186. divisor = {'log': D(1), 'log2': D(2).ln(), 'log10': D(10).ln()}
  187. denominator = dy.ln() if name == 'log_base' else divisor[name]
  188. return (rounded(logarithm / denominator),)
  189. if name in ('sin', 'cos', 'tan', 'sincos'):
  190. if ax == INF:
  191. return (NAN, NAN) if name == 'sincos' else (NAN,)
  192. s, c = sincos(dx)
  193. if name == 'sincos':
  194. return rounded(s), rounded(c)
  195. return (rounded({'sin': s, 'cos': c, 'tan': s / c}[name]),)
  196. if name == 'atan':
  197. return (rounded(atan(dx)),)
  198. if name == 'atan2':
  199. return (rounded(atan2(dx, dy)),) # arguments are (y, x)
  200. if name in ('asin', 'acos'):
  201. if abs(dx) > 1:
  202. return (NAN,)
  203. root = (1 - dx * dx).sqrt()
  204. return (rounded(atan2(dx, root) if name == 'asin' else atan2(root, dx)),)
  205. raise ValueError(name)