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
« prev ^ index » next coverage.py v7.16.2, created at 2026-10-01 13:43 +0000
1"""Trigonometric (Fourier) technology for periodic function approximation.
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:
7 Onefun → Smoothfun → Trigtech
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
14 coeffs[k] = (1/n) * sum_j f(x_j) * exp(-2*pi*i*j*k/n)
15 = (numpy.fft.fft(values) / n)[k]
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.
21Use ``_coeffs_to_plotorder()`` to obtain the human-readable DC-centred ordering
22(equivalent to ``numpy.fft.fftshift``).
24Evaluation
25----------
26Any point x ∈ [-1, 1] is evaluated via the DFT summation formula:
28 f(x) = Σ_k coeffs[k] * exp(i*π*ω_k*(x+1))
30where ω_k = numpy.fft.fftfreq(n)*n gives the integer frequencies in FFT order.
32References:
33----------
34* Trefethen, "Spectral Methods in MATLAB" (SIAM 2000)
35* Chebfun @trigtech (github.com/chebfun/chebfun)
36"""
38import warnings
39from abc import ABC
40from typing import Any, cast
42import matplotlib.pyplot as plt
43import numpy as np
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
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*.
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*.
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``).
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])
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
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])
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
123class Trigtech(Smoothfun, ABC):
124 """Trigonometric (Fourier) function approximation on [-1, 1].
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.
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.
137 Examples:
138 A periodic function needs only a handful of Fourier modes, where the
139 Chebyshev constructor would need many more:
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
150 Fixed-length construction samples on *n* equispaced points:
152 >>> g = Trigtech.initfun_fixedlen(lambda x: np.sin(np.pi * x), 16)
153 >>> g.size
154 16
155 """
157 # ------------------------------------------------------------------
158 # alternative constructors
159 # ------------------------------------------------------------------
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)
170 @classmethod
171 def initempty(cls, *, interval: Any = None) -> "Trigtech":
172 """Initialise an empty Trigtech."""
173 return cls(np.array([]), interval=interval)
175 @classmethod
176 def initidentity(cls, *, interval: Any = None) -> "Trigtech":
177 """Trigtech approximation of the identity f(x) = x on [-1, 1].
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)
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)
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)
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)
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)
216 # ------------------------------------------------------------------
217 # core dunder methods
218 # ------------------------------------------------------------------
220 def __init__(self, coeffs: Any, interval: Any = None) -> None:
221 """Initialise a Trigtech with FFT-order *coeffs* on *interval*.
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)).
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)
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.
239 f(x) = Σ_k coeffs[k] * exp(i*π*ω_k*(x+1))
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.
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()
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
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
267 def __repr__(self) -> str:
268 """Return a concise string representation."""
269 return f"<{self.__class__.__name__}{{{self.size}}}>"
271 # ------------------------------------------------------------------
272 # properties
273 # ------------------------------------------------------------------
275 @property
276 def coeffs(self) -> np.ndarray:
277 """Fourier coefficients in NumPy FFT order (always complex128)."""
278 return self._coeffs
280 @property
281 def interval(self) -> Interval:
282 """Interval that the Trigtech is mapped to."""
283 return self._interval
285 @property
286 def size(self) -> int:
287 """Number of stored Fourier coefficients."""
288 return self._coeffs.size
290 @property
291 def isempty(self) -> bool:
292 """True if the Trigtech has no coefficients."""
293 return self.size == 0
295 @property
296 def iscomplex(self) -> bool:
297 """True if the function is complex-valued (values have a non-negligible imaginary part).
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))
314 @property
315 def isconst(self) -> bool:
316 """True if the Trigtech represents a constant (single coefficient)."""
317 return self.size == 1
319 @property
320 def isperiodic(self) -> bool:
321 """Always True: Trigtech always represents a periodic function."""
322 return True
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())
330 # ------------------------------------------------------------------
331 # utilities
332 # ------------------------------------------------------------------
334 def copy(self) -> "Trigtech":
335 """Return a deep copy."""
336 return self.__class__(self._coeffs.copy(), interval=self._interval.copy())
338 def imag(self) -> "Trigtech":
339 """Return the imaginary part of the function as a real-valued Trigtech.
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)
356 def prolong(self, n: int) -> "Trigtech":
357 """Return a Trigtech of length *n* (truncate or zero-pad in frequency space).
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()
367 centered = np.fft.fftshift(self._coeffs)
368 dc_src = m // 2
369 dc_tgt = n // 2
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)
381 def real(self) -> "Trigtech":
382 """Return the real part of the function as a real-valued Trigtech.
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)
399 def simplify(self) -> "Trigtech":
400 """Truncate high-frequency Fourier coefficients that are below tolerance.
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)
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)
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
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
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)
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)
439 def _coeffs_to_plotorder(self) -> np.ndarray:
440 """Return coefficients in DC-centred (human-readable) order.
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)
447 # ------------------------------------------------------------------
448 # algebra
449 # ------------------------------------------------------------------
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)
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)
485 __truediv__ = __div__
487 @self_empty()
488 def __mul__(self, g: Any) -> "Trigtech":
489 """Multiply this Trigtech by a scalar or another Trigtech.
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)
506 def __neg__(self) -> "Trigtech":
507 """Return the negation."""
508 return self.__class__(-self._coeffs, interval=self._interval)
510 def __pos__(self) -> "Trigtech":
511 """Return self (unary plus)."""
512 return self
514 @self_empty()
515 def __pow__(self, f: Any) -> "Trigtech":
516 """Raise this Trigtech to a power *f* (scalar or Trigtech)."""
518 def powfun(fn: Any, x: Any) -> Any:
519 return fn if np.isscalar(fn) else fn(x)
521 return self.__class__.initfun_adaptive(
522 lambda x: np.power(self(x), powfun(f, x)),
523 interval=self._interval,
524 )
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 )
533 __radd__ = __add__
534 __rmul__ = __mul__
535 __rtruediv__ = __rdiv__
537 def __rsub__(self, f: Any) -> "Trigtech":
538 """Compute f - self."""
539 return cast("Trigtech", -(self - f))
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 )
549 def __sub__(self, f: Any) -> "Trigtech":
550 """Subtract *f* (scalar or Trigtech) from this Trigtech."""
551 return cast("Trigtech", self + (-f))
553 # ------------------------------------------------------------------
554 # rootfinding
555 # ------------------------------------------------------------------
557 def roots(self, sort: bool | None = None) -> np.ndarray:
558 """Find the roots of this Trigtech on [-1, 1].
560 Converts to a Chebyshev representation via re-sampling on Chebyshev
561 points and delegates to the Chebtech colleague-matrix root-finder.
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
569 if self.isempty:
570 return np.array([])
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
582 # ------------------------------------------------------------------
583 # calculus
584 # ------------------------------------------------------------------
586 @self_empty(resultif=0.0)
587 def sum(self) -> Any:
588 """Definite integral of the Trigtech over [-1, 1].
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]))
596 @self_empty()
597 def cumsum(self) -> "Trigtech":
598 """Indefinite integral, zero at x = -1, in Fourier coefficient space.
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.
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
612 int_c = np.zeros(n, dtype=complex)
613 mask = freqs != 0
614 int_c[mask] = c[mask] / (1j * np.pi * freqs[mask])
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)
623 @self_empty()
624 def diff(self) -> "Trigtech":
625 """Derivative via the Fourier multiplier i*π*ω_k.
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)
636 # ------------------------------------------------------------------
637 # static helpers (FFT ↔ values)
638 # ------------------------------------------------------------------
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
647 @staticmethod
648 def _vals2coeffs(vals: Any) -> np.ndarray:
649 """Convert values at equispaced points to FFT coefficients (divided by n).
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.
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)
663 @staticmethod
664 def _coeffs2vals(coeffs: Any) -> np.ndarray:
665 """Convert FFT coefficients (divided by n) to values at equispaced points.
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)
680 # ------------------------------------------------------------------
681 # plotting
682 # ------------------------------------------------------------------
684 def plot(self, ax: Any = None, **kwargs: Any) -> Any:
685 """Plot the Trigtech over [-1, 1].
687 Args:
688 ax: Matplotlib axes. If None, uses the current axes.
689 **kwargs: Forwarded to matplotlib.
691 Returns:
692 The axes on which the plot was drawn.
693 """
694 return plotfun(self, (-1, 1), ax=ax, **kwargs)
696 def plotcoeffs(self, ax: Any = None, **kwargs: Any) -> Any:
697 """Plot the absolute Fourier coefficient magnitudes in DC-centred order.
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.
703 Args:
704 ax: Matplotlib axes. If None, uses the current axes.
705 **kwargs: Forwarded to matplotlib.
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)