ARTICLE DETAIL

建站实战干货

来自一线的建站与推广经验沉淀,每一条都经过真实交付验证。

pyprobml 线性代数实践指南:Probabilistic Machine Learning 第 7 章的 JAX 数值计算全解

2026/9/29 3:26:31 拓冰建站 浏览量
pyprobml 线性代数实践指南:Probabilistic Machine Learning 第 7 章的 JAX 数值计算全解 机器学习深度学习【免费下载链接】pyprobmlPython code for Probabilistic Machine learning book by Kevin Murphy项目地址https://gitcode.com/gh_mirrors/py/pyprobml点击查看免费下载导读本篇技术指南围绕 pyprobml 仓库中《Probabilistic Machine Learning: An Introduction》第 7 章Linear algebra线性代数的配套内容展开以 第 7 章 README 为骨架以该章节的补充教程 linalg.ipynb 为核心主体结合deprecated/scripts/下同名演示脚本的源码实现系统梳理张量基础操作、广播机制、Einstein 求和、特征值分解EVD、奇异值分解SVD、LU/QR/Cholesky 分解、JAX 自动微分以及线性方程组求解等主题。读完本文你将掌握第 7 章全部图表对应的生成代码位置与原理并能直接复用 JAX/NumPy/SciPy 的线性代数 API 解决机器学习建模中的数值计算问题。一、章节概览README 骨架与仓库布局notebooks/book1/07/目录是全书第 7 章的配套 Notebook 集合其 README.md 是一份图号 → Notebook → 图文件的映射索引是理解该章节内容组织方式的入口。目录下实际包含以下文件README.md图号映射表与补充资料索引linalg.ipynb章节补充教程Intro to linear algebra基于 JAX是本章技术主体的最完整载体gaussEvec.ipynb特征向量几何意义演示对应图 7.6height_weight_whiten_plot.ipynb数据标准化与白化对比对应图 7.7svd_image_demo.ipynbSVD 图像低秩近似对应图 7.9、7.10cholesky_demo.ipynb、power_method_demo.ipynb、einsum_demo.ipynb分别对应 Cholesky 分解、幂迭代法、Einstein 求和三个主题README 给出的完整图号映射表如下图中Notebook列为空表示该图由书籍官方图库提供本地仓库没有直接生成它的 Notebook非空列即本仓库中的 Notebook 文件| 图号 | Notebook | 图文件 | |--|--|--| | 7.1 | - | Figure_7.1.png | | 7.2 | - | Figure_7.2_A.png / Figure_7.2_B.png | | 7.3 | - | Figure_7.3_B.png / Figure_7.3_A.png | | 7.4 | - | Figure_7.4.png | | 7.5 | - | Figure_7.5.png | | 7.6 | gaussEvec.ipynb | - | | 7.7 | height_weight_whiten_plot.ipynb | - | | 7.8 | - | Figure_7.8_B.png / Figure_7.8_A.png | | 7.9 | svd_image_demo.ipynb | - | | 7.10 | svd_image_demo.ipynb | - | | 7.11 | - | Figure_7.11_A.png / Figure_7.11_B.png | | 7.12 | - | Figure_7.12.png |此外README 的 Supplementary material 一节登记了本章唯一的补充教程Intro to linear algebra软件栈为JAX即 linalg.ipynb。该 Notebook 的目录结构TOC为Basics → Sparse matrices → Broadcasting → Einstein summation → Eigenvalue decomposition (EVD) → Singular value decomposition (SVD) → Other decompositions → Matrix calculus → Linear systems of equations本文后续小节即沿此脉络展开。值得注意同一批图表的可执行脚本版本保留在 deprecated/scripts/ 下例如 gaussEvec.py、height_weight_whiten_plot.py、svd_image_demo.py、cholesky_demo.py、power_method_demo.py、einsum_demo.py。它们与 Notebook 版逻辑一致是阅读源码实现的首选入口。二、运行环境与依赖linalg.ipynb 的导入代码明确了本章的运行栈import numpy as np import jax import jax.numpy as jnp from jax import grad, hessian, jacfwd, jacrev, jit, vmap from jax.scipy.special import logsumexp import scipy import sklearn import matplotlib.pyplot as plt import seaborn as sns import pandas as pdNumPy提供np.linalg、np.random等基础接口JAX本章所有线性代数核心运算均以jnpjax.numpy演示并利用其grad / hessian / jacfwd / jacrev做矩阵微积分SciPy稀疏矩阵scipy.sparse、LU/QR 分解scipy.linalg、多元高斯采样scipy.stats.multivariate_normal依赖它其他matplotlib 绘图、seaborn 样式、pandas 数据展示、PIL/requests 用于图像读取。Notebook 还通过print(jax version {}.format(jax.__version__))打印 JAX 版本便于排查 API 差异。由于 JAX 是持续演进的框架不同版本间部分 API 存在差异例如 Notebook 中标注np.linalg.norm(A, ord2)、np.linalg.det、np.linalg.cond等当时not supported by jax运行前应确认本机 JAX 版本与 Notebook 编写时代匹配。此外图 7.6/7.9 等 Notebook 在开头会尝试import probml_utils as pml用于pml.savefig保存图片若未安装会通过%pip install -qq githttps://github.com/probml/probml-utils.git自动安装。三、张量与数组基础操作Basics3.1 向量的创建与索引v jnp.array([0, 1, 2]) # 1d 向量 print(v.ndim) # 1 print(v.shape) # (3,) # Python 使用 0 索引而非 1 索引 print(v[0], v[1], v[2]) # 0 1 23.2 二维数组的基本属性A jnp.array([[0, 1, 2], [3, 4, 5]]) print(A.ndim) # 2 print(A.shape) # (2,3) print(A.size) # 6 print(A.T.shape) # (3,2)3.3 向量 ↔ 单行/单列矩阵的等价写法将向量扩展为一行矩阵有四种等价方式Notebook 用断言验证了它们的等价性x jnp.array([1, 2]) # vector X1 jnp.array([x]) # matrix with one row X2 jnp.reshape(x, (1, -1)) X3 x[None, :] X4 x[jnp.newaxis, :] assert jnp.array_equal(X1, X2) assert jnp.array_equal(X1, X3) print(jnp.shape(X1)) # (1,2)扩展为一列矩阵的对称写法为x[:, None]、x[:, jnp.newaxis]、jnp.reshape(x, (-1, 1))、jnp.array([x]).T。在机器学习代码中x[:, None]与x[None, :]是调整维度对齐配合广播最高频的惯用法。3.4 实用构造模式one-hot 编码利用布尔比较 类型转换def one_hot(x, k, dtypejnp.float32): return jnp.array(x[:, None] jnp.arange(k), dtype) x jnp.array([1, 2, 0, 2]) X one_hot(x, 3)按列/按行堆叠重建数组jnp.stack([col0, col1, col2], axis1)等价于原矩阵水平/垂直拼接jnp.concatenate([M, C], axis1)与jnp.hstack([M, C])等价axis0时与jnp.vstack等价注意 JAX 不支持 NumPy 的np.c_语法Notebook 已注明。给数据矩阵加全 1 列线性回归/增广矩阵的经典惯用法X jnp.array([[9, 8, 7], [6, 5, 4]]) N jnp.shape(X)[0] # num. rows X1 jnp.hstack([jnp.ones((N, 1)), X])展平A.ravel()按行拼接得到[0 1 2 3 4 5]。3.5 内存布局row-major 与 column-majorNotebook 专门用一节说明NumPy 数组在内存中按row-major行主序布局即最右侧索引变化最快这与 C、Eigen、PyTorch 一致而 Julia、Matlab、R、Fortran 使用column-major列主序。这一差异对性能有实际影响——在 NumPy/JAX 中书写双重循环时应让内层循环遍历最右侧索引A jnp.reshape(jnp.arange(6), (2, 3)) d1, d2 jnp.shape(A) for i in range(d1): for j in range(d2): # Do something with A[i,j]3.6 张量reshape 与 transpose 的 view 语义T jnp.reshape(jnp.arange(24), (2, 3, 4)) # 0..23 填充 print(jnp.transpose(x, (1, 0, 2)).shape) # (2, 1, 3)Notebook 特别强调transpose只是改变下标 n 维向量 → 一维整数的映射并不会真正搬移内存数据那会很慢它提供的是同一份数据的视图view。这是理解张量运算零开销转置的关键。四、矩阵乘法、外积与范数4.1 矩阵乘法与元素乘的区别C jnp.dot(A, B) # shape (2,4) C2 A.dot(B) C3 A B # 三种写法等价 assert jnp.allclose(C, C2) and jnp.allclose(C, C3)注意A * B是逐元素乘积当形状不兼容时如(2,3) * (3,4)会直接报错不会像某些语言那样自动做矩阵乘。4.2 用点乘实现按行/列求和这是线性代数视角的经典技巧用全 1 矩阵做乘法即可实现求和且与sum(axis...)完全等价XS jnp.dot(jnp.ones((1, 2)), X) # 按列求和 (对行求和) XS2 jnp.sum(X, axis0) assert jnp.allclose(XS, XS2) XS jnp.dot(X, jnp.ones((3, 1))) # 按行求和 XS2 jnp.sum(X, axis1).reshape(-1, 1) assert jnp.allclose(XS, XS2) S1 jnp.dot(jnp.ones((1, 2)), jnp.dot(X, jnp.ones((3, 1))))[0] # 全元素和 S2 jnp.sum(X) assert jnp.allclose(S1, S2)4.3 外积、Kronecker 积与范数A jnp.outer(x, y) # 外积 x y^T jnp.kron(jnp.eye(2), jnp.ones((2, 2))) # Kronecker 积 # 范数 print(jnp.linalg.norm(x, 2) ** 2) # l2 范数平方 sum(x^2) print(jnp.linalg.norm(x, jnp.inf)) # 无穷范数 print(jnp.linalg.norm(A, ordfro)) # Frobenius 范数 # np.linalg.norm(A, ord2) 与 ordnuc 在 Notebook 编写时代未在 JAX 实现print(jnp.trace(A)) # np.linalg.det(A) 与 np.linalg.cond(A) 在 Notebook 编写时代未在 JAX 实现五、稀疏矩阵与特殊结构本章对稀疏结构的演示依赖 SciPy 而非 JAXfrom scipy.sparse import diags A diags([1, 2, 3]) # 三对角/对角稀疏矩阵 print(A.toarray()) from scipy.linalg import block_diag block_diag([2, 3], [[4, 5], [6, 7]]) # 块对角矩阵此外 Notebook 提及**带状对角矩阵band diagonal**可借助bandmat库PyPI 上的bandmat包处理这类结构在状态空间模型、样条平滑中常见。六、广播机制Broadcasting在 NumPy/JAX 中A * B对形状不同的数组做逐元素乘法时会自动广播broadcast即隐式复制某些维度以对齐形状。Notebook 给出了依次应用的转换规则若两个数组维数不同维数少的一侧在左边补 1例如标量变成向量、向量变成单行矩阵若某维度形状不匹配形状为 1 的那一侧在该维度被拉伸复制到另一侧的形状若某维度两边都不为 1 且不相等则报错。两个典型例子也是矩阵乘法与广播的等价对照# 按列缩放等价于右乘对角矩阵 X jnp.reshape(jnp.arange(6), (2, 3)) s jnp.array([1, 2, 3]) XS X * s XS2 jnp.dot(X, jnp.diag(s)) # post-multiply by diagonal assert jnp.allclose(XS, XS2) # 按行缩放等价于左乘对角矩阵 s jnp.array([1, 2]) XS X * jnp.reshape(s, (-1, 1)) XS2 jnp.dot(jnp.diag(s), X) # pre-multiply by diagonal assert jnp.allclose(XS, XS2)理解广播是正确书写向量化代码的前提——它把复制-逐元素运算的循环彻底隐式化既简洁又高效。七、Einstein 求和einsum从记号到代码einsum用inputs - outputs这种下标记号同时表达维度命名与张量缩并tensor contraction出现在输出中未命名的维度会被求和掉。本章提供了从单张量到多张量的完整对照表。7.1 单张量操作A jnp.arange(6).reshape(2, 3) S jnp.arange(9).reshape(3, 3) T np.random.randn(2, 2, 2, 2) assert jnp.allclose(A.T, jnp.einsum(ij-ji, A)) # 转置 assert jnp.allclose(jnp.sum(A), jnp.einsum(ij-, A)) # 全元素和 assert jnp.allclose(jnp.sum(A, axis0), jnp.einsum(ij-j, A)) # 按列和 assert jnp.allclose(jnp.sum(A, axis1), jnp.einsum(ij-i, A)) # 按行和 assert jnp.allclose(jnp.sum(T, axis1), jnp.einsum(ijkl-ikl, T)) assert jnp.allclose(np.diag(S), jnp.einsum(ii-i, S)) # 提取对角 assert jnp.allclose(np.trace(S), jnp.einsum(ii-, S)) # 迹7.2 双张量操作a jnp.arange(3); b jnp.arange(3) A jnp.arange(6).reshape(2, 3) B jnp.arange(15).reshape(3, 5) assert jnp.allclose(jnp.dot(A, b), jnp.einsum(ik,k-i, A, b)) # 矩阵-向量乘 assert jnp.allclose(jnp.dot(A, B), jnp.einsum(ik,kj-ij, A, B)) # 矩阵乘 assert jnp.allclose(jnp.matmul(A, B), jnp.einsum(ik,kj-ij, A, B)) assert jnp.allclose(jnp.dot(a, b), jnp.einsum(i,i-, a, b)) # 内积 assert jnp.allclose(jnp.outer(a, b), jnp.einsum(i,j-ij, a, b)) # 外积 assert jnp.allclose(a * a, jnp.einsum(i,i-i, a, a)) # 元素乘 assert jnp.allclose(CC, jnp.einsum(ijk,ikl-ijl, AA, BB)) # 批量矩阵乘其中批量矩阵乘的ijk,ikl-ijl正是深度学习 batch matmul 的标准写法。7.3 多张量实战句嵌入sentence embeddingNotebook 给出了一个教科书级的复合例子设S_ntk为批量n× 序列位置t× 词 one-hotk的三维张量W_kd为词嵌入矩阵V_dc为分类器线性层则词嵌入 → 句子向量 → logits的整条链路可以一步写成N, C, D, K, T 2, 3, 4, 5, 6 S np.random.randn(N, T, K) W np.random.randn(K, D) V np.random.randn(D, C) Lfast jnp.einsum(ntk,kd,dc-nc, S, W, V)Notebook 用五重 for 循环的暴力实现与Lfast做了allclose断言验证einsum一行等价于整个循环体。7.4 收缩路径优化einsum_path对多张量缩并einsum支持路径优化以控制中间张量大小与 FLOP 数path jnp.einsum_path(ntk,kd,dc-nc, S, W, V, optimizeoptimal)[0] assert jnp.allclose(L, jnp.einsum(ntk,kd,dc-nc, S, W, V, optimizepath))einsum_demo.py 还演示了一个来自 Koller Friedman 的学生网络student network的完整图模型收缩c,dc,gdi,si,lg,jls,hgj-其注释给出了路径优化的量化收益——naive 缩并 FLOP 约 2.734e06optimal优化后约 2.176e03理论加速约 1256 倍greedy策略约 7.101e03约 385 倍最大中间张量 1.25e02 元素。这说明einsum 不只是语法糖配合路径优化本身就是一种性能优化手段。八、特征值分解EVD8.1 对称矩阵的 EVDjnp.linalg.eighnp.random.seed(42) M np.random.randn(4, 4) A M M.T # 构造对称矩阵 assert (A A.T).all() evals, evecs jnp.linalg.eigh(A) # 告诉 JAX 矩阵是对称的 # 按特征值绝对值从大到小排序列随之排序 idx jnp.argsort(jnp.abs(evals))[::-1] evecs evecs[:, idx] evals evals[idx]eigh利用对称性使用更快的专用算法np.linalg.eig则面向一般矩阵。特征向量排序在 PCA、谱方法中是标准化操作。8.2 旋转矩阵的对角化实例Notebook 构造了一个绕 z 轴旋转 45° → 按 diag(1,2,3) 缩放 → 反向旋转 -45°的矩阵A R S R^T再通过 EVD 把这些成分还原出来a (45 / 180) * jnp.pi R jnp.array([[jnp.cos(a), -jnp.sin(a), 0], [jnp.sin(a), jnp.cos(a), 0], [0, 0, 1]]) S jnp.diag(jnp.array([1.0, 2.0, 3.0])) A jnp.dot(jnp.dot(R, S), R.T) # Rotate, scale, then unrotate evals, evecs jnp.linalg.eig(A) idx jnp.argsort(jnp.abs(evals)) # 从小到大 U evecs[:, idx] D jnp.diag(evals[idx]) assert jnp.allclose(A, jnp.dot(U, jnp.dot(D, U.T))) # 特征分解还原8.3 用特征值判定正定性对称矩阵正定当且仅当所有特征值大于 0可直接编码为判定函数def is_symmetric(A): return (A A.T).all() def isposdef(A): if not is_symmetric(A): return False evals, evecs jnp.linalg.eigh(A) return jnp.all(evals 0) np.random.seed(42) M np.random.randn(3, 4) A jnp.dot(M, M.T) # A M M^T 保证半正定 print(isposdef(A))8.4 幂迭代法Power method最大特征值/特征向量对大矩阵用幂迭代逼近主特征向量是内存友好的经典算法linalg.ipynb 与 power_method_demo.py 都给出了实现def power_method(A, max_iter100, tol1e-5): n jnp.shape(A)[0] u np.random.rand(n) converged False iter 0 while (not converged) and (iter max_iter): old_u u u jnp.dot(A, u) u u / norm(u) # 归一化 lam jnp.dot(u, jnp.dot(A, u)) # Rayleigh 商 converged norm(u - old_u) tol iter 1 return lam, u X np.random.randn(10, 5) A jnp.dot(X.T, X) # 半正定矩阵 lam, u power_method(A) # 与完整特征分解对照验证 evals, evecs np.linalg.eig(A) idx np.argsort(np.abs(evals))[::-1] assert np.allclose(evecs[:, 0], u, 1e-3) assert np.allclose(evals[0], lam, 1e-3)power_method_demo.py中的断言证明幂迭代收敛到的主特征向量/特征值与np.linalg.eig排序后的首列/首值一致容差 1e-3。8.5 特征向量的几何意义图 7.6gaussEvec.py 绘制了图 7.6将一个椭圆旋转 30° 后其主轴方向恰好是特征向量 u1、u2半轴长度是特征值平方根 λ1^(1/2)、λ2^(1/2)。代码用matplotlib.transforms.Affine2D().rotate_deg(30)旋转椭圆与箭头并标注u1、u2、λ1^{1/2}、λ2^{1/2}直观展示 EVD 的几何含义——A U Λ U^T把线性变换分解为旋转 → 各轴缩放 → 旋转回来。这是理解协方差矩阵椭圆等高线后续章节大量使用的基石。九、奇异值分解SVD9.1 full_matrices 与瘦身分解A np.random.randn(10, 5) U, S, V jnp.linalg.svd(A, full_matricesFalse) print(U.shape, S.shape, V.shape) # (10,5) (5,) (5,5) U, S, V jnp.linalg.svd(A, full_matricesTrue) print(U.shape, S.shape, V.shape) # (10,10) (5,) (5,5)full_matricesFalse返回瘦身thin/economy分解是实际计算中最常用的形式。9.2 低秩矩阵与数值秩def make_random_low_rank(D, K): A np.zeros((D, D), dtypejnp.float32) for i in range(K): x np.random.randn(D) A A jnp.outer(x, x) # 秩 1 矩阵之和 return A A make_random_low_rank(10, 3) U, S, V jnp.linalg.svd(A, full_matricesFalse) print(jnp.sum(S 1e-5)) # 非零奇异值个数 秩 print(np.linalg.matrix_rank(A)) # 两种方式一致9.3 图像低秩近似图 7.9、7.10svd_image_demo.ipynb 与 svd_image_demo.py 用经典小丑图clown.png从 probml-data 数据仓库获取演示 SVD 低秩压缩def rgb2gray(rgb): # Y 0.2989 R 0.5870 G 0.1140 B return np.dot(rgb[..., :3], [0.2989, 0.5870, 0.1140]) X np.array(img) # 灰度图 r np.linalg.matrix_rank(X) U, sigma, V np.linalg.svd(X, full_matricesTrue) ranks [1, 2, 5, 10, 20, r] # 图 7.9不同秩的重建 for k in ranks: x_hat np.dot(np.dot(U[:, :k], np.diag(sigma[:k])), V[:k, :]) plt.imshow(x_hat, cmapgray) plt.title(rank {}.format(k))同时脚本把图像像素随机打乱后再做 SVD对比奇异值谱图 7.10原始图像的奇异值迅速衰减前 100 个log σ_i呈陡峭下降的红线而打乱后的图像奇异值衰减缓慢——这正是自然图像具有内在低秩结构的可视化证据也是 PCA/降维方法的理论基础。sigma[:k]的前 k 个奇异值即为最优 k 秩近似的能量占比指标。十、其他矩阵分解LU、QR、Cholesky10.1 LU 分解np.random.seed(42) A np.random.randn(5, 5) L, U scipy.linalg.lu(A, True) # permute_lTrue 吸收置换LU 分解是高斯消元法带部分主元的矩阵形式是solve类函数的底层基础见第十二节。10.2 QR 分解经济模式 vs 完整模式A np.random.randn(5, 3) Q, R scipy.linalg.qr(A, modeeconomic) # Q:(5,3) R:(3,3) Q, R scipy.linalg.qr(A, modefull) # Q:(5,5) R:(5,3) assert jnp.allclose(jnp.eye(5), jnp.dot(Q, Q.T), atol1e-3) # Q 正交10.3 Cholesky 分解正定检验与 MVN 采样Cholesky 只对正定矩阵定义因此它本身就是一个正定检验器同时Σ L L^T也是从多元高斯采样的标准途径def isposdef(A): try: _ np.linalg.cholesky(A) return True except np.linalg.LinAlgError: return False np.random.seed(42) A np.random.randn(5, 5) assert not isposdef(A) # 一般矩阵不正定 assert isposdef(np.dot(A, A.T)) # A A^T 正定 def sample_mvn(mu, Sigma, N): L jnp.linalg.cholesky(Sigma) # Sigma L L^T D len(mu) Z np.random.randn(N, D) X jnp.dot(Z, L.T) jnp.reshape(mu, (-1, D)) return X D 5 np.random.seed(42) mu np.random.randn(D) A np.random.randn(D, D) Sigma jnp.dot(A, A.T) X sample_mvn(mu, Sigma, 10000) C np.cov(X, rowvarFalse)cholesky_demo.py 用np.allclose(C, Sigma, 1e-0)验证了采样协方差与目标协方差一致并与scipy.stats.multivariate_normal.rvs的采样结果对照。该白噪声左乘 Cholesky 因子的技巧在贝叶斯推断、重参数化reparameterization trick中反复出现。十一、矩阵微积分JAX 自动微分linalg.ipynb 的 Matrix calculus 一节展示了用 JAX 对三个典型凸函数求梯度、Jacobian、Hessian并逐一与解析结果断言比对。11.1 线性函数多输入、标量输出解析式f(x; a) a^T x∇_x f aDin, Dout 3, 1 np.random.seed(42) a np.random.randn(Dout, Din) def fun1d(x): return jnp.dot(a, x)[0] x np.random.randn(Din) g grad(fun1d)(x) assert jnp.allclose(g, a) # 梯度 系数向量 J jacrev(fun1d)(x) assert jnp.allclose(J, g) # 标量输出时 Jacobian 梯度11.2 线性函数多输入、多输出解析式f(x; A) A xJ A。Notebook 同时用前向模式jacfwd与反向模式jacrev计算并断言二者一致Din, Dout 3, 4 A np.random.randn(Dout, Din) def fun(x): return jnp.dot(A, x) x np.random.randn(Din) Jf jacfwd(fun)(x) Jr jacrev(fun)(x) assert jnp.allclose(Jf, Jr) assert jnp.allclose(Jf, A)11.3 二次型解析式f(x; A) x^T A x∇f (A A^T)x∇²f A A^TD 4 A np.random.randn(D, D) x np.random.randn(D) quadfun lambda x: jnp.dot(x, jnp.dot(A, x)) J jacfwd(quadfun)(x) assert jnp.allclose(J, jnp.dot(A A.T, x)) H1 hessian(quadfun)(x) assert jnp.allclose(H1, A A.T) def my_hessian(fun): return jacfwd(jacrev(fun)) # Hessian 两层 Jacobian 的组合 H2 my_hessian(quadfun)(x) assert jnp.allclose(H1, H2)这里hessian(f) jacfwd(jacrev(f))的组合关系直接点明了 JAX 自动微分 API 的构成逻辑是阅读 JAX 源码jax/experimental、jax顶层 API时值得记住的结构。十二、求解线性方程组12.1 方阵唯一解A jnp.array([[3, 2, -1], [2, -2, 4], [-1, 0.5, -1]]) b jnp.array([1, -2, 0]) x jax.scipy.linalg.solve(A, b) print(jnp.dot(A, x) - b) # 残差 ≈ 0等价地可以显式走 LU 回代路线这正是solve的底层套路L, U jax.scipy.linalg.lu(A, permute_lTrue) y jax.scipy.linalg.solve_triangular(L, b, lowerTrue) x jax.scipy.linalg.solve_triangular(U, y, lowerFalse)12.2 欠定系统最小范数解当方程数 m 未知数 n 时有无穷多解取范数最小的那个最小范数解np.random.seed(42) m, n 3, 4 A np.random.randn(m, n) x np.random.randn(n) b jnp.dot(A, x) x_least_norm scipy.linalg.lstsq(A, b)[0] print(jnp.dot(A, x_least_norm) - b) # 满足约束 print(jnp.linalg.norm(x_least_norm, 2))Notebook 还深入解释了scipy.linalg.lstsq的底层它只是对 LAPACKFortran 编写的 Python 包装LAPACK 提供gelsd默认基于 SVD、gelssSVD、gelsyQR等多种求解方法——这也解释了为什么大量 NumPy/SciPy 函数本质上是历史遗留 Fortran/C 数值库的薄封装。12.3 超定系统最小二乘解的四种算法对照当 m n 时通常无精确解取最小二乘解。Notebook 一次性给出四种实现并断言等价def naive_solve(A, b): # 正规方程数值上最不稳定 return jax.numpy.linalg.inv(A.T A) A.T b def qr_solve(A, b): # QR 分解法 Q, R jnp.linalg.qr(A) Qb jnp.dot(Q.T, b) return jax.scipy.linalg.solve_triangular(R, Qb) def lstsq_solve(A, b): # LAPACK 最小二乘 return scipy.linalg.lstsq(A, b, rcondNone)[0] def pinv_solve(A, b): # 伪逆法 return jnp.dot(jnp.linalg.pinv(A), b) np.random.seed(42) m, n 4, 3 A np.random.randn(m, n) x np.random.randn(n) b jnp.dot(A, x)四种方法结果一致但数值稳定性与计算成本不同正规方程inv(AᵀA)Aᵀb会放大条件数QR 与 SVDlstsq/pinv数值上更稳健是生产代码的首选。这也是线性回归、岭回归实现本仓库如 linreg_poly_ridge.py 等章节脚本的公共数学内核。十三、白化与协方差椭圆图 7.7 背后的实现图 7.7 由 height_weight_whiten_plot.ipynb / height_weight_whiten_plot.py 生成对比四种数据变换Raw原始、Standardized标准化、PCA-whitenedPCA 白化、ZCA-whitenedZCA 白化数据来自身高/体重数据集heightWeight.mat仅取男性样本y_vec 1。三种变换的核心公式如下数据矩阵X协方差ΣΣ E D Eᵀ# 标准化逐列减均值除标准差 xs_mat (x_mat - np.mean(x_mat, axis0)) / np.std(x_mat) # PCA 白化W_pca D^{-1/2} E^T d_vec, e_mat np.linalg.eigh(sigma) d_mat np.diag(d_vec) w_pca_mat np.dot(np.sqrt(np.linalg.inv(d_mat)), e_mat.T) xw_pca_mat np.dot(w_pca_mat, x_mat.T - mu).T # ZCA 白化W_zca E D^{-1/2} E^T保持原始坐标系方向 w_zca_mat np.dot(e_mat, np.dot(np.sqrt(np.linalg.inv(d_mat)), e_mat.T)) xw_zca_mat np.dot(w_zca_mat, x_mat.T - mu).T图中每个子图叠加了协方差椭圆draw_ell用np.linalg.eigh(cov)计算主轴方向与半轴长度半轴放大 5 倍以近似覆盖 95% 数据直观展示标准化只是消除均值/尺度PCA 白化使数据变成各向同性的球形协方差为单位阵但会旋转坐标系ZCA 白化在去相关的同时尽量保留原始坐标方向。这段代码把本章的 EVD、协方差、线性变换三个主题串成了一条完整的应用链。十四、结语与延伸阅读第 7 章的配套内容以 linalg.ipynb 为技术主线覆盖了从数组基础到线性系统求解的完整线性代数工具箱README.md 中的图号映射表则指明了每个概念的几何可视化位置特征向量几何、白化对比、SVD 图像压缩。本章涉及的许多运算在全书后续章节被反复使用EVD / 特征值正定判定 → 协方差矩阵、PCA如 pca_demo.py、高斯分布性质Cholesky 分解 白噪声采样 → 多元高斯采样、变分推断与重参数化SVD 低秩近似 → 矩阵补全、降维、推荐系统最小二乘求解 → 线性回归、岭回归、高斯过程预测JAX 自动微分 → 梯度下降、随机变分推断SVI等所有基于梯度的推断算法。对源码级实现感兴趣的读者可在 deprecated/scripts/ 下找到上述全部演示脚本的可运行版本gaussEvec.py、height_weight_whiten_plot.py、svd_image_demo.py、cholesky_demo.py、power_method_demo.py、einsum_demo.py它们均包含assert自检可作为深入学习与回归验证的起点。赞分享机器学习深度学习【免费下载链接】pyprobmlPython code for Probabilistic Machine learning book by Kevin Murphy项目地址https://gitcode.com/gh_mirrors/py/pyprobml点击查看免费下载相关推荐Home Assistant 使用 Xiaomi Home 集成开启飞利浦护眼台灯 2 护眼模式xiaomi_miio.light_eyecare_mode_on 动作实战Home Assistant 使用 Xiaomi Home 集成开启飞利浦护眼台灯 2 护眼模式xiaomi_miio.light_eyecare_mode_示例工程机器学习深度学习PyWxDump微信聊天记录导出工具为何突然下线了PyWxDump微信聊天记录导出工具为何突然下线了 PyWxDump 是一款微信 PC 端数据库解密与聊天记录导出工具。但项目已于 2025 年 10 月因pyprobml 使用指南搭建环境并运行 Probabilistic Machine Learning 书籍配套 Notebook 复现全部插图pyprobml 使用指南搭建环境并运行 Probabilistic Machine Learning 书籍配套 Notebook 复现全部插图 pyprob机器学习深度学习上一篇3分钟搭好免费虚拟Amiibo系统emuiibo快速上手教程下一篇pot-desktop 生词本完整上手指南划词翻译快速收词创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考