Coverage for src/chebpy/trigtech.py: 100%

323 statements  

« prev     ^ index     » next       coverage.py v7.16.2, created at 2026-10-01 13:43 +0000

1"""Trigonometric (Fourier) technology for periodic function approximation. 

2 

3This module provides the Trigtech class, which represents smooth periodic functions 

4on [-1, 1] using truncated Fourier series. It is the trigonometric analogue of 

5Chebtech and sits in the same class hierarchy: 

6 

7 Onefun → Smoothfun → Trigtech 

8 

9Coefficient storage convention (NumPy-native / FFT order) 

10---------------------------------------------------------- 

11Given n equispaced sample points x_j = -1 + 2j/n (j = 0, …, n-1), the stored 

12coefficients are 

13 

14 coeffs[k] = (1/n) * sum_j f(x_j) * exp(-2*pi*i*j*k/n) 

15 = (numpy.fft.fft(values) / n)[k] 

16 

17This is exactly the output of ``numpy.fft.fft(values) / n``, i.e. NumPy-native 

18(FFT) ordering: DC at index 0, positive frequencies 1 … n//2, then negative 

19frequencies -(n//2)+1 … -1. 

20 

21Use ``_coeffs_to_plotorder()`` to obtain the human-readable DC-centred ordering 

22(equivalent to ``numpy.fft.fftshift``). 

23 

24Evaluation 

25---------- 

26Any point x ∈ [-1, 1] is evaluated via the DFT summation formula: 

27 

28 f(x) = Σ_k coeffs[k] * exp(i*π*ω_k*(x+1)) 

29 

30where ω_k = numpy.fft.fftfreq(n)*n gives the integer frequencies in FFT order. 

31 

32References: 

33---------- 

34* Trefethen, "Spectral Methods in MATLAB" (SIAM 2000) 

35* Chebfun @trigtech (github.com/chebfun/chebfun) 

36""" 

37 

38import warnings 

39from abc import ABC 

40from typing import Any, cast 

41 

42import matplotlib.pyplot as plt 

43import numpy as np 

44 

45from .algorithms import newtonroots, rootsunit 

46from .chebtech import Chebtech 

47from .decorators import self_empty 

48from .exceptions import BadFunLengthArgument 

49from .plotting import plotfun, plotfuncoeffs 

50from .settings import _preferences as prefs 

51from .smoothfun import Smoothfun 

52from .utilities import Interval, coerce_list 

53 

54 

55def _trig_adaptive( 

56 cls: Any, 

57 fun: Any, 

58 hscale: float = 1, 

59 maxpow2: int | None = None, 

60) -> np.ndarray: 

61 """Adaptively determine the Fourier coefficients needed to represent *fun*. 

62 

63 Uses successively finer equispaced grids (sizes 2**k) until the 

64 high-frequency Fourier modes decay below tolerance. Convergence is 

65 assessed via the one-sided symmetric maximum of the DC-centred coefficient 

66 magnitudes: ``abs_sym[k] = max(|c_k|, |c_{-k}|) / vscale``. The series 

67 is considered converged when the Nyquist/highest-frequency mode 

68 ``abs_sym[-1]`` falls below *tol*. 

69 

70 Args: 

71 cls: Trigtech class (provides ``_trigpts`` and ``_vals2coeffs``). 

72 fun: Callable to approximate. 

73 hscale: Horizontal scale for tolerance adjustment. 

74 maxpow2: Maximum power of 2 to try (defaults to ``prefs.maxpow2``). 

75 

76 Returns: 

77 numpy.ndarray: Fourier coefficients in NumPy FFT order. 

78 """ 

79 minpow2 = 3 # start at n = 8 

80 maxpow2 = maxpow2 if maxpow2 is not None else prefs.maxpow2 

81 tol = prefs.eps * max(hscale, 1) 

82 coeffs: np.ndarray = np.array([]) 

83 for k in range(minpow2, max(minpow2, maxpow2) + 1): 

84 n = 2**k 

85 points = cls._trigpts(n) 

86 values = fun(points) 

87 coeffs = cls._vals2coeffs(values) 

88 vscale = float(np.max(np.abs(values))) 

89 if vscale <= tol: 

90 return np.array([0.0]) 

91 

92 # Build one-sided symmetric maximum: 

93 # abs_sym[ki] = max(|c_{ki}|, |c_{-ki}|) / vscale for ki = 0…n//2 

94 centered = np.fft.fftshift(coeffs) 

95 dc_idx = n // 2 

96 abs_sym = np.zeros(dc_idx + 1) 

97 for ki in range(dc_idx + 1): 

98 p = centered[dc_idx + ki] if dc_idx + ki < n else 0.0 

99 q = centered[dc_idx - ki] 

100 abs_sym[ki] = max(abs(p), abs(q)) / vscale 

101 

102 # Convergence: the Nyquist/highest-frequency mode is negligible. 

103 if abs_sym[-1] <= tol: 

104 above = np.where(abs_sym > tol)[0] 

105 if len(above) == 0: # pragma: no cover - defensive, see below 

106 # Defensive: the normalised peak coefficient is >= 1/n >> tol 

107 # whenever vscale > tol, so 'above' is never empty here. 

108 return np.array([0.0]) 

109 max_k = int(above[-1]) # highest significant frequency index 

110 start = dc_idx - max_k 

111 end = dc_idx + max_k + 1 

112 return np.fft.ifftshift(centered[start:end]) 

113 

114 if k == maxpow2: 

115 warnings.warn( 

116 f"The {cls.__name__} constructor did not converge: using {n} points", 

117 stacklevel=3, 

118 ) 

119 break 

120 return coeffs 

121 

122 

123class Trigtech(Smoothfun, ABC): 

124 """Trigonometric (Fourier) function approximation on [-1, 1]. 

125 

126 Represents a smooth periodic function f: [-1, 1] -> R (or C) as a 

127 truncated Fourier series. Coefficients are stored in NumPy FFT order; 

128 see module docstring for the precise convention. 

129 

130 This class is ``ABC`` so that it cannot be instantiated directly—exactly 

131 mirroring Chebtech, which is also abstract (concrete only through the 

132 ``Chebtech`` name used everywhere). In practice ``Trigtech`` is both the 

133 abstract base and the concrete class: it is not further subclassed, but 

134 the ABC marker prevents accidental bare construction without going through 

135 a named constructor. 

136 

137 Examples: 

138 A periodic function needs only a handful of Fourier modes, where the 

139 Chebyshev constructor would need many more: 

140 

141 >>> import numpy as np 

142 >>> f = Trigtech.initfun_adaptive(lambda x: np.cos(np.pi * x)) 

143 >>> f.size 

144 3 

145 >>> bool(abs(f(0.0) - 1.0) < 1e-13) 

146 True 

147 >>> bool(abs(f(1.0) + 1.0) < 1e-13) 

148 True 

149 

150 Fixed-length construction samples on *n* equispaced points: 

151 

152 >>> g = Trigtech.initfun_fixedlen(lambda x: np.sin(np.pi * x), 16) 

153 >>> g.size 

154 16 

155 """ 

156 

157 # ------------------------------------------------------------------ 

158 # alternative constructors 

159 # ------------------------------------------------------------------ 

160 

161 @classmethod 

162 def initconst(cls, c: Any = None, *, interval: Any = None) -> "Trigtech": 

163 """Initialise a Trigtech from a constant *c*.""" 

164 if not np.isscalar(c): 

165 raise ValueError(c) 

166 if isinstance(c, int): 

167 c = float(c) 

168 return cls(np.array([c]), interval=interval) 

169 

170 @classmethod 

171 def initempty(cls, *, interval: Any = None) -> "Trigtech": 

172 """Initialise an empty Trigtech.""" 

173 return cls(np.array([]), interval=interval) 

174 

175 @classmethod 

176 def initidentity(cls, *, interval: Any = None) -> "Trigtech": 

177 """Trigtech approximation of the identity f(x) = x on [-1, 1]. 

178 

179 Note: f(x) = x is *not* periodic on [-1, 1], so this will not converge 

180 to machine precision. It is provided for interface compatibility with 

181 Chebtech; in practice ``Classicfun.initidentity`` is used instead. 

182 """ 

183 interval = interval if interval is not None else prefs.domain 

184 return cls.initfun_adaptive(lambda x: x, interval=interval) 

185 

186 @classmethod 

187 def initfun(cls, fun: Any = None, n: Any = None, *, interval: Any = None) -> "Trigtech": 

188 """Convenience constructor: adaptive if *n* is None, fixed-length otherwise.""" 

189 if n is None: 

190 return cls.initfun_adaptive(fun, interval=interval) 

191 return cls.initfun_fixedlen(fun, n, interval=interval) 

192 

193 @classmethod 

194 def initfun_fixedlen(cls, fun: Any = None, n: Any = None, *, interval: Any = None) -> "Trigtech": 

195 """Initialise a Trigtech from callable *fun* using *n* equispaced points.""" 

196 if n is None: 

197 raise BadFunLengthArgument("initfun_fixedlen requires the n parameter to be specified") # noqa: TRY003 

198 points = cls._trigpts(int(n)) 

199 values = fun(points) 

200 coeffs = cls._vals2coeffs(values) 

201 return cls(coeffs, interval=interval) 

202 

203 @classmethod 

204 def initfun_adaptive(cls, fun: Any = None, *, interval: Any = None) -> "Trigtech": 

205 """Initialise a Trigtech from callable *fun* using the adaptive constructor.""" 

206 interval = interval if interval is not None else prefs.domain 

207 interval = Interval(*interval) 

208 coeffs = _trig_adaptive(cls, fun, hscale=interval.hscale) 

209 return cls(coeffs, interval=interval) 

210 

211 @classmethod 

212 def initvalues(cls, values: Any = None, *, interval: Any = None) -> "Trigtech": 

213 """Initialise a Trigtech from function values at equispaced points.""" 

214 return cls(cls._vals2coeffs(np.asarray(values)), interval=interval) 

215 

216 # ------------------------------------------------------------------ 

217 # core dunder methods 

218 # ------------------------------------------------------------------ 

219 

220 def __init__(self, coeffs: Any, interval: Any = None) -> None: 

221 """Initialise a Trigtech with FFT-order *coeffs* on *interval*. 

222 

223 Coefficients are always stored as complex128. The :attr:`iscomplex` 

224 property returns True only when the function *values* are complex 

225 (i.e., the coefficients do **not** satisfy the conjugate-symmetry 

226 condition C_{n-k} ≈ conj(C_k)). 

227 

228 Args: 

229 coeffs: 1-D array of Fourier coefficients in NumPy FFT order. 

230 interval: Two-element interval [a, b]. Defaults to ``prefs.domain``. 

231 """ 

232 interval = interval if interval is not None else prefs.domain 

233 self._coeffs = np.array(coeffs, dtype=complex) 

234 self._interval = Interval(*interval) 

235 

236 def __call__(self, x: Any, how: str = "fft") -> Any: # noqa: ARG002 (how kept for Chebtech interface parity) 

237 """Evaluate the Trigtech at points *x* via the DFT summation formula. 

238 

239 f(x) = Σ_k coeffs[k] * exp(i*π*ω_k*(x+1)) 

240 

241 where ω_k = fftfreq(n)*n gives integer frequencies in FFT order. 

242 For real-valued functions the imaginary part of the result is discarded. 

243 

244 Args: 

245 x: Evaluation points in [-1, 1]. 

246 how: Ignored; present for interface compatibility with Chebtech. 

247 """ 

248 if self.isempty: 

249 return np.array([]) 

250 scalar = np.isscalar(x) 

251 x = np.atleast_1d(np.asarray(x, dtype=float)).ravel() 

252 

253 if self.isconst: 

254 c0 = self._coeffs[0].real if not self.iscomplex else self._coeffs[0] 

255 out = c0 * np.ones(x.size) 

256 return float(out[0]) if scalar else out 

257 

258 n = self.size 

259 freqs = np.fft.fftfreq(n) * n # [0, 1, …, n//2, -(n//2)+1, …, -1] 

260 # shape: (len(x), n) @ (n,) → (len(x),) 

261 phases = np.exp(1j * np.pi * np.outer(x + 1.0, freqs)) 

262 result = phases @ self._coeffs 

263 if not self.iscomplex: 

264 result = result.real 

265 return float(result[0]) if scalar else result 

266 

267 def __repr__(self) -> str: 

268 """Return a concise string representation.""" 

269 return f"<{self.__class__.__name__}{{{self.size}}}>" 

270 

271 # ------------------------------------------------------------------ 

272 # properties 

273 # ------------------------------------------------------------------ 

274 

275 @property 

276 def coeffs(self) -> np.ndarray: 

277 """Fourier coefficients in NumPy FFT order (always complex128).""" 

278 return self._coeffs 

279 

280 @property 

281 def interval(self) -> Interval: 

282 """Interval that the Trigtech is mapped to.""" 

283 return self._interval 

284 

285 @property 

286 def size(self) -> int: 

287 """Number of stored Fourier coefficients.""" 

288 return self._coeffs.size 

289 

290 @property 

291 def isempty(self) -> bool: 

292 """True if the Trigtech has no coefficients.""" 

293 return self.size == 0 

294 

295 @property 

296 def iscomplex(self) -> bool: 

297 """True if the function is complex-valued (values have a non-negligible imaginary part). 

298 

299 This is determined by checking whether the Fourier coefficients violate 

300 the conjugate-symmetry condition C_{n-k} ≈ conj(C_k) that holds for 

301 every real-valued periodic function. 

302 """ 

303 n = self.size 

304 if n <= 1: 

305 return bool(np.any(np.abs(np.imag(self._coeffs)) > 0)) 

306 abs_max = float(np.max(np.abs(self._coeffs))) 

307 if abs_max == 0.0: 

308 return False 

309 tol = 1e-8 * abs_max 

310 # mirror[k-1] = conj(C_{n-k}) for k = 1,...,n-1 

311 mirror = np.conj(self._coeffs[-1:0:-1]) 

312 return bool(np.any(np.abs(self._coeffs[1:] - mirror) > tol)) 

313 

314 @property 

315 def isconst(self) -> bool: 

316 """True if the Trigtech represents a constant (single coefficient).""" 

317 return self.size == 1 

318 

319 @property 

320 def isperiodic(self) -> bool: 

321 """Always True: Trigtech always represents a periodic function.""" 

322 return True 

323 

324 @property 

325 @self_empty(0.0) 

326 def vscale(self) -> float: 

327 """Estimate the vertical scale (max |f|).""" 

328 return float(np.abs(np.asarray(coerce_list(self.values()))).max()) 

329 

330 # ------------------------------------------------------------------ 

331 # utilities 

332 # ------------------------------------------------------------------ 

333 

334 def copy(self) -> "Trigtech": 

335 """Return a deep copy.""" 

336 return self.__class__(self._coeffs.copy(), interval=self._interval.copy()) 

337 

338 def imag(self) -> "Trigtech": 

339 """Return the imaginary part of the function as a real-valued Trigtech. 

340 

341 For a complex function f(x) = g(x) + i·h(x), the Fourier coefficients 

342 of h(x) are H[k] = (D[k] - conj(D[n-k])) / (2i) for k ≥ 1, 

343 and H[0] = Im(D[0]). 

344 """ 

345 if not self.iscomplex: 

346 return self.initconst(0.0, interval=self._interval) 

347 n = self.size 

348 c = self._coeffs 

349 imag_c = np.zeros(n, dtype=complex) 

350 imag_c[0] = np.imag(c[0]) 

351 if n > 1: 

352 mirror = np.conj(c[-1:0:-1]) # conj(c[n-1]), ..., conj(c[1]) 

353 imag_c[1:] = (c[1:] - mirror) / (2j) 

354 return self.__class__(imag_c, self._interval) 

355 

356 def prolong(self, n: int) -> "Trigtech": 

357 """Return a Trigtech of length *n* (truncate or zero-pad in frequency space). 

358 

359 The operation aligns DC components of the source and target DC-centred 

360 representations, then either pads with zeros (n > m) or slices (n < m). 

361 This correctly handles the asymmetry between even- and odd-length arrays. 

362 """ 

363 m = self.size 

364 if n == m: 

365 return self.copy() 

366 

367 centered = np.fft.fftshift(self._coeffs) 

368 dc_src = m // 2 

369 dc_tgt = n // 2 

370 

371 if n > m: 

372 padded = np.zeros(n, dtype=centered.dtype) 

373 start = dc_tgt - dc_src 

374 padded[start : start + m] = centered 

375 return self.__class__(np.fft.ifftshift(padded), interval=self._interval) 

376 else: 

377 start = dc_src - dc_tgt 

378 truncated = centered[start : start + n] 

379 return self.__class__(np.fft.ifftshift(truncated), interval=self._interval) 

380 

381 def real(self) -> "Trigtech": 

382 """Return the real part of the function as a real-valued Trigtech. 

383 

384 For a complex function f(x) = g(x) + i·h(x), the Fourier coefficients 

385 of g(x) are G[k] = (D[k] + conj(D[n-k])) / 2 for k ≥ 1, 

386 and G[0] = Re(D[0]). 

387 """ 

388 if not self.iscomplex: 

389 return self 

390 n = self.size 

391 c = self._coeffs 

392 real_c = np.zeros(n, dtype=complex) 

393 real_c[0] = np.real(c[0]) 

394 if n > 1: 

395 mirror = np.conj(c[-1:0:-1]) # conj(c[n-1]), ..., conj(c[1]) 

396 real_c[1:] = (c[1:] + mirror) / 2 

397 return self.__class__(real_c, self._interval) 

398 

399 def simplify(self) -> "Trigtech": 

400 """Truncate high-frequency Fourier coefficients that are below tolerance. 

401 

402 Uses the same one-sided symmetric-maximum criterion as the adaptive 

403 constructor: the highest-frequency mode retained is the one where 

404 ``max(|c_k|, |c_{-k}|) / vscale > tol``. 

405 """ 

406 oldlen = len(self._coeffs) 

407 longself = self.prolong(max(17, oldlen)) 

408 n = longself.size 

409 tol = prefs.eps * max(self._interval.hscale, 1) 

410 

411 centered = np.fft.fftshift(longself._coeffs) 

412 dc_idx = n // 2 

413 abs_max = float(np.max(np.abs(centered))) 

414 if abs_max == 0.0: 

415 return self.initconst(0.0, interval=self._interval) 

416 

417 abs_sym = np.zeros(dc_idx + 1) 

418 for ki in range(dc_idx + 1): 

419 p = centered[dc_idx + ki] if dc_idx + ki < n else 0.0 

420 q = centered[dc_idx - ki] 

421 abs_sym[ki] = max(abs(p), abs(q)) / abs_max 

422 

423 above = np.where(abs_sym > tol)[0] 

424 if len(above) == 0: # pragma: no cover - defensive, see below 

425 # Defensive: with abs_max > 0 the normalised peak equals 1 > tol, 

426 # so 'above' always contains at least the peak index. 

427 return self.initconst(0.0, interval=self._interval) 

428 max_k = int(above[-1]) 

429 max_k = min(max_k, oldlen // 2) # don't exceed original size 

430 

431 start = dc_idx - max_k 

432 end = dc_idx + max_k + 1 

433 return self.__class__(np.fft.ifftshift(centered[start:end]), interval=self._interval) 

434 

435 def values(self) -> np.ndarray: 

436 """Function values at the n equispaced points x_j = -1 + 2j/n.""" 

437 return self._coeffs2vals(self._coeffs) 

438 

439 def _coeffs_to_plotorder(self) -> np.ndarray: 

440 """Return coefficients in DC-centred (human-readable) order. 

441 

442 Equivalent to ``numpy.fft.fftshift(self.coeffs)``: 

443 ordering is [c_{-n//2}, …, c_{-1}, c_0, c_1, …, c_{n//2-1}]. 

444 """ 

445 return np.fft.fftshift(self._coeffs) 

446 

447 # ------------------------------------------------------------------ 

448 # algebra 

449 # ------------------------------------------------------------------ 

450 

451 @self_empty() 

452 def __add__(self, f: Any) -> "Trigtech": 

453 """Add a scalar or another Trigtech.""" 

454 cls = self.__class__ 

455 if np.isscalar(f): 

456 dtype: Any = complex if np.iscomplexobj(f) else self._coeffs.dtype 

457 cfs = np.array(self._coeffs, dtype=dtype) 

458 cfs[0] += f # add to DC component 

459 return cls(cfs, interval=self._interval) 

460 if f.isempty: 

461 return cast("Trigtech", f.copy()) 

462 g = self 

463 n, m = g.size, f.size 

464 if n < m: 

465 g = g.prolong(m) 

466 elif m < n: 

467 f = f.prolong(n) 

468 cfs = f.coeffs + g.coeffs 

469 eps = prefs.eps 

470 tol = 0.5 * eps * max(f.vscale, g.vscale) 

471 if np.all(np.abs(cfs) < tol): 

472 return cls.initconst(0.0, interval=self._interval) 

473 return cls(cfs, interval=self._interval) 

474 

475 @self_empty() 

476 def __div__(self, f: Any) -> "Trigtech": 

477 """Divide this Trigtech by a scalar or another Trigtech.""" 

478 cls = self.__class__ 

479 if np.isscalar(f): 

480 return cls(self._coeffs / np.asarray(f), interval=self._interval) 

481 if f.isempty: 

482 return cast("Trigtech", f.copy()) 

483 return cls.initfun_adaptive(lambda x: self(x) / f(x), interval=self._interval) 

484 

485 __truediv__ = __div__ 

486 

487 @self_empty() 

488 def __mul__(self, g: Any) -> "Trigtech": 

489 """Multiply this Trigtech by a scalar or another Trigtech. 

490 

491 Trig-polynomial multiplication is circular convolution in frequency 

492 space. We implement this cleanly by evaluating both on a grid of 

493 size n1 + n2 (sufficient to avoid aliasing), multiplying pointwise, 

494 and taking the FFT. 

495 """ 

496 cls = self.__class__ 

497 if np.isscalar(g): 

498 return cls(g * self._coeffs, interval=self._interval) 

499 if g.isempty: 

500 return cast("Trigtech", g.copy()) 

501 n = self.size + g.size 

502 f_vals = self.prolong(n).values() 

503 g_vals = g.prolong(n).values() 

504 return cls(cls._vals2coeffs(f_vals * g_vals), interval=self._interval) 

505 

506 def __neg__(self) -> "Trigtech": 

507 """Return the negation.""" 

508 return self.__class__(-self._coeffs, interval=self._interval) 

509 

510 def __pos__(self) -> "Trigtech": 

511 """Return self (unary plus).""" 

512 return self 

513 

514 @self_empty() 

515 def __pow__(self, f: Any) -> "Trigtech": 

516 """Raise this Trigtech to a power *f* (scalar or Trigtech).""" 

517 

518 def powfun(fn: Any, x: Any) -> Any: 

519 return fn if np.isscalar(fn) else fn(x) 

520 

521 return self.__class__.initfun_adaptive( 

522 lambda x: np.power(self(x), powfun(f, x)), 

523 interval=self._interval, 

524 ) 

525 

526 def __rdiv__(self, f: Any) -> "Trigtech": 

527 """Compute f / self where *f* is a scalar.""" 

528 return self.__class__.initfun_adaptive( 

529 lambda x: (0.0 * x + f) / self(x), 

530 interval=self._interval, 

531 ) 

532 

533 __radd__ = __add__ 

534 __rmul__ = __mul__ 

535 __rtruediv__ = __rdiv__ 

536 

537 def __rsub__(self, f: Any) -> "Trigtech": 

538 """Compute f - self.""" 

539 return cast("Trigtech", -(self - f)) 

540 

541 @self_empty() 

542 def __rpow__(self, f: Any) -> "Trigtech": 

543 """Compute f ** self.""" 

544 return self.__class__.initfun_adaptive( 

545 lambda x: np.power(f, self(x)), 

546 interval=self._interval, 

547 ) 

548 

549 def __sub__(self, f: Any) -> "Trigtech": 

550 """Subtract *f* (scalar or Trigtech) from this Trigtech.""" 

551 return cast("Trigtech", self + (-f)) 

552 

553 # ------------------------------------------------------------------ 

554 # rootfinding 

555 # ------------------------------------------------------------------ 

556 

557 def roots(self, sort: bool | None = None) -> np.ndarray: 

558 """Find the roots of this Trigtech on [-1, 1]. 

559 

560 Converts to a Chebyshev representation via re-sampling on Chebyshev 

561 points and delegates to the Chebtech colleague-matrix root-finder. 

562 

563 Args: 

564 sort: If True, sort the roots in ascending order. Defaults to 

565 ``prefs.sortroots``. 

566 """ 

567 sort = sort if sort is not None else prefs.sortroots 

568 

569 if self.isempty: 

570 return np.array([]) 

571 

572 # Sample on a Chebyshev grid and fit a Chebtech of the same resolution 

573 n = max(2 * self.size + 1, 33) 

574 cheb_pts = Chebtech._chebpts(n) 

575 vals = self(cheb_pts) 

576 ct = Chebtech(Chebtech._vals2coeffs(vals)) 

577 rts = rootsunit(ct.coeffs) 

578 rts = newtonroots(ct, rts) 

579 rts = np.clip(rts, -1.0, 1.0) 

580 return np.sort(rts) if sort else rts 

581 

582 # ------------------------------------------------------------------ 

583 # calculus 

584 # ------------------------------------------------------------------ 

585 

586 @self_empty(resultif=0.0) 

587 def sum(self) -> Any: 

588 """Definite integral of the Trigtech over [-1, 1]. 

589 

590 Only the DC coefficient contributes: 

591 ∫_{-1}^{1} exp(i*π*k*(x+1)) dx = 0 for k ≠ 0 

592 ∫_{-1}^{1} 1 dx = 2 for k = 0 

593 """ 

594 return 2.0 * float(np.real(self._coeffs[0])) 

595 

596 @self_empty() 

597 def cumsum(self) -> "Trigtech": 

598 """Indefinite integral, zero at x = -1, in Fourier coefficient space. 

599 

600 For mode k ≠ 0: antiderivative coefficient = c_k / (i*π*ω_k) 

601 For mode k = 0: set to the constant needed so that F(-1) = 0. 

602 

603 Note: if the DC component (self.coeffs[0]) is non-zero the true 

604 antiderivative contains a linear trend and is not periodic. We still 

605 return a Trigtech representing the *periodic* part, adjusted so that 

606 the result evaluates to 0 at x = -1. 

607 """ 

608 n = self.size 

609 c = self._coeffs.copy() 

610 freqs = np.fft.fftfreq(n) * n # FFT-order integer frequencies 

611 

612 int_c = np.zeros(n, dtype=complex) 

613 mask = freqs != 0 

614 int_c[mask] = c[mask] / (1j * np.pi * freqs[mask]) 

615 

616 # Enforce F(-1) = 0. 

617 # F(x) = Σ_k int_c[k] * exp(i*π*ω_k*(x+1)) 

618 # At x = -1: exp(i*π*ω_k*0) = 1 for all k, so F(-1) = Σ int_c 

619 # Set int_c[0] so that sum(int_c) = 0. 

620 int_c[0] = -np.sum(int_c[1:]) 

621 return self.__class__(int_c, interval=self._interval) 

622 

623 @self_empty() 

624 def diff(self) -> "Trigtech": 

625 """Derivative via the Fourier multiplier i*π*ω_k. 

626 

627 d/dx [c_k * exp(i*π*ω_k*(x+1))] = i*π*ω_k * c_k * exp(i*π*ω_k*(x+1)) 

628 """ 

629 if self.isconst: 

630 return self.__class__(np.array([0.0 + 0.0j]), interval=self._interval) 

631 n = self.size 

632 freqs = np.fft.fftfreq(n) * n 

633 d_coeffs = (1j * np.pi * freqs) * self._coeffs 

634 return self.__class__(d_coeffs, interval=self._interval) 

635 

636 # ------------------------------------------------------------------ 

637 # static helpers (FFT ↔ values) 

638 # ------------------------------------------------------------------ 

639 

640 @staticmethod 

641 def _trigpts(n: int) -> np.ndarray: 

642 """Return *n* equispaced points on [-1, 1).""" 

643 if n == 0: 

644 return np.array([]) 

645 return -1.0 + 2.0 * np.arange(n) / n 

646 

647 @staticmethod 

648 def _vals2coeffs(vals: Any) -> np.ndarray: 

649 """Convert values at equispaced points to FFT coefficients (divided by n). 

650 

651 Always returns complex128, even for real-valued inputs, because Fourier 

652 coefficients for functions such as sin are purely imaginary and would be 

653 discarded if forced to real. 

654 

655 Inverse of ``_coeffs2vals``. 

656 """ 

657 vals = np.asarray(vals) 

658 n = vals.size 

659 if n == 0: 

660 return np.array([], dtype=complex) 

661 return cast(np.ndarray, np.fft.fft(vals) / n) 

662 

663 @staticmethod 

664 def _coeffs2vals(coeffs: Any) -> np.ndarray: 

665 """Convert FFT coefficients (divided by n) to values at equispaced points. 

666 

667 Inverse of ``_vals2coeffs``. 

668 """ 

669 coeffs = np.asarray(coeffs, dtype=complex) 

670 n = coeffs.size 

671 if n == 0: 

672 return np.array([], dtype=float) 

673 vals = n * np.fft.ifft(coeffs) 

674 # Discard negligible imaginary parts for conjugate-symmetric coefficients 

675 max_real = float(np.max(np.abs(np.real(vals)))) 

676 if float(np.max(np.abs(np.imag(vals)))) < 1e-10 * max(max_real, 1.0): 

677 return np.real(vals) 

678 return cast(np.ndarray, vals) 

679 

680 # ------------------------------------------------------------------ 

681 # plotting 

682 # ------------------------------------------------------------------ 

683 

684 def plot(self, ax: Any = None, **kwargs: Any) -> Any: 

685 """Plot the Trigtech over [-1, 1]. 

686 

687 Args: 

688 ax: Matplotlib axes. If None, uses the current axes. 

689 **kwargs: Forwarded to matplotlib. 

690 

691 Returns: 

692 The axes on which the plot was drawn. 

693 """ 

694 return plotfun(self, (-1, 1), ax=ax, **kwargs) 

695 

696 def plotcoeffs(self, ax: Any = None, **kwargs: Any) -> Any: 

697 """Plot the absolute Fourier coefficient magnitudes in DC-centred order. 

698 

699 Uses ``_coeffs_to_plotorder()`` so the horizontal axis runs from 

700 the most-negative frequency on the left to the most-positive on 

701 the right, with DC in the centre. 

702 

703 Args: 

704 ax: Matplotlib axes. If None, uses the current axes. 

705 **kwargs: Forwarded to matplotlib. 

706 

707 Returns: 

708 The axes on which the plot was drawn. 

709 """ 

710 ax = ax or plt.gca() 

711 return plotfuncoeffs(np.abs(self._coeffs_to_plotorder()), ax=ax, **kwargs)