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

235 statements  

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

1"""Quasimatrix: a matrix with one continuous dimension. 

2 

3A quasimatrix is an inf x n matrix whose columns are Chebfun objects defined on 

4the same domain. This enables continuous analogues of linear algebra operations 

5such as QR factorization, SVD, least-squares, and more. 

6 

7Reference: Trefethen, "Householder triangularization of a quasimatrix," 

8IMA Journal of Numerical Analysis, 30 (2010), 887-897. 

9""" 

10 

11from __future__ import annotations 

12 

13from collections.abc import Iterator 

14from typing import Any, cast 

15 

16import matplotlib.pyplot as plt 

17import numpy as np 

18from matplotlib.axes import Axes 

19from matplotlib.patches import Rectangle 

20 

21from .chebfun import Chebfun 

22from .exceptions import SupportMismatch 

23from .settings import _preferences as prefs 

24 

25 

26class Quasimatrix: 

27 """An inf x n column quasimatrix whose columns are Chebfun objects. 

28 

29 A quasimatrix generalises the idea of a matrix so that one of its 

30 dimensions is continuous. Here the rows are indexed by points in an 

31 interval and the columns are Chebfun objects. 

32 

33 Attributes: 

34 columns: list of Chebfun objects forming the columns. 

35 

36 Examples: 

37 Build a quasimatrix from Chebfun columns. One dimension is 

38 continuous, so the shape is ``(inf, n)``: 

39 

40 >>> import numpy as np 

41 >>> from chebpy import chebfun 

42 >>> x = chebfun("x") 

43 >>> A = Quasimatrix([x, x**2]) 

44 >>> A.shape 

45 (inf, 2) 

46 >>> len(A) 

47 2 

48 

49 Calling evaluates every column at once: 

50 

51 >>> A(0.5).tolist() 

52 [0.5, 0.25] 

53 

54 Matrix-vector products contract the discrete dimension, giving a 

55 Chebfun back: 

56 

57 >>> f = A @ np.array([1.0, 1.0]) 

58 >>> bool(abs(f(0.5) - 0.75) < 1e-13) 

59 True 

60 

61 The continuous analogue of QR yields orthonormal columns: 

62 

63 >>> Q, R = A.qr() 

64 >>> R.shape 

65 (2, 2) 

66 >>> bool(np.allclose(Q.inner(), np.eye(2), atol=1e-10)) 

67 True 

68 """ 

69 

70 # ------------------------------------------------------------------ 

71 # Construction 

72 # ------------------------------------------------------------------ 

73 def __init__(self, columns: list[Any]) -> None: 

74 """Initialise from a list of Chebfun objects, callables, or scalars.""" 

75 cols: list[Chebfun] = [] 

76 for c in columns: 

77 if isinstance(c, Chebfun): 

78 cols.append(c) 

79 elif callable(c): 

80 cols.append(Chebfun.initfun_adaptive(c, cols[0].domain if cols else None)) 

81 else: 

82 # scalar → constant chebfun on the domain of the first column 

83 if cols: 

84 cols.append(Chebfun.initconst(float(c), cols[0].domain)) 

85 else: 

86 cols.append(Chebfun.initconst(float(c), prefs.domain)) 

87 if len(cols) > 1: 

88 # Verify all columns share the same support 

89 ref = cols[0].support 

90 for k, col in enumerate(cols[1:], 1): 

91 if col.support != ref: 

92 msg = f"Column {k} support {col.support} does not match column 0 support {ref}" 

93 raise SupportMismatch(msg) 

94 self.columns: list[Chebfun] = cols 

95 

96 # ------------------------------------------------------------------ 

97 # Properties 

98 # ------------------------------------------------------------------ 

99 @property 

100 def shape(self) -> tuple[float, int]: 

101 """Return (∞, n) where n is the number of columns.""" 

102 return (np.inf, len(self.columns)) 

103 

104 @property 

105 def T(self) -> _TransposedQuasimatrix: 

106 """Return the transpose (an n x inf row quasimatrix).""" 

107 return _TransposedQuasimatrix(self) 

108 

109 @property 

110 def domain(self) -> Any: 

111 """Domain of the quasimatrix columns.""" 

112 if not self.columns: 

113 return None 

114 return self.columns[0].domain 

115 

116 @property 

117 def support(self) -> tuple[float, float]: 

118 """Support interval of the quasimatrix.""" 

119 if not self.columns: 

120 return (0.0, 0.0) 

121 return cast("tuple[float, float]", self.columns[0].support) 

122 

123 @property 

124 def isempty(self) -> bool: 

125 """Return True if the quasimatrix has no columns.""" 

126 return len(self.columns) == 0 

127 

128 # ------------------------------------------------------------------ 

129 # Indexing A[:, k], A(x, k) 

130 # ------------------------------------------------------------------ 

131 def __getitem__(self, key: Any) -> Any: 

132 """Column indexing: ``A[:, k]`` returns column k as a Chebfun.""" 

133 if isinstance(key, tuple): 

134 row, col = key 

135 if isinstance(col, slice): 

136 return Quasimatrix(self.columns[col]) 

137 # A[x, k] - evaluate column k at point x 

138 if isinstance(row, slice) and row == slice(None): 

139 return self.columns[col] 

140 return self.columns[col](row) 

141 # A[k] - return column k 

142 if isinstance(key, (int, np.integer)): 

143 return self.columns[key] 

144 if isinstance(key, slice): 

145 return Quasimatrix(self.columns[key]) 

146 raise TypeError(key) 

147 

148 def __len__(self) -> int: 

149 """Return the number of columns.""" 

150 return len(self.columns) 

151 

152 def __iter__(self) -> Iterator[Chebfun]: 

153 """Iterate over the columns.""" 

154 return iter(self.columns) 

155 

156 # ------------------------------------------------------------------ 

157 # Calling A(x) - evaluate all columns at x, return array 

158 # ------------------------------------------------------------------ 

159 def __call__(self, x: Any) -> np.ndarray: 

160 """Evaluate every column at *x* and return the results as an array. 

161 

162 If *x* is a scalar the result has shape ``(n,)``. 

163 If *x* is an array of length *m* the result has shape ``(m, n)``. 

164 """ 

165 vals = [col(x) for col in self.columns] 

166 return np.column_stack(vals) if np.ndim(x) else np.array(vals) 

167 

168 # ------------------------------------------------------------------ 

169 # Arithmetic 

170 # ------------------------------------------------------------------ 

171 def __matmul__(self, other: Any) -> Any: 

172 """Matrix-vector product: ``A @ c`` returns a Chebfun. 

173 

174 *other* must be a 1-D array-like of length n. 

175 """ 

176 c = np.asarray(other, dtype=float) 

177 if c.ndim != 1 or len(c) != len(self.columns): 

178 msg = f"Cannot multiply {self.shape} quasimatrix by vector of length {len(c)}" 

179 raise ValueError(msg) 

180 result = c[0] * self.columns[0] 

181 for coeff, col in zip(c[1:], self.columns[1:], strict=True): 

182 result = result + coeff * col 

183 return result 

184 

185 def __mul__(self, other: Any) -> Quasimatrix: 

186 """Element-wise scalar multiplication.""" 

187 return Quasimatrix([c * other for c in self.columns]) 

188 

189 def __rmul__(self, other: Any) -> Quasimatrix: 

190 """Right scalar multiplication.""" 

191 return self.__mul__(other) 

192 

193 # ------------------------------------------------------------------ 

194 # Integrals and inner products 

195 # ------------------------------------------------------------------ 

196 def sum(self) -> np.ndarray: 

197 """Definite integral of each column (column sums).""" 

198 return np.array([col.sum() for col in self.columns]) 

199 

200 def inner(self, other: Quasimatrix | None = None) -> np.ndarray: 

201 """Gram matrix ``self.T @ other`` (or ``self.T @ self``).""" 

202 other = other if other is not None else self 

203 m = len(self.columns) 

204 n = len(other.columns) 

205 G = np.empty((m, n)) 

206 for i in range(m): 

207 for j in range(n): 

208 G[i, j] = self.columns[i].dot(other.columns[j]) 

209 return G 

210 

211 # ------------------------------------------------------------------ 

212 # QR factorization (modified Gram-Schmidt) 

213 # ------------------------------------------------------------------ 

214 def qr(self) -> tuple[Quasimatrix, np.ndarray]: 

215 """Compute the reduced QR factorization ``A = Q R``. 

216 

217 Uses modified Gram-Schmidt orthogonalisation in function space. 

218 

219 Returns: 

220 Q: Quasimatrix with orthonormal columns. 

221 R: Upper-triangular n x n NumPy array. 

222 """ 

223 n = len(self.columns) 

224 Q = [col.copy() for col in self.columns] 

225 R = np.zeros((n, n)) 

226 for k in range(n): 

227 for j in range(k): 

228 R[j, k] = Q[j].dot(Q[k]) 

229 Q[k] = Q[k] - R[j, k] * Q[j] 

230 R[k, k] = Q[k].norm(2) 

231 if R[k, k] == 0: 

232 msg = "Rank-deficient quasimatrix: QR factorization failed" 

233 raise np.linalg.LinAlgError(msg) 

234 Q[k] = (1.0 / R[k, k]) * Q[k] 

235 return Quasimatrix(Q), R 

236 

237 # ------------------------------------------------------------------ 

238 # SVD 

239 # ------------------------------------------------------------------ 

240 def svd(self) -> tuple[Quasimatrix, np.ndarray, np.ndarray]: 

241 """Compute the reduced SVD ``A = U S V^T``. 

242 

243 Returns: 

244 U: inf x n quasimatrix with orthonormal columns. 

245 S: 1-D array of singular values (length n). 

246 V: n x n orthogonal NumPy matrix. 

247 """ 

248 Q, R = self.qr() 

249 # Economy SVD of the n x n matrix R 

250 U_r, S, Vt = np.linalg.svd(R, full_matrices=False) 

251 # U = Q @ U_r (linear combinations of orthonormal columns) 

252 U_cols = [] 

253 for j in range(U_r.shape[1]): 

254 U_cols.append(Q @ U_r[:, j]) 

255 return Quasimatrix(U_cols), S, Vt.T # V = Vt.T 

256 

257 # ------------------------------------------------------------------ 

258 # Least-squares (backslash) 

259 # ------------------------------------------------------------------ 

260 def solve(self, f: Chebfun) -> np.ndarray: 

261 r"""Least-squares solution ``c`` to ``A c ~ f``. 

262 

263 Equivalent to MATLAB ``A\f``. Computed via QR factorisation. 

264 """ 

265 Q, R = self.qr() 

266 # b = Q' * f (inner products) 

267 b = np.array([col.dot(f) for col in Q.columns]) 

268 # Solve R c = b (back-substitution) 

269 return np.linalg.solve(R, b) 

270 

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

272 # Norms 

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

274 def norm(self, p: Any = "fro") -> float: 

275 """Compute the norm of the quasimatrix. 

276 

277 Args: 

278 p: Norm type. 

279 - 2: the 2-norm (largest singular value). 

280 - 1: max column 1-norm. 

281 - np.inf: max row-sum := max_x sum_j |A_j(x)|. 

282 - 'fro': Frobenius norm (default). 

283 """ 

284 if p == 2: 

285 _, S, _ = self.svd() 

286 return float(S[0]) 

287 if p == 1: 

288 return float(max(col.norm(1) for col in self.columns)) 

289 if p == np.inf: 

290 abssum = self.columns[0].absolute() 

291 for col in self.columns[1:]: 

292 abssum = abssum + col.absolute() 

293 return float(abssum.norm(np.inf)) 

294 if p == "fro": 

295 _, S, _ = self.svd() 

296 return float(np.sqrt(np.sum(S**2))) 

297 raise ValueError(f"Unsupported norm type: {p}") # noqa: TRY003 

298 

299 # ------------------------------------------------------------------ 

300 # Condition number 

301 # ------------------------------------------------------------------ 

302 def cond(self) -> float: 

303 """2-norm condition number (ratio of largest to smallest singular value).""" 

304 _, S, _ = self.svd() 

305 return float(S[0] / S[-1]) 

306 

307 # ------------------------------------------------------------------ 

308 # Rank 

309 # ------------------------------------------------------------------ 

310 def rank(self, tol: float | None = None) -> int: 

311 """Numerical rank (number of significant singular values).""" 

312 _, S, _ = self.svd() 

313 if tol is None: 

314 tol = max(self.shape[1], 20) * np.finfo(float).eps * S[0] 

315 return int(np.sum(tol < S)) 

316 

317 # ------------------------------------------------------------------ 

318 # Null space 

319 # ------------------------------------------------------------------ 

320 def null(self, tol: float | None = None) -> np.ndarray: 

321 """Orthonormal basis for the null space of the quasimatrix. 

322 

323 Returns an n x k NumPy array whose columns span ``null(A)``. 

324 """ 

325 _, S, V = self.svd() 

326 if tol is None: 

327 tol = max(self.shape[1], 20) * np.finfo(float).eps * S[0] 

328 mask = tol >= S 

329 return V[:, mask] 

330 

331 # ------------------------------------------------------------------ 

332 # Orth (orthonormal basis for range) 

333 # ------------------------------------------------------------------ 

334 def orth(self, tol: float | None = None) -> Quasimatrix: 

335 """Orthonormal basis for the column space (range) of the quasimatrix.""" 

336 U, S, _ = self.svd() 

337 if tol is None: 

338 tol = max(self.shape[1], 20) * np.finfo(float).eps * S[0] 

339 mask = tol < S 

340 return Quasimatrix([U.columns[j] for j in range(len(S)) if mask[j]]) 

341 

342 # ------------------------------------------------------------------ 

343 # Pseudoinverse 

344 # ------------------------------------------------------------------ 

345 def pinv(self) -> _TransposedQuasimatrix: 

346 """Moore-Penrose pseudoinverse (returned as an n x inf row quasimatrix). 

347 

348 ``pinv(A) @ f`` gives the same result as ``A.solve(f)``. 

349 """ 

350 U, S, V = self.svd() 

351 # pinv(A) = V S^{-1} U^T 

352 # The rows of pinv(A) are: sum_k V[i,k] / S[k] * U_k 

353 n = len(S) 

354 pinv_cols: list[Chebfun] = [] 

355 for i in range(n): 

356 col = (V[i, 0] / S[0]) * U.columns[0] 

357 for k in range(1, n): 

358 col = col + (V[i, k] / S[k]) * U.columns[k] 

359 pinv_cols.append(col) 

360 return _TransposedQuasimatrix(Quasimatrix(pinv_cols)) 

361 

362 # ------------------------------------------------------------------ 

363 # Plotting 

364 # ------------------------------------------------------------------ 

365 def plot(self, ax: Axes | None = None, **kwds: Any) -> Axes: 

366 """Plot all columns on the same axes.""" 

367 ax = ax or plt.gca() 

368 for col in self.columns: 

369 col.plot(ax=ax, **kwds) 

370 return ax 

371 

372 def spy(self, ax: Axes | None = None, **kwds: Any) -> Axes: 

373 """Visualise the shape of the quasimatrix. 

374 

375 Draws a rectangle representing the inf x n structure, with a dot for 

376 each column to indicate nonzero content. 

377 """ 

378 ax = ax or plt.gca() 

379 n = len(self.columns) 

380 # Draw the bounding rectangle 

381 rect = Rectangle((0.5, 0.5), n, 10, fill=False, edgecolor="black", linewidth=1.5) 

382 ax.add_patch(rect) 

383 # A dot for each column 

384 for j in range(n): 

385 ax.plot(j + 1, 5.5, "bs", markersize=8, **kwds) 

386 ax.set_xlim(0, n + 1) 

387 ax.set_ylim(0, 11) 

388 ax.set_aspect("equal") 

389 ax.set_xlabel(f"n = {n}") 

390 ax.set_ylabel("∞") 

391 ax.set_xticks(range(1, n + 1)) 

392 ax.set_yticks([]) 

393 return ax 

394 

395 # ------------------------------------------------------------------ 

396 # Representation 

397 # ------------------------------------------------------------------ 

398 def __repr__(self) -> str: 

399 """Return a string representation.""" 

400 n = len(self.columns) 

401 if n == 0: 

402 return "Quasimatrix(empty)" 

403 sup = self.support 

404 return f"Quasimatrix(inf x {n} on [{sup[0]}, {sup[1]}])" 

405 

406 def __str__(self) -> str: 

407 """Return a string representation.""" 

408 return self.__repr__() 

409 

410 

411class _TransposedQuasimatrix: 

412 """An n x inf row quasimatrix (transpose of a column quasimatrix). 

413 

414 This is a thin wrapper that enables ``A.T @ f`` and ``A.T @ B`` 

415 with the correct semantics. 

416 """ 

417 

418 def __init__(self, qm: Quasimatrix) -> None: 

419 """Wrap a column quasimatrix as its transpose.""" 

420 self._qm = qm 

421 

422 @property 

423 def shape(self) -> tuple[int, float]: 

424 """Return (n, inf).""" 

425 return (len(self._qm.columns), np.inf) 

426 

427 @property 

428 def T(self) -> Quasimatrix: 

429 """Return the original column quasimatrix.""" 

430 return self._qm 

431 

432 def __matmul__(self, other: Any) -> Any: 

433 """Compute inner products: ``A.T @ f`` or ``A.T @ B``.""" 

434 if isinstance(other, Quasimatrix): 

435 return self._qm.inner(other) 

436 if isinstance(other, Chebfun): 

437 return np.array([col.dot(other) for col in self._qm.columns]) 

438 raise TypeError(f"Cannot multiply _TransposedQuasimatrix by {type(other)}") # noqa: TRY003 

439 

440 def spy(self, ax: Axes | None = None, **kwds: Any) -> Axes: 

441 """Visualise the shape of the transposed quasimatrix.""" 

442 ax = ax or plt.gca() 

443 n = len(self._qm.columns) 

444 rect = Rectangle((0.5, 0.5), 10, n, fill=False, edgecolor="black", linewidth=1.5) 

445 ax.add_patch(rect) 

446 for j in range(n): 

447 ax.plot(5.5, j + 1, "bs", markersize=8, **kwds) 

448 ax.set_xlim(0, 11) 

449 ax.set_ylim(0, n + 1) 

450 ax.set_aspect("equal") 

451 ax.set_ylabel(f"n = {n}") 

452 ax.set_xlabel("∞") 

453 ax.set_yticks(range(1, n + 1)) 

454 ax.set_xticks([]) 

455 return ax 

456 

457 def __repr__(self) -> str: 

458 """Return a string representation.""" 

459 n = len(self._qm.columns) 

460 if n == 0: 

461 return "_TransposedQuasimatrix(empty)" 

462 sup = self._qm.support 

463 return f"_TransposedQuasimatrix({n}x inf on [{sup[0]}, {sup[1]}])" 

464 

465 

466# ------------------------------------------------------------------ 

467# Module-level convenience functions 

468# ------------------------------------------------------------------ 

469def polyfit(f: Chebfun, n: int) -> Chebfun: 

470 """Least-squares polynomial fit of degree *n* to a Chebfun *f*. 

471 

472 Returns a Chebfun representing the best degree-*n* polynomial 

473 approximation to *f* in the L²-norm. 

474 """ 

475 x = Chebfun.initidentity(f.domain) 

476 cols: list[Chebfun] = [Chebfun.initconst(1.0, f.domain)] 

477 xk = cols[0] 

478 for _ in range(n): 

479 xk = xk * x 

480 cols.append(xk) 

481 A = Quasimatrix(cols) 

482 c = A.solve(f) 

483 return cast("Chebfun", A @ c)