Skip to content

Commit 3763880

Browse files
committed
misc: Fix oversight regarding SparseTimeFunction index errors
1 parent f75f4fc commit 3763880

2 files changed

Lines changed: 15 additions & 1 deletion

File tree

‎devito/types/dense.py‎

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -950,7 +950,9 @@ def _arg_check(self, args, intervals, **kwargs):
950950
for i, s in zip(self.dimensions, data.shape, strict=True):
951951
i._arg_check(args, s, intervals[i])
952952

953-
if args.options['index-mode'] == 'int32' and \
953+
# SparseTimeFunctions use int64 indexing regardless of index mode
954+
if not self.is_SparseTimeFunction and \
955+
args.options['index-mode'] == 'int32' and \
954956
args.options['linearize'] and \
955957
self.is_regular and \
956958
data.size - 1 >= np.iinfo(np.int32).max:

‎tests/test_linearize.py‎

Lines changed: 12 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -149,6 +149,18 @@ def test_interpolation_enforcing_int64_indexing():
149149
assert 'long p_src_stride0' in str(op) # for `src`
150150

151151

152+
def test_interpolation_enforcing_int64_indexing_v2():
153+
grid = Grid(shape=(4, 4))
154+
155+
# Too many elements for int32 (1e10), but should not trigger any errors, as
156+
# will be indexed with int64
157+
rec = SparseTimeFunction(name='rec', grid=grid, npoint=int(1e6), nt=int(1e4))
158+
u = TimeFunction(name="u", grid=grid, time_order=2)
159+
160+
Operator(rec.interpolate(expr=u.forward),
161+
opt=('advanced', {'linearize': True, 'index-mode': 'int32'}))
162+
163+
152164
def test_interpolation_msf():
153165
grid = Grid(shape=(4, 4))
154166

0 commit comments

Comments
 (0)