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
« prev ^ index » next coverage.py v7.16.2, created at 2026-10-01 13:43 +0000
1"""Quasimatrix: a matrix with one continuous dimension.
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.
7Reference: Trefethen, "Householder triangularization of a quasimatrix,"
8IMA Journal of Numerical Analysis, 30 (2010), 887-897.
9"""
11from __future__ import annotations
13from collections.abc import Iterator
14from typing import Any, cast
16import matplotlib.pyplot as plt
17import numpy as np
18from matplotlib.axes import Axes
19from matplotlib.patches import Rectangle
21from .chebfun import Chebfun
22from .exceptions import SupportMismatch
23from .settings import _preferences as prefs
26class Quasimatrix:
27 """An inf x n column quasimatrix whose columns are Chebfun objects.
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.
33 Attributes:
34 columns: list of Chebfun objects forming the columns.
36 Examples:
37 Build a quasimatrix from Chebfun columns. One dimension is
38 continuous, so the shape is ``(inf, n)``:
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
49 Calling evaluates every column at once:
51 >>> A(0.5).tolist()
52 [0.5, 0.25]
54 Matrix-vector products contract the discrete dimension, giving a
55 Chebfun back:
57 >>> f = A @ np.array([1.0, 1.0])
58 >>> bool(abs(f(0.5) - 0.75) < 1e-13)
59 True
61 The continuous analogue of QR yields orthonormal columns:
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 """
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
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))
104 @property
105 def T(self) -> _TransposedQuasimatrix:
106 """Return the transpose (an n x inf row quasimatrix)."""
107 return _TransposedQuasimatrix(self)
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
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)
123 @property
124 def isempty(self) -> bool:
125 """Return True if the quasimatrix has no columns."""
126 return len(self.columns) == 0
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)
148 def __len__(self) -> int:
149 """Return the number of columns."""
150 return len(self.columns)
152 def __iter__(self) -> Iterator[Chebfun]:
153 """Iterate over the columns."""
154 return iter(self.columns)
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.
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)
168 # ------------------------------------------------------------------
169 # Arithmetic
170 # ------------------------------------------------------------------
171 def __matmul__(self, other: Any) -> Any:
172 """Matrix-vector product: ``A @ c`` returns a Chebfun.
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
185 def __mul__(self, other: Any) -> Quasimatrix:
186 """Element-wise scalar multiplication."""
187 return Quasimatrix([c * other for c in self.columns])
189 def __rmul__(self, other: Any) -> Quasimatrix:
190 """Right scalar multiplication."""
191 return self.__mul__(other)
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])
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
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``.
217 Uses modified Gram-Schmidt orthogonalisation in function space.
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
237 # ------------------------------------------------------------------
238 # SVD
239 # ------------------------------------------------------------------
240 def svd(self) -> tuple[Quasimatrix, np.ndarray, np.ndarray]:
241 """Compute the reduced SVD ``A = U S V^T``.
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
257 # ------------------------------------------------------------------
258 # Least-squares (backslash)
259 # ------------------------------------------------------------------
260 def solve(self, f: Chebfun) -> np.ndarray:
261 r"""Least-squares solution ``c`` to ``A c ~ f``.
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)
271 # ------------------------------------------------------------------
272 # Norms
273 # ------------------------------------------------------------------
274 def norm(self, p: Any = "fro") -> float:
275 """Compute the norm of the quasimatrix.
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
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])
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))
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.
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]
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]])
342 # ------------------------------------------------------------------
343 # Pseudoinverse
344 # ------------------------------------------------------------------
345 def pinv(self) -> _TransposedQuasimatrix:
346 """Moore-Penrose pseudoinverse (returned as an n x inf row quasimatrix).
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))
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
372 def spy(self, ax: Axes | None = None, **kwds: Any) -> Axes:
373 """Visualise the shape of the quasimatrix.
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
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]}])"
406 def __str__(self) -> str:
407 """Return a string representation."""
408 return self.__repr__()
411class _TransposedQuasimatrix:
412 """An n x inf row quasimatrix (transpose of a column quasimatrix).
414 This is a thin wrapper that enables ``A.T @ f`` and ``A.T @ B``
415 with the correct semantics.
416 """
418 def __init__(self, qm: Quasimatrix) -> None:
419 """Wrap a column quasimatrix as its transpose."""
420 self._qm = qm
422 @property
423 def shape(self) -> tuple[int, float]:
424 """Return (n, inf)."""
425 return (len(self._qm.columns), np.inf)
427 @property
428 def T(self) -> Quasimatrix:
429 """Return the original column quasimatrix."""
430 return self._qm
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
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
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]}])"
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*.
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)