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

1"""Convolution of :class:`~chebpy.chebfun.Chebfun` objects. 

2 

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`. 

7 

8Two strategies are used: 

9 

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`). 

15 

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""" 

20 

21from __future__ import annotations 

22 

23from collections.abc import Callable 

24from typing import TYPE_CHECKING, Any 

25 

26import numpy as np 

27 

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 

37 

38if TYPE_CHECKING: 

39 from .chebfun import Chebfun 

40 

41 

42def convolve(f: Chebfun, g: Chebfun) -> Chebfun: 

43 """Return the convolution ``h = f ★ g`` as a piecewise Chebfun. 

44 

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() 

50 

51 _reject_unsupported(f, g) 

52 _reject_nonzero_tails(f, g) 

53 

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]) 

57 

58 # General piecewise convolution. 

59 return _piecewise(f, g) 

60 

61 

62def _reject_unsupported(f: Chebfun, g: Chebfun) -> None: 

63 """Reject convolution operands the algorithms cannot handle. 

64 

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. 

70 

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 ) 

89 

90 

91def _reject_nonzero_tails(f: Chebfun, g: Chebfun) -> None: 

92 """Reject convolution when a :class:`CompactFun` piece has a non-zero tail. 

93 

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. 

97 

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 ) 

112 

113 

114def _use_equal_width_fast_path(f: Chebfun, g: Chebfun) -> bool: 

115 """Return True if the fast equal-width Legendre path applies. 

116 

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)) 

130 

131 

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. 

134 

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]) 

142 

143 h = (b - a) / 2.0 # half-width (same for both funs) 

144 

145 leg_f = cheb2leg(f_fun.coeffs) 

146 leg_g = cheb2leg(g_fun.coeffs) 

147 

148 gamma_left, gamma_right = _conv_legendre(leg_f, leg_g) 

149 

150 gamma_left = h * gamma_left 

151 gamma_right = h * gamma_right 

152 

153 cheb_left = leg2cheb(gamma_left) 

154 cheb_right = leg2cheb(gamma_right) 

155 

156 mid = (a + b + c + d) / 2.0 

157 left_interval = Interval(a + c, mid) 

158 right_interval = Interval(mid, b + d) 

159 

160 left_fun = Bndfun(Chebtech(cheb_left), left_interval) 

161 right_fun = Bndfun(Chebtech(cheb_right), right_interval) 

162 

163 return f.__class__([left_fun, right_fun]) 

164 

165 

166def _piecewise(f: Chebfun, g: Chebfun) -> Chebfun: 

167 """General piecewise convolution via Gauss-Legendre quadrature. 

168 

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])) 

181 

182 f_breaks = _effective_breakpoints(f) 

183 g_breaks = _effective_breakpoints(g) 

184 

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] 

191 

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) 

194 

195 

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 

204 

205 

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)``. 

210 

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]) 

217 

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) 

221 

222 f_bps = [float(bp) for bp in f_breaks] 

223 g_bps = [float(bp) for bp in g_breaks] 

224 

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 

246 

247 return conv_eval 

248 

249 

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)``. 

252 

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)) 

261 

262 

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. 

272 

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) 

292 

293 return f.__class__(funs_list)