Coverage for src/chebpy/chebtech.py: 100%
256 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"""Implementation of Chebyshev polynomial technology for function approximation.
3This module provides the Chebtech class, which is an abstract base class for
4representing functions using Chebyshev polynomial expansions. It serves as the
5foundation for the Chebtech class, which uses Chebyshev points of the second kind.
7The Chebtech classes implement core functionality for working with Chebyshev
8expansions, including:
9- Function evaluation using Clenshaw's algorithm or barycentric interpolation
10- Algebraic operations (addition, multiplication, etc.)
11- Calculus operations (differentiation, integration, etc.)
12- Rootfinding
13- Plotting
15These classes are primarily used internally by higher-level classes like Bndfun
16and Chebfun, rather than being used directly by end users.
17"""
19from abc import ABC
20from collections.abc import Callable
21from typing import Any
23import matplotlib.pyplot as plt
24import numpy as np
26from .algorithms import (
27 adaptive,
28 bary,
29 barywts2,
30 chebpts2,
31 clenshaw,
32 coeffmult,
33 coeffs2vals2,
34 newtonroots,
35 rootsunit,
36 standard_chop,
37 vals2coeffs2,
38)
39from .decorators import self_empty
40from .exceptions import BadFunLengthArgument
41from .plotting import plotfun, plotfuncoeffs
42from .settings import _preferences as prefs
43from .smoothfun import Smoothfun
44from .utilities import Interval, coerce_list
47class Chebtech(Smoothfun, ABC):
48 """Abstract base class serving as the template for Chebtech1 and Chebtech subclasses.
50 Chebtech objects always work with first-kind coefficients, so much
51 of the core operational functionality is defined this level.
53 The user will rarely work with these classes directly so we make
54 several assumptions regarding input data types.
56 Examples:
57 The adaptive constructor picks however many coefficients the function
58 needs to reach machine precision on [-1, 1]:
60 >>> import numpy as np
61 >>> f = Chebtech.initfun_adaptive(np.exp)
62 >>> bool(abs(f(0.0) - 1.0) < 1e-14)
63 True
64 >>> bool(abs(f(0.5) - np.exp(0.5)) < 1e-14)
65 True
67 A constant needs exactly one:
69 >>> c = Chebtech.initconst(3.0)
70 >>> c.size
71 1
72 >>> c.isconst
73 True
75 Coefficients are in the T_k basis, so the identity is ``[0, 1]``:
77 >>> Chebtech.initidentity().coeffs.tolist()
78 [0, 1]
79 """
81 @classmethod
82 def initconst(cls, c: Any = None, *, interval: Any = None) -> Any:
83 """Initialise a Chebtech from a constant c."""
84 if not np.isscalar(c):
85 raise ValueError(c)
86 if isinstance(c, int):
87 c = float(c)
88 return cls(np.array([c]), interval=interval)
90 @classmethod
91 def initempty(cls, *, interval: Any = None) -> "Chebtech":
92 """Initialise an empty Chebtech."""
93 return cls(np.array([]), interval=interval)
95 @classmethod
96 def initidentity(cls, *, interval: Any = None) -> "Chebtech":
97 """Chebtech representation of f(x) = x on [-1,1]."""
98 return cls(np.array([0, 1]), interval=interval)
100 @classmethod
101 def initfun(cls, fun: Any = None, n: Any = None, *, interval: Any = None) -> Any:
102 """Convenience constructor to automatically select the adaptive or fixedlen constructor.
104 This constructor automatically selects between the adaptive or fixed-length
105 constructor based on the input arguments passed.
106 """
107 if n is None:
108 return cls.initfun_adaptive(fun, interval=interval)
109 else:
110 return cls.initfun_fixedlen(fun, n, interval=interval)
112 @classmethod
113 def initfun_fixedlen(cls, fun: Any = None, n: Any = None, *, interval: Any = None) -> Any:
114 """Initialise a Chebtech from the callable fun using n degrees of freedom.
116 This constructor creates a Chebtech representation of the function using
117 a fixed number of degrees of freedom specified by n.
118 """
119 if n is None:
120 raise BadFunLengthArgument("n must be specified for fixed-length initialization") # noqa: TRY003
121 points = cls._chebpts(int(n))
122 values = fun(points)
123 coeffs = vals2coeffs2(values)
124 return cls(coeffs, interval=interval)
126 @classmethod
127 def initfun_adaptive(cls, fun: Any = None, *, interval: Any = None) -> Any:
128 """Initialise a Chebtech from the callable fun utilising the adaptive constructor.
130 This constructor uses an adaptive algorithm to determine the appropriate
131 number of degrees of freedom needed to represent the function.
132 """
133 interval = interval if interval is not None else prefs.domain
134 interval = Interval(*interval)
135 coeffs = adaptive(cls, fun, hscale=interval.hscale)
136 return cls(coeffs, interval=interval)
138 @classmethod
139 def initvalues(cls, values: Any = None, *, interval: Any = None) -> Any:
140 """Initialise a Chebtech from an array of values at Chebyshev points."""
141 return cls(cls._vals2coeffs(values), interval=interval)
143 def __init__(self, coeffs: Any, interval: Any = None) -> None:
144 """Initialize a Chebtech object.
146 This method initializes a new Chebtech object with the given coefficients
147 and interval. If no interval is provided, the default interval from
148 preferences is used.
150 Args:
151 coeffs (array-like): The coefficients of the Chebyshev series.
152 interval (array-like, optional): The interval on which the function
153 is defined. Defaults to None, which uses the default interval
154 from preferences.
155 """
156 interval = interval if interval is not None else prefs.domain
157 self._coeffs = np.array(coeffs)
158 self._interval = Interval(*interval)
160 def __call__(self, x: Any, how: str = "clenshaw") -> Any:
161 """Evaluate the Chebtech at the given points.
163 Args:
164 x: Points at which to evaluate the Chebtech.
165 how (str, optional): Method to use for evaluation. Either "clenshaw" or "bary".
166 Defaults to "clenshaw".
168 Returns:
169 The values of the Chebtech at the given points.
171 Raises:
172 ValueError: If the specified method is not supported.
173 """
174 method: dict[str, Callable[[Any], Any]] = {
175 "clenshaw": self.__call__clenshaw,
176 "bary": self.__call__bary,
177 }
178 try:
179 return method[how](x)
180 except KeyError as err:
181 raise ValueError(how) from err
183 def __call__clenshaw(self, x: Any) -> Any:
184 """Evaluate at *x* using Clenshaw recurrence on the coefficients."""
185 return clenshaw(x, self.coeffs)
187 def __call__bary(self, x: Any) -> Any:
188 """Evaluate at *x* using the barycentric interpolation formula."""
189 fk = self.values()
190 xk = self._chebpts(fk.size)
191 vk = self._barywts(fk.size)
192 return bary(x, fk, xk, vk)
194 def __repr__(self) -> str:
195 """Return a string representation of the Chebtech.
197 Returns:
198 str: A string representation of the Chebtech.
199 """
200 out = f"<{self.__class__.__name__}{{{self.size}}}>"
201 return out
203 # ------------
204 # properties
205 # ------------
206 @property
207 def coeffs(self) -> np.ndarray:
208 """Chebyshev expansion coefficients in the T_k basis."""
209 return self._coeffs
211 @property
212 def interval(self) -> Interval:
213 """Interval that Chebtech is mapped to."""
214 return self._interval
216 @property
217 def size(self) -> int:
218 """Return the size of the object."""
219 return self.coeffs.size
221 @property
222 def isempty(self) -> bool:
223 """Return True if the Chebtech is empty."""
224 return self.size == 0
226 @property
227 def iscomplex(self) -> bool:
228 """Determine whether the underlying onefun is complex or real valued."""
229 return self._coeffs.dtype == complex
231 @property
232 def isconst(self) -> bool:
233 """Return True if the Chebtech represents a constant."""
234 return self.size == 1
236 @property
237 @self_empty(0.0)
238 def vscale(self) -> float:
239 """Estimate the vertical scale of a Chebtech."""
240 return float(np.abs(np.asarray(coerce_list(self.values()))).max())
242 # -----------
243 # utilities
244 # -----------
245 def copy(self) -> "Chebtech":
246 """Return a deep copy of the Chebtech."""
247 return self.__class__(self.coeffs.copy(), interval=self.interval.copy())
249 def imag(self) -> Any:
250 """Return the imaginary part of the Chebtech.
252 Returns:
253 Chebtech: A new Chebtech representing the imaginary part of this Chebtech.
254 If this Chebtech is real-valued, returns a zero Chebtech.
255 """
256 if self.iscomplex:
257 return self.__class__(np.imag(self.coeffs), self.interval)
258 else:
259 return self.initconst(0, interval=self.interval)
261 def prolong(self, n: int) -> "Chebtech":
262 """Return a Chebtech of length n.
264 Obtained either by truncating if n < self.size or zero-padding if n > self.size.
265 In all cases a deep copy is returned.
266 """
267 m = self.size
268 ak = self.coeffs
269 cls = self.__class__
270 if n - m < 0:
271 out = cls(ak[:n].copy(), interval=self.interval)
272 elif n - m > 0:
273 out = cls(np.append(ak, np.zeros(n - m)), interval=self.interval)
274 else:
275 out = self.copy()
276 return out
278 def real(self) -> "Chebtech":
279 """Return the real part of the Chebtech.
281 Returns:
282 Chebtech: A new Chebtech representing the real part of this Chebtech.
283 If this Chebtech is already real-valued, returns self.
284 """
285 if self.iscomplex:
286 return self.__class__(np.real(self.coeffs), self.interval)
287 else:
288 return self
290 def simplify(self) -> "Chebtech":
291 """Call standard_chop on the coefficients of self.
293 Returns a Chebtech comprised of a copy of the truncated coefficients.
294 """
295 # coefficients
296 oldlen = len(self.coeffs)
297 longself = self.prolong(max(17, oldlen))
298 cfs = longself.coeffs
299 # scale (decrease) tolerance by hscale
300 tol = prefs.eps * max(self.interval.hscale, 1)
301 # chop
302 npts = standard_chop(cfs, tol=tol)
303 npts = min(oldlen, npts)
304 # construct
305 return self.__class__(cfs[:npts].copy(), interval=self.interval)
307 def values(self) -> np.ndarray:
308 """Function values at Chebyshev points."""
309 return coeffs2vals2(self.coeffs)
311 # ---------
312 # algebra
313 # ---------
314 @self_empty()
315 def __add__(self, f: Any) -> Any:
316 """Add a scalar or another Chebtech to this Chebtech.
318 Args:
319 f: A scalar or another Chebtech to add to this Chebtech.
321 Returns:
322 Chebtech: A new Chebtech representing the sum.
323 """
324 cls = self.__class__
325 if np.isscalar(f):
326 if np.iscomplexobj(f):
327 dtype: Any = complex
328 else:
329 dtype = self.coeffs.dtype
330 cfs = np.array(self.coeffs, dtype=dtype)
331 cfs[0] += f
332 return cls(cfs, interval=self.interval)
333 else:
334 # TODO: is a more general decorator approach better here?
335 # TODO: for constant Chebtech, convert to constant and call __add__ again
336 if f.isempty:
337 return f.copy()
338 g = self
339 n, m = g.size, f.size
340 if n < m:
341 g = g.prolong(m)
342 elif m < n:
343 f = f.prolong(n)
344 cfs = f.coeffs + g.coeffs
346 # check for zero output
347 eps = prefs.eps
348 tol = 0.5 * eps * max([f.vscale, g.vscale])
349 if all(abs(cfs) < tol):
350 return cls.initconst(0.0, interval=self.interval)
351 else:
352 return cls(cfs, interval=self.interval)
354 @self_empty()
355 def __div__(self, f: Any) -> Any:
356 """Divide this Chebtech by a scalar or another Chebtech.
358 Args:
359 f: A scalar or another Chebtech to divide this Chebtech by.
361 Returns:
362 Chebtech: A new Chebtech representing the quotient.
363 """
364 cls = self.__class__
365 if np.isscalar(f):
366 cfs = 1.0 / np.asarray(f) * self.coeffs
367 return cls(cfs, interval=self.interval)
368 else:
369 # TODO: review with reference to __add__
370 if f.isempty:
371 return f.copy()
372 return cls.initfun_adaptive(lambda x: self(x) / f(x), interval=self.interval)
374 __truediv__ = __div__
376 @self_empty()
377 def __mul__(self, g: Any) -> Any:
378 """Multiply this Chebtech by a scalar or another Chebtech.
380 Args:
381 g: A scalar or another Chebtech to multiply this Chebtech by.
383 Returns:
384 Chebtech: A new Chebtech representing the product.
385 """
386 cls = self.__class__
387 if np.isscalar(g):
388 cfs = g * self.coeffs
389 return cls(cfs, interval=self.interval)
390 else:
391 # TODO: review with reference to __add__
392 if g.isempty:
393 return g.copy()
394 f = self
395 n = f.size + g.size - 1
396 f = f.prolong(n)
397 g = g.prolong(n)
398 cfs = coeffmult(f.coeffs, g.coeffs)
399 out = cls(cfs, interval=self.interval)
400 return out
402 def __neg__(self) -> "Chebtech":
403 """Return the negative of this Chebtech.
405 Returns:
406 Chebtech: A new Chebtech representing the negative of this Chebtech.
407 """
408 coeffs = -self.coeffs
409 return self.__class__(coeffs, interval=self.interval)
411 def __pos__(self) -> "Chebtech":
412 """Return this Chebtech (unary positive).
414 Returns:
415 Chebtech: This Chebtech (self).
416 """
417 return self
419 @self_empty()
420 def __pow__(self, f: Any) -> Any:
421 """Raise this Chebtech to a power.
423 Args:
424 f: The exponent, which can be a scalar or another Chebtech.
426 Returns:
427 Chebtech: A new Chebtech representing this Chebtech raised to the power f.
428 """
430 def powfun(fn: Any, x: Any) -> Any:
431 if np.isscalar(fn):
432 return fn
433 else:
434 return fn(x)
436 return self.__class__.initfun_adaptive(lambda x: np.power(self(x), powfun(f, x)), interval=self.interval)
438 def __rdiv__(self, f: Any) -> Any:
439 """Divide a scalar by this Chebtech.
441 This is called when f / self is executed and f is not a Chebtech.
443 Args:
444 f: A scalar to be divided by this Chebtech.
446 Returns:
447 Chebtech: A new Chebtech representing f divided by this Chebtech.
448 """
450 # Executed when __div__(f, self) fails, which is to say whenever f
451 # is not a Chebtech. We proceeed on the assumption f is a scalar.
452 def constfun(x: Any) -> Any:
453 return 0.0 * x + f
455 return self.__class__.initfun_adaptive(lambda x: constfun(x) / self(x), interval=self.interval)
457 __radd__ = __add__
459 def __rsub__(self, f: Any) -> Any:
460 """Subtract this Chebtech from a scalar.
462 This is called when f - self is executed and f is not a Chebtech.
464 Args:
465 f: A scalar from which to subtract this Chebtech.
467 Returns:
468 Chebtech: A new Chebtech representing f minus this Chebtech.
469 """
470 return -(self - f)
472 @self_empty()
473 def __rpow__(self, f: Any) -> Any:
474 """Raise a scalar to the power of this Chebtech.
476 This is called when f ** self is executed and f is not a Chebtech.
478 Args:
479 f: A scalar to be raised to the power of this Chebtech.
481 Returns:
482 Chebtech: A new Chebtech representing f raised to the power of this Chebtech.
483 """
484 return self.__class__.initfun_adaptive(lambda x: np.power(f, self(x)), interval=self.interval)
486 __rtruediv__ = __rdiv__
487 __rmul__ = __mul__
489 def __sub__(self, f: Any) -> Any:
490 """Subtract a scalar or another Chebtech from this Chebtech.
492 Args:
493 f: A scalar or another Chebtech to subtract from this Chebtech.
495 Returns:
496 Chebtech: A new Chebtech representing the difference.
497 """
498 return self + (-f)
500 # -------
501 # roots
502 # -------
503 def roots(self, sort: bool | None = None) -> np.ndarray:
504 """Compute the roots of the Chebtech on [-1,1].
506 Uses the coefficients in the associated Chebyshev series approximation.
507 """
508 sort = sort if sort is not None else prefs.sortroots
509 rts = rootsunit(self.coeffs)
510 rts = newtonroots(self, rts)
511 # fix problems with newton for roots that are numerically very close
512 rts = np.clip(rts, -1, 1) # if newton roots are just outside [-1,1]
513 rts = rts if not sort else np.sort(rts)
514 return rts
516 # ----------
517 # calculus
518 # ----------
519 # Note that function returns 0 for an empty Chebtech object; this is
520 # consistent with numpy, which returns zero for the sum of an empty array
521 @self_empty(resultif=0.0)
522 def sum(self) -> Any:
523 """Definite integral of a Chebtech on the interval [-1,1]."""
524 if self.isconst:
525 out = 2.0 * self(0.0)
526 else:
527 ak = self.coeffs.copy()
528 ak[1::2] = 0
529 kk = np.arange(2, ak.size)
530 ii = np.append([2, 0], 2 / (1 - kk**2))
531 out = (ak * ii).sum()
532 return out
534 @self_empty()
535 def cumsum(self) -> "Chebtech":
536 """Return a Chebtech object representing the indefinite integral.
538 Computes the indefinite integral of a Chebtech on the interval [-1,1].
539 The constant term is chosen such that F(-1) = 0.
540 """
541 n = self.size
542 ak = np.append(self.coeffs, [0, 0])
543 bk = np.zeros(n + 1, dtype=self.coeffs.dtype)
544 rk = np.arange(2, n + 1)
545 bk[2:] = 0.5 * (ak[1:n] - ak[3:]) / rk
546 bk[1] = ak[0] - 0.5 * ak[2]
547 vk = np.ones(n)
548 vk[1::2] = -1
549 bk[0] = (vk * bk[1:]).sum()
550 out = self.__class__(bk, interval=self.interval)
551 return out
553 @self_empty()
554 def diff(self) -> "Chebtech":
555 """Return a Chebtech object representing the derivative.
557 Computes the derivative of a Chebtech on the interval [-1,1].
558 """
559 if self.isconst:
560 out = self.__class__(np.array([0.0]), interval=self.interval)
561 else:
562 n = self.size
563 ak = self.coeffs
564 zk = np.zeros(n - 1, dtype=self.coeffs.dtype)
565 wk = 2 * np.arange(1, n)
566 vk = wk * ak[1:]
567 zk[-1::-2] = vk[-1::-2].cumsum()
568 zk[-2::-2] = vk[-2::-2].cumsum()
569 zk[0] = 0.5 * zk[0]
570 out = self.__class__(zk, interval=self.interval)
571 return out
573 @staticmethod
574 def _chebpts(n: int) -> np.ndarray:
575 """Return n Chebyshev points of the second-kind."""
576 return chebpts2(n)
578 @staticmethod
579 def _barywts(n: int) -> np.ndarray:
580 """Barycentric weights for Chebyshev points of 2nd kind."""
581 return barywts2(n)
583 @staticmethod
584 def _vals2coeffs(vals: Any) -> np.ndarray:
585 """Map function values at Chebyshev points of 2nd kind.
587 Converts values at Chebyshev points of 2nd kind to first-kind Chebyshev polynomial coefficients.
588 """
589 return vals2coeffs2(vals)
591 @staticmethod
592 def _coeffs2vals(coeffs: Any) -> np.ndarray:
593 """Map first-kind Chebyshev polynomial coefficients.
595 Converts first-kind Chebyshev polynomial coefficients to function values at Chebyshev points of 2nd kind.
596 """
597 return coeffs2vals2(coeffs)
599 # ----------
600 # plotting
601 # ----------
602 def plot(self, ax: Any = None, **kwargs: Any) -> Any:
603 """Plot the Chebtech on the interval [-1, 1].
605 Args:
606 ax (matplotlib.axes.Axes, optional): The axes on which to plot. Defaults to None.
607 **kwargs: Additional keyword arguments to pass to the plot function.
609 Returns:
610 matplotlib.lines.Line2D: The line object created by the plot.
611 """
612 return plotfun(self, (-1, 1), ax=ax, **kwargs)
614 def plotcoeffs(self, ax: Any = None, **kwargs: Any) -> Any:
615 """Plot the absolute values of the Chebyshev coefficients.
617 Args:
618 ax (matplotlib.axes.Axes, optional): The axes on which to plot. Defaults to None.
619 **kwargs: Additional keyword arguments to pass to the plot function.
621 Returns:
622 matplotlib.lines.Line2D: The line object created by the plot.
623 """
624 ax = ax or plt.gca()
625 return plotfuncoeffs(abs(self.coeffs), ax=ax, **kwargs)