Skip to content

Commit 2a3e49a

Browse files
committed
compiler: Zero an Array whose out-of-DOMAIN entries are data
1 parent e9db68d commit 2a3e49a

9 files changed

Lines changed: 113 additions & 32 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/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/types/basic.py‎

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -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

‎devito/types/misc.py‎

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -257,6 +257,7 @@ class TempArray(Array):
257257
"""
258258

259259
is_autopaddable = True
260+
_is_zero_init = True
260261

261262
__rkwargs__ = (Array.__rkwargs__ + ('shift',))
262263

‎examples/performance/00_overview.ipynb‎

Lines changed: 26 additions & 18 deletions
Original file line numberDiff line numberDiff line change
@@ -572,13 +572,7 @@
572572
"+ u[t1][x + 4][y + 4][z + 4] = (f[x + 1][y + 1][z + 1]*f[x + 1][y + 1][z + 1])*((-6.66666667e-1F*r0)*(8.33333333e-2F*r0*u[t0][x + 4][y + 1][z + 4] - 6.66666667e-1F*r0*u[t0][x + 4][y + 2][z + 4] + 6.66666667e-1F*r0*u[t0][x + 4][y + 4][z + 4] - 8.33333333e-2F*r0*u[t0][x + 4][y + 5][z + 4]) + (-8.33333333e-2F*r0)*(8.33333333e-2F*r0*u[t0][x + 4][y + 4][z + 4] - 6.66666667e-1F*r0*u[t0][x + 4][y + 5][z + 4] + 6.66666667e-1F*r0*u[t0][x + 4][y + 7][z + 4] - 8.33333333e-2F*r0*u[t0][x + 4][y + 8][z + 4]) + (8.33333333e-2F*r0)*(8.33333333e-2F*r0*u[t0][x + 4][y][z + 4] - 6.66666667e-1F*r0*u[t0][x + 4][y + 1][z + 4] + 6.66666667e-1F*r0*u[t0][x + 4][y + 3][z + 4] - 8.33333333e-2F*r0*u[t0][x + 4][y + 4][z + 4]) + (6.66666667e-1F*r0)*(8.33333333e-2F*r0*u[t0][x + 4][y + 3][z + 4] - 6.66666667e-1F*r0*u[t0][x + 4][y + 4][z + 4] + 6.66666667e-1F*r0*u[t0][x + 4][y + 6][z + 4] - 8.33333333e-2F*r0*u[t0][x + 4][y + 7][z + 4]))*sinf(f[x + 1][y + 1][z + 1]);\n",
573573
" }\n",
574574
" }\n",
575-
" }\n"
576-
]
577-
},
578-
{
579-
"name": "stdout",
580-
"output_type": "stream",
581-
"text": [
575+
" }\n",
582576
"\n"
583577
]
584578
}
@@ -759,14 +753,7 @@
759753
" }\n",
760754
" }\n",
761755
" STOP(section0,timers)\n",
762-
"}"
763-
]
764-
},
765-
{
766-
"name": "stdout",
767-
"output_type": "stream",
768-
"text": [
769-
"\n"
756+
"}\n"
770757
]
771758
}
772759
],
@@ -969,7 +956,14 @@
969956
" }\n",
970957
" }\n",
971958
" STOP(section0,timers)\n",
972-
"}\n"
959+
"}"
960+
]
961+
},
962+
{
963+
"name": "stdout",
964+
"output_type": "stream",
965+
"text": [
966+
"\n"
973967
]
974968
}
975969
],
@@ -1145,6 +1139,7 @@
11451139
"#include \"xmmintrin.h\"\n",
11461140
"#include \"pmmintrin.h\"\n",
11471141
"#include \"omp.h\"\n",
1142+
"#include \"string.h\"\n",
11481143
"\n",
11491144
"struct dataobj\n",
11501145
"{\n",
@@ -1172,6 +1167,7 @@
11721167
" posix_memalign((void**)(&pr2_vec),64,sizeof(float*)*(long)nthreads);\n",
11731168
" float *restrict r0_vec __attribute__ ((aligned (64)));\n",
11741169
" posix_memalign((void**)(&r0_vec),64,sizeof(float)*(long)z_size*(long)y_size*(long)x_size);\n",
1170+
" memset(r0_vec,0,sizeof(float)*(long)z_size*(long)y_size*(long)x_size);\n",
11751171
" #pragma omp parallel num_threads(nthreads)\n",
11761172
" {\n",
11771173
" const int tid = omp_get_thread_num();\n",
@@ -1254,7 +1250,13 @@
12541250
" free(r0_vec);\n",
12551251
"\n",
12561252
" return 0;\n",
1257-
"}\n",
1253+
"}\n"
1254+
]
1255+
},
1256+
{
1257+
"name": "stdout",
1258+
"output_type": "stream",
1259+
"text": [
12581260
"\n"
12591261
]
12601262
}
@@ -1436,6 +1438,7 @@
14361438
"#include \"xmmintrin.h\"\n",
14371439
"#include \"pmmintrin.h\"\n",
14381440
"#include \"omp.h\"\n",
1441+
"#include \"string.h\"\n",
14391442
"\n",
14401443
"struct dataobj\n",
14411444
"{\n",
@@ -1463,6 +1466,7 @@
14631466
" posix_memalign((void**)(&pr2_vec),64,sizeof(float*)*(long)nthreads);\n",
14641467
" float *restrict r0_vec __attribute__ ((aligned (64)));\n",
14651468
" posix_memalign((void**)(&r0_vec),64,sizeof(float)*(long)z_size*(long)y_size*(long)x_size);\n",
1469+
" memset(r0_vec,0,sizeof(float)*(long)z_size*(long)y_size*(long)x_size);\n",
14661470
" #pragma omp parallel num_threads(nthreads)\n",
14671471
" {\n",
14681472
" const int tid = omp_get_thread_num();\n",
@@ -1579,6 +1583,7 @@
15791583
"#include \"xmmintrin.h\"\n",
15801584
"#include \"pmmintrin.h\"\n",
15811585
"#include \"omp.h\"\n",
1586+
"#include \"string.h\"\n",
15821587
"\n",
15831588
"struct dataobj\n",
15841589
"{\n",
@@ -1604,10 +1609,13 @@
16041609
"{\n",
16051610
" float *restrict r0_vec __attribute__ ((aligned (64)));\n",
16061611
" posix_memalign((void**)(&r0_vec),64,sizeof(float)*(long)z_size*(long)y_size*(long)x_size);\n",
1612+
" memset(r0_vec,0,sizeof(float)*(long)z_size*(long)y_size*(long)x_size);\n",
16071613
" float *restrict r3_vec __attribute__ ((aligned (64)));\n",
16081614
" posix_memalign((void**)(&r3_vec),64,sizeof(float)*(long)z_size*(4 + (long)y_size)*(4 + (long)x_size));\n",
1615+
" memset(r3_vec,0,sizeof(float)*(long)z_size*(4 + (long)y_size)*(4 + (long)x_size));\n",
16091616
" float *restrict r4_vec __attribute__ ((aligned (64)));\n",
16101617
" posix_memalign((void**)(&r4_vec),64,sizeof(float)*(long)z_size*(4 + (long)y_size)*(4 + (long)x_size));\n",
1618+
" memset(r4_vec,0,sizeof(float)*(long)z_size*(4 + (long)y_size)*(4 + (long)x_size));\n",
16111619
"\n",
16121620
" float (*restrict f)[f_vec->size[1]][f_vec->size[2]] __attribute__ ((aligned (64))) = (float (*)[f_vec->size[1]][f_vec->size[2]]) f_vec->data;\n",
16131621
" float (*restrict r0)[y_size][z_size] __attribute__ ((aligned (64))) = (float (*)[y_size][z_size]) r0_vec;\n",
@@ -1731,7 +1739,7 @@
17311739
"name": "python",
17321740
"nbconvert_exporter": "python",
17331741
"pygments_lexer": "ipython3",
1734-
"version": "3.13.11"
1742+
"version": "3.13.15"
17351743
}
17361744
},
17371745
"nbformat": 4,

‎tests/test_gpu_openmp.py‎

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -214,11 +214,13 @@ def test_array_rw(self):
214214

215215
op = Operator(eqn, language='openmp')
216216

217-
assert len(op.body.allocs) == 1
217+
assert len(op.body.allocs) == 2
218218
assert str(op.body.allocs[0]) ==\
219219
('float * r0_vec = (float *)'
220220
'omp_target_alloc(x_size*y_size*z_size*sizeof(float),'
221221
'omp_get_default_device());')
222+
assert str(op.body.allocs[1]) ==\
223+
'init0(x_M,x_m,y_M,y_m,z_M,z_m,r0_vec,y_size,z_size);'
222224
assert len(op.body.maps) == 2
223225
assert all('r0' not in str(i) for i in op.body.maps)
224226

‎tests/test_operator.py‎

Lines changed: 20 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1394,6 +1394,26 @@ def test_conditional_declarations(self):
13941394
assert i[0].is_Expression
13951395
assert i[0].expr.rhs is init_value
13961396

1397+
def test_zero_init_array(self):
1398+
"""
1399+
An Array whose entries outside the DOMAIN are data, rather than
1400+
scratch, is zeroed right after being allocated.
1401+
"""
1402+
grid = Grid(shape=(4, 4))
1403+
1404+
class ZeroInitArray(Array):
1405+
_is_zero_init = True
1406+
1407+
a = ZeroInitArray(name='a', dimensions=grid.dimensions,
1408+
dtype=grid.dtype, space='local')
1409+
b = Array(name='b', dimensions=grid.dimensions, dtype=grid.dtype,
1410+
space='local')
1411+
1412+
f = Function(name='f', grid=grid)
1413+
1414+
assert 'memset(a' in str(Operator(Eq(f, a.indexify())))
1415+
assert 'memset(b' not in str(Operator(Eq(f, b.indexify())))
1416+
13971417
def test_nested_scalar_assigns(self):
13981418
grid = Grid(shape=(4, 4))
13991419

0 commit comments

Comments
 (0)