Skip to content

Commit cda0311

Browse files
committed
compiler: precompute interp pos/weights
2 parents e401545 + 67066e7 commit cda0311

9 files changed

Lines changed: 238 additions & 55 deletions

File tree

‎.github/workflows/pytest-core-mpi.yaml‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -22,7 +22,7 @@ jobs:
2222
runs-on: ubuntu-22.04
2323
strategy:
2424
matrix:
25-
python-version: ['3.10', '3.11']
25+
python-version: ['3.11', '3.14']
2626

2727
env:
2828
DEVITO_LANGUAGE: "openmp"

‎.github/workflows/pytest-core-nompi.yaml‎

Lines changed: 4 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -39,7 +39,7 @@ jobs:
3939
pytest-ubuntu-py311-gcc10-noomp,
4040
pytest-ubuntu-py311-gcc11-cxxnoomp,
4141
pytest-ubuntu-py312-gcc12-cxxomp,
42-
pytest-ubuntu-py312-gcc13-omp,
42+
pytest-ubuntu-py314-gcc13-omp,
4343
pytest-ubuntu-py313-gcc14-omp
4444
]
4545
set: [base, adjoint]
@@ -73,8 +73,8 @@ jobs:
7373
language: "C"
7474
sympy: "1.14"
7575

76-
- name: pytest-ubuntu-py312-gcc13-omp
77-
python-version: '3.12'
76+
- name: pytest-ubuntu-py314-gcc13-omp
77+
python-version: '3.14'
7878
os: ubuntu-24.04
7979
arch: "gcc-13"
8080
language: "openmp"
@@ -134,7 +134,7 @@ jobs:
134134
fi
135135
id: set-tests
136136

137-
- name: Set pip flags for latest python (3.12)
137+
- name: Set pip flags for Python 3.12+
138138
run: |
139139
ver="${{ matrix.python-version }}"
140140
major=${ver%%.*}

‎devito/operations/interpolators.py‎

Lines changed: 182 additions & 31 deletions
Original file line numberDiff line numberDiff line change
@@ -15,8 +15,10 @@
1515
from devito.finite_differences.elementary import floor
1616
from devito.logger import warning
1717
from devito.symbolics import INT, retrieve_function_carriers, retrieve_functions
18-
from devito.tools import Pickable, as_tuple, filter_ordered, flatten, memoized_meth
19-
from devito.types import Eq, Evaluable, Inc, SubFunction, Symbol
18+
from devito.tools import (
19+
Pickable, as_fp64_decimal, as_tuple, filter_ordered, flatten, memoized_meth
20+
)
21+
from devito.types import CustomDimension, Eq, Evaluable, Inc, SubFunction, Symbol
2022
from devito.types.utils import DimensionTuple
2123

2224
__all__ = ['LinearInterpolator', 'PrecomputedInterpolator', 'SincInterpolator']
@@ -510,46 +512,195 @@ def _inject(self, field, expr, implicit_dims=None):
510512
return filter_ordered(temps) + eqns
511513

512514

515+
def _shift_tag(shifts):
516+
"""Suffix used to distinguish per-staggering table names ("_s10", ...)."""
517+
if not shifts or not any(shifts):
518+
return ''
519+
return '_s' + ''.join('1' if s else '0' for s in shifts)
520+
521+
522+
def _shift_values(shifts, grid, spacing):
523+
"""Physical half-cell offsets for each grid dim, as fp64."""
524+
if not shifts:
525+
return np.zeros(grid.dim, dtype=np.float64)
526+
subs = {d.spacing: float(h)
527+
for d, h in zip(grid.dimensions, spacing, strict=True)}
528+
return np.array([float(sympy.sympify(s).xreplace(subs)) for s in shifts])
529+
530+
531+
class _HostTable(SubFunction):
532+
"""Base for a SubFunction populated on the host by the linear interpolator
533+
at arg-prep time. Data honors runtime grid overrides (spacing/origin) and
534+
the parent SparseFunction's scattered coordinates."""
535+
536+
def _arg_apply(self, *args, **kwargs):
537+
return
538+
539+
def _emit(self, coords, spacing, origin): # pragma: no cover - abstract
540+
raise NotImplementedError
541+
542+
def _arg_defaults(self, alias=None, metadata=None, estimate_memory=False):
543+
args = dict(self.dimensions[-1]._arg_defaults(_min=0,
544+
size=self.shape[-1]))
545+
key = alias or self
546+
if estimate_memory:
547+
args[key.name] = self
548+
else:
549+
spacing, origin = _resolved_geometry(self._parent.grid, {})
550+
args[key.name] = self._emit(self._parent.coordinates.data,
551+
spacing, origin)
552+
return args
553+
554+
def _arg_values(self, estimate_memory=False, **kwargs):
555+
# Route through the parent so its dim defaults (and any override)
556+
# are collected; then overlay this table's data honoring runtime
557+
# coord / spacing / origin overrides.
558+
values = super()._arg_values(estimate_memory=estimate_memory, **kwargs)
559+
if estimate_memory:
560+
return values
561+
args = kwargs.get('args')
562+
coords = (args.get(self._parent.coordinates.name)
563+
if args is not None else None)
564+
if coords is None:
565+
coords = self._parent.coordinates.data
566+
spacing, origin = _resolved_geometry(self._parent.grid, kwargs)
567+
values[self.name] = self._emit(coords, spacing, origin)
568+
return values
569+
570+
571+
class Gridpoints(_HostTable):
572+
"""int32 cell indices per sparse point, shape ``(npoint, ndim)``."""
573+
574+
__rkwargs__ = _HostTable.__rkwargs__ + ('shifts',)
575+
576+
def __init_finalize__(self, *args, **kwargs):
577+
self._shifts = kwargs.pop('shifts', None)
578+
super().__init_finalize__(*args, **kwargs)
579+
580+
@property
581+
def shifts(self):
582+
return self._shifts
583+
584+
def _emit(self, coords, spacing, origin):
585+
return _cell_indices(coords, self._parent, self._shifts,
586+
spacing, origin)
587+
588+
589+
class Coeffs(_HostTable):
590+
"""Per-dim `(1 - frac, frac)` interpolation weights, shape ``(npoint, 2)``."""
591+
592+
__rkwargs__ = _HostTable.__rkwargs__ + ('shifts', 'dim_index')
593+
594+
def __init_finalize__(self, *args, **kwargs):
595+
self._shifts = kwargs.pop('shifts', None)
596+
self._dim_index = kwargs.pop('dim_index')
597+
super().__init_finalize__(*args, **kwargs)
598+
599+
@property
600+
def shifts(self):
601+
return self._shifts
602+
603+
@property
604+
def dim_index(self):
605+
return self._dim_index
606+
607+
def _emit(self, coords, spacing, origin):
608+
return _linear_weights(coords, self._parent, self._shifts,
609+
self._dim_index, self.dtype, spacing, origin)
610+
611+
612+
def _resolved_geometry(grid, kwargs):
613+
"""Fp64 (spacing, origin) tuple honoring runtime `h_x`/`o_x`/... overrides."""
614+
spacing = np.array([as_fp64_decimal(kwargs.get(s.name, v)) for s, v
615+
in zip(grid.spacing_symbols, grid.spacing, strict=True)])
616+
origin = np.array([as_fp64_decimal(kwargs.get(o.name, v)) for o, v
617+
in zip(grid.origin_symbols, grid.origin, strict=True)])
618+
return spacing, origin
619+
620+
621+
def _positions_fp64(coords, sfunc, shifts, spacing, origin):
622+
c64 = np.asarray(coords, dtype=np.float64)
623+
return (c64 - origin - _shift_values(shifts, sfunc.grid, spacing)) / spacing
624+
625+
626+
def _cell_indices(coords, sfunc, shifts, spacing, origin):
627+
return np.floor(_positions_fp64(coords, sfunc, shifts, spacing,
628+
origin)).astype(np.int32)
629+
630+
631+
def _linear_weights(coords, sfunc, shifts, j, dtype, spacing, origin):
632+
pos = _positions_fp64(coords, sfunc, shifts, spacing, origin)
633+
frac = pos[:, j] - np.floor(pos[:, j])
634+
data = np.empty((pos.shape[0], 2), dtype=dtype)
635+
data[:, 0] = 1.0 - frac
636+
data[:, 1] = frac
637+
return data
638+
639+
513640
class LinearInterpolator(WeightedInterpolator):
514641
"""
515-
Concrete implementation of WeightedInterpolator implementing a Linear interpolation
516-
scheme, i.e. Bilinear for 2D and Trilinear for 3D problems.
642+
Linear (bilinear/trilinear) interpolator.
517643
518-
Parameters
519-
----------
520-
sfunction: The SparseFunction that this Interpolator operates on.
644+
Gridpoints and per-dim `(1-frac, frac)` weights are precomputed on the
645+
host in fp64 (see `_arg_defaults`) and passed to the kernel as int32/fp
646+
SubFunctions. The generated C only indexes those tables and never sees
647+
`(c-o)/h` or `floor` on fp32.
521648
"""
522649

523650
_name = 'linear'
524651

652+
@cached_property
653+
def _coeff_dtype(self):
654+
# Weights are real even for complex-valued sparse fields.
655+
dtype = np.dtype(self.sfunction.dtype)
656+
if np.issubdtype(dtype, np.complexfloating):
657+
return np.finfo(dtype).dtype.type
658+
return dtype.type
659+
525660
@memoized_meth
526-
def _weights(self, subdomain=None, shifts=None):
527-
rdim = self._rdim(subdomain=subdomain, shifts=shifts)
528-
c = [(1 - p) * (1 - r) + p * r
529-
for (p, d, r) in zip(self._point_symbols(shifts), self._gdims, rdim,
530-
strict=True)]
531-
return Mul(*c)
661+
def _gridpoints(self, shifts=None):
662+
name = f'gp{self.sfunction.name}{_shift_tag(shifts)}'
663+
sfdim = self.sfunction._sparse_dim
664+
ddim = CustomDimension(f'{name}d', 0, self.grid.dim - 1,
665+
self.grid.dim, sfdim)
666+
return Gridpoints(name=name, dtype=np.int32,
667+
shape=(self.sfunction.npoint, self.grid.dim),
668+
dimensions=(sfdim, ddim), space_order=0,
669+
alias=self.sfunction.alias,
670+
parent=self.sfunction, shifts=shifts)
532671

533672
@memoized_meth
534-
def _point_symbols(self, shifts=None):
535-
"""Symbol for coordinate value in each Dimension of the point."""
536-
dtype = self.sfunction.coordinates.dtype
537-
symbols = []
538-
for d in self.grid.dimensions:
539-
if shifts and shifts[self.grid.dimensions.index(d)] != 0:
540-
symbols.append(Symbol(name=f'p{d}_s1', dtype=dtype))
541-
else:
542-
symbols.append(Symbol(name=f'p{d}', dtype=dtype))
543-
return DimensionTuple(*symbols, getters=self.grid.dimensions)
673+
def _coeffs(self, shifts=None):
674+
sfdim = self.sfunction._sparse_dim
675+
tag = _shift_tag(shifts)
676+
return tuple(
677+
Coeffs(name=f'w{self.sfunction.name}{d.name}{tag}',
678+
dtype=self._coeff_dtype,
679+
shape=(self.sfunction.npoint, 2),
680+
dimensions=(sfdim, r), space_order=0,
681+
alias=self.sfunction.alias,
682+
parent=self.sfunction, shifts=shifts, dim_index=i)
683+
for i, (d, r) in enumerate(zip(self._gdims, self._cdim, strict=True))
684+
)
685+
686+
def _positions(self, implicit_dims, shifts=None):
687+
gp = self._gridpoints(shifts=shifts)
688+
ddim = gp.dimensions[-1]
689+
return [Eq(p, gp._subs(ddim, di), implicit_dims=implicit_dims)
690+
for (di, p) in enumerate(
691+
self.sfunction._pos_symbols(shifts=shifts))]
544692

545693
def _coeff_temps(self, implicit_dims, shifts=None):
546-
# Positions
547-
pmap = self.sfunction._position_map(shifts=shifts)
548-
psyms = self._point_symbols(shifts)
549-
poseq = [Eq(psyms[d], pos - floor(pos),
550-
implicit_dims=implicit_dims)
551-
for (d, pos) in zip(self._gdims, pmap.keys(), strict=True)]
552-
return poseq
694+
return []
695+
696+
@memoized_meth
697+
def _weights(self, subdomain=None, shifts=None):
698+
rdims = self._rdim(subdomain=subdomain, shifts=shifts)
699+
coeffs = self._coeffs(shifts=shifts)
700+
return Mul(*[
701+
w._subs(rd, rd - rd.parent.symbolic_min)
702+
for (rd, w) in zip(rdims, coeffs, strict=True)
703+
])
553704

554705

555706
class PrecomputedInterpolator(WeightedInterpolator):
@@ -634,7 +785,7 @@ def _weights(self, subdomain=None, shifts=None):
634785
for (rd, w) in zip(rdims, self.interpolation_coeffs, strict=True)
635786
])
636787

637-
def _arg_defaults(self, coords=None, sfunc=None):
788+
def _arg_defaults(self, coords=None, sfunc=None, origin=None):
638789
args = {}
639790
b = self._b_table[self.r]
640791
b0 = i0(b)

‎devito/tools/dtypes_lowering.py‎

Lines changed: 11 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -16,7 +16,17 @@
1616
'dtype_to_cstr', 'dtype_to_ctype', 'infer_datasize', 'dtype_to_mpitype',
1717
'dtype_len', 'ctypes_to_cstr', 'c_restrict_void_p', 'ctypes_vector_mapper',
1818
'is_external_ctype', 'infer_dtype', 'extract_dtype', 'CustomDtype',
19-
'mpi4py_mapper']
19+
'mpi4py_mapper', 'as_fp64_decimal']
20+
21+
22+
def as_fp64_decimal(v):
23+
"""
24+
fp64 value of ``v`` matching its shortest round-tripping decimal.
25+
For an `np.float32` this recovers the decimal the user wrote (e.g.
26+
``np.float32(0.1)`` -> ``0.1`` exact in fp64) rather than the widened
27+
fp32 bit pattern (``0.10000000149...``).
28+
"""
29+
return np.float64(np.format_float_positional(v, unique=True, trim='0'))
2030

2131

2232
# *** Custom np.dtypes

‎devito/types/sparse.py‎

Lines changed: 6 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -760,8 +760,12 @@ def _arg_values(self, estimate_memory=False, **kwargs):
760760
values = new._arg_defaults(alias=self,
761761
estimate_memory=estimate_memory).reduce_all()
762762
else:
763-
# We've been provided a pure-data replacement (array)
764-
values = {}
763+
# Pure-data replacement (ndarray). Re-derive full defaults so
764+
# any interpolator-owned SubFunctions get rebuilt alongside
765+
# the scattered data.
766+
values = self._arg_defaults(
767+
alias=self, estimate_memory=estimate_memory
768+
).reduce_all()
765769
for k, v in self._dist_scatter(data=new).items():
766770
values[k.name] = v
767771
for i, s in zip(k.indices, v.shape, strict=True):
@@ -995,7 +999,6 @@ def _arg_defaults(self, alias=None, estimate_memory=False):
995999
defaults = super()._arg_defaults(alias=alias, estimate_memory=estimate_memory)
9961000
if estimate_memory:
9971001
return defaults
998-
9991002
key = alias or self
10001003
coords = defaults.get(key.coordinates.name, key.coordinates.data)
10011004
defaults.update(key.interpolator._arg_defaults(coords=coords,

‎pyproject.toml‎

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -12,7 +12,7 @@ description = "Finite Difference DSL for symbolic computation."
1212
license = { file = "LICENSE.md" }
1313
readme = "README.md"
1414
keywords = ["finite-difference", "DSL", "symbolic", "jit", "devito"]
15-
requires-python = ">=3.10,<3.14"
15+
requires-python = ">=3.10,<3.15"
1616
authors = [
1717
{ name = "Imperial College London", email = "g.gorman@imperial.ac.uk" },
1818
{ name = "Fabio Luporini", email = "fabio@devitocodes.com" },
@@ -32,6 +32,8 @@ classifiers = [
3232
"Programming Language :: Python :: 3.10",
3333
"Programming Language :: Python :: 3.11",
3434
"Programming Language :: Python :: 3.12",
35+
"Programming Language :: Python :: 3.13",
36+
"Programming Language :: Python :: 3.14",
3537
"Programming Language :: Python :: 3 :: Only",
3638
"Topic :: Scientific/Engineering",
3739
"Topic :: Scientific/Engineering :: Mathematics",

‎requirements-optional.txt‎

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1,7 +1,7 @@
11
matplotlib<3.10.10
22
pillow>11,<12.2.1
3-
pyrevolve==2.2.7
4-
scipy<1.15.4
3+
pyrevolve==2.2.8
4+
scipy>=1.8.0,<1.18.1
55
distributed<2026.7.2
66
click<9.0
77
cloudpickle<3.1.3

‎requirements-testing.txt‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -3,7 +3,7 @@ pytest-runner<6.0.2
33
pytest-cov<7.1.1
44
flake8-pyproject>=1.2.3,<1.2.5
55
nbval<0.11.1
6-
scipy<1.15.4
6+
scipy>=1.8.0,<1.18.1
77
pooch<1.9.1
88
click<9.0
99
cloudpickle<3.1.3

0 commit comments

Comments
 (0)