Merge pull request #692 from skirpichev/ctx_fp-stdlib

Change FPContext to use more functions from the stdlib
This commit is contained in:
Sergey B Kirpichev
2023-05-19 18:03:00 +03:00
committed by GitHub
3 changed files with 69 additions and 218 deletions
+5 -5
View File
@@ -255,14 +255,14 @@ Although mpmath is generally designed for arbitrary-precision arithmetic, many o
To take advantage of this feature, simply use the ``fp`` prefix, i.e. write ``fp.func`` instead of ``func`` or ``mp.func``::
>>> u = fp.erfc(2.5)
>>> u = fp.erfc(0.5)
>>> print(u)
0.000406952017444959
0.4795001221869535
>>> type(u)
<class 'float'>
>>> mp.dps = 15
>>> print(mp.erfc(2.5))
0.000406952017444959
>>> mp.dps = 16
>>> print(mp.erfc(0.5))
0.4795001221869535
>>> fp.matrix([[1,2],[3,4]]) ** 2
matrix(
[['7.0', '10.0'],
+44 -47
View File
@@ -1,6 +1,8 @@
import cmath
import math
import sys
from . import function_docs, libmp, math2
from . import function_docs, libfp, libmp
from .ctx_base import StandardBaseContext
from .libmp import int_types, mpf_bernoulli, to_float
@@ -15,7 +17,7 @@ class FPContext(StandardBaseContext):
StandardBaseContext.__init__(ctx)
# Override SpecialFunctions implementation
ctx.loggamma = math2.loggamma
ctx.loggamma = libfp.loggamma
ctx._bernoulli_cache = {}
ctx.pretty = False
@@ -35,10 +37,10 @@ class FPContext(StandardBaseContext):
zero = 0.0
one = 1.0
eps = math2.EPS
inf = math2.INF
ninf = math2.NINF
nan = math2.NAN
eps = sys.float_info.epsilon
inf = libfp.INF
ninf = -math.inf
nan = math.nan
j = 1j
# Called by SpecialFunctions.__init__()
@@ -61,15 +63,15 @@ class FPContext(StandardBaseContext):
cache[n] = to_float(mpf_bernoulli(n, 53, 'n'), strict=True)
return cache[n]
pi = math2.pi
e = math2.e
euler = math2.euler
pi = libfp.pi
e = math.e
euler = libfp.euler
sqrt2 = 1.4142135623730950488
sqrt5 = 2.2360679774997896964
phi = 1.6180339887498948482
ln2 = 0.69314718055994530942
ln10 = 2.302585092994045684
euler = 0.57721566490153286061
euler = libfp.euler
catalan = 0.91596559417721901505
khinchin = 2.6854520010653064453
apery = 1.2020569031595942854
@@ -84,7 +86,7 @@ class FPContext(StandardBaseContext):
return x != x
def isinf(ctx, x):
return abs(x) == math2.INF
return abs(x) == libfp.INF
def isnormal(ctx, x):
if x:
@@ -107,37 +109,33 @@ class FPContext(StandardBaseContext):
except:
return complex(x)
power = staticmethod(math2.pow)
sqrt = staticmethod(math2.sqrt)
exp = staticmethod(math2.exp)
ln = log = staticmethod(math2.log)
cos = staticmethod(math2.cos)
sin = staticmethod(math2.sin)
tan = staticmethod(math2.tan)
cos_sin = staticmethod(math2.cos_sin)
acos = staticmethod(math2.acos)
asin = staticmethod(math2.asin)
atan = staticmethod(math2.atan)
cosh = staticmethod(math2.cosh)
sinh = staticmethod(math2.sinh)
tanh = staticmethod(math2.tanh)
gamma = staticmethod(math2.gamma)
rgamma = staticmethod(math2.rgamma)
fac = factorial = staticmethod(math2.factorial)
floor = staticmethod(math2.floor)
ceil = staticmethod(math2.ceil)
cospi = staticmethod(math2.cospi)
sinpi = staticmethod(math2.sinpi)
cbrt = staticmethod(math2.cbrt)
_nthroot = staticmethod(math2.nthroot)
_ei = staticmethod(math2.ei)
_e1 = staticmethod(math2.e1)
_zeta = _zeta_int = staticmethod(math2.zeta)
# XXX: math2
def arg(ctx, z):
z = complex(z)
return math.atan2(z.imag, z.real)
power = staticmethod(libfp.pow)
sqrt = staticmethod(libfp.sqrt)
exp = staticmethod(libfp.exp)
ln = log = staticmethod(libfp.log)
cos = staticmethod(libfp.cos)
sin = staticmethod(libfp.sin)
tan = staticmethod(libfp.tan)
cos_sin = staticmethod(libfp.cos_sin)
acos = staticmethod(libfp.acos)
asin = staticmethod(libfp.asin)
atan = staticmethod(libfp.atan)
cosh = staticmethod(libfp.cosh)
sinh = staticmethod(libfp.sinh)
tanh = staticmethod(libfp.tanh)
gamma = staticmethod(libfp.gamma)
rgamma = staticmethod(libfp.rgamma)
fac = factorial = staticmethod(libfp.factorial)
floor = staticmethod(libfp.floor)
ceil = staticmethod(libfp.ceil)
cospi = staticmethod(libfp.cospi)
sinpi = staticmethod(libfp.sinpi)
cbrt = staticmethod(libfp.cbrt)
_nthroot = staticmethod(libfp.nthroot)
_ei = staticmethod(libfp.ei)
_e1 = staticmethod(libfp.e1)
_zeta = _zeta_int = staticmethod(libfp.zeta)
arg = staticmethod(cmath.phase)
def expj(ctx, x):
return ctx.exp(ctx.j*x)
@@ -202,8 +200,7 @@ class FPContext(StandardBaseContext):
if k > maxterms:
raise ctx.NoConvergence
def atan2(ctx, x, y):
return math.atan2(x, y)
atan2 = staticmethod(math.atan2)
def psi(ctx, m, z):
m = int(m)
@@ -211,7 +208,7 @@ class FPContext(StandardBaseContext):
return ctx.digamma(z)
return (-1)**(m+1) * ctx.fac(m) * ctx.zeta(m+1, z)
digamma = staticmethod(math2.digamma)
digamma = staticmethod(libfp.digamma)
def harmonic(ctx, x):
x = ctx.convert(x)
@@ -228,8 +225,8 @@ class FPContext(StandardBaseContext):
import random
return random.random()
_erf = staticmethod(math2.erf)
_erfc = staticmethod(math2.erfc)
_erf = staticmethod(math.erf)
_erfc = staticmethod(math.erfc)
def sum_accurately(ctx, terms, check_step=1):
s = ctx.zero
+20 -166
View File
@@ -1,6 +1,6 @@
"""
This module complements the math and cmath builtin modules by providing
fast machine precision versions of some additional functions (gamma, ...)
fast machine precision versions of some additional functions (ei, e1, ...)
and wrapping math/cmath functions so that they can be called with either
real or complex arguments.
"""
@@ -10,26 +10,13 @@ import math
import cmath
# Irrational (?) constants
pi = 3.1415926535897932385
e = 2.7182818284590452354
sqrt2 = 1.4142135623730950488
sqrt5 = 2.2360679774997896964
phi = 1.6180339887498948482
ln2 = 0.69314718055994530942
ln10 = 2.302585092994045684
pi = math.pi
euler = 0.57721566490153286061
catalan = 0.91596559417721901505
khinchin = 2.6854520010653064453
apery = 1.2020569031595942854
logpi = 1.1447298858494001741
def _mathfun_real(f_real, f_complex):
def f(x, **kwargs):
if type(x) is float:
return f_real(x)
if type(x) is complex:
return f_complex(x)
try:
x = float(x)
return f_real(x)
@@ -142,16 +129,6 @@ def _cospi_complex(z):
cospi = _mathfun_real(_cospi_real, _cospi_complex)
sinpi = _mathfun_real(_sinpi_real, _sinpi_complex)
def tanpi(x):
try:
return sinpi(x) / cospi(x)
except OverflowError:
if complex(x).imag > 10:
return 1j
if complex(x).imag < 10:
return -1j
raise
def cotpi(x):
try:
return cospi(x) / sinpi(x)
@@ -162,10 +139,7 @@ def cotpi(x):
return 1j
raise
INF = 1e300*1e300
NINF = -INF
NAN = INF-INF
EPS = 2.2204460492503131e-16
INF = math.inf
_exact_gamma = (INF, 1.0, 1.0, 2.0, 6.0, 24.0, 120.0, 720.0, 5040.0, 40320.0,
362880.0, 3628800.0, 39916800.0, 479001600.0, 6227020800.0, 87178291200.0,
@@ -180,28 +154,9 @@ _lanczos_p = (0.99999999999980993, 676.5203681218851, -1259.1392167224028,
771.32342877765313, -176.61502916214059, 12.507343278686905,
-0.13857109526572012, 9.9843695780195716e-6, 1.5056327351493116e-7)
def _gamma_real(x):
_intx = int(x)
if _intx == x:
if _intx <= 0:
#return (-1)**_intx * INF
raise ZeroDivisionError("gamma function pole")
if _intx <= _max_exact_gamma:
return _exact_gamma[_intx]
if x < 0.5:
# TODO: sinpi
return pi / (_sinpi_real(x)*_gamma_real(1-x))
else:
x -= 1.0
r = _lanczos_p[0]
for i in range(1, _lanczos_g+2):
r += _lanczos_p[i]/(x+i)
t = x + _lanczos_g + 0.5
return 2.506628274631000502417 * t**(x+0.5) * math.exp(-t) * r
def _gamma_complex(x):
if not x.imag:
return complex(_gamma_real(x.real))
return complex(math.gamma(x.real))
if x.real < 0.5:
# TODO: sinpi
return pi / (_sinpi_complex(x)*_gamma_complex(1-x))
@@ -213,7 +168,7 @@ def _gamma_complex(x):
t = x + _lanczos_g + 0.5
return 2.506628274631000502417 * t**(x+0.5) * cmath.exp(-t) * r
gamma = _mathfun_real(_gamma_real, _gamma_complex)
gamma = _mathfun_real(math.gamma, _gamma_complex)
def rgamma(x):
try:
@@ -224,11 +179,6 @@ def rgamma(x):
def factorial(x):
return gamma(x+1.0)
def arg(x):
if type(x) is float:
return math.atan2(0.0,x)
return math.atan2(x.imag,x.real)
# XXX: broken for negatives
def loggamma(x):
if type(x) not in (float, complex):
@@ -331,105 +281,13 @@ def _digamma_complex(x):
digamma = _mathfun_real(_digamma_real, _digamma_complex)
# TODO: could implement complex erf and erfc here. Need
# to find an accurate method (avoiding cancellation)
# for approx. 1 < abs(x) < 9.
_erfc_coeff_P = [
1.0000000161203922312,
2.1275306946297962644,
2.2280433377390253297,
1.4695509105618423961,
0.66275911699770787537,
0.20924776504163751585,
0.045459713768411264339,
0.0063065951710717791934,
0.00044560259661560421715][::-1]
_erfc_coeff_Q = [
1.0000000000000000000,
3.2559100272784894318,
4.9019435608903239131,
4.4971472894498014205,
2.7845640601891186528,
1.2146026030046904138,
0.37647108453729465912,
0.080970149639040548613,
0.011178148899483545902,
0.00078981003831980423513][::-1]
def _polyval(coeffs, x):
p = coeffs[0]
for c in coeffs[1:]:
p = c + x*p
return p
def _erf_taylor(x):
# Taylor series assuming 0 <= x <= 1
x2 = x*x
s = t = x
n = 1
while abs(t) > 1e-17:
t *= x2/n
s -= t/(n+n+1)
n += 1
t *= x2/n
s += t/(n+n+1)
n += 1
return 1.1283791670955125739*s
def _erfc_mid(x):
# Rational approximation assuming 0 <= x <= 9
return exp(-x*x)*_polyval(_erfc_coeff_P,x)/_polyval(_erfc_coeff_Q,x)
def _erfc_asymp(x):
# Asymptotic expansion assuming x >= 9
x2 = x*x
v = exp(-x2)/x*0.56418958354775628695
r = t = 0.5 / x2
s = 1.0
for n in range(1,22,4):
s -= t
t *= r * (n+2)
s += t
t *= r * (n+4)
if abs(t) < 1e-17:
break
return s * v
def erf(x):
"""
erf of a real number.
"""
x = float(x)
if x != x:
return x
if x < 0.0:
return -erf(-x)
if x >= 1.0:
if x >= 6.0:
return 1.0
return 1.0 - _erfc_mid(x)
return _erf_taylor(x)
def erfc(x):
"""
erfc of a real number.
"""
x = float(x)
if x != x:
return x
if x < 0.0:
if x < -6.0:
return 2.0
return 2.0-erfc(-x)
if x > 9.0:
return _erfc_asymp(x)
if x >= 1.0:
return _erfc_mid(x)
return 1.0 - _erf_taylor(x)
gauss42 = [\
gauss42 = [
(0.99839961899006235, 0.0041059986046490839),
(-0.99839961899006235, 0.0041059986046490839),
(0.9915772883408609, 0.009536220301748501),
@@ -523,14 +381,12 @@ def ei_taylor(z, _e1=False):
return s
def ei(z, _e1=False):
typez = type(z)
if typez not in (float, complex):
try:
z = float(z)
typez = float
except (TypeError, ValueError):
z = complex(z)
typez = complex
try:
z = float(z)
typez = float
except (TypeError, ValueError):
z = complex(z)
typez = complex
if not z:
return -INF
absz = abs(z)
@@ -560,22 +416,20 @@ def ei(z, _e1=False):
return ref
def e1(z):
try:
z = float(z)
typez = float
except (TypeError, ValueError):
z = complex(z)
typez = complex
# hack to get consistent signs if the imaginary part if 0
# and signed
typez = type(z)
if type(z) not in (float, complex):
try:
z = float(z)
typez = float
except (TypeError, ValueError):
z = complex(z)
typez = complex
if typez is complex and not z.imag:
z = complex(z.real, 0.0)
# end hack
return -ei(-z, _e1=True)
_zeta_int = [\
_zeta_int = [
-0.5,
0.0,
1.6449340668482264365,1.2020569031595942854,1.0823232337111381915,
@@ -637,7 +491,7 @@ def zeta(s):
if not (n % 2):
return 0.0
if s <= 0.0:
return 2.**s*pi**(s-1)*_sinpi_real(0.5*s)*_gamma_real(1-s)*zeta(1-s)
return 2.**s*pi**(s-1)*_sinpi_real(0.5*s)*math.gamma(1-s)*zeta(1-s)
if s <= 2.0:
if s <= 1.0:
return _polyval(_zeta_0,s)/(s-1)