Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
100 changes: 97 additions & 3 deletions devito/finite_differences/differentiable.py
Original file line number Diff line number Diff line change
Expand Up @@ -21,8 +21,8 @@
from devito.finite_differences.tools import coeff_priority, make_shift_x0
from devito.logger import warning
from devito.tools import (
Tag, as_tuple, extract_dtype, filter_ordered, flatten, frozendict, infer_dtype,
is_integer, is_number, memoized_func, split
Pickable, Tag, as_tuple, extract_dtype, filter_ordered, flatten, frozendict,
infer_dtype, is_integer, is_number, memoized_func, split
)
from devito.types import Array, DimensionTuple, Evaluable, StencilDimension
from devito.types.basic import AbstractFunction, Indexed
Expand All @@ -35,6 +35,7 @@
'Imag',
'IndexDerivative',
'IndexDerivativeProperty',
'LocalSum',
'Real',
'Weights',
]
Expand Down Expand Up @@ -940,13 +941,106 @@ def _evaluate(self, **kwargs):
terms.append(expr.xreplace(mapper))
return sum(terms)

@property
def bound_symbols(self):
return set(self.dimensions)

@property
def free_symbols(self):
return super().free_symbols - set(self.dimensions)
return super().free_symbols - self.bound_symbols

func = DifferentiableOp._rebuild


class LocalSum(IndexSum, Pickable):

"""
A zero-initialized sum over guarded local dimensions.

`cdims` are guarded ConditionalDimensions, retained with their original
parents and conditions. `dimensions` exposes the parent iteration dimensions.
Masked points contribute zero. The sum remains symbolic until Cluster lowering
chooses its implementation.

Examples
--------
For bilinear interpolation, `posx` and `posy` are the grid indices of sparse
point `p`, and `wx` and `wy` hold its interpolation weights::

i = CustomDimension('i', 0, 1, 2)
j = CustomDimension('j', 0, 1, 2)
ci = ConditionalDimension('i', i, indirect=True,
condition=And(posx + i >= x_m, posx + i <= x_M))
cj = ConditionalDimension('j', j, indirect=True,
condition=And(posy + j >= y_m, posy + j <= y_M))
value = LocalSum(
wx[p, ci]*wy[p, cj]*f[posx + ci, posy + cj],
cdims=(ci, cj)
)
Eq(rcv[p], value)

The scalar lowering has the following semantics (pseudocode)::

acc = 0
for i in range(2):
for j in range(2):
if x_m <= posx + i <= x_M and y_m <= posy + j <= y_M:
acc += wx[p, i]*wy[p, j]*f[posx + i, posy + j]
rcv[p] = acc

The guarded indices and their parents are local to the sum; `p` remains an
outer iteration dimension.
If every tap is masked, `rcv[p]` receives zero.
"""

__rargs__ = ('expr',)
__rkwargs__ = ('cdims', 'dtype')

def __new__(cls, expr, cdims=(), dtype=None, **kwargs):
obj = sympy.Expr.__new__(cls, expr)

obj._expr = expr
obj._cdims = as_tuple(cdims)
obj._dtype = dtype

return obj

def _hashable_content(self):
return super()._hashable_content() + (self.cdims, self.dtype)

@property
def cdims(self):
return self._cdims

@cached_property
def dtype(self):
if self._dtype is None:
return extract_dtype(self.expr)
return self._dtype

@cached_property
def dimensions(self):
return tuple(d.parent for d in self.cdims)

@cached_property
def conditionals(self):
return frozendict({d: d.condition for d in self.cdims})

@property
def bound_symbols(self):
return super().bound_symbols | set(self.cdims)

@property
def free_symbols(self):
symbols = self.expr.free_symbols.union(*[d.free_symbols for d in self.cdims])
return symbols - self.bound_symbols

def _evaluate(self, **kwargs):
return self._rebuild(*self._evaluate_args(**kwargs))

__reduce_ex__ = Pickable.__reduce_ex__


class WeightsIndexed(Indexed):

@property
Expand Down
11 changes: 11 additions & 0 deletions devito/ir/clusters/cluster.py
Original file line number Diff line number Diff line change
Expand Up @@ -148,6 +148,13 @@ def dist_dimensions(self):
ret.update(f._dist_dimensions)
return frozenset(ret)

@cached_property
def local_sums(self):
"""
The local sums in equation order, retaining occurrences across equations.
"""
return tuple(s for e in self.exprs for s in e.local_sums)

@cached_property
def scope(self):
return Scope(self.exprs)
Expand Down Expand Up @@ -674,6 +681,10 @@ def rebuild(self, **kwargs):
def exprs(self):
return flatten(c.exprs for c in self)

@cached_property
def local_sums(self):
return tuple(s for c in self for s in c.local_sums)

@cached_property
def scope(self):
return Scope(exprs=self.exprs)
Expand Down
15 changes: 10 additions & 5 deletions devito/ir/equations/algorithms.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,7 +3,7 @@

from devito.data.allocators import DataReference
from devito.finite_differences.differentiable import diff2sympy
from devito.ir.support import GuardFactor
from devito.ir.support import GuardFactor, bounded
from devito.logger import warning
from devito.symbolics import (
IntDiv, retrieve_dimensions, retrieve_functions, retrieve_indexed, uxreplace
Expand All @@ -28,12 +28,17 @@ def dimension_sort(expr):
appear within Indexeds.
"""

# Bound Dimensions do not impact the order of the enclosing iteration space
bound = bounded(expr)

def handle_indexed(indexed):
relation = []
for i in indexed.indices:
try:
# Assume it's an AffineIndexAccessFunction...
relation.append(i.d)
# It may contain only a scalar offset and bound stencil indices
if i.d:
relation.append(i.d)
except AttributeError:
# It's not! Maybe there are some nested Indexeds (e.g., the
# situation is A[B[i]])
Expand All @@ -46,9 +51,9 @@ def handle_indexed(indexed):
# what the user is attempting to do
relation.extend(filter_sorted(i.atoms(Dimension)))

# StencilDimensions are lowered subsequently through special compiler
# Bound Dimensions are lowered subsequently through special compiler
# passes, so they can be ignored here
relation = tuple(d for d in relation if not d.is_Stencil)
relation = tuple(d for d in relation if d not in bound)

return relation

Expand All @@ -61,7 +66,7 @@ def handle_indexed(indexed):
relations.add(expr.implicit_dims)

# Add in leftover free dimensions (not an Indexed' index)
extra = set(retrieve_dimensions(expr, deep=True))
extra = set(retrieve_dimensions(expr, deep=True)) - bound

# Add in pure data dimensions (e.g., those accessed only via explicit values,
# such as A[3])
Expand Down
18 changes: 14 additions & 4 deletions devito/ir/equations/equation.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,14 +4,15 @@
import numpy as np
import sympy

from devito.finite_differences.differentiable import diff2sympy
from devito.finite_differences.differentiable import LocalSum, diff2sympy
from devito.ir.equations.algorithms import dimension_sort, generate_conditionals
from devito.ir.support import (
Interval, IntervalGroup, IterationSpace, Stencil, detect_accesses
Interval, IntervalGroup, IterationSpace, Stencil, bounded, detect_accesses
)
from devito.symbolics import limits_mapper, retrieve_accesses
from devito.symbolics import limits_mapper, retrieve_accesses, search
from devito.tools import (
Pickable, Tag, as_hashable, filter_sorted, frozendict, reuse_if_unchanged
Pickable, Tag, as_hashable, filter_ordered, filter_sorted, frozendict,
reuse_if_unchanged
)
from devito.types import Eq, Inc, ReduceMax, ReduceMin, ReduceMinMax

Expand Down Expand Up @@ -50,6 +51,11 @@ def ispace(self):
def dimensions(self):
return set(self.ispace.dimensions)

@cached_property
def local_sums(self):
"""The local sums in dependency order, with nested sums first."""
return tuple(filter_ordered(search(self, LocalSum, mode='all')))

@property
def implicit_dims(self):
return self._implicit_dims
Expand Down Expand Up @@ -273,11 +279,15 @@ def __new__(cls, *args, **kwargs):
# Analyze the expression
accesses = detect_accesses(expr)
dimensions = Stencil.union(*accesses.values())
bound = bounded(expr)

# Separate out the SubIterators from the main iteration Dimensions, that
# is those which define an actual iteration space
iterators = {}
for d in dimensions:
if d in bound:
# Local sum dimensions belong to the sum's own iteration space
continue
if d.is_SubIterator:
iterators.setdefault(d.root, set()).add(d)
elif d.is_Conditional:
Expand Down
16 changes: 14 additions & 2 deletions devito/ir/support/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,7 +2,7 @@
from contextlib import suppress
from itertools import product

from devito.finite_differences import IndexDerivative
from devito.finite_differences.differentiable import IndexDerivative, IndexSum
from devito.symbolics import retrieve_indexed, search
from devito.tools import DefaultOrderedDict, as_tuple, filter_sorted, split
from devito.types import (
Expand All @@ -13,6 +13,7 @@
'AccessMode',
'IMask',
'Stencil',
'bounded',
'detect_accesses',
'erange',
'extrema',
Expand Down Expand Up @@ -231,7 +232,18 @@ def pull_dims(exprs, flag=True):
return dims


# *** Utility functions for expressions that potentially contain StencilDimensions
# *** Utility functions for bound and unbound Dimensions


def bounded(expr):
"""
Retrieve all Dimensions bound by symbolic sums in `expr`.
"""
sums = search(expr, IndexSum, mode='unique', deep=True)
dims = set().union(*(i.bound_symbols for i in sums))

return dims - expr.free_symbols


def unbounded(expr):
"""
Expand Down
30 changes: 11 additions & 19 deletions devito/operations/interpolators.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,14 +11,14 @@
except ImportError:
from numpy import i0

from devito.finite_differences.differentiable import Mul
from devito.finite_differences.differentiable import LocalSum, Mul
from devito.finite_differences.elementary import floor
from devito.logger import warning
from devito.symbolics import INT, retrieve_function_carriers, retrieve_functions
from devito.tools import (
Pickable, as_fp64_decimal, as_list, as_tuple, filter_ordered, flatten, memoized_meth
)
from devito.types import CustomDimension, Eq, Evaluable, Inc, SubFunction, Symbol
from devito.types import CustomDimension, Eq, Evaluable, Inc, SubFunction
from devito.types.utils import DimensionTuple

__all__ = ['LinearInterpolator', 'NearestInterpolator',
Expand Down Expand Up @@ -453,19 +453,14 @@ def _interp_idx(self, variables, implicit_dims=None, subdomain=None,

return idx_subs, temps

def _local_accumulator(self, expr, idx_subs, implicit_dims=None, subdomain=None):
def _local_accumulator(self, expr, idx_subs, subdomain=None):
"""
Generate a local accumulator for the interpolation/injection operation.
Represent the local sum of weighted interpolation contributions.
"""
# Accumulate point-wise contributions into a temporary
rhs = Symbol(name=f'sum{self.sfunction.name}', dtype=self.sfunction.dtype)
summands = [Eq(rhs, 0., implicit_dims=implicit_dims)]
# Substitute coordinate base symbols into the interpolation coefficients
weights = self._weights(subdomain=subdomain)
summands.extend([Inc(rhs, (weights * expr).xreplace(idx_subs),
implicit_dims=implicit_dims)])

return summands, rhs
rdims = self._rdim(subdomain=subdomain)
summand = (weights * expr).xreplace(idx_subs)
return LocalSum(summand, cdims=rdims, dtype=self.sfunction.dtype)

@check_radius
@check_coords
Expand Down Expand Up @@ -539,16 +534,13 @@ def _interpolate(self, expr, increment=False, self_subs=None, implicit_dims=None
idx_subs, temps = self._interp_idx(variables, implicit_dims=implicit_dims,
subdomain=subdomain)

# Local scalar for accumulation over radius
summands, rhs = self._local_accumulator(expr, idx_subs,
implicit_dims=implicit_dims,
subdomain=subdomain)
rhs = self._local_accumulator(expr, idx_subs, subdomain=subdomain)
# Write/Incr `self`
lhs = self.sfunction.subs(self_subs)
ecls = Inc if increment else Eq
last = [ecls(lhs, rhs, implicit_dims=implicit_dims)]

return temps + summands + last
return temps + last

def _inject(self, field, expr, increment=True, implicit_dims=None):
"""
Expand Down Expand Up @@ -785,8 +777,8 @@ class NearestInterpolator(LinearInterpolator):

_name = 'nearest'

def _local_accumulator(self, expr, idx_subs, implicit_dims=None, subdomain=None):
return [], expr.xreplace(idx_subs)
def _local_accumulator(self, expr, idx_subs, subdomain=None):
return expr.xreplace(idx_subs)

@memoized_meth
def _rdim(self, subdomain=None, shifts=None):
Expand Down
4 changes: 3 additions & 1 deletion devito/operator/operator.py
Original file line number Diff line number Diff line change
Expand Up @@ -33,7 +33,8 @@
from devito.parameters import configuration
from devito.passes import (
Graph, error_mapper, finalize_args, generate_implicit, generate_macros, is_on_device,
lower_dtypes, lower_index_derivatives, minimize_symbols, optimize_pows, unevaluate
lower_dtypes, lower_index_derivatives, lower_local_sums, minimize_symbols,
optimize_pows, unevaluate
)
from devito.symbolics import estimate_cost, subs_op_args
from devito.tools import (
Expand Down Expand Up @@ -423,6 +424,7 @@ def _lower_clusters(cls, expressions, profiler=None, **kwargs):
clusters = generate_implicit(clusters)

# Lower all remaining high order symbolic objects
clusters = lower_local_sums(clusters, **kwargs)
clusters = lower_index_derivatives(clusters, **kwargs)

# Turn pows into multiplications. This must happen as late as possible
Expand Down
1 change: 1 addition & 0 deletions devito/passes/clusters/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,4 +9,5 @@
from .implicit import * # noqa
from .misc import * # noqa
from .derivatives import * # noqa
from .localsum import * # noqa
from .unevaluate import * # noqa
4 changes: 2 additions & 2 deletions devito/passes/clusters/cse.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,7 +11,7 @@
# Moved in 1.13
from sympy.core.basic import ordering_of_classes

from devito.finite_differences.differentiable import IndexDerivative
from devito.finite_differences.differentiable import IndexSum
from devito.ir import Cluster, Scope, cluster_pass
from devito.symbolics import (
DefFunction, Reserved, estimate_cost, q_leaf, q_terminal, search
Expand Down Expand Up @@ -427,7 +427,7 @@ def _(expr):
return {}


@_catch.register(IndexDerivative)
@_catch.register(IndexSum)
def _(expr):
"""
Handler for symbol-binding objects. There can be many of them and therefore
Expand Down
Loading
Loading