|
11 | 11 | except ImportError: |
12 | 12 | from numpy import i0 |
13 | 13 |
|
14 | | -from devito.finite_differences.differentiable import Mul |
| 14 | +from devito.finite_differences.differentiable import Mul, SparseLocalSum |
15 | 15 | from devito.finite_differences.elementary import floor |
16 | 16 | from devito.logger import warning |
17 | 17 | from devito.symbolics import INT, retrieve_function_carriers, retrieve_functions |
18 | 18 | from devito.tools import ( |
19 | 19 | Pickable, as_fp64_decimal, as_list, as_tuple, filter_ordered, flatten, memoized_meth |
20 | 20 | ) |
21 | | -from devito.types import CustomDimension, Eq, Evaluable, Inc, SubFunction, Symbol |
| 21 | +from devito.types import ( |
| 22 | + CustomDimension, Eq, Evaluable, Inc, StencilDimension, SubFunction |
| 23 | +) |
22 | 24 | from devito.types.utils import DimensionTuple |
23 | 25 |
|
24 | 26 | __all__ = ['LinearInterpolator', 'NearestInterpolator', |
@@ -453,19 +455,21 @@ def _interp_idx(self, variables, implicit_dims=None, subdomain=None, |
453 | 455 |
|
454 | 456 | return idx_subs, temps |
455 | 457 |
|
456 | | - def _local_accumulator(self, expr, idx_subs, implicit_dims=None, subdomain=None): |
| 458 | + def _local_accumulator(self, expr, idx_subs, subdomain=None): |
457 | 459 | """ |
458 | | - Generate a local accumulator for the interpolation/injection operation. |
| 460 | + Represent the local sum of weighted interpolation contributions. |
459 | 461 | """ |
460 | | - # Accumulate point-wise contributions into a temporary |
461 | | - rhs = Symbol(name=f'sum{self.sfunction.name}', dtype=self.sfunction.dtype) |
462 | | - summands = [Eq(rhs, 0., implicit_dims=implicit_dims)] |
463 | | - # Substitute coordinate base symbols into the interpolation coefficients |
464 | 462 | weights = self._weights(subdomain=subdomain) |
465 | | - summands.extend([Inc(rhs, (weights * expr).xreplace(idx_subs), |
466 | | - implicit_dims=implicit_dims)]) |
467 | | - |
468 | | - return summands, rhs |
| 463 | + rdims = self._rdim(subdomain=subdomain) |
| 464 | + # Bound stencil indices are independent of the outer sparse-point dimension |
| 465 | + dimensions = tuple(StencilDimension(d.name, d.parent.symbolic_min, |
| 466 | + d.parent.symbolic_max) for d in rdims) |
| 467 | + aliases = {a: i for d, i in zip(rdims, dimensions, strict=True) |
| 468 | + for a in (d, d.parent)} |
| 469 | + summand = (weights * expr).xreplace(idx_subs).xreplace(aliases) |
| 470 | + predicates = tuple(d.condition.xreplace(aliases) for d in rdims) |
| 471 | + return SparseLocalSum(summand, dimensions, self.sfunction.dtype, |
| 472 | + predicates=predicates) |
469 | 473 |
|
470 | 474 | @check_radius |
471 | 475 | @check_coords |
@@ -539,16 +543,13 @@ def _interpolate(self, expr, increment=False, self_subs=None, implicit_dims=None |
539 | 543 | idx_subs, temps = self._interp_idx(variables, implicit_dims=implicit_dims, |
540 | 544 | subdomain=subdomain) |
541 | 545 |
|
542 | | - # Local scalar for accumulation over radius |
543 | | - summands, rhs = self._local_accumulator(expr, idx_subs, |
544 | | - implicit_dims=implicit_dims, |
545 | | - subdomain=subdomain) |
| 546 | + rhs = self._local_accumulator(expr, idx_subs, subdomain=subdomain) |
546 | 547 | # Write/Incr `self` |
547 | 548 | lhs = self.sfunction.subs(self_subs) |
548 | 549 | ecls = Inc if increment else Eq |
549 | 550 | last = [ecls(lhs, rhs, implicit_dims=implicit_dims)] |
550 | 551 |
|
551 | | - return temps + summands + last |
| 552 | + return temps + last |
552 | 553 |
|
553 | 554 | def _inject(self, field, expr, increment=True, implicit_dims=None): |
554 | 555 | """ |
@@ -785,8 +786,8 @@ class NearestInterpolator(LinearInterpolator): |
785 | 786 |
|
786 | 787 | _name = 'nearest' |
787 | 788 |
|
788 | | - def _local_accumulator(self, expr, idx_subs, implicit_dims=None, subdomain=None): |
789 | | - return [], expr.xreplace(idx_subs) |
| 789 | + def _local_accumulator(self, expr, idx_subs, subdomain=None): |
| 790 | + return expr.xreplace(idx_subs) |
790 | 791 |
|
791 | 792 | @memoized_meth |
792 | 793 | def _rdim(self, subdomain=None, shifts=None): |
|
0 commit comments