diff --git a/docs/sphinx/api-interpolation.md b/docs/sphinx/api-interpolation.md new file mode 100644 index 00000000..cbdbef7f --- /dev/null +++ b/docs/sphinx/api-interpolation.md @@ -0,0 +1,10 @@ +# Interpolation Functions + +```{eval-rst} +.. currentmodule:: array_api_extra +.. autosummary:: + :nosignatures: + :toctree: generated + + interp +``` diff --git a/docs/sphinx/api-reference.md b/docs/sphinx/api-reference.md index 7cc326d7..b18ac21b 100644 --- a/docs/sphinx/api-reference.md +++ b/docs/sphinx/api-reference.md @@ -7,6 +7,7 @@ api-creation.md api-elementwise.md api-indexing.md api-inspection.md +api-interpolation.md api-linalg.md api-manipulation.md api-searching.md diff --git a/meson.build b/meson.build index ebc124e3..cb9fb91c 100644 --- a/meson.build +++ b/meson.build @@ -14,6 +14,7 @@ sources = { 'src/array_api_extra/_creation.py', 'src/array_api_extra/_elementwise.py', 'src/array_api_extra/_indexing.py', + 'src/array_api_extra/_interpolation.py', 'src/array_api_extra/_lazy.py', 'src/array_api_extra/_linalg.py', 'src/array_api_extra/_manipulation.py', @@ -29,6 +30,7 @@ sources = { 'src/array_api_extra/_agnostic/_elementwise.py', 'src/array_api_extra/_agnostic/_indexing.py', 'src/array_api_extra/_agnostic/_inspection.py', + 'src/array_api_extra/_agnostic/_interpolation.py', 'src/array_api_extra/_agnostic/_linalg.py', 'src/array_api_extra/_agnostic/_manipulation.py', 'src/array_api_extra/_agnostic/_searching.py', diff --git a/src/array_api_extra/__init__.py b/src/array_api_extra/__init__.py index c670fe44..4d94e2f8 100644 --- a/src/array_api_extra/__init__.py +++ b/src/array_api_extra/__init__.py @@ -7,6 +7,7 @@ from ._creation import create_diagonal, one_hot from ._elementwise import deg2rad, isclose, nan_to_num, rad2deg, sinc from ._indexing import diag_indices, tril_indices, triu_indices, unravel_index +from ._interpolation import interp from ._lazy import lazy_apply from ._linalg import kron from ._manipulation import atleast_nd, broadcast_shapes, expand_dims, pad @@ -31,6 +32,7 @@ "deg2rad", "diag_indices", "expand_dims", + "interp", "isclose", "isin", "kron", diff --git a/src/array_api_extra/_agnostic/__init__.py b/src/array_api_extra/_agnostic/__init__.py index 8af17333..44e62e0e 100644 --- a/src/array_api_extra/_agnostic/__init__.py +++ b/src/array_api_extra/_agnostic/__init__.py @@ -5,6 +5,7 @@ _elementwise, _indexing, _inspection, + _interpolation, _linalg, _manipulation, _searching, @@ -18,6 +19,7 @@ "_elementwise", "_indexing", "_inspection", + "_interpolation", "_linalg", "_manipulation", "_searching", diff --git a/src/array_api_extra/_agnostic/_interpolation.py b/src/array_api_extra/_agnostic/_interpolation.py new file mode 100644 index 00000000..76b40f49 --- /dev/null +++ b/src/array_api_extra/_agnostic/_interpolation.py @@ -0,0 +1,331 @@ +"""Array-agnostic implementations for interpolation functions.""" + +import math +from typing import cast + +from .._lib import _compat +from .._lib._typing import Array, ArrayNamespace, DType + +__all__ = ["interp"] + + +def _safe_difference(x1: Array, x2: Array, /, *, xp: ArrayNamespace) -> Array: + """Subtract finite arrays with IEEE overflow results but without warnings.""" + zero = xp.zeros_like(x1) + maximum = xp.asarray( + xp.finfo(x1.dtype).max, dtype=x1.dtype, device=_compat.device(x1) + ) + negative_x2 = xp.where(x2 < 0, x2, zero) + positive_x2 = xp.where(x2 > 0, x2, zero) + positive_overflow = (x1 > 0) & (x2 < 0) & (x1 > maximum + negative_x2) + negative_overflow = (x1 < 0) & (x2 > 0) & (x1 < -maximum + positive_x2) + overflow = positive_overflow | negative_overflow + + out = xp.where(overflow, zero, x1) - xp.where(overflow, zero, x2) + infinity = xp.asarray(math.inf, dtype=x1.dtype, device=_compat.device(x1)) + out = xp.where(positive_overflow, infinity, out) + return xp.where(negative_overflow, -infinity, out) + + +def _safe_divide( + numerator: Array, denominator: Array, /, *, xp: ArrayNamespace +) -> Array: + """Divide finite arrays with IEEE overflow results but without warnings.""" + zero = xp.zeros_like(numerator) + one = xp.ones_like(denominator) + maximum = xp.asarray( + xp.finfo(numerator.dtype).max, + dtype=numerator.dtype, + device=_compat.device(numerator), + ) + absolute_denominator = xp.abs(denominator) + small_denominator = absolute_denominator < 1 + overflow_threshold = maximum * xp.where( + small_denominator, absolute_denominator, zero + ) + overflow = ( + (numerator != 0) + & (denominator != 0) + & small_denominator + & (xp.abs(numerator) > overflow_threshold) + ) + + out = xp.where(overflow, zero, numerator) / xp.where(overflow, one, denominator) + infinity = xp.asarray( + math.inf, dtype=numerator.dtype, device=_compat.device(numerator) + ) + positive_overflow = overflow & ((numerator > 0) == (denominator > 0)) + out = xp.where(positive_overflow, infinity, out) + return xp.where(overflow & ~positive_overflow, -infinity, out) + + +def _safe_multiply( + x1: Array, x2: Array, /, *, zero_times_infinity: float, xp: ArrayNamespace +) -> Array: + """Multiply arrays without warnings from overflow or zero times infinity.""" + zero = xp.zeros_like(x1) + one = xp.ones_like(x2) + maximum = xp.asarray( + xp.finfo(x1.dtype).max, dtype=x1.dtype, device=_compat.device(x1) + ) + absolute_x1 = xp.abs(x1) + absolute_x2 = xp.abs(x2) + large_x2 = absolute_x2 > 1 + overflow_threshold = maximum / xp.where(large_x2, absolute_x2, one) + overflow = ( + xp.isfinite(x1) + & xp.isfinite(x2) + & large_x2 + & (absolute_x1 > overflow_threshold) + ) + zero_inf = ((x1 == 0) & xp.isinf(x2)) | (xp.isinf(x1) & (x2 == 0)) + suppressed = overflow | zero_inf + + out = xp.where(suppressed, zero, x1) * xp.where(suppressed, zero, x2) + infinity = xp.asarray(math.inf, dtype=x1.dtype, device=_compat.device(x1)) + positive_overflow = overflow & ((x1 > 0) == (x2 > 0)) + out = xp.where(positive_overflow, infinity, out) + out = xp.where(overflow & ~positive_overflow, -infinity, out) + zero_inf_value = xp.asarray( + zero_times_infinity, dtype=x1.dtype, device=_compat.device(x1) + ) + return xp.where(zero_inf, zero_inf_value, out) + + +def _interp_component( + x: Array, + x_lo: Array, + x_hi: Array, + y_lo: Array, + y_hi: Array, + /, + *, + inactive: Array, + reciprocal_first: bool = False, + xp: ArrayNamespace, +) -> Array: + """Interpolate one real component without invalid arithmetic.""" + coordinate_infinite = xp.isinf(x_lo) | xp.isinf(x_hi) + values_finite = xp.isfinite(y_lo) & xp.isfinite(y_hi) + regular = ~inactive & ~coordinate_infinite & values_finite + + safe_x = xp.where(regular, x, xp.zeros_like(x)) + safe_x_lo = xp.where(regular, x_lo, xp.zeros_like(x_lo)) + safe_x_hi = xp.where(regular, x_hi, xp.ones_like(x_hi)) + safe_y_lo = xp.where(regular, y_lo, xp.zeros_like(y_lo)) + safe_y_hi = xp.where(regular, y_hi, xp.zeros_like(y_hi)) + device = _compat.device(y_lo) + nan = xp.asarray(math.nan, dtype=y_lo.dtype, device=device) + equal_values = y_lo == y_hi + + coordinate_difference = _safe_difference(safe_x_hi, safe_x_lo, xp=xp) + value_difference = _safe_difference(safe_y_hi, safe_y_lo, xp=xp) + indeterminate_slope = xp.isinf(coordinate_difference) & xp.isinf(value_difference) + safe_value_difference = xp.where( + indeterminate_slope, xp.zeros_like(value_difference), value_difference + ) + safe_coordinate_difference = xp.where( + indeterminate_slope, + xp.ones_like(coordinate_difference), + coordinate_difference, + ) + if reciprocal_first: + inverse_coordinate_difference = _safe_divide( + xp.ones_like(safe_coordinate_difference), + safe_coordinate_difference, + xp=xp, + ) + slope = _safe_multiply( + safe_value_difference, + inverse_coordinate_difference, + zero_times_infinity=math.nan, + xp=xp, + ) + else: + slope = _safe_divide(safe_value_difference, safe_coordinate_difference, xp=xp) + slope = xp.where(indeterminate_slope, nan, slope) + + left_delta = _safe_difference(safe_x, safe_x_lo, xp=xp) + left_invalid = (slope == 0) & xp.isinf(left_delta) + left_out = ( + xp.where(left_invalid, xp.zeros_like(slope), slope) + * xp.where(left_invalid, xp.zeros_like(left_delta), left_delta) + + safe_y_lo + ) + retry = left_invalid | xp.isnan(left_out) + + right_delta = _safe_difference(safe_x, safe_x_hi, xp=xp) + right_invalid = (slope == 0) & xp.isinf(right_delta) + right_out = ( + xp.where(right_invalid, xp.zeros_like(slope), slope) + * xp.where(right_invalid, xp.zeros_like(right_delta), right_delta) + + safe_y_hi + ) + right_failed = right_invalid | xp.isnan(right_out) + out = xp.where(retry, right_out, left_out) + out = xp.where(retry & right_failed, nan, out) + out = xp.where(retry & right_failed & equal_values, safe_y_lo, out) + + finite_coordinates_nonfinite_values = ~coordinate_infinite & ~values_finite + nonfinite_value_out = xp.where( + xp.isfinite(y_lo), + y_hi, + xp.where(xp.isfinite(y_hi), y_lo, xp.where(equal_values, y_lo, nan)), + ) + out = xp.where(finite_coordinates_nonfinite_values, nonfinite_value_out, out) + + only_lo_coordinate_infinite = xp.isinf(x_lo) & ~xp.isinf(x_hi) + only_hi_coordinate_infinite = ~xp.isinf(x_lo) & xp.isinf(x_hi) + finite_y_lo = xp.where(values_finite, y_lo, xp.zeros_like(y_lo)) + finite_y_hi = xp.where(values_finite, y_hi, xp.zeros_like(y_hi)) + infinite_coordinate_value_overflow = values_finite & xp.isinf( + _safe_difference(finite_y_hi, finite_y_lo, xp=xp) + ) + infinite_coordinate_out = xp.where( + infinite_coordinate_value_overflow, + nan, + xp.where( + values_finite & only_lo_coordinate_infinite, + y_hi, + xp.where( + values_finite & only_hi_coordinate_infinite, + y_lo, + xp.where(equal_values, y_lo, nan), + ), + ), + ) + return xp.where(coordinate_infinite, infinite_coordinate_out, out) + + +def _combine_complex( + real: Array, imag: Array, /, *, dtype: DType, xp: ArrayNamespace +) -> Array: + """Combine real components without multiplying zero by an infinity.""" + device = _compat.device(real) + real_complex = xp.astype(real, dtype) + finite_imag = xp.where(xp.isfinite(imag), imag, xp.zeros_like(imag)) + out = real_complex + xp.astype(finite_imag, dtype) * 1j + + positive_inf = xp.asarray(complex(0.0, math.inf), dtype=dtype, device=device) + negative_inf = xp.asarray(complex(0.0, -math.inf), dtype=dtype, device=device) + imaginary_nan = xp.asarray(complex(0.0, math.nan), dtype=dtype, device=device) + out = xp.where(xp.isinf(imag) & (imag > 0), real_complex + positive_inf, out) + out = xp.where(xp.isinf(imag) & (imag < 0), real_complex + negative_inf, out) + return xp.where(xp.isnan(imag), real_complex + imaginary_nan, out) + + +def interp( + x: Array, + x_points: Array, + values: Array, + /, + *, + left: Array | None, + right: Array | None, + period: float | None, + xp: ArrayNamespace, +) -> Array: + # numpydoc ignore=PR01,RT01 + """See docstring in `array_api_extra._interpolation`.""" + if period is not None: + x = x % period + x_points = x_points % period + order = xp.argsort(x_points, stable=True) + x_points = xp.take(x_points, order, axis=0) + if xp.isdtype(values.dtype, "complex floating"): + values = _combine_complex( + xp.take(xp.real(values), order, axis=0), + xp.take(xp.imag(values), order, axis=0), + dtype=values.dtype, + xp=xp, + ) + else: + values = xp.take(values, order, axis=0) + x_points = xp.concat( + (x_points[-1:] - period, x_points, x_points[:1] + period), axis=0 + ) + values = xp.concat((values[-1:], values, values[:1]), axis=0) + + x_shape = x.shape + x_flat = xp.reshape(x, (-1,)) + n_points = cast(int, x_points.shape[0]) + + if period is None and n_points == 1: + out = xp.broadcast_to(values[0], x_flat.shape) + left_array = values[0] if left is None else left + right_array = values[-1] if right is None else right + out = xp.where(x_flat < x_points[0], left_array, out) + out = xp.where(x_flat > x_points[-1], right_array, out) + return xp.reshape(out, x_shape) + + right_indices = xp.searchsorted(x_points, x_flat, side="right") + exact_indices = xp.clip(right_indices - 1, 0, n_points - 1) + exact_x = xp.take(x_points, exact_indices, axis=0) + exact = x_flat == exact_x + + interval_indices = xp.clip(right_indices - 1, 0, n_points - 2) + x_lo = xp.take(x_points, interval_indices, axis=0) + x_hi = xp.take(x_points, interval_indices + 1, axis=0) + + below = x_flat < x_points[0] + above = x_flat > x_points[-1] + query_nan = xp.isnan(x_flat) + inactive = exact | below | above | query_nan | (x_lo == x_hi) + + if xp.isdtype(values.dtype, "complex floating"): + values_real = xp.real(values) + values_imag = xp.imag(values) + exact_real = xp.take(values_real, exact_indices, axis=0) + exact_imag = xp.take(values_imag, exact_indices, axis=0) + y_lo_real = xp.take(values_real, interval_indices, axis=0) + y_hi_real = xp.take(values_real, interval_indices + 1, axis=0) + y_lo_imag = xp.take(values_imag, interval_indices, axis=0) + y_hi_imag = xp.take(values_imag, interval_indices + 1, axis=0) + out_real = _interp_component( + x_flat, + x_lo, + x_hi, + y_lo_real, + y_hi_real, + inactive=inactive, + reciprocal_first=True, + xp=xp, + ) + out_imag = _interp_component( + x_flat, + x_lo, + x_hi, + y_lo_imag, + y_hi_imag, + inactive=inactive, + reciprocal_first=True, + xp=xp, + ) + left_array = values[0] if left is None else left + right_array = values[-1] if right is None else right + out_real = xp.where(exact, exact_real, out_real) + out_imag = xp.where(exact, exact_imag, out_imag) + out_real = xp.where(below, xp.real(left_array), out_real) + out_imag = xp.where(below, xp.imag(left_array), out_imag) + out_real = xp.where(above, xp.real(right_array), out_real) + out_imag = xp.where(above, xp.imag(right_array), out_imag) + nan = xp.asarray(math.nan, dtype=out_real.dtype, device=_compat.device(values)) + out_real = xp.where(query_nan, nan, out_real) + out_imag = xp.where(query_nan, xp.zeros_like(out_imag), out_imag) + out = _combine_complex(out_real, out_imag, dtype=values.dtype, xp=xp) + else: + exact_y = xp.take(values, exact_indices, axis=0) + y_lo = xp.take(values, interval_indices, axis=0) + y_hi = xp.take(values, interval_indices + 1, axis=0) + out = _interp_component( + x_flat, x_lo, x_hi, y_lo, y_hi, inactive=inactive, xp=xp + ) + out = xp.where(exact, exact_y, out) + left_array = values[0] if left is None else left + right_array = values[-1] if right is None else right + out = xp.where(below, left_array, out) + out = xp.where(above, right_array, out) + nan = xp.asarray(math.nan, dtype=values.dtype, device=_compat.device(values)) + out = xp.where(query_nan, nan, out) + + return xp.reshape(out, x_shape) diff --git a/src/array_api_extra/_interpolation.py b/src/array_api_extra/_interpolation.py new file mode 100644 index 00000000..7d5217b7 --- /dev/null +++ b/src/array_api_extra/_interpolation.py @@ -0,0 +1,326 @@ +"""Delegation layer for interpolation functions.""" + +import math +from numbers import Number, Real + +from . import _agnostic +from ._lib import _compat, _helpers +from ._lib._typing import Array, ArrayNamespace, Device, DType + +__all__ = ["interp"] + + +def _is_python_real_scalar(x: object, /) -> bool: + if x is True or x is False: + return False + return _helpers.is_python_scalar(x) and isinstance(x, Real) + + +def _is_real_scalar(x: object, /) -> bool: + if x is True or x is False: + return False + return isinstance(x, Real) + + +def _same_namespace(xp1: ArrayNamespace, xp2: ArrayNamespace, /) -> bool: + predicates = ( + _compat.is_array_api_strict_namespace, + _compat.is_cupy_namespace, + _compat.is_dask_namespace, + _compat.is_jax_namespace, + _compat.is_numpy_namespace, + _compat.is_pydata_sparse_namespace, + _compat.is_torch_namespace, + ) + return xp1 is xp2 or any( + predicate(xp1) and predicate(xp2) for predicate in predicates + ) + + +def _require_dtype(xp: ArrayNamespace, name: str, /, *, device: Device) -> DType: + try: + dtype = xp.__array_namespace_info__().dtypes(device=device)[name] + except (AttributeError, KeyError, TypeError) as error: + msg = f"`interp` requires {name} support on the selected device." + raise TypeError(msg) from error + return dtype + + +def _astype_required( + x: Array, dtype: DType, /, *, name: str, xp: ArrayNamespace +) -> Array: + out = xp.astype(x, dtype, copy=False) + if out.dtype != dtype: + msg = f"`interp` requires {dtype!s} support for `{name}`." + raise TypeError(msg) + return out + + +def _validate_bound( + bound: Array | complex | None, + /, + *, + name: str, + values_are_complex: bool, + xp: ArrayNamespace, +) -> None: + if bound is None: + return + + if _compat.is_array_api_obj(bound): + if bound.ndim != 0: + msg = f"`{name}` must be a numerical scalar or a 0-dimensional array." + raise ValueError(msg) + if not xp.isdtype( + bound.dtype, ("integral", "real floating", "complex floating") + ): + msg = f"`{name}` must have a real or complex numeric dtype." + raise TypeError(msg) + bound_is_complex = xp.isdtype(bound.dtype, "complex floating") + else: + if bound is True or bound is False or not isinstance(bound, Number): + msg = f"`{name}` must be a numerical scalar or a 0-dimensional array." + raise TypeError(msg) + bound_is_complex = isinstance(bound, complex) and not isinstance(bound, Real) + + if bound_is_complex and not values_are_complex: + msg = f"`{name}` must be real when `values` is real-valued." + raise TypeError(msg) + + +def _as_bound_array( + bound: Array | complex | None, + /, + *, + dtype: DType, + device: Device, + name: str, + xp: ArrayNamespace, +) -> Array | None: + if bound is None: + return None + if _compat.is_array_api_obj(bound): + return _astype_required(bound, dtype, name=name, xp=xp) + out = xp.asarray(bound, dtype=dtype, device=device) + if out.dtype != dtype: + msg = f"`interp` requires {dtype!s} support for `{name}`." + raise TypeError(msg) + return out + + +def interp( + x: Array | float, + x_points: Array, + values: Array, + /, + *, + left: Array | complex | None = None, + right: Array | complex | None = None, + period: float | Real | None = None, + xp: ArrayNamespace | None = None, +) -> Array: + """ + One-dimensional piecewise linear interpolation. + + Evaluate the piecewise linear function defined by sample coordinates + `x_points` and sample `values` at the query coordinates `x`. + + Parameters + ---------- + x : Array or real scalar + Query coordinates. Arrays may have any shape and must have an integral or + real floating-point dtype. A scalar must be a Python real scalar. Boolean + values are not supported. + x_points : Array + One-dimensional, nonempty sample coordinates with an integral or real + floating-point dtype. When `period` is ``None``, coordinates must be in + non-decreasing order; this precondition is not checked. NaN sample + coordinates are not supported. + values : Array + One-dimensional sample values with a real or complex numeric dtype. Its + length must match `x_points`. + left : numerical scalar or 0-dimensional Array, optional + Value returned for queries below the first sample coordinate. By default, + the first element of `values` is used. Ignored, without validation, when + `period` is provided. + right : numerical scalar or 0-dimensional Array, optional + Value returned for queries above the last sample coordinate. By default, + the last element of `values` is used. Ignored, without validation, when + `period` is provided. + period : real scalar, optional + Period for the sample and query coordinates. It must remain finite and + nonzero when converted to ``float64``. A negative value is treated as its + absolute value. When provided, coordinates are normalized to the period + and samples are sorted by normalized coordinate. + xp : array_namespace, optional + The standard-compatible namespace for the array arguments. Default: infer. + + Returns + ------- + Array + Interpolated values with the same shape as `x`. A scalar query produces a + 0-dimensional array. Coordinates are evaluated in ``float64``. Real sample + values produce ``float64`` output and complex sample values produce + ``complex128`` output. + + Notes + ----- + Without `period`, an exact repeated sample coordinate uses the last corresponding + value. Ties between coordinates that become equal after periodic normalization + are backend-dependent. The selected backend and device must support ``float64`` + and, for complex `values`, ``complex128``. Native NumPy and CuPy implementations + are used when available; other namespaces use the array-agnostic implementation. + + Examples + -------- + >>> import array_api_extra as xpx + >>> import array_api_strict as xp + >>> x_points = xp.asarray([0.0, 1.0, 2.0]) + >>> values = xp.asarray([0.0, 10.0, 20.0]) + >>> xpx.interp(xp.asarray([0.5, 1.5]), x_points, values, xp=xp) + Array([ 5., 15.], dtype=array_api_strict.float64) + """ + if not _compat.is_array_api_obj(x_points): + msg = "`x_points` must be an array." + raise TypeError(msg) + if not _compat.is_array_api_obj(values): + msg = "`values` must be an array." + raise TypeError(msg) + + x_is_scalar = _is_python_real_scalar(x) + x_is_array = _compat.is_array_api_obj(x) + if not x_is_scalar and not x_is_array: + msg = "`x` must be an array or a Python real scalar." + raise TypeError(msg) + + if period is not None: + if not _is_real_scalar(period): + msg = "`period` must be a finite real scalar or None." + raise TypeError(msg) + try: + period = float(abs(period)) + except (OverflowError, ValueError) as error: + msg = "`period` must be representable as a finite float." + raise ValueError(msg) from error + if not math.isfinite(period): + msg = "`period` must be finite after conversion to float." + raise ValueError(msg) + if period == 0: + msg = "`period` must be nonzero after conversion to float." + raise ValueError(msg) + + namespace_args: list[Array] = [x_points, values] + if _compat.is_array_api_obj(x): + namespace_args.append(x) + if period is None: + if _compat.is_array_api_obj(left): + namespace_args.append(left) + if _compat.is_array_api_obj(right): + namespace_args.append(right) + inferred_xp = _compat.array_namespace(*namespace_args) + if xp is None: + xp = inferred_xp + elif not _same_namespace(xp, inferred_xp): + msg = "`xp` must match the namespace of the array arguments." + raise TypeError(msg) + + if x_points.ndim != 1: + msg = "`x_points` must be one-dimensional." + raise ValueError(msg) + if values.ndim != 1: + msg = "`values` must be one-dimensional." + raise ValueError(msg) + (n_points,) = _helpers.eager_shape(x_points) + (n_values,) = _helpers.eager_shape(values) + if n_points == 0: + msg = "`x_points` and `values` must be nonempty." + raise ValueError(msg) + if n_points != n_values: + msg = "`x_points` and `values` must have the same length." + raise ValueError(msg) + + if not xp.isdtype(x_points.dtype, ("integral", "real floating")): + msg = "`x_points` must have an integral or real floating-point dtype." + raise TypeError(msg) + if _compat.is_array_api_obj(x) and not xp.isdtype( + x.dtype, ("integral", "real floating") + ): + msg = "`x` must have an integral or real floating-point dtype." + raise TypeError(msg) + if not xp.isdtype(values.dtype, ("integral", "real floating", "complex floating")): + msg = "`values` must have a real or complex numeric dtype." + raise TypeError(msg) + + values_are_complex = xp.isdtype(values.dtype, "complex floating") + if period is None: + _validate_bound(left, name="left", values_are_complex=values_are_complex, xp=xp) + _validate_bound( + right, name="right", values_are_complex=values_are_complex, xp=xp + ) + + arrays: list[Array] = [x_points, values] + if _compat.is_array_api_obj(x): + arrays.append(x) + if period is None: + if _compat.is_array_api_obj(left): + arrays.append(left) + if _compat.is_array_api_obj(right): + arrays.append(right) + device = _compat.device(x_points) + if any(_compat.device(array) != device for array in arrays[1:]): + msg = "All array arguments must be on the same device." + raise ValueError(msg) + + coordinate_dtype = _require_dtype(xp, "float64", device=device) + value_dtype = _require_dtype( + xp, "complex128" if values_are_complex else "float64", device=device + ) + x_points = _astype_required(x_points, coordinate_dtype, name="x_points", xp=xp) + if _compat.is_array_api_obj(x): + x_array = _astype_required(x, coordinate_dtype, name="x", xp=xp) + else: + x_array = xp.asarray(x, dtype=coordinate_dtype, device=device) + if x_array.dtype != coordinate_dtype: + msg = "`interp` requires float64 support for `x`." + raise TypeError(msg) + values = _astype_required(values, value_dtype, name="values", xp=xp) + + if period is None: + left_array = _as_bound_array( + left, dtype=value_dtype, device=device, name="left", xp=xp + ) + right_array = _as_bound_array( + right, dtype=value_dtype, device=device, name="right", xp=xp + ) + else: + left_array = right_array = None + + native_complex_bounds_unsupported = ( + _compat.is_numpy_namespace(xp) + and values_are_complex + and (left_array is not None or right_array is not None) + ) + if ( + _compat.is_numpy_namespace(xp) or _compat.is_cupy_namespace(xp) + ) and not native_complex_bounds_unsupported: + out = xp.interp( + x_array, + x_points, + values, + left=left_array, + right=right_array, + period=period, + ) + if x_array.ndim == 0: + out = xp.asarray(out, dtype=value_dtype, device=device) + return out + + return _agnostic._interpolation.interp( + x_array, + x_points, + values, + left=left_array, + right=right_array, + period=period, + xp=xp, + ) diff --git a/tests/main/meson.build b/tests/main/meson.build index 720251f7..300d57ae 100644 --- a/tests/main/meson.build +++ b/tests/main/meson.build @@ -8,6 +8,7 @@ py.install_sources([ 'test_helpers.py', 'test_indexing.py', 'test_inspection.py', + 'test_interpolation.py', 'test_lazy.py', 'test_linalg.py', 'test_manipulation.py', diff --git a/tests/main/test_interpolation.py b/tests/main/test_interpolation.py new file mode 100644 index 00000000..33bbc270 --- /dev/null +++ b/tests/main/test_interpolation.py @@ -0,0 +1,505 @@ +from fractions import Fraction +from typing import Any, Literal + +import numpy as np +import pytest + +from array_api_extra import interp as xpx_interp +from array_api_extra._agnostic._interpolation import interp as agnostic_interp +from array_api_extra._lib._backends import Backend +from array_api_extra._lib._compat import array_namespace, is_array_api_obj +from array_api_extra._lib._compat import device as get_device +from array_api_extra._lib._typing import Array, ArrayNamespace, Device +from array_api_extra.testing import assert_close, assert_equal + +Implementation = Literal["public", "agnostic"] + +implementations = pytest.mark.parametrize("implementation", ["public", "agnostic"]) + + +def _interp( + implementation: Implementation, + x: Any, + x_points: Array, + values: Array, + /, + *, + xp: ArrayNamespace, + left: Any = None, + right: Any = None, + period: Any = None, +) -> Array: + if implementation == "public": + return xpx_interp(x, x_points, values, left=left, right=right, period=period) + + coordinate_device = get_device(x_points) + if is_array_api_obj(x): + x = xp.astype(x, xp.float64, copy=False) + else: + x = xp.asarray(x, dtype=xp.float64, device=coordinate_device) + x_points = xp.astype(x_points, xp.float64, copy=False) + + value_dtype = ( + xp.complex128 if xp.isdtype(values.dtype, "complex floating") else xp.float64 + ) + values = xp.astype(values, value_dtype, copy=False) + if period is not None: + period = float(abs(period)) + left = right = None + else: + if left is not None: + left = ( + xp.astype(left, value_dtype, copy=False) + if is_array_api_obj(left) + else xp.asarray(left, dtype=value_dtype, device=get_device(values)) + ) + if right is not None: + right = ( + xp.astype(right, value_dtype, copy=False) + if is_array_api_obj(right) + else xp.asarray(right, dtype=value_dtype, device=get_device(values)) + ) + return agnostic_interp( + x, + x_points, + values, + left=left, + right=right, + period=period, + xp=xp, + ) + + +@pytest.mark.skip_xp_backend(Backend.SPARSE, reason="no searchsorted") +@pytest.mark.skip_xp_backend(Backend.MPARRAY, reason="no searchsorted") +class TestInterp: + @implementations + def test_finite_values_and_shapes( + self, xp: ArrayNamespace, implementation: Implementation + ): + x_points = xp.asarray([0, 1, 2], dtype=xp.float32) + values = xp.asarray([0, 10, 20], dtype=xp.float32) + + actual = _interp( + implementation, + xp.asarray([[-1, 0.5], [1.5, 3]], dtype=xp.float32), + x_points, + values, + xp=xp, + ) + expected = xp.asarray([[0, 5], [15, 20]], dtype=xp.float64) + assert_close(actual, expected) + assert array_namespace(actual) == array_namespace(x_points) + + array_scalar = xp.asarray(0.5, dtype=xp.float64) + queries = (0.5, array_scalar) if implementation == "public" else (array_scalar,) + for x in queries: + actual = _interp(implementation, x, x_points, values, xp=xp) + assert actual.shape == () + assert_equal(actual, xp.asarray(5, dtype=xp.float64)) + + actual = _interp( + implementation, + xp.asarray([], dtype=xp.float64), + x_points, + values, + xp=xp, + ) + assert_equal(actual, xp.asarray([], dtype=xp.float64)) + + @pytest.mark.parametrize( + "values", + [[0.0, 1.0], [0.0 + 0.0j, 1.0 + 2.0j]], + ids=["real", "complex"], + ) + @implementations + def test_reference_narrow_interval( + self, + xp: ArrayNamespace, + implementation: Implementation, + values: list[float] | list[complex], + ): + x_points = np.asarray([0.0, 1e-20]) + x = np.asarray([0.0, 2.5e-21, 1e-20]) + expected = np.interp(x, x_points, values) + dtype = xp.complex128 if np.iscomplexobj(expected) else xp.float64 + + actual = _interp( + implementation, + xp.asarray(x, dtype=xp.float64), + xp.asarray(x_points, dtype=xp.float64), + xp.asarray(values, dtype=dtype), + xp=xp, + ) + assert_close(actual, xp.asarray(expected, dtype=dtype)) + + @implementations + def test_reference_extreme_finite_coordinates( + self, xp: ArrayNamespace, implementation: Implementation + ): + x_points = np.asarray([-1e308, 1e308]) + x = np.asarray([5e307]) + values = np.asarray([0.0, 1.0]) + with np.errstate(over="ignore", invalid="ignore"): + expected = np.interp(x, x_points, values) + + actual = _interp( + implementation, + xp.asarray(x, dtype=xp.float64), + xp.asarray(x_points, dtype=xp.float64), + xp.asarray(values, dtype=xp.float64), + xp=xp, + ) + assert_equal(actual, xp.asarray(expected, dtype=xp.float64)) + + @pytest.mark.parametrize("values", [[0.0, 1.0], [0.0, 1e-320]]) + @implementations + def test_reference_subnormal_interval( + self, + xp: ArrayNamespace, + implementation: Implementation, + values: list[float], + ): + x_points = np.asarray([0.0, 1e-320]) + x = np.asarray([5e-321]) + expected = np.interp(x, x_points, values) + + actual = _interp( + implementation, + xp.asarray(x, dtype=xp.float64), + xp.asarray(x_points, dtype=xp.float64), + xp.asarray(values, dtype=xp.float64), + xp=xp, + ) + assert_equal(actual, xp.asarray(expected, dtype=xp.float64)) + + @pytest.mark.parametrize( + "end_value", [complex(1e-320, 1e-320), complex(1e-320, 0.0)] + ) + @implementations + def test_reference_subnormal_complex_slope( + self, + xp: ArrayNamespace, + implementation: Implementation, + end_value: complex, + ): + x_points = np.asarray([0.0, 1e-320]) + x = np.asarray([5e-321]) + values = np.asarray([0j, end_value]) + expected = np.interp(x, x_points, values) + + actual = _interp( + implementation, + xp.asarray(x, dtype=xp.float64), + xp.asarray(x_points, dtype=xp.float64), + xp.asarray(values, dtype=xp.complex128), + left=0, + xp=xp, + ) + assert_equal(xp.real(actual), xp.asarray(expected.real, dtype=xp.float64)) + assert_equal(xp.imag(actual), xp.asarray(expected.imag, dtype=xp.float64)) + + @pytest.mark.parametrize( + ("x_points", "x"), + [([-np.inf, 0.0], [-1.0]), ([0.0, np.inf], [1.0])], + ids=["left", "right"], + ) + @pytest.mark.parametrize("complex_values", [False, True], ids=["real", "complex"]) + @implementations + def test_reference_infinite_coordinate_value_overflow( + self, + xp: ArrayNamespace, + implementation: Implementation, + x_points: list[float], + x: list[float], + complex_values: bool, + ): + values = np.asarray([-1e308, 1e308]) + dtype = xp.float64 + if complex_values: + values = values + values * 1j + dtype = xp.complex128 + expected = np.interp(x, x_points, values) + + actual = _interp( + implementation, + xp.asarray(x, dtype=xp.float64), + xp.asarray(x_points, dtype=xp.float64), + xp.asarray(values, dtype=dtype), + xp=xp, + ) + if complex_values: + assert_equal(xp.real(actual), xp.asarray(expected.real, dtype=xp.float64)) + assert_equal(xp.imag(actual), xp.asarray(expected.imag, dtype=xp.float64)) + else: + assert_equal(actual, xp.asarray(expected, dtype=xp.float64)) + + # NumPy recommends strictly increasing sample coordinates. array-api-extra + # deliberately defines these repeated-knot cases, including which value wins at + # the exact knot, rather than exposing searchsorted clipping or a 0/0 division. + @pytest.mark.parametrize( + ("x_points", "values", "x", "expected"), + [ + ([0, 0, 1], [1, 2, 4], [-0.5, 0, 0.5], [1, 2, 3]), + ([0, 1, 1, 2], [0, 10, 20, 30], [0.5, 1, 1.5], [5, 20, 25]), + ([0, 1, 2, 2], [0, 10, 20, 30], [1.5, 2, 2.5], [15, 30, 30]), + ], + ids=["leading", "middle", "trailing"], + ) + @implementations + def test_repeated_coordinates( + self, + xp: ArrayNamespace, + implementation: Implementation, + x_points: list[int], + values: list[int], + x: list[float], + expected: list[float], + ): + actual = _interp( + implementation, + xp.asarray(x, dtype=xp.float64), + xp.asarray(x_points, dtype=xp.float64), + xp.asarray(values, dtype=xp.float64), + xp=xp, + ) + assert_close(actual, xp.asarray(expected, dtype=xp.float64)) + + @implementations + def test_singleton_and_nan_query( + self, xp: ArrayNamespace, implementation: Implementation + ): + x = xp.asarray([np.nan, -3, 99], dtype=xp.float64) + actual = _interp( + implementation, + x, + xp.asarray([1], dtype=xp.float64), + xp.asarray([7], dtype=xp.float64), + xp=xp, + ) + assert_equal(actual, xp.asarray([7, 7, 7], dtype=xp.float64)) + + actual = _interp( + implementation, + xp.asarray([np.nan], dtype=xp.float64), + xp.asarray([0, 1], dtype=xp.float64), + xp.asarray([1, 2], dtype=xp.float64), + xp=xp, + ) + assert_equal(actual, xp.asarray([np.nan], dtype=xp.float64)) + + @implementations + def test_complex_nan_query( + self, xp: ArrayNamespace, implementation: Implementation + ): + x = np.asarray([np.nan]) + x_points = np.asarray([0.0, 1.0]) + values = np.asarray([1 + 2j, 3 + 4j]) + expected = np.interp(x, x_points, values) + + actual = _interp( + implementation, + xp.asarray(x, dtype=xp.float64), + xp.asarray(x_points, dtype=xp.float64), + xp.asarray(values, dtype=xp.complex128), + xp=xp, + ) + assert_equal(xp.real(actual), xp.asarray(expected.real, dtype=xp.float64)) + assert_equal(xp.imag(actual), xp.asarray(expected.imag, dtype=xp.float64)) + with pytest.raises(AssertionError): + assert_equal(xp.imag(actual), xp.asarray([123.0], dtype=xp.float64)) + + @implementations + def test_left_right_and_complex_values( + self, xp: ArrayNamespace, implementation: Implementation + ): + x_points = xp.asarray([0, 1, 2], dtype=xp.float64) + values = xp.asarray([0, 2 + 4j, 4], dtype=xp.complex128) + actual = _interp( + implementation, + xp.asarray([-1, 0.5, 1.5, 3], dtype=xp.float64), + x_points, + values, + left=xp.asarray(1 - 1j, dtype=xp.complex128), + right=5 + 2j, + xp=xp, + ) + expected = xp.asarray([1 - 1j, 1 + 2j, 3 + 2j, 5 + 2j], dtype=xp.complex128) + assert_close(actual, expected) + + @implementations + def test_periodic_unsorted_negative_period_and_ignored_fills( + self, xp: ArrayNamespace, implementation: Implementation + ): + x = xp.asarray([-1, 0, 1, 2, 3, 4, 5], dtype=xp.float64) + x_points = xp.asarray([3, 1], dtype=xp.float64) + values = xp.asarray([30, 10], dtype=xp.float64) + + actual = _interp( + implementation, + x, + x_points, + values, + left="ignored", + right=False, + period=np.float64(-4.0), + xp=xp, + ) + expected = xp.asarray([30, 20, 10, 20, 30, 20, 10], dtype=xp.float64) + assert_close(actual, expected) + + assert_equal(x, xp.asarray([-1, 0, 1, 2, 3, 4, 5], dtype=xp.float64)) + assert_equal(x_points, xp.asarray([3, 1], dtype=xp.float64)) + assert_equal(values, xp.asarray([30, 10], dtype=xp.float64)) + + @implementations + def test_integer_inputs_and_large_coordinates( + self, xp: ArrayNamespace, implementation: Implementation + ): + actual = _interp( + implementation, + xp.asarray([0, 1, 2], dtype=xp.int64), + xp.asarray([0, 2], dtype=xp.int64), + xp.asarray([1, 5], dtype=xp.int64), + xp=xp, + ) + assert_close(actual, xp.asarray([1, 3, 5], dtype=xp.float64)) + + base = 2**24 + actual = _interp( + implementation, + xp.asarray( + [base, base + 0.5, base + 1, base + 1.5, base + 2], + dtype=xp.float64, + ), + xp.asarray([base, base + 1, base + 2], dtype=xp.int64), + xp.asarray([0, 10, 20], dtype=xp.int64), + xp=xp, + ) + assert_close(actual, xp.asarray([0, 5, 10, 15, 20], dtype=xp.float64)) + + @implementations + def test_nonfinite_values(self, xp: ArrayNamespace, implementation: Implementation): + x_points = xp.asarray([0, 1, 2], dtype=xp.float64) + actual = _interp( + implementation, + x_points, + x_points, + xp.asarray([-np.inf, 2, np.inf], dtype=xp.float64), + xp=xp, + ) + assert_equal(actual, xp.asarray([-np.inf, 2, np.inf], dtype=xp.float64)) + + actual = _interp( + implementation, + xp.asarray([0.5], dtype=xp.float64), + xp.asarray([0, 1], dtype=xp.float64), + xp.asarray([np.inf, np.inf], dtype=xp.float64), + xp=xp, + ) + assert_equal(actual, xp.asarray([np.inf], dtype=xp.float64)) + + actual = _interp( + implementation, + xp.asarray([0, 0.5, 1], dtype=xp.float64), + xp.asarray([0, 1], dtype=xp.float64), + xp.asarray([np.nan, 2], dtype=xp.float64), + xp=xp, + ) + assert_equal(actual, xp.asarray([np.nan, np.nan, 2], dtype=xp.float64)) + + actual = _interp( + implementation, + xp.asarray([0.5], dtype=xp.float64), + xp.asarray([0, 1], dtype=xp.float64), + xp.asarray([complex(np.inf, 0), complex(np.inf, 2)], dtype=xp.complex128), + xp=xp, + ) + assert_equal(actual, xp.asarray([complex(np.inf, 1)], dtype=xp.complex128)) + + @implementations + @pytest.mark.skip_xp_backend( + Backend.TORCH, reason="device='meta' does not support searchsorted" + ) + def test_device( + self, + xp: ArrayNamespace, + device: Device, + implementation: Implementation, + ): + actual = _interp( + implementation, + xp.asarray([0.5], dtype=xp.float64, device=device), + xp.asarray([0, 1], dtype=xp.float64, device=device), + xp.asarray([0, 2], dtype=xp.float64, device=device), + xp=xp, + ) + assert get_device(actual) == device + + @pytest.mark.skip_xp_backend(Backend.NUMPY_READONLY, reason="xp=xp") + def test_xp_keyword(self, xp: ArrayNamespace): + actual = xpx_interp( + xp.asarray([0.5], dtype=xp.float64), + xp.asarray([0, 1], dtype=xp.float64), + xp.asarray([0, 2], dtype=xp.float64), + xp=xp, + ) + assert_equal(actual, xp.asarray([1], dtype=xp.float64)) + + def test_shape_and_empty_validation(self, xp: ArrayNamespace): + x = xp.asarray([0.5], dtype=xp.float64) + x_points = xp.asarray([0, 1], dtype=xp.float64) + values = xp.asarray([0, 1], dtype=xp.float64) + + with pytest.raises(ValueError, match="one-dimensional"): + _ = xpx_interp(x, xp.reshape(x_points, (1, 2)), values) + with pytest.raises(ValueError, match="one-dimensional"): + _ = xpx_interp(x, x_points, xp.reshape(values, (1, 2))) + with pytest.raises(ValueError, match="same length"): + _ = xpx_interp(x, x_points, values[:1]) + with pytest.raises(ValueError, match="nonempty"): + _ = xpx_interp(x, x_points[:0], values[:0]) + + def test_type_validation(self, xp: ArrayNamespace): + x = xp.asarray([0.5], dtype=xp.float64) + x_points = xp.asarray([0, 1], dtype=xp.float64) + values = xp.asarray([0, 1], dtype=xp.float64) + + with pytest.raises(TypeError): + _ = xpx_interp([0.5], x_points, values) # type: ignore[arg-type] # pyright: ignore[reportArgumentType] + with pytest.raises(TypeError): + _ = xpx_interp(x, [0, 1], values) # type: ignore[arg-type] # pyright: ignore[reportArgumentType] + with pytest.raises(TypeError): + _ = xpx_interp(x, x_points, [0, 1]) # type: ignore[arg-type] # pyright: ignore[reportArgumentType] + with pytest.raises(TypeError): + _ = xpx_interp(xp.asarray([True]), x_points, values) + with pytest.raises(TypeError): + _ = xpx_interp(x, xp.asarray([False, True]), values) + with pytest.raises(TypeError): + _ = xpx_interp(x, x_points, xp.asarray([False, True])) + with pytest.raises(TypeError): + _ = xpx_interp(True, x_points, values) + with pytest.raises(TypeError): + _ = xpx_interp(1j, x_points, values) # type: ignore[arg-type] # pyright: ignore[reportArgumentType] + + def test_bound_and_period_validation(self, xp: ArrayNamespace): + x = xp.asarray([0.5], dtype=xp.float64) + x_points = xp.asarray([0, 1], dtype=xp.float64) + values = xp.asarray([0, 1], dtype=xp.float64) + + with pytest.raises(ValueError, match="numerical scalar"): + _ = xpx_interp(x, x_points, values, left=xp.asarray([0])) + with pytest.raises(TypeError): + _ = xpx_interp(x, x_points, values, right="bad") # type: ignore[arg-type] # pyright: ignore[reportArgumentType] + with pytest.raises(TypeError): + _ = xpx_interp(x, x_points, values, left=1j) + + for period in (0.0, np.inf, -np.inf, np.nan, Fraction(1, 10**400)): + with pytest.raises(ValueError, match=r"must be (finite|nonzero)"): + _ = xpx_interp(x, x_points, values, period=period) + with pytest.raises(ValueError, match="representable as a finite float"): + _ = xpx_interp(x, x_points, values, period=Fraction(10**400, 1)) + with pytest.raises(TypeError): + _ = xpx_interp(x, x_points, values, period=True) + with pytest.raises(TypeError): + _ = xpx_interp(x, x_points, values, period=xp.asarray(4.0)) diff --git a/tests/main/test_public_api.py b/tests/main/test_public_api.py index 617cdc27..4ad11488 100644 --- a/tests/main/test_public_api.py +++ b/tests/main/test_public_api.py @@ -6,6 +6,7 @@ _creation, _elementwise, _indexing, + _interpolation, _lazy, _linalg, _manipulation, @@ -23,6 +24,7 @@ def test_all_contains_all_public_functions(): _creation, _elementwise, _indexing, + _interpolation, _lazy, _linalg, _manipulation, @@ -34,6 +36,7 @@ def test_all_contains_all_public_functions(): _agnostic._elementwise, _agnostic._indexing, _agnostic._inspection, + _agnostic._interpolation, _agnostic._linalg, _agnostic._manipulation, _agnostic._searching,