Skip to content

Commit cb8288c

Browse files
authored
Merge pull request #3016 from devitocodes/fix-fd-gather-priority
Fix fd gather priority
2 parents e875143 + 07a38b0 commit cb8288c

2 files changed

Lines changed: 46 additions & 3 deletions

File tree

‎devito/finite_differences/differentiable.py‎

Lines changed: 17 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -22,7 +22,7 @@
2222
from devito.logger import warning
2323
from devito.tools import (
2424
as_tuple, extract_dtype, filter_ordered, flatten, frozendict, infer_dtype, is_integer,
25-
is_number, split
25+
is_number, memoized_func, split
2626
)
2727
from devito.types import Array, DimensionTuple, Evaluable, StencilDimension
2828
from devito.types.basic import AbstractFunction, Indexed
@@ -489,6 +489,19 @@ def has_free(self, *patterns):
489489
return all(i in self.free_symbols for i in patterns)
490490

491491

492+
@memoized_func(scope='build')
493+
def deep_priority(expr):
494+
"""
495+
The highest `_fd_priority` among the Functions inside `expr`.
496+
497+
`expr._fd_priority` does not give this: an `Add` or `Mul` falls back on a
498+
generic value, so `mu*tau_xx` reports .75 rather than `tau_xx`'s 2.1.
499+
"""
500+
prio = getattr(expr, '_fd_priority', 0)
501+
return max([prio] + [deep_priority(i)
502+
for i in getattr(expr, '_args_diff', ())])
503+
504+
492505
def highest_priority(diff_op, candidates=None):
493506
"""
494507
The Function whose location a product should be evaluated at.
@@ -504,7 +517,9 @@ def highest_priority(diff_op, candidates=None):
504517
# We also need to make sure that the object with the largest
505518
# set of dimensions is used when multiple ones with the same
506519
# priority appear
507-
prio = lambda x: (getattr(x, '_fd_priority', 0), len(x.dimensions))
520+
# `deep_priority`, not `_fd_priority`: an `Add` or `Mul` reports a generic
521+
# fallback, so operands would tie and `sorted` would pick on argument order
522+
prio = lambda x: (deep_priority(x), len(x.dimensions))
508523
prio_func = sorted(args_diff, key=prio, reverse=True)[0]
509524

510525
# The highest priority must be a Function

‎tests/test_differentiable.py‎

Lines changed: 29 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -6,7 +6,8 @@
66

77
from devito import NODE, Differentiable, Eq, Function, Grid, Operator
88
from devito.finite_differences.differentiable import (
9-
Add, EvalDerivative, Mul, Pow, SafeInv, diffify, interp_for_fd
9+
Add, EvalDerivative, Mul, Pow, SafeInv, deep_priority, diffify, highest_priority,
10+
interp_for_fd
1011
)
1112

1213

@@ -224,6 +225,33 @@ def test_mul_three_funcs(self, interp_mode, targets):
224225
assert b.name in evaluated_str
225226
assert c.name in evaluated_str
226227

228+
def test_mul_gather_priority(self):
229+
"""
230+
The gather lands on the highest-priority Function inside the operands.
231+
232+
Without `deep_priority` both operands report the `Differentiable`
233+
fallback, tie, and the winner comes down to SymPy's argument order.
234+
"""
235+
grid = Grid((11, 11))
236+
funcs = self._all_funcs(grid)
237+
node, fx, fy, fxy = (funcs['node'], funcs['x'],
238+
funcs['y'], funcs['xy'])
239+
240+
# Two sums, only the first carrying a NODE Function
241+
with_node = node * fxy + fxy
242+
staggered = fx + fy
243+
244+
assert deep_priority(with_node) == node._fd_priority
245+
assert deep_priority(staggered) == fx._fd_priority
246+
assert deep_priority(with_node) > deep_priority(staggered)
247+
248+
# Whichever way the product sorts, NODE wins
249+
for prod in (with_node * staggered, staggered * with_node):
250+
assert highest_priority(
251+
prod, candidates=[with_node, staggered]
252+
) is node
253+
assert prod.indices_ref == node.indices_ref
254+
227255
@pytest.mark.parametrize('interp_mode', ['direct', 'symmetric'])
228256
@pytest.mark.parametrize('targets', [
229257
('node', 'x', 'xy'),

0 commit comments

Comments
 (0)