Skip to content

Commit b330c14

Browse files
authored
Merge pull request #3012 from devitocodes/fix-staggered-sum-indices
api: Fix staggered sum indices
2 parents 4109b58 + e68a40c commit b330c14

19 files changed

Lines changed: 301 additions & 56 deletions

File tree

‎devito/core/gpu.py‎

Lines changed: 4 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -162,7 +162,10 @@ def wrapper(expressions, mode='default', options=None, **kwargs1):
162162
# small kernels typically generated by recursive compilation
163163
par_tile0 = options0['par-tile']
164164
par_tile = options.get('par-tile')
165-
if par_tile0 and par_tile:
165+
if par_tile is False:
166+
# The caller explicitly opted out of tiling
167+
options = {**options0, **options, 'par-tile': ParTile(None)}
168+
elif par_tile0 and par_tile:
166169
options = {**options0, **options, 'par-tile': par_tile}
167170
elif par_tile0:
168171
par_tile = ParTile(par_tile0.default, default=par_tile0.default)

‎devito/finite_differences/derivative.py‎

Lines changed: 7 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -14,7 +14,7 @@
1414
from devito.warnings import warn
1515

1616
from .differentiable import Add, Differentiable, Mul, diffify, interp_for_fd
17-
from .finite_difference import cross_derivative, generic_derivative
17+
from .finite_difference import cross_derivative, generic_derivative, indices_at
1818
from .rsfd import d45
1919
from .tools import direct, transpose
2020

@@ -575,6 +575,12 @@ def _eval_fd(self, expr, **kwargs):
575575
shited derivative.
576576
- 4: Apply substitutions.
577577
"""
578+
# Differentiation is linear, and a sum of terms at different staggered
579+
# locations must use it: `Add` reports its first argument's location,
580+
# so `x0` would shift the other terms off the point they sat at.
581+
if expr.is_Add and any(len(indices_at(expr, d)) > 1 for d in self.dims):
582+
return expr.func(*[self._eval_fd(a, **kwargs) for a in expr.args])
583+
578584
# Step 1: Evaluate non-derivative x0. We currently enforce a simple 2nd order
579585
# interpolation to avoid very expensive finite differences on top of it
580586
x0_deriv = self._filter_dims(self.x0)

‎devito/finite_differences/finite_difference.py‎

Lines changed: 26 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -100,6 +100,29 @@ def cross_derivative(expr, dims, fd_order, deriv_order, x0=None, side=None, **kw
100100
return expr
101101

102102

103+
def indices_at(expr, dim):
104+
"""
105+
The locations `expr`'s terms sit at along `dim`.
106+
107+
Terms with no location of their own, a scalar say, contribute none.
108+
"""
109+
indices = set()
110+
for i in (expr.args if expr.is_Add else (expr,)):
111+
try:
112+
indices.add(i.indices_ref[dim])
113+
except (AttributeError, KeyError, IndexError, TypeError):
114+
continue
115+
return indices
116+
117+
118+
def index_at(expr, dim):
119+
"""
120+
Where `expr` sits along `dim`, or None if it does not say.
121+
"""
122+
indices = indices_at(expr, dim)
123+
return indices.pop() if len(indices) == 1 else None
124+
125+
103126
@check_input
104127
def generic_derivative(expr, dim, fd_order, deriv_order, matvec=direct, x0=None,
105128
coefficients='taylor', expand=True, weights=None, side=None):
@@ -139,8 +162,9 @@ def generic_derivative(expr, dim, fd_order, deriv_order, matvec=direct, x0=None,
139162
if deriv_order == 1 and fd_order == 2 and side is None:
140163
fd_order = 1
141164

142-
# Zeroth order derivative is just the expression itself if not shifted
143-
if deriv_order == 0 and not x0:
165+
# Zeroth order is the identity when `expr` already sits at `x0`, not a
166+
# stencil centred there.
167+
if deriv_order == 0 and (not x0 or index_at(expr, dim) == x0.get(dim)):
144168
return expr
145169

146170
# Enforce stable time coefficients

‎devito/ir/clusters/algorithms.py‎

Lines changed: 4 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -742,11 +742,10 @@ def _normalize_reductions_dense(cluster, mapper, sregistry, platform):
742742
elif rhs in mapper:
743743
# Seen this RHS already, so reuse the Array that was created for it
744744
processed.append(e.func(lhs, mapper[rhs].indexify()))
745-
elif rf and rf.is_Array and sum(flatten(rf._size_nodomain)) == 0:
746-
# Special case: the RHS is an Array with no halo/padding, meaning
747-
# that the written data values are contiguous in memory, hence
748-
# we can simply reuse the Array itself as we're already in the
749-
# desired memory layout
745+
elif rf and rf.is_Array and rf._is_reduction_ready:
746+
# Special case: the RHS is an Array whose written data values are
747+
# contiguous in memory, hence we can simply reuse the Array
748+
# itself as we're already in the desired memory layout
750749
processed.append(e)
751750
else:
752751
name = sregistry.make_name()

‎devito/ir/iet/visitors.py‎

Lines changed: 20 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -7,6 +7,7 @@
77
import ctypes
88
from collections import OrderedDict
99
from collections.abc import Callable, Generator, Iterable, Iterator, Sequence
10+
from contextlib import suppress
1011
from itertools import chain, groupby
1112
from typing import Any, Generic, TypeVar
1213

@@ -17,8 +18,8 @@
1718
from devito.exceptions import CompilationError
1819
from devito.ir.cgen.printer import get_printer
1920
from devito.ir.iet.nodes import (
20-
BlankLine, Call, Expression, ExpressionBundle, Iteration, Lambda, ListMajor, Node,
21-
Section, _same_as_before
21+
BlankLine, Call, Definition, Expression, ExpressionBundle, Iteration, Lambda,
22+
ListMajor, Node, Section, _same_as_before
2223
)
2324
from devito.ir.support.space import Backward
2425
from devito.symbolics import (
@@ -1275,6 +1276,23 @@ def visit_Call(self, o: Call, **kwargs) -> Iterator[ApplicationType]:
12751276
except (AttributeError, TypeError):
12761277
yield from self._visit(i)
12771278

1279+
def visit_Definition(self, o: Definition, **kwargs) -> Iterator[ApplicationType]:
1280+
# The defined object carries expressions in its constructor arguments
1281+
# and in its initializer, both of which end up in the generated code
1282+
f = o.function
1283+
if f.is_LocalObject:
1284+
candidates = (*f.cargs, f.initvalue)
1285+
elif f.is_Array:
1286+
candidates = as_tuple(f.initvalue)
1287+
else:
1288+
return
1289+
1290+
for i in candidates:
1291+
# Not everything in there is a symbolic expression, e.g. a plain
1292+
# number, a string, or nothing at all
1293+
with suppress(AttributeError, TypeError):
1294+
yield from i.find(self.match)
1295+
12781296

12791297
class IsPerfectIteration(Visitor):
12801298

‎devito/passes/iet/definitions.py‎

Lines changed: 49 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -20,7 +20,7 @@
2020
VOID, Byref, DefFunction, FieldFromPointer, IndexedPointer, ListInitializer, SizeOf,
2121
as_long, pow_to_mul, unevaluate
2222
)
23-
from devito.tools import as_list, as_mapper, as_tuple, filter_sorted, flatten
23+
from devito.tools import as_list, as_mapper, as_tuple, filter_sorted, flatten, is_integer
2424
from devito.types import (
2525
Array, ComponentAccess, CustomDimension, DeviceMap, DeviceRM, Dimension, Eq, Symbol,
2626
size_t
@@ -91,6 +91,27 @@ def __init__(self, rcompile=None, sregistry=None, platform=None,
9191
self.sregistry = sregistry
9292
self.platform = platform
9393

94+
# Off inside the recursive compilation of a zero-init itself, which
95+
# would otherwise ask for a zero-init of its own, ad infinitum
96+
self.zero_init = (options or {}).get('zero-init', True)
97+
98+
def _zero_init(self, obj, storage):
99+
"""
100+
The nodes zeroing `obj` upfront, if it asks for it, plus the efuncs
101+
they call, if any.
102+
"""
103+
if not (obj._is_zero_init and self.zero_init):
104+
return (), ()
105+
106+
return self._make_zero_init(obj, storage)
107+
108+
def _make_zero_init(self, obj, storage):
109+
"""How to zero `obj`'s whole allocation, padding included."""
110+
storage.include(self.langbb['header-memcpy'])
111+
nbytes = SizeOf(obj._C_typedata)*as_long(obj.size)
112+
113+
return (self.langbb['host-memset'](obj._C_symbol, 0, nbytes),), ()
114+
94115
def _alloc_object_on_low_lat_mem(self, site, obj, storage):
95116
"""
96117
Allocate a LocalObject in the low latency memory.
@@ -172,11 +193,13 @@ def _alloc_host_array_on_high_bw_mem(self, site, obj, storage, *args):
172193
memptr = VOID(Byref(obj._C_symbol), '**')
173194
alignment = obj._data_alignment
174195
nbytes = SizeOf(obj._C_typedata)*as_long(obj.size)
175-
alloc = self.langbb['host-alloc'](memptr, alignment, nbytes)
196+
zeroing, efuncs = self._zero_init(obj, storage)
197+
allocs = [decl, self.langbb['host-alloc'](memptr, alignment, nbytes),
198+
*zeroing]
176199

177200
free = self.langbb['host-free'](obj._C_symbol)
178201

179-
storage.update(obj, site, allocs=(decl, alloc), frees=free)
202+
storage.update(obj, site, allocs=tuple(allocs), frees=free, efuncs=efuncs)
180203

181204
def _alloc_local_array_on_high_bw_mem(self, site, obj, storage, *args):
182205
"""
@@ -568,7 +591,7 @@ def __init__(self, options=None, **kwargs):
568591
self.gpu_create = options['gpu-create']
569592
self.gpu_place_transfers = options.get('place-transfers')
570593

571-
super().__init__(**kwargs)
594+
super().__init__(options=options, **kwargs)
572595

573596
def _alloc_local_array_on_high_bw_mem(self, site, obj, storage):
574597
"""
@@ -579,11 +602,22 @@ def _alloc_local_array_on_high_bw_mem(self, site, obj, storage):
579602
dofree = self.langbb['device-free']
580603

581604
nbytes = SizeOf(obj._C_typedata)*obj.size
582-
init = doalloc(nbytes, deviceid, retobj=obj)
605+
606+
zeroing, efuncs = self._zero_init(obj, storage)
607+
allocs = [doalloc(nbytes, deviceid, retobj=obj), *zeroing]
583608

584609
free = dofree(obj._C_name, deviceid)
585610

586-
storage.update(obj, site, allocs=init, frees=free)
611+
storage.update(obj, site, allocs=tuple(allocs), frees=free, efuncs=efuncs)
612+
613+
def _make_zero_init(self, obj, storage):
614+
# No language here has a device-side memset, so use a kernel. It gains
615+
# nothing from tiling, and nvc++ trips over the padded loop bounds
616+
# when it is asked to tile them
617+
efuncs, init = make_zero_init(obj, self.rcompile, self.sregistry,
618+
options={'par-tile': False})
619+
620+
return (init,), efuncs
587621

588622
def _map_array_on_high_bw_mem(self, site, obj, storage):
589623
"""
@@ -702,18 +736,20 @@ def process(self, graph):
702736
self.place_casts(graph)
703737

704738

705-
def make_zero_init(obj, rcompile, sregistry):
739+
def make_zero_init(obj, rcompile, sregistry, options=None):
706740
cdims = []
707-
for d, (h0, h1), s in zip(
708-
obj.dimensions, obj._size_halo, obj.symbolic_shape, strict=True
741+
for d, (h0, h1), (_, p1), s in zip(
742+
obj.dimensions, obj._size_halo, obj._size_padding, obj.symbolic_shape,
743+
strict=True
709744
):
710745
if d.is_NonlinearDerived:
711-
assert h0 == h1 == 0
746+
assert h0 == h1
712747
m = 0
713748
M = s - 1
714749
else:
715750
m = d.symbolic_min - h0
716-
M = d.symbolic_max + h1
751+
# Object needing padding zeroing need symbolic padding
752+
M = d.symbolic_max + h1 + (0 if is_integer(p1) else p1)
717753
cdims.append(CustomDimension(name=d.name, parent=d,
718754
symbolic_min=m, symbolic_max=M))
719755

@@ -722,7 +758,8 @@ def make_zero_init(obj, rcompile, sregistry):
722758
else:
723759
eqns = [Eq(obj[cdims], 0)]
724760

725-
irs, byproduct = rcompile(eqns)
761+
irs, byproduct = rcompile(eqns, options={'zero-init': False,
762+
**(options or {})})
726763

727764
init = irs.iet.body.body[0]
728765

‎devito/passes/iet/languages/C.py‎

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -56,6 +56,8 @@ class CBB(LangBB):
5656
Call('free', (i,)),
5757
'host-free-pin': lambda i:
5858
Call('free', (i,)),
59+
'host-memset': lambda i, j, k:
60+
Call('memset', (i, j, k)),
5961
'alloc-global-symbol': lambda i, j, k:
6062
Call('memcpy', (i, j, k))
6163
}

‎devito/passes/iet/languages/CXX.py‎

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -141,6 +141,8 @@ class CXXBB(LangBB):
141141
Call('free', (i,)),
142142
'host-free-pin': lambda i:
143143
Call('free', (i,)),
144+
'host-memset': lambda i, j, k:
145+
Call('memset', (i, j, k)),
144146
'alloc-global-symbol': lambda i, j, k:
145147
Call('memcpy', (i, j, k))
146148
}

‎devito/passes/iet/misc.py‎

Lines changed: 3 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -259,8 +259,9 @@ def _(expr, langbb, printer):
259259

260260
@_lower_macro_math.register(RoundUp)
261261
def _(expr, langbb, printer):
262-
return (('ROUND_UP(a,b)',
263-
'((((a)%(b)) == 0) ? (a) : ((a) + (b) - ((a)%(b))))'),), {}
262+
# Branchless: a ternary makes for a poor loop bound, and it crashes the
263+
# nvc++ frontend outright when the loop is collapsed and optimized
264+
return (('ROUND_UP(a,b)', '((((a) + (b) - 1)/(b))*(b))'),), {}
264265

265266

266267
@iet_pass

‎devito/types/basic.py‎

Lines changed: 15 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -16,7 +16,7 @@
1616
from devito.data import default_allocator
1717
from devito.parameters import configuration
1818
from devito.tools import (
19-
CustomDtype, Pickable, as_tuple, dtype_to_ctype, frozendict, memoized_meth,
19+
CustomDtype, Pickable, as_tuple, dtype_to_ctype, flatten, frozendict, memoized_meth,
2020
sympy_mutex
2121
)
2222
from devito.types.args import ArgProvider
@@ -714,6 +714,12 @@ class AbstractFunction(sympy.Function, Basic, Pickable, Evaluable):
714714
effect if autopadding is disabled, which is the default behavior.
715715
"""
716716

717+
_is_zero_init = False
718+
"""
719+
Whether the entries outside `self`'s DOMAIN carry meaningful data rather
720+
than scratch, in which case the whole allocation must be zeroed upfront.
721+
"""
722+
717723
__rkwargs__ = ('name', 'dtype', 'grid', 'halo', 'ghost',
718724
'alias', 'space', 'function', 'is_transient', 'avg_mode')
719725

@@ -1351,6 +1357,14 @@ def _size_nodomain(self):
13511357

13521358
return DimensionTuple(*sizes, getters=self.dimensions, left=left, right=right)
13531359

1360+
@property
1361+
def _is_reduction_ready(self):
1362+
"""
1363+
True if a reduction over `self` may run over the whole allocated data
1364+
rather than over the DOMAIN alone, False otherwise.
1365+
"""
1366+
return not sum(flatten(self._size_nodomain))
1367+
13541368
@cached_property
13551369
def _size_ghost(self):
13561370
"""

0 commit comments

Comments
 (0)