|
15 | 15 | from devito.finite_differences.elementary import floor |
16 | 16 | from devito.logger import warning |
17 | 17 | 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 |
20 | 22 | from devito.types.utils import DimensionTuple |
21 | 23 |
|
22 | 24 | __all__ = ['LinearInterpolator', 'PrecomputedInterpolator', 'SincInterpolator'] |
@@ -510,46 +512,195 @@ def _inject(self, field, expr, implicit_dims=None): |
510 | 512 | return filter_ordered(temps) + eqns |
511 | 513 |
|
512 | 514 |
|
| 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 | + |
513 | 640 | class LinearInterpolator(WeightedInterpolator): |
514 | 641 | """ |
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. |
517 | 643 |
|
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. |
521 | 648 | """ |
522 | 649 |
|
523 | 650 | _name = 'linear' |
524 | 651 |
|
| 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 | + |
525 | 660 | @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) |
532 | 671 |
|
533 | 672 | @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))] |
544 | 692 |
|
545 | 693 | 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 | + ]) |
553 | 704 |
|
554 | 705 |
|
555 | 706 | class PrecomputedInterpolator(WeightedInterpolator): |
@@ -634,7 +785,7 @@ def _weights(self, subdomain=None, shifts=None): |
634 | 785 | for (rd, w) in zip(rdims, self.interpolation_coeffs, strict=True) |
635 | 786 | ]) |
636 | 787 |
|
637 | | - def _arg_defaults(self, coords=None, sfunc=None): |
| 788 | + def _arg_defaults(self, coords=None, sfunc=None, origin=None): |
638 | 789 | args = {} |
639 | 790 | b = self._b_table[self.r] |
640 | 791 | b0 = i0(b) |
|
0 commit comments