Coverage for src/chebpy/_convolution.py: 100%
131 statements
« prev ^ index » next coverage.py v7.16.2, created at 2026-10-01 13:43 +0000
« prev ^ index » next coverage.py v7.16.2, created at 2026-10-01 13:43 +0000
1"""Convolution of :class:`~chebpy.chebfun.Chebfun` objects.
3This module implements ``h = f ★ g`` for piecewise Chebfuns. It is kept
4separate from :mod:`chebpy.chebfun` so the (sizeable) convolution algorithm
5does not bloat the main class; :meth:`chebpy.chebfun.Chebfun.conv` is a thin
6wrapper around :func:`convolve`.
8Two strategies are used:
10* When both inputs are single-piece :class:`~chebpy.bndfun.Bndfun` funs of
11 equal width, the fast Hale-Townsend Legendre convolution is used
12 (:func:`_equal_width_pair`).
13* Otherwise each output sub-interval is built adaptively using
14 Gauss-Legendre quadrature (:func:`_piecewise`).
16The algorithm is based on:
17 N. Hale and A. Townsend, "An algorithm for the convolution of Legendre
18 series", SIAM J. Sci. Comput., 36(3), A1207-A1220, 2014.
19"""
21from __future__ import annotations
23from collections.abc import Callable
24from typing import TYPE_CHECKING, Any
26import numpy as np
28from .algorithms import _conv_legendre, cheb2leg, leg2cheb
29from .bndfun import Bndfun
30from .chebtech import Chebtech
31from .compactfun import CompactFun
32from .exceptions import DivergentIntegralError
33from .fun import Fun
34from .singfun import Singfun
35from .trigtech import Trigtech
36from .utilities import Interval
38if TYPE_CHECKING:
39 from .chebfun import Chebfun
42def convolve(f: Chebfun, g: Chebfun) -> Chebfun:
43 """Return the convolution ``h = f ★ g`` as a piecewise Chebfun.
45 See :meth:`chebpy.chebfun.Chebfun.conv` for the full description; this is
46 the implementation behind that method.
47 """
48 if f.isempty or g.isempty:
49 return f.__class__.initempty()
51 _reject_unsupported(f, g)
52 _reject_nonzero_tails(f, g)
54 # Fast path: both single-piece with equal-width finite domains.
55 if _use_equal_width_fast_path(f, g):
56 return _equal_width_pair(f, f.funs[0], g.funs[0])
58 # General piecewise convolution.
59 return _piecewise(f, g)
62def _reject_unsupported(f: Chebfun, g: Chebfun) -> None:
63 """Reject convolution operands the algorithms cannot handle.
65 Both the Hale-Townsend Legendre algorithm and the Gauss-Legendre fallback
66 assume Chebyshev coefficients on an affine map, so :class:`Trigtech`-backed
67 pieces (Fourier coefficients; periodic convolution is a distinct
68 ``circconv`` operation) and :class:`~chebpy.singfun.Singfun` pieces
69 (non-affine clustering map) are refused with a clear error.
71 Raises:
72 NotImplementedError: If either operand contains a Trigtech- or
73 Singfun-backed piece.
74 """
75 pieces = (*f.funs, *g.funs)
76 if any(isinstance(piece.onefun, Trigtech) for piece in pieces):
77 raise NotImplementedError(
78 "conv() is not supported for trigfun (Trigtech-backed) inputs. "
79 "Aperiodic convolution and periodic (circular) convolution are "
80 "distinct operations; a dedicated circconv() for trigfuns is "
81 "not yet implemented."
82 )
83 if any(isinstance(piece, Singfun) for piece in pieces):
84 raise NotImplementedError(
85 "conv() is not supported for Chebfuns containing Singfun pieces "
86 "(functions with endpoint singularities represented by a non-affine "
87 "clustering map)."
88 )
91def _reject_nonzero_tails(f: Chebfun, g: Chebfun) -> None:
92 """Reject convolution when a :class:`CompactFun` piece has a non-zero tail.
94 Convolution of a function with a non-zero asymptotic limit diverges on an
95 unbounded interval, so refuse early with a clear error pointing the user at
96 the algebraic-closure escape hatch.
98 Raises:
99 DivergentIntegralError: If any CompactFun piece of either operand has a
100 non-zero ``tail_left`` or ``tail_right``.
101 """
102 for label, h in (("self", f), ("other", g)):
103 for piece in h.funs:
104 if isinstance(piece, CompactFun) and (piece.tail_left != 0.0 or piece.tail_right != 0.0):
105 raise DivergentIntegralError( # noqa: TRY003
106 f"Convolution requires both operands to decay to zero at "
107 f"±inf; got tail_left={piece.tail_left}, "
108 f"tail_right={piece.tail_right} for {label}. Consider "
109 f"subtracting a matched sigmoid first so the residual "
110 f"has zero tails, then convolving the residual."
111 )
114def _use_equal_width_fast_path(f: Chebfun, g: Chebfun) -> bool:
115 """Return True if the fast equal-width Legendre path applies.
117 The fast path requires both operands to be single finite :class:`Bndfun`
118 pieces of equal width. :class:`CompactFun` pieces (with possibly infinite
119 logical support) always take the general piecewise path so the output is
120 wrapped correctly.
121 """
122 if any(isinstance(piece, CompactFun) for piece in (*f.funs, *g.funs)):
123 return False
124 if f.funs.size != 1 or g.funs.size != 1:
125 return False
126 f_fun, g_fun = f.funs[0], g.funs[0]
127 f_w = float(f_fun.support[1]) - float(f_fun.support[0])
128 g_w = float(g_fun.support[1]) - float(g_fun.support[0])
129 return bool(np.isclose(f_w, g_w))
132def _equal_width_pair(f: Chebfun, f_fun: Any, g_fun: Any) -> Chebfun:
133 """Convolve two single Bndfuns of equal width using the fast algorithm.
135 Uses the Hale-Townsend Legendre convolution. The two funs may be on
136 different intervals as long as they have the same width.
137 """
138 a = float(f_fun.support[0])
139 b = float(f_fun.support[1])
140 c = float(g_fun.support[0])
141 d = float(g_fun.support[1])
143 h = (b - a) / 2.0 # half-width (same for both funs)
145 leg_f = cheb2leg(f_fun.coeffs)
146 leg_g = cheb2leg(g_fun.coeffs)
148 gamma_left, gamma_right = _conv_legendre(leg_f, leg_g)
150 gamma_left = h * gamma_left
151 gamma_right = h * gamma_right
153 cheb_left = leg2cheb(gamma_left)
154 cheb_right = leg2cheb(gamma_right)
156 mid = (a + b + c + d) / 2.0
157 left_interval = Interval(a + c, mid)
158 right_interval = Interval(mid, b + d)
160 left_fun = Bndfun(Chebtech(cheb_left), left_interval)
161 right_fun = Bndfun(Chebtech(cheb_right), right_interval)
163 return f.__class__([left_fun, right_fun])
166def _piecewise(f: Chebfun, g: Chebfun) -> Chebfun:
167 """General piecewise convolution via Gauss-Legendre quadrature.
169 The breakpoints of the result are the sorted, unique pairwise sums of the
170 breakpoints of ``f`` and ``g``. On each sub-interval the convolution
171 integral is smooth, so we construct it adaptively. When either input
172 contains :class:`CompactFun` pieces, the corresponding ``±inf`` breakpoints
173 are replaced with the numerical-support bounds for the purposes of
174 integration; the outermost output pieces are then wrapped as
175 :class:`CompactFun` so the result preserves the unbounded logical support.
176 """
177 f_logical_breaks = np.array(f.breakpoints, dtype=float)
178 g_logical_breaks = np.array(g.breakpoints, dtype=float)
179 left_inf = (not np.isfinite(f_logical_breaks[0])) or (not np.isfinite(g_logical_breaks[0]))
180 right_inf = (not np.isfinite(f_logical_breaks[-1])) or (not np.isfinite(g_logical_breaks[-1]))
182 f_breaks = _effective_breakpoints(f)
183 g_breaks = _effective_breakpoints(g)
185 # Output breakpoints: all pairwise sums, uniquified and coalesced.
186 out_breaks = np.unique(np.add.outer(f_breaks, g_breaks).ravel())
187 hscl = max(abs(out_breaks[0]), abs(out_breaks[-1]), 1.0)
188 tol = 10.0 * np.finfo(float).eps * hscl
189 mask = np.concatenate(([True], np.diff(out_breaks) > tol))
190 out_breaks = out_breaks[mask]
192 conv_eval = _make_evaluator(f, g, f_breaks, g_breaks)
193 return _build_pieces(f, out_breaks, conv_eval, left_inf=left_inf, right_inf=right_inf)
196def _effective_breakpoints(h: Chebfun) -> np.ndarray:
197 """Return ``h``'s breakpoints with ±inf replaced by numerical-support bounds."""
198 bps = np.array(h.breakpoints, dtype=float)
199 if not np.isfinite(bps[0]) and isinstance(h.funs[0], CompactFun):
200 bps[0] = float(h.funs[0].numerical_support[0])
201 if not np.isfinite(bps[-1]) and isinstance(h.funs[-1], CompactFun):
202 bps[-1] = float(h.funs[-1].numerical_support[1])
203 return bps
206def _make_evaluator(
207 f: Chebfun, g: Chebfun, f_breaks: np.ndarray, g_breaks: np.ndarray
208) -> Callable[[np.ndarray], np.ndarray]:
209 """Build the Gauss-Legendre quadrature evaluator for ``(f ★ g)``.
211 The integrand is broken at the breakpoints of ``f`` and the shifted
212 breakpoints of ``g`` so it is polynomial on each sub-interval; the
213 quadrature order is chosen to integrate that polynomial exactly.
214 """
215 f_a, f_b = float(f_breaks[0]), float(f_breaks[-1])
216 g_c, g_d = float(g_breaks[0]), float(g_breaks[-1])
218 max_deg = max(fun.size for fun in f.funs) + max(fun.size for fun in g.funs)
219 n_quad = max(int(np.ceil((max_deg + 1) / 2)), 16)
220 quad_nodes, quad_weights = np.polynomial.legendre.leggauss(n_quad)
222 f_bps = [float(bp) for bp in f_breaks]
223 g_bps = [float(bp) for bp in g_breaks]
225 def conv_eval(x: np.ndarray) -> np.ndarray:
226 """Evaluate (f ★ g)(x) via Gauss-Legendre quadrature."""
227 x = np.atleast_1d(np.asarray(x, dtype=float))
228 result = np.zeros(x.shape)
229 for idx in range(x.size):
230 xi = x[idx]
231 t_lo = max(f_a, xi - g_d)
232 t_hi = min(f_b, xi - g_c)
233 if t_hi <= t_lo:
234 continue
235 inner = _subinterval_nodes(xi, t_lo, t_hi, f_bps, g_bps)
236 total = 0.0
237 for j in range(len(inner) - 1):
238 a_int, b_int = inner[j], inner[j + 1]
239 hw = (b_int - a_int) / 2.0
240 mid = (a_int + b_int) / 2.0
241 nodes = hw * quad_nodes + mid
242 wts = hw * quad_weights
243 total += np.dot(wts, f(nodes) * g(xi - nodes))
244 result[idx] = total
245 return result
247 return conv_eval
250def _subinterval_nodes(xi: float, t_lo: float, t_hi: float, f_bps: list[float], g_bps: list[float]) -> list[float]:
251 """Return the sorted integration sub-interval boundaries in ``(t_lo, t_hi)``.
253 Breaks at the breakpoints of ``f`` and the shifted breakpoints of ``g``
254 that fall strictly inside ``(t_lo, t_hi)``, keeping the integrand polynomial
255 on each resulting sub-interval.
256 """
257 inner = [t_lo, t_hi]
258 inner.extend(bp for bp in f_bps if t_lo < bp < t_hi)
259 inner.extend(xi - bp for bp in g_bps if t_lo < xi - bp < t_hi)
260 return sorted(set(inner))
263def _build_pieces(
264 f: Chebfun,
265 out_breaks: np.ndarray,
266 conv_eval: Callable[[np.ndarray], np.ndarray],
267 *,
268 left_inf: bool,
269 right_inf: bool,
270) -> Chebfun:
271 """Build one fun per output sub-interval, wrapping the unbounded ends.
273 Interior pieces are finite :class:`Bndfun`; the outermost pieces are wrapped
274 as :class:`CompactFun` when the corresponding logical edge is ``±inf`` so
275 the result preserves the unbounded logical support.
276 """
277 n_pieces = len(out_breaks) - 1
278 funs_list: list[Fun] = []
279 for i in range(n_pieces):
280 a_storage = float(out_breaks[i])
281 b_storage = float(out_breaks[i + 1])
282 interval = Interval(a_storage, b_storage)
283 bnd = Bndfun.initfun_adaptive(conv_eval, interval)
284 wrap_left = i == 0 and left_inf
285 wrap_right = i == n_pieces - 1 and right_inf
286 if wrap_left or wrap_right:
287 a_logical = -np.inf if wrap_left else a_storage
288 b_logical = np.inf if wrap_right else b_storage
289 funs_list.append(CompactFun(bnd.onefun, interval, logical_interval=(a_logical, b_logical)))
290 else:
291 funs_list.append(bnd)
293 return f.__class__(funs_list)