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

1"""Implementation of Chebyshev polynomial technology for function approximation. 

2 

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. 

6 

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 

14 

15These classes are primarily used internally by higher-level classes like Bndfun 

16and Chebfun, rather than being used directly by end users. 

17""" 

18 

19from abc import ABC 

20from collections.abc import Callable 

21from typing import Any 

22 

23import matplotlib.pyplot as plt 

24import numpy as np 

25 

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 

45 

46 

47class Chebtech(Smoothfun, ABC): 

48 """Abstract base class serving as the template for Chebtech1 and Chebtech subclasses. 

49 

50 Chebtech objects always work with first-kind coefficients, so much 

51 of the core operational functionality is defined this level. 

52 

53 The user will rarely work with these classes directly so we make 

54 several assumptions regarding input data types. 

55 

56 Examples: 

57 The adaptive constructor picks however many coefficients the function 

58 needs to reach machine precision on [-1, 1]: 

59 

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 

66 

67 A constant needs exactly one: 

68 

69 >>> c = Chebtech.initconst(3.0) 

70 >>> c.size 

71 1 

72 >>> c.isconst 

73 True 

74 

75 Coefficients are in the T_k basis, so the identity is ``[0, 1]``: 

76 

77 >>> Chebtech.initidentity().coeffs.tolist() 

78 [0, 1] 

79 """ 

80 

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) 

89 

90 @classmethod 

91 def initempty(cls, *, interval: Any = None) -> "Chebtech": 

92 """Initialise an empty Chebtech.""" 

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

94 

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) 

99 

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. 

103 

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) 

111 

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. 

115 

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) 

125 

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. 

129 

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) 

137 

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) 

142 

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

144 """Initialize a Chebtech object. 

145 

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. 

149 

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) 

159 

160 def __call__(self, x: Any, how: str = "clenshaw") -> Any: 

161 """Evaluate the Chebtech at the given points. 

162 

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

167 

168 Returns: 

169 The values of the Chebtech at the given points. 

170 

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 

182 

183 def __call__clenshaw(self, x: Any) -> Any: 

184 """Evaluate at *x* using Clenshaw recurrence on the coefficients.""" 

185 return clenshaw(x, self.coeffs) 

186 

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) 

193 

194 def __repr__(self) -> str: 

195 """Return a string representation of the Chebtech. 

196 

197 Returns: 

198 str: A string representation of the Chebtech. 

199 """ 

200 out = f"<{self.__class__.__name__}{{{self.size}}}>" 

201 return out 

202 

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 

210 

211 @property 

212 def interval(self) -> Interval: 

213 """Interval that Chebtech is mapped to.""" 

214 return self._interval 

215 

216 @property 

217 def size(self) -> int: 

218 """Return the size of the object.""" 

219 return self.coeffs.size 

220 

221 @property 

222 def isempty(self) -> bool: 

223 """Return True if the Chebtech is empty.""" 

224 return self.size == 0 

225 

226 @property 

227 def iscomplex(self) -> bool: 

228 """Determine whether the underlying onefun is complex or real valued.""" 

229 return self._coeffs.dtype == complex 

230 

231 @property 

232 def isconst(self) -> bool: 

233 """Return True if the Chebtech represents a constant.""" 

234 return self.size == 1 

235 

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

241 

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

248 

249 def imag(self) -> Any: 

250 """Return the imaginary part of the Chebtech. 

251 

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) 

260 

261 def prolong(self, n: int) -> "Chebtech": 

262 """Return a Chebtech of length n. 

263 

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 

277 

278 def real(self) -> "Chebtech": 

279 """Return the real part of the Chebtech. 

280 

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 

289 

290 def simplify(self) -> "Chebtech": 

291 """Call standard_chop on the coefficients of self. 

292 

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) 

306 

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

308 """Function values at Chebyshev points.""" 

309 return coeffs2vals2(self.coeffs) 

310 

311 # --------- 

312 # algebra 

313 # --------- 

314 @self_empty() 

315 def __add__(self, f: Any) -> Any: 

316 """Add a scalar or another Chebtech to this Chebtech. 

317 

318 Args: 

319 f: A scalar or another Chebtech to add to this Chebtech. 

320 

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 

345 

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) 

353 

354 @self_empty() 

355 def __div__(self, f: Any) -> Any: 

356 """Divide this Chebtech by a scalar or another Chebtech. 

357 

358 Args: 

359 f: A scalar or another Chebtech to divide this Chebtech by. 

360 

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) 

373 

374 __truediv__ = __div__ 

375 

376 @self_empty() 

377 def __mul__(self, g: Any) -> Any: 

378 """Multiply this Chebtech by a scalar or another Chebtech. 

379 

380 Args: 

381 g: A scalar or another Chebtech to multiply this Chebtech by. 

382 

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 

401 

402 def __neg__(self) -> "Chebtech": 

403 """Return the negative of this Chebtech. 

404 

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) 

410 

411 def __pos__(self) -> "Chebtech": 

412 """Return this Chebtech (unary positive). 

413 

414 Returns: 

415 Chebtech: This Chebtech (self). 

416 """ 

417 return self 

418 

419 @self_empty() 

420 def __pow__(self, f: Any) -> Any: 

421 """Raise this Chebtech to a power. 

422 

423 Args: 

424 f: The exponent, which can be a scalar or another Chebtech. 

425 

426 Returns: 

427 Chebtech: A new Chebtech representing this Chebtech raised to the power f. 

428 """ 

429 

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

431 if np.isscalar(fn): 

432 return fn 

433 else: 

434 return fn(x) 

435 

436 return self.__class__.initfun_adaptive(lambda x: np.power(self(x), powfun(f, x)), interval=self.interval) 

437 

438 def __rdiv__(self, f: Any) -> Any: 

439 """Divide a scalar by this Chebtech. 

440 

441 This is called when f / self is executed and f is not a Chebtech. 

442 

443 Args: 

444 f: A scalar to be divided by this Chebtech. 

445 

446 Returns: 

447 Chebtech: A new Chebtech representing f divided by this Chebtech. 

448 """ 

449 

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 

454 

455 return self.__class__.initfun_adaptive(lambda x: constfun(x) / self(x), interval=self.interval) 

456 

457 __radd__ = __add__ 

458 

459 def __rsub__(self, f: Any) -> Any: 

460 """Subtract this Chebtech from a scalar. 

461 

462 This is called when f - self is executed and f is not a Chebtech. 

463 

464 Args: 

465 f: A scalar from which to subtract this Chebtech. 

466 

467 Returns: 

468 Chebtech: A new Chebtech representing f minus this Chebtech. 

469 """ 

470 return -(self - f) 

471 

472 @self_empty() 

473 def __rpow__(self, f: Any) -> Any: 

474 """Raise a scalar to the power of this Chebtech. 

475 

476 This is called when f ** self is executed and f is not a Chebtech. 

477 

478 Args: 

479 f: A scalar to be raised to the power of this Chebtech. 

480 

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) 

485 

486 __rtruediv__ = __rdiv__ 

487 __rmul__ = __mul__ 

488 

489 def __sub__(self, f: Any) -> Any: 

490 """Subtract a scalar or another Chebtech from this Chebtech. 

491 

492 Args: 

493 f: A scalar or another Chebtech to subtract from this Chebtech. 

494 

495 Returns: 

496 Chebtech: A new Chebtech representing the difference. 

497 """ 

498 return self + (-f) 

499 

500 # ------- 

501 # roots 

502 # ------- 

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

504 """Compute the roots of the Chebtech on [-1,1]. 

505 

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 

515 

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 

533 

534 @self_empty() 

535 def cumsum(self) -> "Chebtech": 

536 """Return a Chebtech object representing the indefinite integral. 

537 

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 

552 

553 @self_empty() 

554 def diff(self) -> "Chebtech": 

555 """Return a Chebtech object representing the derivative. 

556 

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 

572 

573 @staticmethod 

574 def _chebpts(n: int) -> np.ndarray: 

575 """Return n Chebyshev points of the second-kind.""" 

576 return chebpts2(n) 

577 

578 @staticmethod 

579 def _barywts(n: int) -> np.ndarray: 

580 """Barycentric weights for Chebyshev points of 2nd kind.""" 

581 return barywts2(n) 

582 

583 @staticmethod 

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

585 """Map function values at Chebyshev points of 2nd kind. 

586 

587 Converts values at Chebyshev points of 2nd kind to first-kind Chebyshev polynomial coefficients. 

588 """ 

589 return vals2coeffs2(vals) 

590 

591 @staticmethod 

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

593 """Map first-kind Chebyshev polynomial coefficients. 

594 

595 Converts first-kind Chebyshev polynomial coefficients to function values at Chebyshev points of 2nd kind. 

596 """ 

597 return coeffs2vals2(coeffs) 

598 

599 # ---------- 

600 # plotting 

601 # ---------- 

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

603 """Plot the Chebtech on the interval [-1, 1]. 

604 

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. 

608 

609 Returns: 

610 matplotlib.lines.Line2D: The line object created by the plot. 

611 """ 

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

613 

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

615 """Plot the absolute values of the Chebyshev coefficients. 

616 

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. 

620 

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)