图片解析应用
You can not select more than 25 topics Topics must start with a letter or number, can include dashes ('-') and can be up to 35 characters long.

394 lines
14 KiB

  1. from sympy.codegen import Assignment
  2. from sympy.codegen.ast import none
  3. from sympy.codegen.cfunctions import expm1, log1p
  4. from sympy.codegen.scipy_nodes import cosm1
  5. from sympy.codegen.matrix_nodes import MatrixSolve
  6. from sympy.core import Expr, Mod, symbols, Eq, Le, Gt, zoo, oo, Rational, Pow
  7. from sympy.core.numbers import pi
  8. from sympy.core.singleton import S
  9. from sympy.functions import acos, KroneckerDelta, Piecewise, sign, sqrt
  10. from sympy.logic import And, Or
  11. from sympy.matrices import SparseMatrix, MatrixSymbol, Identity
  12. from sympy.printing.pycode import (
  13. MpmathPrinter, PythonCodePrinter, pycode, SymPyPrinter
  14. )
  15. from sympy.printing.numpy import NumPyPrinter, SciPyPrinter
  16. from sympy.testing.pytest import raises, skip
  17. from sympy.tensor import IndexedBase
  18. from sympy.external import import_module
  19. from sympy.functions.special.gamma_functions import loggamma
  20. from sympy.parsing.latex import parse_latex
  21. x, y, z = symbols('x y z')
  22. p = IndexedBase("p")
  23. def test_PythonCodePrinter():
  24. prntr = PythonCodePrinter()
  25. assert not prntr.module_imports
  26. assert prntr.doprint(x**y) == 'x**y'
  27. assert prntr.doprint(Mod(x, 2)) == 'x % 2'
  28. assert prntr.doprint(-Mod(x, y)) == '-(x % y)'
  29. assert prntr.doprint(Mod(-x, y)) == '(-x) % y'
  30. assert prntr.doprint(And(x, y)) == 'x and y'
  31. assert prntr.doprint(Or(x, y)) == 'x or y'
  32. assert not prntr.module_imports
  33. assert prntr.doprint(pi) == 'math.pi'
  34. assert prntr.module_imports == {'math': {'pi'}}
  35. assert prntr.doprint(x**Rational(1, 2)) == 'math.sqrt(x)'
  36. assert prntr.doprint(sqrt(x)) == 'math.sqrt(x)'
  37. assert prntr.module_imports == {'math': {'pi', 'sqrt'}}
  38. assert prntr.doprint(acos(x)) == 'math.acos(x)'
  39. assert prntr.doprint(Assignment(x, 2)) == 'x = 2'
  40. assert prntr.doprint(Piecewise((1, Eq(x, 0)),
  41. (2, x>6))) == '((1) if (x == 0) else (2) if (x > 6) else None)'
  42. assert prntr.doprint(Piecewise((2, Le(x, 0)),
  43. (3, Gt(x, 0)), evaluate=False)) == '((2) if (x <= 0) else'\
  44. ' (3) if (x > 0) else None)'
  45. assert prntr.doprint(sign(x)) == '(0.0 if x == 0 else math.copysign(1, x))'
  46. assert prntr.doprint(p[0, 1]) == 'p[0, 1]'
  47. assert prntr.doprint(KroneckerDelta(x,y)) == '(1 if x == y else 0)'
  48. assert prntr.doprint((2,3)) == "(2, 3)"
  49. assert prntr.doprint([2,3]) == "[2, 3]"
  50. def test_PythonCodePrinter_standard():
  51. prntr = PythonCodePrinter()
  52. assert prntr.standard == 'python3'
  53. raises(ValueError, lambda: PythonCodePrinter({'standard':'python4'}))
  54. def test_MpmathPrinter():
  55. p = MpmathPrinter()
  56. assert p.doprint(sign(x)) == 'mpmath.sign(x)'
  57. assert p.doprint(Rational(1, 2)) == 'mpmath.mpf(1)/mpmath.mpf(2)'
  58. assert p.doprint(S.Exp1) == 'mpmath.e'
  59. assert p.doprint(S.Pi) == 'mpmath.pi'
  60. assert p.doprint(S.GoldenRatio) == 'mpmath.phi'
  61. assert p.doprint(S.EulerGamma) == 'mpmath.euler'
  62. assert p.doprint(S.NaN) == 'mpmath.nan'
  63. assert p.doprint(S.Infinity) == 'mpmath.inf'
  64. assert p.doprint(S.NegativeInfinity) == 'mpmath.ninf'
  65. assert p.doprint(loggamma(x)) == 'mpmath.loggamma(x)'
  66. def test_NumPyPrinter():
  67. from sympy.core.function import Lambda
  68. from sympy.matrices.expressions.adjoint import Adjoint
  69. from sympy.matrices.expressions.diagonal import (DiagMatrix, DiagonalMatrix, DiagonalOf)
  70. from sympy.matrices.expressions.funcmatrix import FunctionMatrix
  71. from sympy.matrices.expressions.hadamard import HadamardProduct
  72. from sympy.matrices.expressions.kronecker import KroneckerProduct
  73. from sympy.matrices.expressions.special import (OneMatrix, ZeroMatrix)
  74. from sympy.abc import a, b
  75. p = NumPyPrinter()
  76. assert p.doprint(sign(x)) == 'numpy.sign(x)'
  77. A = MatrixSymbol("A", 2, 2)
  78. B = MatrixSymbol("B", 2, 2)
  79. C = MatrixSymbol("C", 1, 5)
  80. D = MatrixSymbol("D", 3, 4)
  81. assert p.doprint(A**(-1)) == "numpy.linalg.inv(A)"
  82. assert p.doprint(A**5) == "numpy.linalg.matrix_power(A, 5)"
  83. assert p.doprint(Identity(3)) == "numpy.eye(3)"
  84. u = MatrixSymbol('x', 2, 1)
  85. v = MatrixSymbol('y', 2, 1)
  86. assert p.doprint(MatrixSolve(A, u)) == 'numpy.linalg.solve(A, x)'
  87. assert p.doprint(MatrixSolve(A, u) + v) == 'numpy.linalg.solve(A, x) + y'
  88. assert p.doprint(ZeroMatrix(2, 3)) == "numpy.zeros((2, 3))"
  89. assert p.doprint(OneMatrix(2, 3)) == "numpy.ones((2, 3))"
  90. assert p.doprint(FunctionMatrix(4, 5, Lambda((a, b), a + b))) == \
  91. "numpy.fromfunction(lambda a, b: a + b, (4, 5))"
  92. assert p.doprint(HadamardProduct(A, B)) == "numpy.multiply(A, B)"
  93. assert p.doprint(KroneckerProduct(A, B)) == "numpy.kron(A, B)"
  94. assert p.doprint(Adjoint(A)) == "numpy.conjugate(numpy.transpose(A))"
  95. assert p.doprint(DiagonalOf(A)) == "numpy.reshape(numpy.diag(A), (-1, 1))"
  96. assert p.doprint(DiagMatrix(C)) == "numpy.diagflat(C)"
  97. assert p.doprint(DiagonalMatrix(D)) == "numpy.multiply(D, numpy.eye(3, 4))"
  98. # Workaround for numpy negative integer power errors
  99. assert p.doprint(x**-1) == 'x**(-1.0)'
  100. assert p.doprint(x**-2) == 'x**(-2.0)'
  101. expr = Pow(2, -1, evaluate=False)
  102. assert p.doprint(expr) == "2**(-1.0)"
  103. assert p.doprint(S.Exp1) == 'numpy.e'
  104. assert p.doprint(S.Pi) == 'numpy.pi'
  105. assert p.doprint(S.EulerGamma) == 'numpy.euler_gamma'
  106. assert p.doprint(S.NaN) == 'numpy.nan'
  107. assert p.doprint(S.Infinity) == 'numpy.PINF'
  108. assert p.doprint(S.NegativeInfinity) == 'numpy.NINF'
  109. def test_issue_18770():
  110. numpy = import_module('numpy')
  111. if not numpy:
  112. skip("numpy not installed.")
  113. from sympy.functions.elementary.miscellaneous import (Max, Min)
  114. from sympy.utilities.lambdify import lambdify
  115. expr1 = Min(0.1*x + 3, x + 1, 0.5*x + 1)
  116. func = lambdify(x, expr1, "numpy")
  117. assert (func(numpy.linspace(0, 3, 3)) == [1.0, 1.75, 2.5 ]).all()
  118. assert func(4) == 3
  119. expr1 = Max(x**2, x**3)
  120. func = lambdify(x,expr1, "numpy")
  121. assert (func(numpy.linspace(-1, 2, 4)) == [1, 0, 1, 8] ).all()
  122. assert func(4) == 64
  123. def test_SciPyPrinter():
  124. p = SciPyPrinter()
  125. expr = acos(x)
  126. assert 'numpy' not in p.module_imports
  127. assert p.doprint(expr) == 'numpy.arccos(x)'
  128. assert 'numpy' in p.module_imports
  129. assert not any(m.startswith('scipy') for m in p.module_imports)
  130. smat = SparseMatrix(2, 5, {(0, 1): 3})
  131. assert p.doprint(smat) == \
  132. 'scipy.sparse.coo_matrix(([3], ([0], [1])), shape=(2, 5))'
  133. assert 'scipy.sparse' in p.module_imports
  134. assert p.doprint(S.GoldenRatio) == 'scipy.constants.golden_ratio'
  135. assert p.doprint(S.Pi) == 'scipy.constants.pi'
  136. assert p.doprint(S.Exp1) == 'numpy.e'
  137. def test_pycode_reserved_words():
  138. s1, s2 = symbols('if else')
  139. raises(ValueError, lambda: pycode(s1 + s2, error_on_reserved=True))
  140. py_str = pycode(s1 + s2)
  141. assert py_str in ('else_ + if_', 'if_ + else_')
  142. def test_issue_20762():
  143. antlr4 = import_module("antlr4")
  144. if not antlr4:
  145. skip('antlr not installed.')
  146. # Make sure pycode removes curly braces from subscripted variables
  147. expr = parse_latex(r'a_b \cdot b')
  148. assert pycode(expr) == 'a_b*b'
  149. expr = parse_latex(r'a_{11} \cdot b')
  150. assert pycode(expr) == 'a_11*b'
  151. def test_sqrt():
  152. prntr = PythonCodePrinter()
  153. assert prntr._print_Pow(sqrt(x), rational=False) == 'math.sqrt(x)'
  154. assert prntr._print_Pow(1/sqrt(x), rational=False) == '1/math.sqrt(x)'
  155. prntr = PythonCodePrinter({'standard' : 'python3'})
  156. assert prntr._print_Pow(sqrt(x), rational=True) == 'x**(1/2)'
  157. assert prntr._print_Pow(1/sqrt(x), rational=True) == 'x**(-1/2)'
  158. prntr = MpmathPrinter()
  159. assert prntr._print_Pow(sqrt(x), rational=False) == 'mpmath.sqrt(x)'
  160. assert prntr._print_Pow(sqrt(x), rational=True) == \
  161. "x**(mpmath.mpf(1)/mpmath.mpf(2))"
  162. prntr = NumPyPrinter()
  163. assert prntr._print_Pow(sqrt(x), rational=False) == 'numpy.sqrt(x)'
  164. assert prntr._print_Pow(sqrt(x), rational=True) == 'x**(1/2)'
  165. prntr = SciPyPrinter()
  166. assert prntr._print_Pow(sqrt(x), rational=False) == 'numpy.sqrt(x)'
  167. assert prntr._print_Pow(sqrt(x), rational=True) == 'x**(1/2)'
  168. prntr = SymPyPrinter()
  169. assert prntr._print_Pow(sqrt(x), rational=False) == 'sympy.sqrt(x)'
  170. assert prntr._print_Pow(sqrt(x), rational=True) == 'x**(1/2)'
  171. def test_frac():
  172. from sympy.functions.elementary.integers import frac
  173. expr = frac(x)
  174. prntr = NumPyPrinter()
  175. assert prntr.doprint(expr) == 'numpy.mod(x, 1)'
  176. prntr = SciPyPrinter()
  177. assert prntr.doprint(expr) == 'numpy.mod(x, 1)'
  178. prntr = PythonCodePrinter()
  179. assert prntr.doprint(expr) == 'x % 1'
  180. prntr = MpmathPrinter()
  181. assert prntr.doprint(expr) == 'mpmath.frac(x)'
  182. prntr = SymPyPrinter()
  183. assert prntr.doprint(expr) == 'sympy.functions.elementary.integers.frac(x)'
  184. class CustomPrintedObject(Expr):
  185. def _numpycode(self, printer):
  186. return 'numpy'
  187. def _mpmathcode(self, printer):
  188. return 'mpmath'
  189. def test_printmethod():
  190. obj = CustomPrintedObject()
  191. assert NumPyPrinter().doprint(obj) == 'numpy'
  192. assert MpmathPrinter().doprint(obj) == 'mpmath'
  193. def test_codegen_ast_nodes():
  194. assert pycode(none) == 'None'
  195. def test_issue_14283():
  196. prntr = PythonCodePrinter()
  197. assert prntr.doprint(zoo) == "math.nan"
  198. assert prntr.doprint(-oo) == "float('-inf')"
  199. def test_NumPyPrinter_print_seq():
  200. n = NumPyPrinter()
  201. assert n._print_seq(range(2)) == '(0, 1,)'
  202. def test_issue_16535_16536():
  203. from sympy.functions.special.gamma_functions import (lowergamma, uppergamma)
  204. a = symbols('a')
  205. expr1 = lowergamma(a, x)
  206. expr2 = uppergamma(a, x)
  207. prntr = SciPyPrinter()
  208. assert prntr.doprint(expr1) == 'scipy.special.gamma(a)*scipy.special.gammainc(a, x)'
  209. assert prntr.doprint(expr2) == 'scipy.special.gamma(a)*scipy.special.gammaincc(a, x)'
  210. prntr = NumPyPrinter()
  211. assert "Not supported" in prntr.doprint(expr1)
  212. assert "Not supported" in prntr.doprint(expr2)
  213. prntr = PythonCodePrinter()
  214. assert "Not supported" in prntr.doprint(expr1)
  215. assert "Not supported" in prntr.doprint(expr2)
  216. def test_Integral():
  217. from sympy.functions.elementary.exponential import exp
  218. from sympy.integrals.integrals import Integral
  219. single = Integral(exp(-x), (x, 0, oo))
  220. double = Integral(x**2*exp(x*y), (x, -z, z), (y, 0, z))
  221. indefinite = Integral(x**2, x)
  222. evaluateat = Integral(x**2, (x, 1))
  223. prntr = SciPyPrinter()
  224. assert prntr.doprint(single) == 'scipy.integrate.quad(lambda x: numpy.exp(-x), 0, numpy.PINF)[0]'
  225. assert prntr.doprint(double) == 'scipy.integrate.nquad(lambda x, y: x**2*numpy.exp(x*y), ((-z, z), (0, z)))[0]'
  226. raises(NotImplementedError, lambda: prntr.doprint(indefinite))
  227. raises(NotImplementedError, lambda: prntr.doprint(evaluateat))
  228. prntr = MpmathPrinter()
  229. assert prntr.doprint(single) == 'mpmath.quad(lambda x: mpmath.exp(-x), (0, mpmath.inf))'
  230. assert prntr.doprint(double) == 'mpmath.quad(lambda x, y: x**2*mpmath.exp(x*y), (-z, z), (0, z))'
  231. raises(NotImplementedError, lambda: prntr.doprint(indefinite))
  232. raises(NotImplementedError, lambda: prntr.doprint(evaluateat))
  233. def test_fresnel_integrals():
  234. from sympy.functions.special.error_functions import (fresnelc, fresnels)
  235. expr1 = fresnelc(x)
  236. expr2 = fresnels(x)
  237. prntr = SciPyPrinter()
  238. assert prntr.doprint(expr1) == 'scipy.special.fresnel(x)[1]'
  239. assert prntr.doprint(expr2) == 'scipy.special.fresnel(x)[0]'
  240. prntr = NumPyPrinter()
  241. assert "Not supported" in prntr.doprint(expr1)
  242. assert "Not supported" in prntr.doprint(expr2)
  243. prntr = PythonCodePrinter()
  244. assert "Not supported" in prntr.doprint(expr1)
  245. assert "Not supported" in prntr.doprint(expr2)
  246. prntr = MpmathPrinter()
  247. assert prntr.doprint(expr1) == 'mpmath.fresnelc(x)'
  248. assert prntr.doprint(expr2) == 'mpmath.fresnels(x)'
  249. def test_beta():
  250. from sympy.functions.special.beta_functions import beta
  251. expr = beta(x, y)
  252. prntr = SciPyPrinter()
  253. assert prntr.doprint(expr) == 'scipy.special.beta(x, y)'
  254. prntr = NumPyPrinter()
  255. assert prntr.doprint(expr) == 'math.gamma(x)*math.gamma(y)/math.gamma(x + y)'
  256. prntr = PythonCodePrinter()
  257. assert prntr.doprint(expr) == 'math.gamma(x)*math.gamma(y)/math.gamma(x + y)'
  258. prntr = PythonCodePrinter({'allow_unknown_functions': True})
  259. assert prntr.doprint(expr) == 'math.gamma(x)*math.gamma(y)/math.gamma(x + y)'
  260. prntr = MpmathPrinter()
  261. assert prntr.doprint(expr) == 'mpmath.beta(x, y)'
  262. def test_airy():
  263. from sympy.functions.special.bessel import (airyai, airybi)
  264. expr1 = airyai(x)
  265. expr2 = airybi(x)
  266. prntr = SciPyPrinter()
  267. assert prntr.doprint(expr1) == 'scipy.special.airy(x)[0]'
  268. assert prntr.doprint(expr2) == 'scipy.special.airy(x)[2]'
  269. prntr = NumPyPrinter()
  270. assert "Not supported" in prntr.doprint(expr1)
  271. assert "Not supported" in prntr.doprint(expr2)
  272. prntr = PythonCodePrinter()
  273. assert "Not supported" in prntr.doprint(expr1)
  274. assert "Not supported" in prntr.doprint(expr2)
  275. def test_airy_prime():
  276. from sympy.functions.special.bessel import (airyaiprime, airybiprime)
  277. expr1 = airyaiprime(x)
  278. expr2 = airybiprime(x)
  279. prntr = SciPyPrinter()
  280. assert prntr.doprint(expr1) == 'scipy.special.airy(x)[1]'
  281. assert prntr.doprint(expr2) == 'scipy.special.airy(x)[3]'
  282. prntr = NumPyPrinter()
  283. assert "Not supported" in prntr.doprint(expr1)
  284. assert "Not supported" in prntr.doprint(expr2)
  285. prntr = PythonCodePrinter()
  286. assert "Not supported" in prntr.doprint(expr1)
  287. assert "Not supported" in prntr.doprint(expr2)
  288. def test_numerical_accuracy_functions():
  289. prntr = SciPyPrinter()
  290. assert prntr.doprint(expm1(x)) == 'numpy.expm1(x)'
  291. assert prntr.doprint(log1p(x)) == 'numpy.log1p(x)'
  292. assert prntr.doprint(cosm1(x)) == 'scipy.special.cosm1(x)'