Skip to content

Commit c1d4a35

Browse files
authored
Merge pull request #3002 from devitocodes/patch-sympy-tensor-args
dsl: misc patches from recent updates (interp, sympy args)
2 parents bb78b3e + 8f12021 commit c1d4a35

6 files changed

Lines changed: 105 additions & 22 deletions

File tree

‎devito/types/basic.py‎

Lines changed: 16 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -4,6 +4,7 @@
44
from contextlib import contextmanager, suppress
55
from ctypes import POINTER, Structure, _Pointer, c_char, c_char_p
66
from functools import cached_property, reduce
7+
from numbers import Number
78
from operator import mul
89

910
import numpy as np
@@ -1539,7 +1540,10 @@ def _new(cls, *args, **kwargs):
15391540
# Filter grid and dimensions
15401541
grid, dimensions = newobj._infer_dims()
15411542
if grid is None and dimensions is None:
1542-
return sympy.ImmutableDenseMatrix(*args)
1543+
# Downgrade to a plain Matrix, reusing the representation rather
1544+
# than rebuilding from `args`, as the latter would sympify the
1545+
# entries with sympy's `sympify` instead of `cls._sympify`
1546+
return sympy.ImmutableDenseMatrix._fromrep(newobj._rep)
15431547
# Initialized with constructed object
15441548
newobj.__init_finalize__(newobj.rows, newobj.cols, newobj.flat(),
15451549
grid=grid, dimensions=dimensions)
@@ -1581,13 +1585,18 @@ def __subfunc_setup__(cls, *args, **kwargs):
15811585
@classmethod
15821586
def _sympify(cls, arg):
15831587
# This is used internally by sympy to process arguments at rebuilt. And since
1584-
# some of our properties are non-sympyfiable we need to have a fallback.
1585-
# `strict` so that strings are left alone rather than parsed into Symbols,
1586-
# while plain numbers are turned into `Expr` as sympy expects (a Matrix
1587-
# holding non-`Expr` entries, such as a plain `int` 0, is deprecated)
1588+
# some of our properties are non-sympyfiable we need to have a fallback
1589+
if isinstance(arg, Number):
1590+
# Plain numbers must be sympified, as sympy assigns the `EXRAW` domain
1591+
# to a Matrix holding non-`Expr` entries such as a plain `int` 0
1592+
return sympy.sympify(arg)
15881593
try:
1589-
return sympy.sympify(arg, strict=True)
1590-
except sympy.SympifyError:
1594+
# Pure sympy object
1595+
return arg._sympy_()
1596+
except AttributeError:
1597+
# Anything else, such as a `Staggering`, is passed through untouched.
1598+
# Note that sympifying is not an option here, as it would convert
1599+
# away the type, `Staggering` being a `tuple` for example
15911600
return arg
15921601

15931602
@classmethod

‎devito/types/sparse.py‎

Lines changed: 9 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -1008,20 +1008,24 @@ def _arg_defaults(self, alias=None, estimate_memory=False):
10081008
if estimate_memory:
10091009
return defaults
10101010
key = alias or self
1011-
coords = defaults.get(key.coordinates.name, key.coordinates.data)
1011+
coords = defaults.get(key.coordinates.name, self.coordinates.data)
10121012
defaults.update(key.interpolator._arg_defaults(coords=coords,
1013-
sfunc=key))
1013+
sfunc=self))
10141014
return defaults
10151015

10161016
def _arg_values(self, estimate_memory=False, **kwargs):
10171017
values = super()._arg_values(estimate_memory=estimate_memory, **kwargs)
10181018
if estimate_memory:
10191019
return values
10201020

1021-
# Resolve the runtime grid origin (honours `o_x`/`o_y`/... overrides)
1022-
# and hand it to the interpolator so tables reflect the actual frame
1023-
# of reference used by the kernel.
1021+
# `super` has already tabulated through `_arg_defaults`, in the frame
1022+
# of whichever object supplied the runtime values. Only an explicit
1023+
# `o_x`/`o_y`/... override moves that frame again, and the tables then
1024+
# have to be rebuilt against it.
10241025
onames = [o.name for o in self.grid.origin_symbols]
1026+
if not any(n in kwargs for n in onames):
1027+
return values
1028+
10251029
origin = tuple(kwargs.get(n, o) for n, o in
10261030
zip(onames, self.grid.origin, strict=True))
10271031
coords = values.get(self.coordinates.name, self.coordinates.data)

‎devito/types/tensor.py‎

Lines changed: 6 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -24,12 +24,15 @@ def staggering(stagg, i, j, d, dims):
2424
if stagg is None:
2525
# No input
2626
return NODE if i == j else (d, dims[j])
27+
elif isinstance(stagg, MatrixBase):
28+
# From rebuild/tensor property. Indexed as a sympy Matrix. Note that this
29+
# may be a plain Matrix rather than an AbstractTensor, as rebuilding a
30+
# tensor component-wise downgrades it when the components aren't Devito
31+
# objects, which is the case for a Matrix of `Staggering`
32+
return stagg[i, j]
2733
elif isinstance(stagg, (tuple, list)):
2834
# User input as list or tuple
2935
return stagg[i][j]
30-
elif isinstance(stagg, AbstractTensor):
31-
# From rebuild/tensor property. Indexed as a sympy Matrix
32-
return stagg[i, j]
3336

3437

3538
class TensorFunction(AbstractTensor):

‎examples/seismic/tti/operators.py‎

Lines changed: 23 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -62,7 +62,7 @@ def trig_func(model):
6262
return costheta, sintheta
6363

6464

65-
def Gzz_centered(model, field):
65+
def Gzz_centered(model, field, b=None):
6666
"""
6767
3D rotated second order derivative in the direction z.
6868
@@ -72,12 +72,17 @@ def Gzz_centered(model, field):
7272
Physical parameters model structure.
7373
field : Function
7474
Input for which the derivative is computed.
75+
b : Function, optional
76+
Buoyancy to build the operator with, defaulting to the model's. Since
77+
the operator is linear in it, passing a perturbation here gives the
78+
derivative of the operator with respect to the buoyancy in that
79+
direction.
7580
7681
Returns
7782
-------
7883
Rotated second order derivative w.r.t. z.
7984
"""
80-
b = getattr(model, 'b', 1)
85+
b = getattr(model, 'b', 1) if b is None else b
8186
costheta, sintheta, cosphi, sinphi = trig_func(model)
8287

8388
order1 = field.space_order // 2
@@ -99,7 +104,7 @@ def Gzz_centered(model, field):
99104
return Gzz
100105

101106

102-
def Gzz_centered_2d(model, field):
107+
def Gzz_centered_2d(model, field, b=None):
103108
"""
104109
2D rotated second order derivative in the direction z.
105110
@@ -109,12 +114,17 @@ def Gzz_centered_2d(model, field):
109114
Physical parameters model structure.
110115
field : Function
111116
Input for which the derivative is computed.
117+
b : Function, optional
118+
Buoyancy to build the operator with, defaulting to the model's. Since
119+
the operator is linear in it, passing a perturbation here gives the
120+
derivative of the operator with respect to the buoyancy in that
121+
direction.
112122
113123
Returns
114124
-------
115125
Rotated second order derivative w.r.t. z.
116126
"""
117-
b = getattr(model, 'b', 1)
127+
b = getattr(model, 'b', 1) if b is None else b
118128
costheta, sintheta = trig_func(model)
119129

120130
order1 = field.space_order // 2
@@ -133,7 +143,7 @@ def Gzz_centered_2d(model, field):
133143

134144

135145
# Centered case produces directly Gxx + Gyy
136-
def Gh_centered(model, field):
146+
def Gh_centered(model, field, b=None):
137147
"""
138148
Sum of the 3D rotated second order derivative in the direction x and y.
139149
As the Laplacian is rotation invariant, it is computed as the conventional
@@ -146,13 +156,19 @@ def Gh_centered(model, field):
146156
Physical parameters model structure.
147157
field : Function
148158
Input field.
159+
b : Function, optional
160+
Buoyancy to build the operator with, defaulting to the model's. See
161+
:func:`Gzz_centered`.
149162
150163
Returns
151164
-------
152165
Sum of the 3D rotated second order derivative in the direction x and y.
153166
"""
154-
Gzz = Gzz_centered(model, field) if model.dim == 3 else Gzz_centered_2d(model, field)
155-
b = getattr(model, 'b', None)
167+
b = getattr(model, 'b', None) if b is None else b
168+
if model.dim == 3: # noqa: SIM108
169+
Gzz = Gzz_centered(model, field, b=b)
170+
else:
171+
Gzz = Gzz_centered_2d(model, field, b=b)
156172
if b is not None:
157173
_diff = lambda f, d: getattr(f, f'd{d.name}')
158174
so = field.space_order // 2

‎tests/test_interpolation.py‎

Lines changed: 31 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1077,6 +1077,37 @@ def test_position(self, shape):
10771077

10781078
assert(np.allclose(rec.data, rec1.data, atol=1e-5))
10791079

1080+
@pytest.mark.parametrize('interpolation,r', [('linear', 1), ('sinc', 4)])
1081+
def test_position_override_grid(self, interpolation, r):
1082+
"""
1083+
Inject through an Operator built on a grid whose origin differs from
1084+
the one it is applied to, as when an Operator compiled against one
1085+
model is applied to another. The point must land where the runtime
1086+
origin puts it, not the compile-time one.
1087+
"""
1088+
shape, spacing, coord = (41, 41), (10., 10.), 120.
1089+
extent = tuple((s - 1) * h for s, h in zip(shape, spacing, strict=True))
1090+
kw = dict(interpolation=interpolation, r=r)
1091+
1092+
def setup(origin):
1093+
grid = Grid(shape=shape, extent=extent, origin=origin)
1094+
u = TimeFunction(name='u', grid=grid, space_order=8)
1095+
src = SparseTimeFunction(name='src', grid=grid, npoint=1, nt=2, **kw)
1096+
src.coordinates.data[0, :] = coord
1097+
src.data[:] = 1.
1098+
return u, src
1099+
1100+
u_build, src_build = setup((0., 0.))
1101+
op = Operator(src_build.inject(field=u_build.forward, expr=src_build))
1102+
1103+
shift = -100.
1104+
u, src = setup((shift, shift))
1105+
op.apply(time_M=0, u=u, src=src)
1106+
1107+
expected = tuple(int((coord - shift) / h) for h in spacing)
1108+
peak = np.unravel_index(np.argmax(np.abs(u.data)), u.data.shape)[1:]
1109+
assert peak == expected
1110+
10801111
def test_sparse_first(self):
10811112
"""
10821113
Tests custom sprase function with sparse dimension as first index.

‎tests/test_tensors.py‎

Lines changed: 20 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -12,6 +12,7 @@
1212
)
1313
from devito.symbolics import retrieve_derivatives
1414
from devito.types import NODE
15+
from devito.types.utils import Staggering
1516

1617

1718
def dimify(dimensions):
@@ -528,6 +529,25 @@ def test_diag_sympified_zeros(func1):
528529
assert all(isinstance(c, sympy.Expr) for c in f2.flat())
529530

530531

532+
@pytest.mark.parametrize('func1', [TensorFunction, TensorTimeFunction,
533+
VectorFunction, VectorTimeFunction])
534+
def test_staggered_attribute_roundtrip(func1):
535+
"""
536+
Accessing an attribute rebuilds the tensor component-wise, which must not
537+
sympify a `Staggering` away, otherwise it can no longer be fed back as the
538+
`staggered` kwarg.
539+
"""
540+
grid = Grid(tuple([5]*3))
541+
f1 = func1(name="f1", grid=grid, time_order=1)
542+
543+
stagg = f1.staggered
544+
assert all(isinstance(s, Staggering) for s in stagg.flat())
545+
546+
f2 = func1(name="f2", grid=grid, time_order=1, staggered=stagg)
547+
assert all(c1.staggered == c2.staggered
548+
for c1, c2 in zip(f1.flat(), f2.flat(), strict=True))
549+
550+
531551
def test_non_expr_components():
532552
"""
533553
A tensor may legitimately hold non-`Expr` components, which sympy deprecates

0 commit comments

Comments
 (0)