Skip to content

Commit 6209d52

Browse files
committed
sympy: remove tensor mul now fixed in sympy
1 parent 9adbe93 commit 6209d52

5 files changed

Lines changed: 10 additions & 53 deletions

File tree

devito/finite_differences/derivative.py

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -6,7 +6,7 @@
66
from devito.finite_differences.finite_difference import (generic_derivative,
77
first_derivative,
88
cross_derivative)
9-
from devito.finite_differences.differentiable import (Differentiable, EvalDiffDerivative,
9+
from devito.finite_differences.differentiable import (Differentiable, EvalDerivative,
1010
diffify)
1111
from devito.finite_differences.tools import direct, transpose
1212
from devito.tools import as_mapper, as_tuple, filter_ordered, frozendict
@@ -334,7 +334,7 @@ def _eval_fd(self, expr):
334334
- 3: Evaluate remaining terms (as `g` may need to be evaluated
335335
at a different point).
336336
- 4: Apply substitutions.
337-
- 5: Cast to an object of type `EvalDiffDerivative` so that we know
337+
- 5: Cast to an object of type `EvalDerivative` so that we know
338338
the argument stems from a `Derivative. This may be useful for
339339
later compilation passes.
340340
"""
@@ -362,6 +362,6 @@ def _eval_fd(self, expr):
362362

363363
# Step 5: Cast to EvaluatedDerivative
364364
assert res.is_Add
365-
res = EvalDiffDerivative(*res.args, evaluate=False)
365+
res = EvalDerivative(*res.args, evaluate=False)
366366

367367
return res

devito/finite_differences/differentiable.py

Lines changed: 4 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -10,7 +10,7 @@
1010
from devito.finite_differences.tools import make_shift_x0
1111
from devito.logger import warning
1212
from devito.tools import filter_ordered, flatten
13-
from devito.types.lazy import Evaluable, EvalDerivative
13+
from devito.types.lazy import Evaluable
1414
from devito.types.utils import DimensionTuple
1515

1616
__all__ = ['Differentiable']
@@ -405,8 +405,8 @@ class Mod(DifferentiableOp, sympy.Mod):
405405
__new__ = DifferentiableOp.__new__
406406

407407

408-
class EvalDiffDerivative(DifferentiableOp, EvalDerivative):
409-
__sympy_class__ = EvalDerivative
408+
class EvalDerivative(DifferentiableOp, sympy.Add):
409+
__sympy_class__ = sympy.Add
410410
__new__ = DifferentiableOp.__new__
411411

412412

@@ -466,7 +466,7 @@ def _(obj):
466466
@_cls.register(Mul)
467467
@_cls.register(Pow)
468468
@_cls.register(Mod)
469-
@_cls.register(EvalDiffDerivative)
469+
@_cls.register(EvalDerivative)
470470
def _(obj):
471471
return obj.__class__
472472

devito/types/lazy.py

Lines changed: 0 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,3 @@
1-
import sympy
2-
31
__all__ = ['Evaluable']
42

53

@@ -49,15 +47,3 @@ def evaluate(self):
4947
args = self._evaluate_args()
5048
evaluate = not all(i is j for i, j in zip(args, self.args))
5149
return self.func(*args, evaluate=evaluate)
52-
53-
54-
# Custom SymPy types used upon evaluation
55-
56-
57-
class EvalDerivative(sympy.Add):
58-
59-
"""
60-
A sympy.Add representing an expanded finite-difference Derivative.
61-
"""
62-
63-
pass

devito/types/tensor.py

Lines changed: 0 additions & 29 deletions
Original file line numberDiff line numberDiff line change
@@ -2,8 +2,6 @@
22
from cached_property import cached_property
33

44
import numpy as np
5-
from sympy import Dummy
6-
from sympy.core.decorators import call_highest_priority
75
from sympy.core.sympify import converter as sympify_converter
86

97
from devito.finite_differences import Differentiable
@@ -123,33 +121,6 @@ def __subfunc_setup__(cls, *args, **kwargs):
123121
funcs = funcs.tolist()
124122
return funcs
125123

126-
@call_highest_priority('__mul__')
127-
def __rmul__(self, other):
128-
"""
129-
Prevents Functions to be interpreted as matrices in 2D becaue of `.shape`
130-
TODO: Remove after sympy 1.7
131-
"""
132-
if getattr(other, 'is_DiscreteFunction', False):
133-
tmp_func = Dummy(other.name)
134-
simplemul = super(TensorFunction, self).__rmul__(tmp_func)
135-
return simplemul.subs(tmp_func, other)
136-
# honest sympy matrices defer to their class's routine
137-
if getattr(other, 'is_Matrix', False):
138-
return self._eval_matrix_rmul(other)
139-
return super(TensorFunction, self).__rmul__(other)
140-
141-
@call_highest_priority('__rmul__')
142-
def __mul__(self, other):
143-
"""
144-
Prevents Functions to be interpreted as matrices in 2D becaue of `.shape`
145-
TODO: Remove after sympy 1.7
146-
"""
147-
if getattr(other, 'is_DiscreteFunction', False):
148-
tmp_func = Dummy(other.name)
149-
simplemul = super(TensorFunction, self).__mul__(tmp_func)
150-
return simplemul.subs(tmp_func, other)
151-
return super(TensorFunction, self).__mul__(other)
152-
153124
def __getattr__(self, name):
154125
"""
155126
Try calling a dynamically created FD shortcut.

tests/test_derivatives.py

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -5,7 +5,7 @@
55
from devito import (Grid, Function, TimeFunction, Eq, Operator, NODE,
66
ConditionalDimension, left, right, centered, div, grad)
77
from devito.finite_differences import Derivative, Differentiable
8-
from devito.finite_differences.differentiable import EvalDiffDerivative
8+
from devito.finite_differences.differentiable import EvalDerivative
99
from devito.symbolics import indexify, retrieve_indexed
1010

1111
_PRECISION = 9
@@ -195,7 +195,7 @@ def test_derivatives_space(self, derivative, dim, order):
195195

196196
s_expr = u.diff(dim).as_finite_difference(indices).evalf(_PRECISION)
197197
assert(simplify(expr - s_expr) == 0) # Symbolic equality
198-
assert type(expr) == EvalDiffDerivative
198+
assert type(expr) == EvalDerivative
199199
expr1 = s_expr.func(*expr.args)
200200
assert(expr1 == s_expr) # Exact equality
201201

@@ -215,7 +215,7 @@ def test_second_derivatives_space(self, derivative, dim, order):
215215
indices = [(dim + i * dim.spacing) for i in range(-width, width + 1)]
216216
s_expr = u.diff(dim, dim).as_finite_difference(indices).evalf(_PRECISION)
217217
assert(simplify(expr - s_expr) == 0) # Symbolic equality
218-
assert type(expr) == EvalDiffDerivative
218+
assert type(expr) == EvalDerivative
219219
expr1 = s_expr.func(*expr.args)
220220
assert(expr1 == s_expr) # Exact equality
221221

0 commit comments

Comments
 (0)