feat: 完整中文翻译 maths-cs-ai-compendium(数学·计算机科学·AI 知识大全)
翻译自英文原版 maths-cs-ai-compendium,共 20 章全部完成。 第01章 向量 | 第02章 矩阵 | 第03章 微积分 第04章 统计学 | 第05章 概率论 | 第06章 机器学习 第07章 计算语言学 | 第08章 计算机视觉 | 第09章 音频与语音 第10章 多模态学习 | 第11章 自主系统 | 第12章 图神经网络 第13章 计算与操作系统 | 第14章 数据结构与算法 第15章 生产级软件工程 | 第16章 SIMD与GPU编程 第17章 AI推理 | 第18章 ML系统设计 第19章 应用人工智能 | 第20章 前沿人工智能 翻译说明: - 所有数学公式 $...$ / $$...$$、代码块、图片引用完整保留 - mkdocs.yml 配置中文导航 + language: zh - README.md 已翻译为中文(兼 docs/index.md) - docs/ 目录包含指向各章文件的 symlink - 约 29,000 行中文内容,排除 .cache/ 构建缓存
This commit is contained in:
@@ -0,0 +1,381 @@
|
||||
# 经典机器学习
|
||||
|
||||
*经典机器学习算法通过数据学习模式而无需显式编程,使用闭式解或启发式搜索而非梯度下降。本文涵盖朴素贝叶斯、k-NN、决策树、随机森林、支持向量机、k-means聚类和主成分分析*
|
||||
|
||||
- 机器学习是研究算法通过从数据中学习来提升其在某项任务上表现的学科,而非通过显式规则编程。与其编写"如果收入 > 50k 且年龄 < 30 则批准贷款",不如将数千条历史贷款决策交给算法,让它自行找出模式。
|
||||
|
||||
- 存在三大范式。**监督学习**使用带标签数据,即每个输入都有已知的正确输出。算法学习从输入到输出的映射。**无监督学习**处理未标签数据,试图发现隐藏结构,如聚类或压缩表示。**强化学习**通过试错学习,根据在环境中采取的动作接收奖励或惩罚(在第04篇中介绍)。
|
||||
|
||||
- 在监督学习中,**分类**预测离散类别(垃圾邮件或非垃圾邮件,猫或狗),而**回归**预测连续值(房价、明天温度)。边界并不总是清晰:逻辑回归虽然名为"回归",但实际上执行分类任务。
|
||||
|
||||
- 概率模型中的一个关键区分是**生成式 vs 判别式**。生成模型学习联合分布 $P(x, y)$,这意味着它理解数据本身的生成方式。它能产生新样本。判别模型直接学习 $P(y \mid x)$,仅关注类别之间的边界。朴素贝叶斯是生成式的;逻辑回归(第02篇)是判别式的。生成模型更灵活但更难训练好;判别模型在数据充足时通常给出更好的分类准确率。
|
||||
|
||||
- **朴素贝叶斯**是最简单且最有效的分类器之一。它直接应用贝叶斯定理(来自第05章):
|
||||
|
||||
$$P(C_k \mid x) = \frac{P(x \mid C_k) \, P(C_k)}{P(x)}$$
|
||||
|
||||
- "朴素"之处在于一个强烈的独立性假设:它假设给定类别后每个特征相互独立。如果你正在将电子邮件分类为垃圾邮件,朴素贝叶斯假设一旦你知道邮件是垃圾邮件,单词"免费"的出现告诉你关于单词"赢家"是否出现的信息为零。这在现实中几乎从不成立,但分类器仍然出奇地好用。
|
||||
|
||||
- 由于 $P(x)$ 对所有类别都一样,分类简化为选择最大化分子的类别:
|
||||
|
||||
$$\hat{y} = \arg\max_{k} \; P(C_k) \prod_{i=1}^{n} P(x_i \mid C_k)$$
|
||||
|
||||
- 先验 $P(C_k)$ 就是每个类别中训练样本的比例。似然 $P(x_i \mid C_k)$ 取决于特征的类型,从而产生三种常见变体。
|
||||
|
||||
- **多项式朴素贝叶斯**专为计数数据设计,如文档中的词频。每个特征 $x_i$ 表示单词 $i$ 出现的次数,似然遵循多项分布。这是文本分类、情感分析和垃圾邮件过滤的标准选择。
|
||||
|
||||
- **高斯朴素贝叶斯**假设每个特征在每个类别内服从正态分布。你从训练数据中估计特征 $i$ 对类别 $k$ 的均值 $\mu_{ik}$ 和方差 $\sigma_{ik}^2$,然后计算:
|
||||
|
||||
$$P(x_i \mid C_k) = \frac{1}{\sqrt{2\pi\sigma_{ik}^2}} \exp\!\left(-\frac{(x_i - \mu_{ik})^2}{2\sigma_{ik}^2}\right)$$
|
||||
|
||||
- 当特征为连续测量值时,如身高、体重或传感器读数,这是自然的选择。
|
||||
|
||||

|
||||
|
||||
- **伯努利朴素贝叶斯**对二元特征建模:每个特征要么存在(1)要么不存在(0)。你不再统计单词出现的次数,而是只跟踪它是否出现。这适用于短文本或二元特征向量。
|
||||
|
||||
- 一个实际问题是,当某个特征值在训练数据中从未与某个类别一起出现时,似然变为零,由于所有概率相乘,整个后验概率也归零。**拉普拉斯平滑**通过为每个特征-类别组合添加一个小计数(通常为1)来解决这个问题:
|
||||
|
||||
$$P(x_i \mid C_k) = \frac{\text{count}(x_i, C_k) + \alpha}{\text{count}(C_k) + \alpha \cdot V}$$
|
||||
|
||||
- 这里 $\alpha$ 是平滑参数(通常为1),$V$ 是该特征的可能取值数量。这确保了任何概率永远不会精确为零。
|
||||
|
||||
- **决策树**采用了一种完全不同的方法。它不是计算概率,而是通过一系列的"是/否"问题来划分特征空间。想象"二十问"游戏:每一步,你问一个最能缩小可能性范围的问题。
|
||||
|
||||
- 树从根节点开始,包含所有训练样本。在每个内部节点,它选择一个特征和一个阈值进行分裂(例如,"年龄 < 30?")。样本根据答案向左或向右流动。这一过程递归进行直到叶节点,叶节点中存放预测结果:分类任务中的多数类别,或回归任务中的均值。
|
||||
|
||||

|
||||
|
||||
- 关键问题是:应该选择哪个特征进行分裂?你希望分裂产生最"纯"的子节点,即大多数样本属于同一类别。衡量不纯度的两种常用指标是**基尼不纯度**和**熵**。
|
||||
|
||||
- **基尼不纯度**衡量的是如果按照该节点中的分布标记,随机选择的样本被错误分类的概率:
|
||||
|
||||
$$\text{Gini}(S) = 1 - \sum_{k=1}^{K} p_k^2$$
|
||||
|
||||
- 如果节点完全纯(全部属于一个类别),基尼值为0。如果类别完全平衡(比如两类各占50%),基尼值达到最大值0.5。
|
||||
|
||||
- **熵**(来自第05章的信息论部分)衡量平均惊讶程度:
|
||||
|
||||
$$H(S) = -\sum_{k=1}^{K} p_k \log_2 p_k$$
|
||||
|
||||
- 纯节点的熵为0。完全平衡的二元节点的熵为1比特。实际上,基尼和熵产生的树非常相似;基尼计算稍快,因为它避免了对数运算。
|
||||
|
||||
- **信息增益**是由一次分裂带来的不纯度降低。对于将集合 $S$ 划分为子集 $S_L$ 和 $S_R$ 的分裂:
|
||||
|
||||
$$\text{IG}(S, \text{split}) = H(S) - \frac{|S_L|}{|S|} H(S_L) - \frac{|S_R|}{|S|} H(S_R)$$
|
||||
|
||||
- 算法在每一节点贪心地选择信息增益最高的分裂。这是一种局部最优策略,而非全局最优,但在实践中效果很好。
|
||||
|
||||
- **回归树**工作原理相同,但叶子预测连续值(到达该叶子的样本的均值),分裂准则使用方差减少而非基尼或熵。
|
||||
|
||||
- 如果不加约束,决策树会一直分裂直到每个叶子都纯,本质上是在记忆训练数据。这是严重的过拟合。**剪枝**用于应对这一问题。预剪枝在树生长之前设置限制:最大深度、每个叶子的最少样本数、或进行分裂的最小信息增益。后剪枝先生长完整树,然后移除那些不能提升验证集性能的分支。
|
||||
|
||||
- 单个决策树易于解释,但往往不稳定:数据的微小变化可能导致完全不同的树。**集成方法**组合多个模型,以获得比任何单个模型更好的预测结果。
|
||||
|
||||
- 核心思想是"群众智慧"。如果你问100个平庸的分类器然后进行多数投票,只要各个分类器做出一定程度上独立的错误,集成结果可以非常出色。
|
||||
|
||||
- **Bagging**(自助汇聚法)在数据的不同随机子集上训练多个模型,采用有放回抽样(bootstrap样本)。每个模型大约看到原始数据的63%。在预测时,你对输出取平均(回归)或进行多数投票(分类)。由于每个模型看到不同的数据,它们犯不同的错误,平均操作抵消了大部分方差。
|
||||
|
||||
- **随机森林**是将bagging应用于决策树并增加一个额外技巧:在每个分裂处,树只考虑一个随机的特征子集(通常是从 $d$ 个总特征中选 $\sqrt{d}$ 个)。这进一步去除了树之间的相关性,使集成更强大。随机森林是整个机器学习中最可靠的现成分类器之一。
|
||||
|
||||

|
||||
|
||||
- **Boosting**采取了相反的哲学。它不是独立地训练模型,而是顺序地训练,每个新模型专注于之前模型分类错误的样本。
|
||||
|
||||
- **AdaBoost**(自适应提升)为每个训练样本维护一个权重。最初所有权重相等。训练一个弱学习器(通常是深度很浅的决策树,称为"桩")后,被错误分类的样本获得更高的权重,因此下一个学习器更加关注它们。最终预测是所有学习器的加权投票,表现更好的学习器拥有更大的发言权:
|
||||
|
||||
$$H(x) = \text{sign}\!\left(\sum_{t=1}^{T} \alpha_t \, h_t(x)\right)$$
|
||||
|
||||
- 学习器 $t$ 的权重 $\alpha_t$ 取决于其错误率 $\epsilon_t$:
|
||||
|
||||
$$\alpha_t = \frac{1}{2} \ln\!\left(\frac{1 - \epsilon_t}{\epsilon_t}\right)$$
|
||||
|
||||
- 错误率低的学习器获得大的正权重;表现与随机水平持平($\epsilon = 0.5$)的学习器获得零权重。
|
||||
|
||||
- **梯度提升**推广了这一思想。不同于重新加权样本,每个新模型被训练来预测当前集成整体的残差误差(损失函数的负梯度)。对于平方误差损失,残差就是预测值与目标值之间的差值。基于决策树的梯度提升(GBDT)是结构化数据竞赛中许多获胜方案背后的方法(XGBoost、LightGBM、CatBoost是流行的实现)。
|
||||
|
||||
- 关键对比:bagging降低**方差**(通过平均消除噪声),而boosting降低**偏差**(纠正系统性错误)。Bagging在个别模型过拟合时效果最好;boosting在模型欠拟合时效果最好。
|
||||
|
||||
- 转向无监督学习,**K-Means聚类**是最简单且使用最广泛的聚类算法。给定 $n$ 个数据点和目标聚类数 $K$,它通过最小化每个点到其聚类中心的距离总和,将每个点分配给 $K$ 个组之一。
|
||||
|
||||
- 算法交替进行两个步骤。首先,将每个点**分配**到最近的中心点。其次,将每个中心点**更新**为分配给它的所有点的均值。重复直到分配不再变化。这保证收敛,因为每一步总簇内距离都会减小(或保持不变)。
|
||||
|
||||

|
||||
|
||||
- 形式上,K-Means最小化簇内平方和,称为**惯性**:
|
||||
|
||||
$$J = \sum_{k=1}^{K} \sum_{x \in C_k} \|x - \mu_k\|^2$$
|
||||
|
||||
- 其中 $\mu_k$ 是簇 $C_k$ 的中心点。
|
||||
|
||||
- K-Means对初始化敏感。糟糕的起始中心点可能导致较差的局部最小值。**K-Means++** 初始化策略首先随机选择一个中心点,然后每个后续中心点的选择概率与其距离最近现有中心点的平方距离成正比。这分散了初始中心点,几乎总是能给出更好的结果。
|
||||
|
||||
- 如何选择 $K$?两种常用工具。**肘部法**绘制惯性随 $K$ 变化的曲线,寻找"肘部"——增加更多簇不再显著帮助的点。**轮廓系数**衡量一个点与其自身簇的相似度相对于最近其他簇的相似度,范围从-1(错误簇)到+1(良好聚类)。所有点的平均轮廓系数给出了聚类质量的整体衡量。
|
||||
|
||||
- K-Means有局限性:它假设大致相等大小的球形簇,并且它做出"硬"分配(每个点恰好属于一个簇)。**高斯混合模型(GMM)** 放松了这两个限制。
|
||||
|
||||
- GMM将数据建模为 $K$ 个高斯分布的混合,每个分布有自己的均值 $\mu_k$、协方差 $\Sigma_k$ 和混合权重 $\pi_k$(所有权重之和为1):
|
||||
|
||||
$$P(x) = \sum_{k=1}^{K} \pi_k \, \mathcal{N}(x \mid \mu_k, \Sigma_k)$$
|
||||
|
||||
- 不同于硬分配,每个点得到一个**软分配**:它属于每个簇的概率(称为"责任")。位于两个高斯边界附近的点可能是60%属于簇A,40%属于簇B。
|
||||
|
||||
- GMM使用**期望-最大化(EM)算法**进行拟合,该算法交替两个步骤,与K-Means非常类似。**E步**计算责任:对于每个点,它来自每个高斯的概率是多少?**M步**更新参数:给定责任,最佳的均值、协方差和混合权重是什么?EM保证每次迭代增加数据似然,并收敛到局部最大值。
|
||||
|
||||
- K-Means实际上是GMM的EM算法的一个特例:它对应于具有相等协方差的球形高斯和硬(0/1)责任分配。
|
||||
|
||||
- **支持向量机(SVM)** 从几何视角处理分类问题。给定两个线性可分的类别,存在无限多个超平面可以将它们分开。SVM找到**最大间隔**的那个——超平面与每个类别最近数据点之间的最大可能间隙。
|
||||
|
||||
- 最近的点,即恰好位于间隔边缘的点,称为**支持向量**。它们是定义决策边界唯一重要的点;你可以移除所有其他训练点,仍然得到相同的超平面。
|
||||
|
||||

|
||||
|
||||
- 对于线性分类器 $f(x) = w \cdot x + b$,找到最大间隔等价于求解:
|
||||
|
||||
$$\min_{w, b} \; \frac{1}{2}\|w\|^2 \quad \text{subject to} \quad y_i(w \cdot x_i + b) \geq 1 \; \text{for all } i$$
|
||||
|
||||
- 这是一个凸二次规划问题,因此有唯一的全局解(无需担心局部最小值)。
|
||||
|
||||
- 真实数据很少完美可分。**软间隔SVM** 通过引入松弛变量 $\xi_i \geq 0$ 允许一些点违反间隔:
|
||||
|
||||
$$\min_{w, b, \xi} \; \frac{1}{2}\|w\|^2 + C \sum_{i=1}^{n} \xi_i \quad \text{subject to} \quad y_i(w \cdot x_i + b) \geq 1 - \xi_i$$
|
||||
|
||||
- 超参数 $C$ 控制权衡:大的 $C$ 对错误分类施加高惩罚(更紧的拟合,有过拟合风险),小的 $C$ 允许更多违规(更宽的间隔,更强的正则化)。
|
||||
|
||||
- SVM最强大的特性是**核技巧**。许多在原始特征空间中不是线性可分的数据集,在映射到高维空间后变得可分。核技巧让你能够在那个高维空间中计算点积,而无需显式计算变换。
|
||||
|
||||
- 核函数 $K(x_i, x_j) = \phi(x_i) \cdot \phi(x_j)$ 替换SVM优化中的每个点积。最流行的核是**径向基函数(RBF)核**:
|
||||
|
||||
$$K(x_i, x_j) = \exp\!\left(-\gamma \|x_i - x_j\|^2\right)$$
|
||||
|
||||
- RBF核隐式地将数据映射到无限维空间。参数 $\gamma$ 控制单个训练点的影响范围:大的 $\gamma$ 意味着每个点只影响其紧邻区域(过拟合风险),小的 $\gamma$ 给出更平滑的边界。
|
||||
|
||||
- 其他常见核包括多项式核 $K(x_i, x_j) = (x_i \cdot x_j + c)^d$ 和线性核 $K(x_i, x_j) = x_i \cdot x_j$(即没有任何变换的标准SVM)。
|
||||
|
||||
- 实际上,带RBF核的SVM在深度学习出现之前是主导分类器。它们在中小规模数据集上仍然表现良好,特别是当特征数量相对于样本数量较大时。
|
||||
|
||||
- SVM与第02章(矩阵)的联系很深。优化通常以其对偶形式求解,其中解仅依赖于训练样本之间的点积——这正是使核技巧成为可能的原因。整个算法以内积和线性代数的语言运作。
|
||||
|
||||
- 汇总经典ML工具箱:
|
||||
|
||||
| 算法 | 类型 | 关键优势 | 关键劣势 |
|
||||
|---|---|---|---|
|
||||
| 朴素贝叶斯 | 监督(生成式) | 快速,少量数据即可工作 | 独立性假设 |
|
||||
| 决策树 | 监督 | 可解释 | 容易过拟合 |
|
||||
| 随机森林 | 监督(集成) | 稳健,超参数少 | 可解释性较差 |
|
||||
| 梯度提升 | 监督(集成) | 表格数据上的最优水平 | 较慢,调参更多 |
|
||||
| K-Means | 无监督(聚类) | 简单,可扩展 | 假设球形簇 |
|
||||
| GMM | 无监督(聚类) | 软分配,形状灵活 | 对初始化敏感 |
|
||||
| SVM | 监督 | 高维有效 | 大数据集上慢 |
|
||||
|
||||
## 编程任务(在CoLab或笔记本中完成)
|
||||
|
||||
1. 从头实现高斯朴素贝叶斯。在合成二维数据(两个类别)上训练并可视化决策边界。与scikit-learn的实现进行比较。
|
||||
```python
|
||||
import jax.numpy as jnp
|
||||
import matplotlib.pyplot as plt
|
||||
from sklearn.datasets import make_classification
|
||||
|
||||
# 生成合成数据
|
||||
X, y = make_classification(n_samples=300, n_features=2, n_redundant=0,
|
||||
n_informative=2, n_clusters_per_class=1, random_state=42)
|
||||
X, y = jnp.array(X), jnp.array(y)
|
||||
|
||||
# 从头拟合高斯朴素贝叶斯
|
||||
classes = jnp.unique(y)
|
||||
params = {}
|
||||
for c in classes:
|
||||
c = int(c)
|
||||
mask = y == c
|
||||
X_c = X[mask]
|
||||
params[c] = {
|
||||
'mean': jnp.mean(X_c, axis=0),
|
||||
'var': jnp.var(X_c, axis=0),
|
||||
'prior': jnp.sum(mask) / len(y)
|
||||
}
|
||||
|
||||
def gaussian_log_likelihood(x, mean, var):
|
||||
return -0.5 * jnp.sum(jnp.log(2 * jnp.pi * var) + (x - mean)**2 / var)
|
||||
|
||||
def predict(X):
|
||||
preds = []
|
||||
for x in X:
|
||||
log_posts = []
|
||||
for c in [0, 1]:
|
||||
log_post = jnp.log(params[c]['prior']) + gaussian_log_likelihood(
|
||||
x, params[c]['mean'], params[c]['var'])
|
||||
log_posts.append(log_post)
|
||||
preds.append(jnp.argmax(jnp.array(log_posts)))
|
||||
return jnp.array(preds)
|
||||
|
||||
# 决策边界可视化
|
||||
xx, yy = jnp.meshgrid(jnp.linspace(X[:,0].min()-1, X[:,0].max()+1, 200),
|
||||
jnp.linspace(X[:,1].min()-1, X[:,1].max()+1, 200))
|
||||
grid = jnp.column_stack([xx.ravel(), yy.ravel()])
|
||||
zz = predict(grid).reshape(xx.shape)
|
||||
|
||||
plt.figure(figsize=(8, 6))
|
||||
plt.contourf(xx, yy, zz, alpha=0.3, cmap='coolwarm')
|
||||
plt.scatter(X[y==0, 0], X[y==0, 1], c='#3498db', label='Class 0', edgecolors='k', s=20)
|
||||
plt.scatter(X[y==1, 0], X[y==1, 1], c='#e74c3c', label='Class 1', edgecolors='k', s=20)
|
||||
plt.title("Gaussian Naive Bayes Decision Boundary")
|
||||
plt.legend()
|
||||
plt.grid(alpha=0.3)
|
||||
plt.show()
|
||||
|
||||
accuracy = jnp.mean(predict(X) == y)
|
||||
print(f"Training accuracy: {accuracy:.2%}")
|
||||
```
|
||||
|
||||
2. 构建一个使用基尼不纯度进行分裂的决策树。实现单个节点的分裂逻辑,并展示信息增益如何选择最佳特征和阈值。
|
||||
```python
|
||||
import jax.numpy as jnp
|
||||
|
||||
def gini_impurity(y):
|
||||
"""计算标签数组的基尼不纯度。"""
|
||||
classes, counts = jnp.unique(y, return_counts=True)
|
||||
probs = counts / len(y)
|
||||
return 1.0 - jnp.sum(probs ** 2)
|
||||
|
||||
def information_gain(y, left_mask):
|
||||
"""通过布尔掩码将y分裂为左/右后的信息增益。"""
|
||||
parent_gini = gini_impurity(y)
|
||||
left_y, right_y = y[left_mask], y[~left_mask]
|
||||
n = len(y)
|
||||
if len(left_y) == 0 or len(right_y) == 0:
|
||||
return 0.0
|
||||
child_gini = (len(left_y)/n) * gini_impurity(left_y) + \
|
||||
(len(right_y)/n) * gini_impurity(right_y)
|
||||
return float(parent_gini - child_gini)
|
||||
|
||||
def best_split(X, y):
|
||||
"""找到最大化信息增益的特征和阈值。"""
|
||||
best_ig, best_feat, best_thresh = -1, None, None
|
||||
for feat in range(X.shape[1]):
|
||||
thresholds = jnp.unique(X[:, feat])
|
||||
for thresh in thresholds:
|
||||
mask = X[:, feat] <= float(thresh)
|
||||
ig = information_gain(y, mask)
|
||||
if ig > best_ig:
|
||||
best_ig, best_feat, best_thresh = ig, feat, float(thresh)
|
||||
return best_feat, best_thresh, best_ig
|
||||
|
||||
# 示例:合成数据
|
||||
from sklearn.datasets import make_classification
|
||||
X, y = make_classification(n_samples=100, n_features=4, n_redundant=0, random_state=0)
|
||||
X, y = jnp.array(X), jnp.array(y)
|
||||
|
||||
feat, thresh, ig = best_split(X, y)
|
||||
print(f"Best split: feature {feat}, threshold {thresh:.3f}, info gain {ig:.4f}")
|
||||
print(f"Parent Gini: {gini_impurity(y):.4f}")
|
||||
mask = X[:, feat] <= thresh
|
||||
print(f"Left Gini: {gini_impurity(y[mask]):.4f} ({int(jnp.sum(mask))} samples)")
|
||||
print(f"Right Gini: {gini_impurity(y[~mask]):.4f} ({int(jnp.sum(~mask))} samples)")
|
||||
```
|
||||
|
||||
3. 从头实现带K-Means++初始化的K-Means。对合成数据集进行聚类并可视化每次迭代的簇。
|
||||
```python
|
||||
import jax
|
||||
import jax.numpy as jnp
|
||||
import matplotlib.pyplot as plt
|
||||
from sklearn.datasets import make_blobs
|
||||
|
||||
# 生成合成簇
|
||||
X, y_true = make_blobs(n_samples=300, centers=4, cluster_std=0.8, random_state=42)
|
||||
X = jnp.array(X)
|
||||
|
||||
def kmeans_plus_plus_init(X, K, key):
|
||||
"""K-Means++初始化。"""
|
||||
n = X.shape[0]
|
||||
idx = jax.random.randint(key, (), 0, n)
|
||||
centroids = [X[idx]]
|
||||
for _ in range(1, K):
|
||||
dists = jnp.min(jnp.stack([jnp.sum((X - c)**2, axis=1) for c in centroids]), axis=0)
|
||||
probs = dists / jnp.sum(dists)
|
||||
key, subkey = jax.random.split(key)
|
||||
idx = jax.random.choice(subkey, n, p=probs)
|
||||
centroids.append(X[idx])
|
||||
return jnp.stack(centroids)
|
||||
|
||||
def kmeans(X, K, max_iters=20, key=jax.random.PRNGKey(0)):
|
||||
centroids = kmeans_plus_plus_init(X, K, key)
|
||||
history = [centroids]
|
||||
for _ in range(max_iters):
|
||||
# 分配步骤
|
||||
dists = jnp.stack([jnp.sum((X - c)**2, axis=1) for c in centroids])
|
||||
labels = jnp.argmin(dists, axis=0)
|
||||
# 更新步骤
|
||||
new_centroids = jnp.stack([
|
||||
jnp.mean(X[labels == k], axis=0) for k in range(K)
|
||||
])
|
||||
history.append(new_centroids)
|
||||
if jnp.allclose(centroids, new_centroids):
|
||||
break
|
||||
centroids = new_centroids
|
||||
return labels, centroids, history
|
||||
|
||||
K = 4
|
||||
labels, centroids, history = kmeans(X, K)
|
||||
|
||||
# 绘制最终结果
|
||||
colors = ['#3498db', '#e74c3c', '#27ae60', '#9b59b6']
|
||||
plt.figure(figsize=(8, 6))
|
||||
for k in range(K):
|
||||
mask = labels == k
|
||||
plt.scatter(X[mask, 0], X[mask, 1], c=colors[k], s=20, alpha=0.6)
|
||||
plt.scatter(centroids[k, 0], centroids[k, 1], c=colors[k], marker='X',
|
||||
s=200, edgecolors='k', linewidths=1.5)
|
||||
plt.title(f"K-Means Clustering (K={K}, {len(history)-1} iterations)")
|
||||
plt.grid(alpha=0.3)
|
||||
plt.show()
|
||||
|
||||
# 计算惯性
|
||||
inertia = sum(jnp.sum((X[labels == k] - centroids[k])**2) for k in range(K))
|
||||
print(f"Final inertia: {inertia:.2f}")
|
||||
```
|
||||
|
||||
4. 演示核技巧。通过比较核矩阵与多项式核的显式特征映射,展示RBF核如何在高维空间中计算点积。
|
||||
```python
|
||||
import jax.numpy as jnp
|
||||
|
||||
# 简单2D数据
|
||||
X = jnp.array([[1.0, 2.0], [3.0, 4.0], [5.0, 6.0]])
|
||||
|
||||
# 多项式核:K(x,y) = (x·y + 1)^2
|
||||
def poly_kernel(X, degree=2, c=1.0):
|
||||
return (X @ X.T + c) ** degree
|
||||
|
||||
# 2D的显式二次特征映射:(1, sqrt(2)*x1, sqrt(2)*x2, x1^2, x2^2, sqrt(2)*x1*x2)
|
||||
def poly_features(X):
|
||||
x1, x2 = X[:, 0], X[:, 1]
|
||||
return jnp.column_stack([
|
||||
jnp.ones(len(X)),
|
||||
jnp.sqrt(2) * x1,
|
||||
jnp.sqrt(2) * x2,
|
||||
x1 ** 2,
|
||||
x2 ** 2,
|
||||
jnp.sqrt(2) * x1 * x2
|
||||
])
|
||||
|
||||
K_trick = poly_kernel(X)
|
||||
phi = poly_features(X)
|
||||
K_explicit = phi @ phi.T
|
||||
|
||||
print("Kernel trick (polynomial degree 2):")
|
||||
print(K_trick)
|
||||
print("\nExplicit feature map dot products:")
|
||||
print(K_explicit)
|
||||
print(f"\nMatrices match: {jnp.allclose(K_trick, K_explicit)}")
|
||||
|
||||
# RBF核:不存在有限的显式映射
|
||||
def rbf_kernel(X, gamma=0.5):
|
||||
sq_dists = jnp.sum(X**2, axis=1, keepdims=True) + \
|
||||
jnp.sum(X**2, axis=1) - 2 * X @ X.T
|
||||
return jnp.exp(-gamma * sq_dists)
|
||||
|
||||
K_rbf = rbf_kernel(X)
|
||||
print("\nRBF kernel matrix:")
|
||||
print(K_rbf)
|
||||
print("Diagonal is always 1 (a point is identical to itself)")
|
||||
print("Off-diagonal entries decay with distance")
|
||||
```
|
||||
@@ -0,0 +1,408 @@
|
||||
# 梯度机器学习
|
||||
|
||||
*基于梯度的学习通过沿着损失曲面的斜率迭代优化模型参数。本文涵盖线性回归、逻辑回归、softmax分类、梯度下降变体、正则化(L1/L2)和偏差-方差权衡*
|
||||
|
||||
- 第01篇中的经典方法使用巧妙的启发式或闭式解。本文涵盖通过沿着梯度学习、在损失曲面上小步下坡直到找到好参数的算法。基于梯度的学习是从线性回归到最大神经网络的一切背后的引擎。
|
||||
|
||||
- **线性回归**是最简单的基于梯度的模型,它也有闭式解,这使其成为完美的起点。模型是一条直线(或更高维的超平面):
|
||||
|
||||
$$\hat{y} = w \cdot x + b = \sum_{i=1}^{d} w_i x_i + b$$
|
||||
|
||||
- 用矩阵符号(来自第02章),如果我们将所有训练输入堆叠为矩阵 $X$ 的行,并通过追加一列1将偏置吸收到 $w$ 中,这就变成了 $\hat{y} = Xw$。
|
||||
|
||||
- 目标是最小化**均方误差(MSE)**,即预测值与实际值之间平均平方差:
|
||||
|
||||
$$\mathcal{L}(w) = \frac{1}{n} \sum_{i=1}^{n} (y_i - \hat{y}_i)^2 = \frac{1}{n} \|y - Xw\|^2$$
|
||||
|
||||
- 为什么采用平方误差?它有概率论上的依据:如果你假设目标值由 $y = Xw + \epsilon$ 生成,其中 $\epsilon \sim \mathcal{N}(0, \sigma^2)$,那么最大化数据的高斯似然(第05章)等价于最小化MSE。平方误差还比小错误更严厉地惩罚大错误,这通常是可取的。
|
||||
|
||||

|
||||
|
||||
- 由于MSE是 $w$ 的二次函数,它具有唯一的全局最小值,我们可以通过解析方法找到。求导、设为零并求解,得到**正规方程**:
|
||||
|
||||
$$w^{*} = (X^T X)^{-1} X^T y$$
|
||||
|
||||
- 这直接使用了第02章的矩阵逆运算。表达式 $X^T X$ 是一个 $d \times d$ 矩阵(其中 $d$ 是特征数量),$X^T y$ 是一个 $d$ 维向量。正规方程一次性给出精确的最优权重。
|
||||
|
||||
- 正规方程何时失效?当 $X^T X$ 奇异(不可逆)时,这发生在特征线性相关或特征数量多于样本数量($d > n$)的情况下。在这些情况下,你需要正则化(后续介绍)或梯度下降。
|
||||
|
||||
- **逻辑回归**将线性模型适用于二元分类。我们不预测连续值,而是想要一个介于0和1之间的概率。**Sigmoid函数**将所有实数压缩到这个范围内:
|
||||
|
||||
$$\sigma(z) = \frac{1}{1 + e^{-z}}$$
|
||||
|
||||
- 模型计算 $z = w \cdot x + b$(线性得分,与线性回归相同),然后将其通过sigmoid:$\hat{y} = \sigma(w \cdot x + b)$。输出 $\hat{y}$ 被解释为 $P(y = 1 \mid x)$。
|
||||
|
||||

|
||||
|
||||
- Sigmoid具有良好的性质:$\sigma(0) = 0.5$,$\sigma(z) \to 1$ 当 $z \to \infty$,$\sigma(z) \to 0$ 当 $z \to -\infty$,且其导数具有优雅的形式 $\sigma'(z) = \sigma(z)(1 - \sigma(z))$。
|
||||
|
||||
- 逻辑回归的损失函数是**二元交叉熵(BCE)**,直接来自于伯努利似然(第05章):
|
||||
|
||||
$$\mathcal{L} = -\frac{1}{n} \sum_{i=1}^{n} \left[ y_i \log(\hat{y}_i) + (1 - y_i) \log(1 - \hat{y}_i) \right]$$
|
||||
|
||||
- 当真实标签为1时,只有第一项起作用,它惩罚过低的预测。当真实标签为0时,只有第二项起作用,它惩罚过高的预测。对数使得对于自信的错误预测,惩罚极其陡峭:当真实标签为1时预测0.01,代价远高于预测0.4。
|
||||
|
||||
- 与线性回归的MSE不同,BCE最小化权重没有闭式解。我们需要一种迭代方法:**梯度下降**。
|
||||
|
||||
- 梯度下降的直觉很简单:想象你身处大雾中的丘陵地带(损失曲面)。你看不到全局最小值,但可以感受到脚下的坡度。你向下坡走一步,再次感受坡度,然后重复。最终你到达一个山谷。
|
||||
|
||||
$$w \leftarrow w - \eta \frac{\partial \mathcal{L}}{\partial w}$$
|
||||
|
||||
- 学习率 $\eta$ 控制你的步长。太大则越过山谷,来回弹跳而不收敛。太小则缓慢前行,可能陷入局部最小值。
|
||||
|
||||

|
||||
|
||||
- 梯度 $\frac{\partial \mathcal{L}}{\partial w}$ 是一个指向最陡上升方向的向量。我们减去它是因为想向下坡走。这是第03章中的链式法则应用于损失函数。
|
||||
|
||||
- **批量梯度下降**每一步使用整个训练集计算梯度。这给出精确梯度,但当 $n$ 很大时计算代价高昂。
|
||||
|
||||
- **随机梯度下降(SGD)** 每一步使用单个随机样本。梯度带有噪声(它从一个样本估计真实梯度),但每一步非常快。噪声实际上可以帮助逃离浅的局部极小值。
|
||||
|
||||
- **小批量梯度下降**折中:每一步使用 $B$ 个样本的批次(通常为32、64或256)。这平衡了计算效率(对批次的向量化操作)与梯度质量。几乎所有深度学习都使用小批量SGD。
|
||||
|
||||
- **反向传播**是我们实际计算具有许多参数的模型(如神经网络)中梯度的方法。它是第03章的链式法则通过计算图系统化地应用。
|
||||
|
||||
- 任何模型都可以表示为操作的有向无环图:输入流入,乘以权重,加在一起,通过非线性函数传递,最终产生损失值。**前向传播**通过让数据从输入到输出流经此图来计算输出(和损失)。
|
||||
|
||||
- **反向传播**反向流动梯度。从损失开始,你使用每个节点的链式法则计算损失相对于每个中间值的变化。如果 $L$ 依赖于 $z$,而 $z$ 依赖于 $w$,则:
|
||||
|
||||
$$\frac{\partial L}{\partial w} = \frac{\partial L}{\partial z} \cdot \frac{\partial z}{\partial w}$$
|
||||
|
||||
- 每个节点只需要知道自己的局部导数和从上方流入的梯度。这使得反向传播模块化且高效:代价大约是前向传播的两倍(一次前向,一次反向)。
|
||||
|
||||
- 原始SGD有一个问题:它在陡峭曲率方向上振荡,而在平坦方向上进展缓慢。**优化器**通过根据梯度历史调整步长来改进这一点。
|
||||
|
||||
- **带动量的SGD**维护过去梯度的运行平均值(指数移动平均,来自第04章)。这平滑了振荡并加速了沿一致方向的进展:
|
||||
|
||||
$$v_t = \beta v_{t-1} + (1 - \beta) \nabla \mathcal{L}$$
|
||||
$$w \leftarrow w - \eta \, v_t$$
|
||||
|
||||
- 想象一个滚下山的球:动量让它沿一致方向积累速度并抑制侧向抖动。典型值为 $\beta = 0.9$。
|
||||
|
||||
- **内斯特罗夫加速梯度(NAG)** 是一个小巧但巧妙的调整:不在当前位置计算梯度,而是在"前瞻"位置 $w - \eta \beta v_{t-1}$ 计算梯度。这一修正步骤减少了过冲:
|
||||
|
||||
$$v_t = \beta \, v_{t-1} + \nabla \mathcal{L}(w - \eta \beta \, v_{t-1})$$
|
||||
$$w \leftarrow w - \eta \, v_t$$
|
||||
|
||||
- **Adagrad** 为每个参数调整学习率。接收大梯度的参数获得较小的学习率,反之亦然。它累积平方梯度:
|
||||
|
||||
$$G_t = G_{t-1} + g_t^2, \quad w \leftarrow w - \frac{\eta}{\sqrt{G_t + \epsilon}} g_t$$
|
||||
|
||||
- 问题在于:$G_t$ 只增不减,因此有效学习率单调递减,最终变得太小而无法学习任何东西。
|
||||
|
||||
- **RMSprop** 通过使用平方梯度的指数移动平均而非求和来修复此问题,使得近期梯度比早期梯度更重要:
|
||||
|
||||
$$s_t = \beta \, s_{t-1} + (1 - \beta) g_t^2, \quad w \leftarrow w - \frac{\eta}{\sqrt{s_t + \epsilon}} g_t$$
|
||||
|
||||
- **Adam**(自适应矩估计)结合了动量和RMSprop。它同时维护一阶矩估计(梯度的均值,像动量)和二阶矩估计(平方梯度的均值,像RMSprop):
|
||||
|
||||
$$m_t = \beta_1 m_{t-1} + (1 - \beta_1) g_t$$
|
||||
$$v_t = \beta_2 v_{t-1} + (1 - \beta_2) g_t^2$$
|
||||
|
||||
- 由于 $m_t$ 和 $v_t$ 初始化为零,它们在早期步骤中有偏近于零。偏差修正解决了这个问题:
|
||||
|
||||
$$\hat{m}_t = \frac{m_t}{1 - \beta_1^t}, \quad \hat{v}_t = \frac{v_t}{1 - \beta_2^t}$$
|
||||
$$w \leftarrow w - \frac{\eta}{\sqrt{\hat{v}_t} + \epsilon} \hat{m}_t$$
|
||||
|
||||

|
||||
|
||||
- 默认超参数($\beta_1 = 0.9$, $\beta_2 = 0.999$, $\epsilon = 10^{-8}$)在广泛的问题上表现良好,这就是为什么Adam是大多数深度学习工作中的默认优化器。
|
||||
|
||||
- **AdamW** 将权重衰减与梯度更新解耦。标准L2正则化和权重衰减对于SGD是等价的,但对于Adam则不然。AdamW直接将权重衰减应用于参数,而不是将 $\lambda w$ 加到梯度上。这带来了更好的泛化性能,现在是Transformer训练的标准:
|
||||
|
||||
$$w \leftarrow w - \eta \left( \frac{\hat{m}_t}{\sqrt{\hat{v}_t} + \epsilon} + \lambda \, w \right)$$
|
||||
|
||||
- **LION**(演化符号动量)是通过程序搜索发现的新优化器。它只使用动量更新的符号(而不是幅度),使得每次更新的尺度均匀。LION比Adam使用更少的内存(没有二阶矩缓冲区),并且在许多任务上可以匹配或超越Adam:
|
||||
|
||||
$$w \leftarrow w - \eta \cdot \text{sign}(\beta_1 \, m_{t-1} + (1 - \beta_1) \, g_t)$$
|
||||
$$m_t = \beta_2 \, m_{t-1} + (1 - \beta_2) \, g_t$$
|
||||
|
||||
- **Muon**(动量 + 正交化)应用内斯特罗夫动量,然后使用Newton-Schulz迭代对更新矩阵进行正交化,该迭代近似极分解。得到的更新方向位于Stiefel流形上,每次更新在所有奇异方向上具有大致相等的幅度,防止任何单一方向主导。这消除了对自适应二阶矩估计(如Adam的 $v_t$ 缓冲区)的需求,减少了内存使用。Muon在Transformer训练中表现出色,通常以更快的收敛速度达到与AdamW相当的质量,尤其适用于注意力矩阵和MLP权重矩阵。嵌入层和输出层通常仍由AdamW处理。
|
||||
|
||||
$$G_t = \text{NesterovMomentum}(\nabla \mathcal{L})$$
|
||||
$$U_t = \text{NewtonSchulz}(G_t) \approx G_t (G_t^T G_t)^{-1/2}$$
|
||||
$$W \leftarrow W - \eta \, U_t$$
|
||||
|
||||
- Newton-Schulz迭代通过重复 $X_{k+1} = \frac{1}{2} X_k (3I - X_k^T X_k)$ 几个步骤(通常5-10步)来计算正交因子。这避免了完整SVD的计算代价,同时提供了良好的近似。
|
||||
|
||||

|
||||
|
||||

|
||||
|
||||
- 除了MSE和BCE之外,还有几种常用的**损失函数**。
|
||||
|
||||
- **平均绝对误差(MAE)**,或L1损失,取绝对差的平均值:$\frac{1}{n}\sum|y_i - \hat{y}_i|$。它对异常值比MSE更鲁棒,因为它不对大误差进行平方。
|
||||
|
||||
- **Huber损失**结合了两者的优点:对于小误差表现像MSE(平滑,易于优化),对于大误差表现像MAE(对异常值鲁棒)。它有一个控制过渡的阈值 $\delta$。
|
||||
|
||||
- **分类交叉熵(CCE)** 将BCE推广到多个类别。如果 $\hat{y}_k$ 是类别 $k$ 的预测概率,真实类别为 $c$:
|
||||
|
||||
$$\mathcal{L} = -\log(\hat{y}_c)$$
|
||||
|
||||
- 这只是正确类别的负对数概率。最小化交叉熵等价于最大化似然,这联系到第05章的信息论:交叉熵衡量当你使用预测分布代替真实分布时需要多少额外比特。
|
||||
|
||||
- **Hinge损失** 被SVM使用:$\mathcal{L} = \max(0, 1 - y \cdot f(x))$。它只惩罚在间隔错误一侧或间隔内的预测。一旦一个点被足够置信地正确分类,损失为零。
|
||||
|
||||
- **正则化**通过添加对复杂模型的惩罚来防止过拟合。正则化后的损失为:
|
||||
|
||||
$$\mathcal{L}_{\text{reg}} = \mathcal{L}_{\text{data}} + \lambda \, R(w)$$
|
||||
|
||||
- **L2正则化**(Ridge,权重衰减)惩罚平方权重之和:$R(w) = \|w\|^2 = \sum w_i^2$。它阻止任何单个权重变得过大,有效地将所有权重向零收缩,但很少使它们精确为零。
|
||||
|
||||
- **L1正则化**(Lasso)惩罚绝对权重之和:$R(w) = \|w\|_1 = \sum |w_i|$。它鼓励稀疏性,将许多权重驱动到精确为零,实现自动特征选择。
|
||||
|
||||
- **弹性网络** 结合了两者:$R(w) = \alpha \|w\|_1 + (1 - \alpha) \|w\|^2$,融合了稀疏性和收缩。
|
||||
|
||||
- 有一个优美的贝叶斯解释(来自第05章)。L2正则化等价于在权重上放置高斯先验并寻找MAP估计。L1正则化对应于拉普拉斯先验。正则化强度 $\lambda$ 控制你相对于数据信任先验的程度。
|
||||
|
||||
- **评估指标**告诉你模型是否真正有效。对于回归,MSE和MAE是标准指标。对于分类,情况更为微妙。
|
||||
|
||||
- **混淆矩阵**是一个二元分类的四格表:
|
||||
- 真正例(TP):预测为正,实际为正
|
||||
- 假正例(FP):预测为正,实际为负
|
||||
- 真负例(TN):预测为负,实际为负
|
||||
- 假负例(FN):预测为负,实际为正
|
||||
|
||||
- **准确率** = $\frac{TP + TN}{TP + TN + FP + FN}$ 在类别不平衡时可能具有误导性。如果99%的电子邮件不是垃圾邮件,一个总是预测"非垃圾邮件"的模型有99%的准确率,但没有用处。
|
||||
|
||||
- **精确率** = $\frac{TP}{TP + FP}$ 回答:在所有预测为正的样本中,有多少实际为正?高精确率意味着误报少。
|
||||
|
||||
- **召回率**(敏感度)= $\frac{TP}{TP + FN}$ 回答:在所有实际为正的样本中,你捕获了多少?高召回率意味着漏检少。
|
||||
|
||||
- **F1分数** = $\frac{2 \cdot \text{precision} \cdot \text{recall}}{\text{precision} + \text{recall}}$ 是精确率和召回率的调和平均数,平衡了两者。
|
||||
|
||||
- **ROC曲线**绘制了真正率(召回率)对假正率($\frac{FP}{FP + TN}$)随分类阈值从0到1变化的曲线。完美分类器紧贴左上角。**AUC**(ROC曲线下面积)用一个数字概括性能:1.0为完美,0.5为随机猜测。
|
||||
|
||||
- **交叉验证**提供了更可靠的泛化性能估计。在 $k$ 折交叉验证中,你将数据分成 $k$ 份,在 $k-1$ 份上训练,在剩余一份上测试,然后轮换。所有 $k$ 折的平均测试性能就是你的估计。这使用了所有数据进行训练和测试(只是不在同一时间),在数据稀缺时尤为宝贵。
|
||||
|
||||
- **偏差-方差权衡**(来自第04章)是ML中的基本张力。模型期望误差分解为:
|
||||
|
||||
$$\text{Error} = \text{Bias}^2 + \text{Variance} + \text{Irreducible Noise}$$
|
||||
|
||||
- **偏差**是错误假设带来的系统性误差(例如,用直线拟合曲线数据)。**方差**是对训练数据波动的敏感度(例如,20次多项式拟合噪声)。简单模型具有高偏差和低方差;复杂模型具有低偏差和高方差。最优在两者之间。
|
||||
|
||||
- **学习率调度**在训练期间调整 $\eta$。常见策略:
|
||||
- 步长衰减:每 $N$ 个epoch将 $\eta$ 乘以一个因子(如0.1)
|
||||
- 余弦退火:按照余弦曲线从初始值平滑降低 $\eta$ 到接近零
|
||||
- 预热:从一个非常小的 $\eta$ 开始,在前几千步线性增加,然后衰减。这防止了大的初始梯度破坏训练稳定性
|
||||
- 1cycle:一个先升后降的余弦周期,可以带来更快的收敛
|
||||
|
||||
- **超参数调优**是找到学习率、批量大小、正则化强度和其他不由梯度下降学习的设置的良好值的过程。常用方法:
|
||||
- 网格搜索:在预定义的网格上尝试每一种组合(穷举但代价高)
|
||||
- 随机搜索:随机采样组合,通常更高效,因为并非所有超参数同等重要
|
||||
- 贝叶斯优化:构建目标函数的模型并智能选择下一个要尝试的超参数
|
||||
- **ASHA**(异步连续减半算法):使用小预算并行运行许多试验,然后将最有希望的提升到更大预算,同时及早终止其余试验。它结合了早停的高效性和大规模并行性——不是运行100次完整的训练,而是廉价地启动所有100次,在每级保留前四分之一,只有少数运行到完成。这是现代大规模调优框架(如Ray Tune)的骨干。
|
||||
|
||||
- **无调度学习**完全消除了对学习率调度的需求。它不是在固定曲线上衰减 $\eta$,而是维护两个序列:一个缓慢移动的迭代平均值 $z_t$(收敛到最优值)和一个快速探索的迭代 $y_t$(在其上评估梯度)。最终输出是平均序列,被证明在事后能匹配最佳调度的收敛速度。这完全消除了调度作为一个超参数——你只需设置基础学习率,优化器处理其余部分。SGD和Adam的无调度变体已被证明能达到或超越其经过调度的对应版本。
|
||||
|
||||
## 编程任务(在CoLab或笔记本中完成)
|
||||
|
||||
1. 使用正规方程和梯度下降两种方法实现线性回归。比较求解结果,并绘制GD损失随迭代的收敛曲线。
|
||||
```python
|
||||
import jax
|
||||
import jax.numpy as jnp
|
||||
import matplotlib.pyplot as plt
|
||||
|
||||
# 生成合成数据:y = 3x + 2 + noise
|
||||
key = jax.random.PRNGKey(42)
|
||||
n = 100
|
||||
X = jax.random.uniform(key, (n, 1), minval=0, maxval=10)
|
||||
y = 3 * X[:, 0] + 2 + jax.random.normal(key, (n,)) * 1.5
|
||||
|
||||
# 添加偏置列
|
||||
X_b = jnp.column_stack([X, jnp.ones(n)])
|
||||
|
||||
# 正规方程
|
||||
w_exact = jnp.linalg.solve(X_b.T @ X_b, X_b.T @ y)
|
||||
print(f"Normal equation: w={w_exact[0]:.4f}, b={w_exact[1]:.4f}")
|
||||
|
||||
# 梯度下降
|
||||
w_gd = jnp.zeros(2)
|
||||
lr = 0.005
|
||||
losses = []
|
||||
for step in range(500):
|
||||
pred = X_b @ w_gd
|
||||
error = pred - y
|
||||
loss = jnp.mean(error ** 2)
|
||||
losses.append(float(loss))
|
||||
grad = (2 / n) * X_b.T @ error
|
||||
w_gd = w_gd - lr * grad
|
||||
|
||||
print(f"Gradient descent: w={w_gd[0]:.4f}, b={w_gd[1]:.4f}")
|
||||
|
||||
fig, axes = plt.subplots(1, 2, figsize=(12, 4))
|
||||
axes[0].scatter(X[:, 0], y, s=15, alpha=0.5, color='#3498db')
|
||||
axes[0].plot([0, 10], [w_exact[1], w_exact[0]*10 + w_exact[1]], color='#e74c3c', linewidth=2)
|
||||
axes[0].set_title("Linear Regression Fit")
|
||||
axes[0].set_xlabel("x"); axes[0].set_ylabel("y")
|
||||
|
||||
axes[1].plot(losses, color='#27ae60', linewidth=1.5)
|
||||
axes[1].set_title("GD Loss Convergence")
|
||||
axes[1].set_xlabel("Step"); axes[1].set_ylabel("MSE")
|
||||
axes[1].set_yscale('log')
|
||||
plt.tight_layout()
|
||||
plt.show()
|
||||
```
|
||||
|
||||
2. 从头实现带梯度下降的逻辑回归。在二维数据集上训练并可视化学习到的决策边界。
|
||||
```python
|
||||
import jax
|
||||
import jax.numpy as jnp
|
||||
import matplotlib.pyplot as plt
|
||||
from sklearn.datasets import make_moons
|
||||
|
||||
# 生成数据
|
||||
X, y = make_moons(n_samples=300, noise=0.2, random_state=42)
|
||||
X, y = jnp.array(X), jnp.array(y, dtype=jnp.float32)
|
||||
|
||||
def sigmoid(z):
|
||||
return 1 / (1 + jnp.exp(-z))
|
||||
|
||||
# 添加偏置列
|
||||
X_b = jnp.column_stack([X, jnp.ones(len(X))])
|
||||
w = jnp.zeros(3)
|
||||
lr = 0.5
|
||||
losses = []
|
||||
|
||||
for step in range(2000):
|
||||
z = X_b @ w
|
||||
pred = sigmoid(z)
|
||||
# BCE损失
|
||||
loss = -jnp.mean(y * jnp.log(pred + 1e-8) + (1 - y) * jnp.log(1 - pred + 1e-8))
|
||||
losses.append(float(loss))
|
||||
# 梯度
|
||||
grad = X_b.T @ (pred - y) / len(y)
|
||||
w = w - lr * grad
|
||||
|
||||
# 决策边界
|
||||
xx, yy = jnp.meshgrid(jnp.linspace(-2, 3, 200), jnp.linspace(-1.5, 2, 200))
|
||||
grid = jnp.column_stack([xx.ravel(), yy.ravel(), jnp.ones(xx.size)])
|
||||
zz = sigmoid(grid @ w).reshape(xx.shape)
|
||||
|
||||
plt.figure(figsize=(8, 6))
|
||||
plt.contourf(xx, yy, zz, levels=[0, 0.5, 1], alpha=0.3, colors=['#e74c3c', '#3498db'])
|
||||
plt.contour(xx, yy, zz, levels=[0.5], colors='#9b59b6', linewidths=2)
|
||||
plt.scatter(X[y==0, 0], X[y==0, 1], c='#e74c3c', s=15, label='Class 0')
|
||||
plt.scatter(X[y==1, 0], X[y==1, 1], c='#3498db', s=15, label='Class 1')
|
||||
plt.title("Logistic Regression Decision Boundary")
|
||||
plt.legend()
|
||||
plt.grid(alpha=0.3)
|
||||
plt.show()
|
||||
```
|
||||
|
||||
3. 在二维二次曲面上比较优化器的轨迹。从相同的起点运行SGD、SGD+Momentum和Adam,绘制它们的路径。
|
||||
```python
|
||||
import jax
|
||||
import jax.numpy as jnp
|
||||
import matplotlib.pyplot as plt
|
||||
|
||||
# 拉长的二次曲面:L(w1, w2) = 0.5*w1^2 + 10*w2^2
|
||||
def loss_fn(w):
|
||||
return 0.5 * w[0]**2 + 10 * w[1]**2
|
||||
|
||||
grad_fn = jax.grad(loss_fn)
|
||||
|
||||
def run_sgd(w0, lr=0.05, steps=80):
|
||||
w = w0.copy()
|
||||
path = [w.copy()]
|
||||
for _ in range(steps):
|
||||
g = grad_fn(w)
|
||||
w = w - lr * g
|
||||
path.append(w.copy())
|
||||
return jnp.stack(path)
|
||||
|
||||
def run_momentum(w0, lr=0.05, beta=0.9, steps=80):
|
||||
w, v = w0.copy(), jnp.zeros(2)
|
||||
path = [w.copy()]
|
||||
for _ in range(steps):
|
||||
g = grad_fn(w)
|
||||
v = beta * v + (1 - beta) * g
|
||||
w = w - lr * v
|
||||
path.append(w.copy())
|
||||
return jnp.stack(path)
|
||||
|
||||
def run_adam(w0, lr=0.05, b1=0.9, b2=0.999, eps=1e-8, steps=80):
|
||||
w, m, v = w0.copy(), jnp.zeros(2), jnp.zeros(2)
|
||||
path = [w.copy()]
|
||||
for t in range(1, steps + 1):
|
||||
g = grad_fn(w)
|
||||
m = b1 * m + (1 - b1) * g
|
||||
v = b2 * v + (1 - b2) * g**2
|
||||
m_hat = m / (1 - b1**t)
|
||||
v_hat = v / (1 - b2**t)
|
||||
w = w - lr * m_hat / (jnp.sqrt(v_hat) + eps)
|
||||
path.append(w.copy())
|
||||
return jnp.stack(path)
|
||||
|
||||
w0 = jnp.array([8.0, 3.0])
|
||||
sgd_path = run_sgd(w0)
|
||||
mom_path = run_momentum(w0)
|
||||
adam_path = run_adam(w0)
|
||||
|
||||
# 绘图
|
||||
fig, ax = plt.subplots(figsize=(8, 6))
|
||||
w1 = jnp.linspace(-10, 10, 100)
|
||||
w2 = jnp.linspace(-4, 4, 100)
|
||||
W1, W2 = jnp.meshgrid(w1, w2)
|
||||
L = 0.5 * W1**2 + 10 * W2**2
|
||||
ax.contour(W1, W2, L, levels=20, cmap='Greys', alpha=0.4)
|
||||
ax.plot(sgd_path[:,0], sgd_path[:,1], 'o-', color='#3498db', markersize=2, linewidth=1, label='SGD')
|
||||
ax.plot(mom_path[:,0], mom_path[:,1], 'o-', color='#27ae60', markersize=2, linewidth=1, label='Momentum')
|
||||
ax.plot(adam_path[:,0], adam_path[:,1], 'o-', color='#e74c3c', markersize=2, linewidth=1, label='Adam')
|
||||
ax.plot(0, 0, 'k*', markersize=15, label='Minimum')
|
||||
ax.set_xlabel('w₁'); ax.set_ylabel('w₂')
|
||||
ax.set_title("Optimizer Trajectories on Elongated Quadratic")
|
||||
ax.legend()
|
||||
plt.grid(alpha=0.3)
|
||||
plt.show()
|
||||
```
|
||||
|
||||
4. 展示L1与L2正则化对权重稀疏性的影响。使用两种惩罚训练线性回归,并比较得到的权重向量。
|
||||
```python
|
||||
import jax
|
||||
import jax.numpy as jnp
|
||||
import matplotlib.pyplot as plt
|
||||
|
||||
# 合成数据:20个特征中只有前3个是相关的
|
||||
key = jax.random.PRNGKey(0)
|
||||
n, d = 200, 20
|
||||
w_true = jnp.zeros(d).at[:3].set(jnp.array([3.0, -2.0, 1.5]))
|
||||
X = jax.random.normal(key, (n, d))
|
||||
y = X @ w_true + 0.5 * jax.random.normal(key, (n,))
|
||||
|
||||
def train_ridge(X, y, lam=1.0, lr=0.01, steps=2000):
|
||||
"""通过GD进行L2正则化线性回归。"""
|
||||
w = jnp.zeros(X.shape[1])
|
||||
for _ in range(steps):
|
||||
pred = X @ w
|
||||
grad = (2/len(y)) * X.T @ (pred - y) + 2 * lam * w
|
||||
w = w - lr * grad
|
||||
return w
|
||||
|
||||
def train_lasso(X, y, lam=1.0, lr=0.01, steps=2000):
|
||||
"""通过近端GD进行L1正则化线性回归。"""
|
||||
w = jnp.zeros(X.shape[1])
|
||||
for _ in range(steps):
|
||||
pred = X @ w
|
||||
grad = (2/len(y)) * X.T @ (pred - y)
|
||||
w = w - lr * grad
|
||||
# 软阈值(L1的近端算子)
|
||||
w = jnp.sign(w) * jnp.maximum(jnp.abs(w) - lr * lam, 0)
|
||||
return w
|
||||
|
||||
w_l2 = train_ridge(X, y, lam=0.1)
|
||||
w_l1 = train_lasso(X, y, lam=0.1)
|
||||
|
||||
fig, axes = plt.subplots(1, 3, figsize=(14, 4))
|
||||
axes[0].bar(range(d), w_true, color='#333', alpha=0.7)
|
||||
axes[0].set_title("True Weights"); axes[0].set_xlabel("Feature")
|
||||
axes[1].bar(range(d), w_l2, color='#3498db', alpha=0.7)
|
||||
axes[1].set_title("L2 (Ridge): shrinks all"); axes[1].set_xlabel("Feature")
|
||||
axes[2].bar(range(d), w_l1, color='#e74c3c', alpha=0.7)
|
||||
axes[2].set_title("L1 (Lasso): zeros out irrelevant"); axes[2].set_xlabel("Feature")
|
||||
plt.tight_layout()
|
||||
plt.show()
|
||||
|
||||
print(f"L2 non-zero weights: {int(jnp.sum(jnp.abs(w_l2) > 0.01))}/{d}")
|
||||
print(f"L1 non-zero weights: {int(jnp.sum(jnp.abs(w_l1) > 0.01))}/{d}")
|
||||
```
|
||||
@@ -0,0 +1,354 @@
|
||||
# 深度学习
|
||||
|
||||
*深度学习堆叠非线性层来构建层次化表示,自动将原始输入转换为有用的特征。本文涵盖MLP、激活函数、反向传播、CNN、RNN、LSTM、注意力机制、Transformer、GAN、VAE、扩散模型和归一化技术*
|
||||
|
||||
- 什么使网络"深"?浅网络只有一个隐藏层;深网络有许多层。深度让网络构建层次化表示,早期层学习简单特征(边缘、音调),后期层将它们组合成复杂概念(人脸、句子)。这种组合性正是深度学习力量的来源。
|
||||
|
||||
- 最简单的深度网络是**多层感知器(MLP)**,也称为全连接或密集网络。每层计算:
|
||||
|
||||
$$h = \sigma(Wx + b)$$
|
||||
|
||||
- 这里 $W$ 是权重矩阵(第02章),$b$ 是偏置向量,$\sigma$ 是非线性激活函数。一层的输出成为下一层的输入。没有非线性,堆叠层将毫无意义:$W_2(W_1 x) = (W_2 W_1)x$,这只是另一个线性变换。这正是第02章中的矩阵乘法塌缩。
|
||||
|
||||
- **激活函数**引入使深度有意义的非线性。
|
||||
|
||||
- **ReLU**(修正线性单元):$\text{ReLU}(x) = \max(0, x)$。它是使用最广泛的激活函数。计算速度快,正输入不饱和,并产生稀疏激活(许多神经元输出精确为零)。缺点:负输入的神经元总是输出零,如果它们永久卡在那里,就会"死亡"并停止学习。
|
||||
|
||||
- **Sigmoid**:$\sigma(x) = \frac{1}{1+e^{-x}}$,将输入压缩到 $(0, 1)$。适用于二元分类的输出层,但在隐藏层中有问题,因为当输入远离零时梯度消失(曲线几乎平坦)。
|
||||
|
||||
- **Tanh**:$\tanh(x) = \frac{e^x - e^{-x}}{e^x + e^{-x}}$,压缩到 $(-1, 1)$。零中心(不同于sigmoid),有助于梯度流动,但在极端值处仍存在梯度消失问题。
|
||||
|
||||
- **GELU**(高斯误差线性单元):$\text{GELU}(x) = x \cdot \Phi(x)$,其中 $\Phi$ 是标准正态CDF。它是ReLU的平滑近似,允许微小的负值通过。GELU是GPT和BERT中的默认选择。
|
||||
|
||||
- **Swish**:$\text{Swish}(x) = x \cdot \sigma(x)$,另一种平滑门控。实际使用中与GELU类似。
|
||||
|
||||

|
||||
|
||||
- 一个具有 $d_{\text{in}}$ 个输入和 $d_{\text{out}}$ 个输出的密集层有 $d_{\text{in}} \times d_{\text{out}} + d_{\text{out}}$ 个参数(权重加偏置)。矩阵乘法 $Wx$ 就是第02章中的矩阵-向量乘法。在批处理设置中,输入是形状为 $(B, d_{\text{in}})$ 的矩阵 $X$,输出是形状为 $(B, d_{\text{out}})$ 的 $XW^T + b$。
|
||||
|
||||
- **万能近似定理**指出,一个具有足够神经元的隐藏层可以在紧致域上以任意精度逼近任何连续函数。这听起来似乎深度无关紧要,但关键在于"足够的神经元"。实际上,深层网络可以用指数级少于浅层网络的参数来表示相同的函数。深度带来的是效率,而不仅仅是表达能力。
|
||||
|
||||
- 随着网络变深,出现两种梯度病理。**梯度消失**:当梯度通过许多层时(通过链式法则,第03章),它们被乘以许多因子。如果这些因子都小于1(如sigmoid和tanh饱和时发生的情况),梯度呈指数级缩小趋近于零。早期层几乎无法学习。**梯度爆炸**:如果因子都大于1,梯度呈指数级增长,导致数值溢出和训练不稳定。
|
||||
|
||||
- 梯度消失/爆炸的解决方案:
|
||||
- 使用ReLU或GELU激活函数(正输入时梯度为1,无饱和)
|
||||
- 仔细的权重初始化
|
||||
- 归一化层
|
||||
- 残差连接(跳跃连接)
|
||||
- 梯度裁剪(针对梯度爆炸):将梯度范数限制在最大值
|
||||
|
||||
- **权重初始化**很重要,因为它决定了训练开始时激活值和梯度的尺度。如果权重太大,激活值爆炸;太小,它们消失。
|
||||
|
||||
- **Xavier (Glorot) 初始化**从方差为 $\frac{2}{d_{\text{in}} + d_{\text{out}}}$ 的分布中设置权重。这假设使用线性或tanh激活函数时,能使激活值的方差在各层大致保持恒定。
|
||||
|
||||
- **He (Kaiming) 初始化**使用方差 $\frac{2}{d_{\text{in}}}$,针对ReLU激活函数校准(由于ReLU将半数激活值置零,需要双倍方差来补偿)。
|
||||
|
||||
- **归一化层**通过确保每层的输入具有一致的统计特性(大致零均值、单位方差)来稳定训练。
|
||||
|
||||
- **批归一化(BatchNorm)** 在批次维度上进行归一化:对于每个通道/特征,计算小批次中所有样本的均值和方差,然后归一化。它添加了可学习的尺度($\gamma$)和偏移($\beta$)参数,以便网络在需要时撤销归一化:
|
||||
|
||||
$$\hat{x} = \frac{x - \mu_B}{\sqrt{\sigma_B^2 + \epsilon}}, \quad y = \gamma \hat{x} + \beta$$
|
||||
|
||||
- BatchNorm有一个问题:它依赖于批量大小。当批次非常小时,统计数据有噪声。在推理时,使用运行平均值而非批次统计,这造成了训练/测试不一致。
|
||||
|
||||
- **层归一化(LayerNorm)** 对每个单独样本在特征维度上进行归一化。它不依赖于批次中的其他样本,使其成为Transformer和循环网络的标准选择。
|
||||
|
||||
- **实例归一化** 对每个样本和每个通道独立地在空间维度上进行归一化。在风格迁移中很流行。
|
||||
|
||||
- **组归一化** 将通道分成组并在每个组内进行归一化。它是LayerNorm和InstanceNorm之间的折中。
|
||||
|
||||

|
||||
|
||||
- **Dropout** 是一种正则化技术,在训练期间随机将一部分 $p$ 的神经元置零。这迫使网络不依赖任何单个神经元,鼓励冗余表示。测试时,所有神经元都被激活。**逆置Dropout** 在训练期间将激活值缩放 $\frac{1}{1-p}$,以便测试时无需缩放。这是标准实现。
|
||||
|
||||
- **卷积神经网络(CNN)** 利用了空间结构。卷积层不是将每个输入连接到每个输出(如密集层),而是在输入上滑动一个小滤波器(核),在每个位置计算点积。相同的滤波器权重在所有位置共享,这大大减少了参数并内建了平移不变性。
|
||||
|
||||
- 二维输入与大小为 $k \times k$ 的滤波器 $K$ 的**卷积操作**:
|
||||
|
||||
$$(\text{input} * K)[i,j] = \sum_{m=0}^{k-1} \sum_{n=0}^{k-1} \text{input}[i+m, j+n] \cdot K[m, n]$$
|
||||
|
||||

|
||||
|
||||
- 输出大小取决于三个超参数。**步幅**控制滤波器在位置之间移动多少像素(步幅2使空间维度减半)。**填充**在输入边界周围添加零("same"填充保持空间大小,"valid"填充不填充)。输出大小公式:$\text{out} = \lfloor (\text{in} - k + 2p) / s \rfloor + 1$。
|
||||
|
||||
- **池化**层对特征图进行下采样。最大池化取每个窗口中的最大值;平均池化取均值。池化在保留最重要信息的同时减少空间维度。
|
||||
|
||||
- **扩张卷积** 在滤波器元素之间插入间隙,增加感受野而不增加参数。扩张率为2意味着3x3滤波器覆盖5x5区域。
|
||||
|
||||
- **1x1卷积** 是使用1x1滤波器的卷积。它们不查看空间邻居;而是跨通道混合信息。可以将其视为在每个空间位置应用密集层。用于廉价地改变通道数。
|
||||
|
||||
- **跳跃连接**(残差连接)让输入绕过一层或多层:$\text{output} = F(x) + x$。该层只需学习残差 $F(x) = \text{output} - x$,当最优变换接近恒等映射时这更容易。ResNet(残差网络)使用这一技巧堆叠超过100层,解决了更深的网络表现比浅层网络更差的退化问题。
|
||||
|
||||
- CNN构建了一个**特征层次结构**。早期层检测边缘和纹理。中间层将这些组合成部件(眼睛、轮子)。后期层识别整个物体。每层的感受野(它"看到"的输入区域)随深度增加。
|
||||
|
||||
- **嵌入**将离散的标记(单词、字符、物品ID)映射到密集向量。嵌入层只是一个查找表:一个形状为(词汇表大小,嵌入维度)的矩阵 $E$。查找标记 $i$ 意味着选择 $E$ 的第 $i$ 行。这等价于乘以one-hot向量,这只是矩阵-向量乘法的一个特例(第02章)。嵌入在训练期间学习,因此相似的标记最终具有相似的向量。
|
||||
|
||||
- **分词**是将原始文本转换为标记序列的过程。词级分词按空格分割,但无法处理未见过的词。**子词分词**(BPE、WordPiece、SentencePiece)将文本分解为频繁的子词单元,平衡词汇表大小和覆盖率。单词"unhappiness"可能变成["un", "happiness"]或["un", "happ", "iness"]。
|
||||
|
||||
- **循环神经网络(RNN)** 一次处理一个序列元素,维护一个向前传递信息的隐藏状态:
|
||||
|
||||
$$h_t = \tanh(W_h h_{t-1} + W_x x_t + b)$$
|
||||
|
||||
- 隐藏状态 $h_t$ 是网络到时间 $t$ 为止所看到内容的压缩摘要。相同的权重 $W_h$ 和 $W_x$ 在所有时间步共享(权重共享,如同CNN共享空间权重)。
|
||||
|
||||
- 原始RNN在长序列上存在梯度消失问题:从步骤 $t$ 到步骤 $t - k$ 的梯度信号经过 $k$ 次与 $W_h$ 的乘法,呈指数级缩小(或爆炸)。
|
||||
|
||||
- **LSTM**(长短时记忆网络)通过引入一个独立的细胞状态 $c_t$ 来解决这一问题,该状态以最小干扰流过时间。三个门控制哪些信息进入、离开和持续存在:
|
||||
|
||||
- **遗忘门**决定从细胞状态中擦除什么:$f_t = \sigma(W_f [h_{t-1}, x_t] + b_f)$
|
||||
- **输入门**决定写入什么新信息:$i_t = \sigma(W_i [h_{t-1}, x_t] + b_i)$,候选值 $\tilde{c}_t = \tanh(W_c [h_{t-1}, x_t] + b_c)$
|
||||
- 细胞状态更新:$c_t = f_t \odot c_{t-1} + i_t \odot \tilde{c}_t$
|
||||
- **输出门**决定暴露什么:$o_t = \sigma(W_o [h_{t-1}, x_t] + b_o)$,$h_t = o_t \odot \tanh(c_t)$
|
||||
|
||||

|
||||
|
||||
- 细胞状态像传送带一样工作:信息可以不变地流过许多时间步(遗忘门保持接近1),这解决了长距离依赖的梯度消失问题。
|
||||
|
||||
- **GRU**(门控循环单元)通过将细胞状态和隐藏状态合并为一个,并使用两个门(更新门和重置门)代替三个门来简化LSTM。GRU参数更少,通常表现与LSTM相当。
|
||||
|
||||
- RNN(包括LSTM)的根本限制是顺序处理:必须按顺序处理标记1、标记2、标记3。这阻止了并行化并造成信息瓶颈,因为所有上下文必须通过固定大小的隐藏状态。
|
||||
|
||||
- **注意力机制**解决了这两个问题。注意力机制不是将整个输入压缩为固定向量,而是让模型回顾所有输入位置并决定哪些位置与当前输出相关。
|
||||
|
||||
- 现代公式使用**查询、键和值(Q, K, V)**。将其想象为图书馆搜索:你有一个查询(你在找什么)、键(每本书的标签)和值(实际书籍内容)。你将查询与所有键比较,以确定检索哪些值。
|
||||
|
||||
- **缩放点积注意力**:
|
||||
|
||||
$$\text{Attention}(Q, K, V) = \text{softmax}\!\left(\frac{QK^T}{\sqrt{d_k}}\right) V$$
|
||||
|
||||
- $QK^T$ 计算每个查询和每个键之间的相似度。这是矩阵乘法(第02章),其中的条目是点积,衡量余弦相似度(第01章)。除以 $\sqrt{d_k}$ 防止点积变得太大(这会使softmax饱和并产生接近one-hot分布,导致梯度消失)。Softmax将相似度转换为概率分布。乘以 $V$ 产生值的加权组合。
|
||||
|
||||
- **多头注意力**运行 $h$ 个并行的注意力操作,每个使用不同的Q、K、V学习投影。这让模型同时从不同的表示子空间关注信息。一个头可能关注句法关系,而另一个关注语义关系。输出被拼接并投影:
|
||||
|
||||
$$\text{MultiHead}(Q, K, V) = \text{Concat}(\text{head}_1, \ldots, \text{head}_h) W^O$$
|
||||
|
||||
- **Transformer**架构(Vaswani等人,2017)完全由注意力和前馈层构建,没有循环。编码器块重复:多头自注意力、加法和层归一化、前馈网络、加法和层归一化。解码器块添加了掩码自注意力(防止模型看到未来的标记)和关注编码器输出的交叉注意力层。
|
||||
|
||||

|
||||
|
||||
- **位置编码**是必需的,因为注意力是排列等变的,意味着它将输入视为集合而非序列。没有位置信息,"猫坐在垫子上"和"垫子坐在猫上"将是相同的。原始Transformer使用正弦位置编码:
|
||||
|
||||
$$PE_{(pos, 2i)} = \sin\!\left(\frac{pos}{10000^{2i/d}}\right), \quad PE_{(pos, 2i+1)} = \cos\!\left(\frac{pos}{10000^{2i/d}}\right)$$
|
||||
|
||||
- 每个位置获得一个唯一的向量,模型可以用来区分位置。现代模型通常使用学习的位置嵌入或相对位置编码(RoPE、ALiBi)代替。
|
||||
|
||||
- Transformer并行处理所有标记(自注意力矩阵 $QK^T$ 在一次矩阵乘法中计算),这使得它们在现代硬件上比RNN训练更快。权衡是自注意力在序列长度上是 $O(n^2)$(每个标记关注每个其他标记),而RNN是 $O(n)$。这就是为什么长上下文模型需要特殊的注意力变体(稀疏注意力、线性注意力、Flash Attention)。
|
||||
|
||||
- **视觉Transformer(ViT)** 通过将图像分割为固定大小的块(如16x16),将每个块展平为向量,并将这些块视为标记序列,将Transformer应用于图像。一个可学习的[CLS]标记被前置,其最终表示用于分类。尽管没有卷积的归纳偏置,ViT在足够数据上训练时可以匹配或超越CNN。
|
||||
|
||||
- **MLP-Mixer** 是一种更简单的架构,用MLP替代了注意力和卷积。它在"标记混合"MLP(跨空间位置应用)和"通道混合"MLP(跨特征应用)之间交替。它的表现具有竞争力,表明现代架构的关键洞察不是注意力本身,而是跨标记和特征的高效信息混合。
|
||||
|
||||
- **自编码器**通过训练网络重构自身输入来学习压缩表示。编码器将输入映射到低维瓶颈(潜码),解码器将其映射回来:
|
||||
|
||||
$$z = f_{\text{enc}}(x), \quad \hat{x} = f_{\text{dec}}(z), \quad \mathcal{L} = \|x - \hat{x}\|^2$$
|
||||
|
||||
- 瓶颈迫使网络学习最重要的特征。自编码器用于降维、去噪(在噪声输入上训练,重构干净输出)和异常检测(高重构误差表明输入异常)。
|
||||
|
||||
- **变分自编码器(VAE)** 增加了概率的变体。编码器不是编码到单个点 $z$,而是输出分布的参数(高斯的均值 $\mu$ 和方差 $\sigma^2$)。潜码从此分布中采样:$z = \mu + \sigma \odot \epsilon$,其中 $\epsilon \sim \mathcal{N}(0, I)$。这个**重参数化技巧**使采样可微,梯度可以流过。
|
||||
|
||||
- VAE损失有两个项:
|
||||
|
||||
$$\mathcal{L} = \underbrace{\|x - \hat{x}\|^2}_{\text{reconstruction}} + \underbrace{D_{\text{KL}}(q(z|x) \| p(z))}_{\text{regularisation}}$$
|
||||
|
||||
- KL散度项(来自第05章)将学习到的后验 $q(z|x)$ 推向先验 $p(z) = \mathcal{N}(0, I)$,确保潜空间平滑且结构良好。然后你可以从先验中采样并解码以生成新数据。这就是使VAE成为生成模型的原因。
|
||||
|
||||
## 编程任务(在CoLab或笔记本中完成)
|
||||
|
||||
1. 在JAX中从头构建一个简单的MLP。在二维分类问题(如同心圆)上训练并可视化决策边界。
|
||||
```python
|
||||
import jax
|
||||
import jax.numpy as jnp
|
||||
import matplotlib.pyplot as plt
|
||||
from sklearn.datasets import make_circles
|
||||
|
||||
# 数据
|
||||
X, y = make_circles(n_samples=500, noise=0.1, factor=0.5, random_state=42)
|
||||
X, y = jnp.array(X), jnp.array(y, dtype=jnp.float32)
|
||||
|
||||
# 初始化一个2层MLP:2 -> 16 -> 16 -> 1
|
||||
def init_params(key):
|
||||
k1, k2, k3 = jax.random.split(key, 3)
|
||||
return {
|
||||
'W1': jax.random.normal(k1, (2, 16)) * 0.5,
|
||||
'b1': jnp.zeros(16),
|
||||
'W2': jax.random.normal(k2, (16, 16)) * 0.5,
|
||||
'b2': jnp.zeros(16),
|
||||
'W3': jax.random.normal(k3, (16, 1)) * 0.5,
|
||||
'b3': jnp.zeros(1),
|
||||
}
|
||||
|
||||
def forward(params, x):
|
||||
h = jnp.maximum(0, x @ params['W1'] + params['b1']) # ReLU
|
||||
h = jnp.maximum(0, h @ params['W2'] + params['b2']) # ReLU
|
||||
logit = (h @ params['W3'] + params['b3']).squeeze()
|
||||
return jax.nn.sigmoid(logit)
|
||||
|
||||
def loss_fn(params, X, y):
|
||||
pred = forward(params, X)
|
||||
return -jnp.mean(y * jnp.log(pred + 1e-7) + (1 - y) * jnp.log(1 - pred + 1e-7))
|
||||
|
||||
grad_fn = jax.jit(jax.grad(loss_fn))
|
||||
params = init_params(jax.random.PRNGKey(0))
|
||||
lr = 0.1
|
||||
|
||||
for step in range(2000):
|
||||
grads = grad_fn(params, X, y)
|
||||
params = {k: params[k] - lr * grads[k] for k in params}
|
||||
|
||||
# 绘制决策边界
|
||||
xx, yy = jnp.meshgrid(jnp.linspace(-2, 2, 200), jnp.linspace(-2, 2, 200))
|
||||
grid = jnp.column_stack([xx.ravel(), yy.ravel()])
|
||||
zz = forward(params, grid).reshape(xx.shape)
|
||||
|
||||
plt.figure(figsize=(7, 6))
|
||||
plt.contourf(xx, yy, zz, levels=[0, 0.5, 1], alpha=0.3, colors=['#e74c3c', '#3498db'])
|
||||
plt.scatter(X[y==0,0], X[y==0,1], c='#e74c3c', s=10, label='Class 0')
|
||||
plt.scatter(X[y==1,0], X[y==1,1], c='#3498db', s=10, label='Class 1')
|
||||
plt.title("MLP Decision Boundary on Concentric Circles")
|
||||
plt.legend(); plt.grid(alpha=0.3); plt.show()
|
||||
|
||||
acc = jnp.mean((forward(params, X) > 0.5) == y)
|
||||
print(f"Accuracy: {acc:.2%}")
|
||||
```
|
||||
|
||||
2. 从头实现一维卷积。将简单的边缘检测滤波器应用于信号,并与内置的 `jnp.convolve` 进行比较。
|
||||
```python
|
||||
import jax.numpy as jnp
|
||||
import matplotlib.pyplot as plt
|
||||
|
||||
def conv1d(signal, kernel):
|
||||
"""从头实现一维卷积(valid模式)。"""
|
||||
n, k = len(signal), len(kernel)
|
||||
output = jnp.zeros(n - k + 1)
|
||||
for i in range(n - k + 1):
|
||||
output = output.at[i].set(jnp.sum(signal[i:i+k] * kernel))
|
||||
return output
|
||||
|
||||
# 创建一个带有阶跃函数的信号
|
||||
t = jnp.linspace(0, 4, 200)
|
||||
signal = jnp.where(t < 1, 0.0, jnp.where(t < 2, 1.0, jnp.where(t < 3, 0.5, 1.5)))
|
||||
|
||||
# 边缘检测核
|
||||
edge_kernel = jnp.array([-1.0, 0.0, 1.0])
|
||||
|
||||
# 我们的实现 vs 内置函数
|
||||
our_output = conv1d(signal, edge_kernel)
|
||||
jnp_output = jnp.convolve(signal, edge_kernel, mode='valid')
|
||||
|
||||
fig, axes = plt.subplots(3, 1, figsize=(10, 6), sharex=True)
|
||||
axes[0].plot(t, signal, color='#3498db', linewidth=1.5)
|
||||
axes[0].set_title("Original Signal"); axes[0].set_ylabel("Value")
|
||||
|
||||
axes[1].plot(t[:len(our_output)], our_output, color='#e74c3c', linewidth=1.5)
|
||||
axes[1].set_title("After Edge Detection (our conv1d)"); axes[1].set_ylabel("Value")
|
||||
|
||||
axes[2].plot(t[:len(jnp_output)], jnp_output, color='#27ae60', linewidth=1.5, linestyle='--')
|
||||
axes[2].set_title("After Edge Detection (jnp.convolve)"); axes[2].set_ylabel("Value")
|
||||
axes[2].set_xlabel("t")
|
||||
|
||||
plt.tight_layout(); plt.show()
|
||||
print(f"Outputs match: {jnp.allclose(our_output, jnp_output)}")
|
||||
```
|
||||
|
||||
3. 从头实现缩放点积注意力。为一个小例子计算注意力权重,并将注意力矩阵可视化为热力图。
|
||||
```python
|
||||
import jax
|
||||
import jax.numpy as jnp
|
||||
import matplotlib.pyplot as plt
|
||||
|
||||
def scaled_dot_product_attention(Q, K, V):
|
||||
"""缩放点积注意力。"""
|
||||
d_k = Q.shape[-1]
|
||||
scores = Q @ K.T / jnp.sqrt(d_k)
|
||||
weights = jax.nn.softmax(scores, axis=-1)
|
||||
output = weights @ V
|
||||
return output, weights
|
||||
|
||||
# 示例:4个标记,嵌入维度8
|
||||
key = jax.random.PRNGKey(42)
|
||||
k1, k2, k3 = jax.random.split(key, 3)
|
||||
seq_len, d_model = 4, 8
|
||||
|
||||
Q = jax.random.normal(k1, (seq_len, d_model))
|
||||
K = jax.random.normal(k2, (seq_len, d_model))
|
||||
V = jax.random.normal(k3, (seq_len, d_model))
|
||||
|
||||
output, weights = scaled_dot_product_attention(Q, K, V)
|
||||
|
||||
print(f"Q shape: {Q.shape}")
|
||||
print(f"Attention weights shape: {weights.shape}")
|
||||
print(f"Output shape: {output.shape}")
|
||||
print(f"\nAttention weights (rows sum to 1):")
|
||||
print(weights)
|
||||
print(f"Row sums: {weights.sum(axis=-1)}")
|
||||
|
||||
# 可视化注意力
|
||||
fig, ax = plt.subplots(figsize=(5, 4))
|
||||
im = ax.imshow(weights, cmap='Blues', vmin=0, vmax=1)
|
||||
ax.set_xlabel("Key position"); ax.set_ylabel("Query position")
|
||||
ax.set_title("Attention Weights")
|
||||
tokens = ['tok 0', 'tok 1', 'tok 2', 'tok 3']
|
||||
ax.set_xticks(range(4)); ax.set_xticklabels(tokens)
|
||||
ax.set_yticks(range(4)); ax.set_yticklabels(tokens)
|
||||
for i in range(4):
|
||||
for j in range(4):
|
||||
ax.text(j, i, f"{weights[i,j]:.2f}", ha='center', va='center', fontsize=10)
|
||||
plt.colorbar(im); plt.tight_layout(); plt.show()
|
||||
```
|
||||
|
||||
4. 构建一个简单的自编码器,通过一维瓶颈压缩二维数据并重建。可视化潜空间和重建结果。
|
||||
```python
|
||||
import jax
|
||||
import jax.numpy as jnp
|
||||
import matplotlib.pyplot as plt
|
||||
from sklearn.datasets import make_moons
|
||||
|
||||
# 数据
|
||||
X, _ = make_moons(n_samples=500, noise=0.05, random_state=42)
|
||||
X = jnp.array(X)
|
||||
|
||||
# 自编码器:2 -> 8 -> 1 -> 8 -> 2
|
||||
def init_ae(key):
|
||||
k1, k2, k3, k4 = jax.random.split(key, 4)
|
||||
return {
|
||||
'enc_W1': jax.random.normal(k1, (2, 8)) * 0.5, 'enc_b1': jnp.zeros(8),
|
||||
'enc_W2': jax.random.normal(k2, (8, 1)) * 0.5, 'enc_b2': jnp.zeros(1),
|
||||
'dec_W1': jax.random.normal(k3, (1, 8)) * 0.5, 'dec_b1': jnp.zeros(8),
|
||||
'dec_W2': jax.random.normal(k4, (8, 2)) * 0.5, 'dec_b2': jnp.zeros(2),
|
||||
}
|
||||
|
||||
def encode(p, x):
|
||||
h = jnp.tanh(x @ p['enc_W1'] + p['enc_b1'])
|
||||
return h @ p['enc_W2'] + p['enc_b2']
|
||||
|
||||
def decode(p, z):
|
||||
h = jnp.tanh(z @ p['dec_W1'] + p['dec_b1'])
|
||||
return h @ p['dec_W2'] + p['dec_b2']
|
||||
|
||||
def ae_loss(p, X):
|
||||
z = encode(p, X)
|
||||
X_hat = decode(p, z)
|
||||
return jnp.mean((X - X_hat) ** 2)
|
||||
|
||||
grad_fn = jax.jit(jax.grad(ae_loss))
|
||||
params = init_ae(jax.random.PRNGKey(0))
|
||||
lr = 0.01
|
||||
|
||||
for step in range(3000):
|
||||
grads = grad_fn(params, X)
|
||||
params = {k: params[k] - lr * grads[k] for k in params}
|
||||
|
||||
z = encode(params, X)
|
||||
X_hat = decode(params, z)
|
||||
|
||||
fig, axes = plt.subplots(1, 2, figsize=(12, 5))
|
||||
axes[0].scatter(X[:,0], X[:,1], c=z.squeeze(), cmap='viridis', s=10)
|
||||
axes[0].set_title("Original Data (coloured by latent code)")
|
||||
axes[1].scatter(X_hat[:,0], X_hat[:,1], c=z.squeeze(), cmap='viridis', s=10)
|
||||
axes[1].set_title("Reconstruction from 1D bottleneck")
|
||||
for ax in axes:
|
||||
ax.set_aspect('equal'); ax.grid(alpha=0.3)
|
||||
plt.tight_layout(); plt.show()
|
||||
|
||||
print(f"Reconstruction MSE: {ae_loss(params, X):.4f}")
|
||||
```
|
||||
@@ -0,0 +1,353 @@
|
||||
# 强化学习
|
||||
|
||||
*强化学习通过试错法最大化累积奖励来训练智能体做出序列决策。本文件涵盖MDP、价值函数、贝尔曼方程、Q学习、策略梯度、演员-评论家方法、PPO和RLHF——这些是游戏智能体和语言模型对齐背后的框架。*
|
||||
|
||||
- 监督学习需要标注数据。无监督学习在无标注数据中发现模式。**强化学习(RL)** 与两者都不同:智能体通过与环境的交互、采取行动和接收奖励来学习。没有正确的标签;智能体必须通过试错来发现好的行为。
|
||||
|
||||
- 想象教狗一个新把戏。你不会给它展示一个正确行为的数据集。相反,它尝试各种动作,你对好的行为给予奖励,随着时间的推移它明白了你想要什么。RL将这个形式化。
|
||||
|
||||
- RL设置包含五个核心组件。**智能体(agent)** 是学习者和决策者。**环境(environment)** 是智能体之外与之交互的一切。在每个时间步,智能体观察一个**状态(state)** $s_t$,选择一个**动作(action)** $a_t$,接收一个**奖励(reward)** $r_t$,并转移到新状态 $s_{t+1}$。智能体的目标是最大化其随时间收集的总奖励。
|
||||
|
||||

|
||||
|
||||
- **策略(policy)** $\pi$ 是智能体的策略:从状态到动作的映射。确定性策略对每个状态给出一个动作:$a = \pi(s)$。随机策略给出动作上的概率分布:$\pi(a \mid s)$。RL的目标是找到最优策略,即最大化期望累积奖励的策略。
|
||||
|
||||
- RL的数学框架是**马尔可夫决策过程(MDP)**,由元组 $(S, A, P, R, \gamma)$ 定义:一组状态 $S$,一组动作 $A$,转移概率 $P(s' \mid s, a)$,奖励函数 $R(s, a)$,以及折扣因子 $\gamma$。
|
||||
|
||||
- **马尔可夫性质**(来自第05章)指出未来仅取决于当前状态,而不是如何到达那里的历史:$P(s_{t+1} \mid s_t, a_t, s_{t-1}, \ldots) = P(s_{t+1} \mid s_t, a_t)$。这意味着状态包含了做出决策所需的全部信息。
|
||||
|
||||
- **折扣因子** $\gamma \in [0, 1)$ 决定了智能体对未来奖励相对于即时奖励的重视程度。从时间 $t$ 开始的折扣回报为:
|
||||
|
||||
$$G_t = r_t + \gamma r_{t+1} + \gamma^2 r_{t+2} + \cdots = \sum_{k=0}^{\infty} \gamma^k r_{t+k}$$
|
||||
|
||||
- 当 $\gamma = 0$ 时,智能体完全短视,只关心下一个奖励。当 $\gamma$ 接近1时,智能体具有长远眼光。折扣因子还确保了求和收敛(如果奖励有界),这对数学上的良定义性很重要。
|
||||
|
||||
- **价值函数**估计处于某个状态(或在某个状态下采取某个动作)有多好。**状态价值函数** $V^\pi(s)$ 是从状态 $s$ 开始并按照策略 $\pi$ 行动所获得的期望回报:
|
||||
|
||||
$$V^\pi(s) = \mathbb{E}_\pi \left[ G_t \mid s_t = s \right]$$
|
||||
|
||||
- **动作价值函数** $Q^\pi(s, a)$ 是从状态 $s$ 开始,采取动作 $a$,然后按照 $\pi$ 行动所获得的期望回报:
|
||||
|
||||
$$Q^\pi(s, a) = \mathbb{E}_\pi \left[ G_t \mid s_t = s, a_t = a \right]$$
|
||||
|
||||
- 两者关系:$V^\pi(s) = \sum_a \pi(a \mid s) \, Q^\pi(s, a)$。状态价值是动作价值按策略加权的平均值。
|
||||
|
||||
- **贝尔曼方程**表达了递归关系:一个状态的价值等于即时奖励加上下一个状态的折扣价值。对于状态价值函数:
|
||||
|
||||
$$V^\pi(s) = \sum_a \pi(a \mid s) \sum_{s'} P(s' \mid s, a) \left[ R(s, a) + \gamma \, V^\pi(s') \right]$$
|
||||
|
||||
- 对于最优价值函数 $V^{*}(s)$,智能体总是选择最佳动作:
|
||||
|
||||
$$V^{*}(s) = \max_a \sum_{s'} P(s' \mid s, a) \left[ R(s, a) + \gamma \, V^{*}(s') \right]$$
|
||||
|
||||
- 类似地,$Q^{*}$ 的**贝尔曼最优方程**为:
|
||||
|
||||
$$Q^{*}(s, a) = \sum_{s'} P(s' \mid s, a) \left[ R(s, a) + \gamma \max_{a'} Q^{*}(s', a') \right]$$
|
||||
|
||||
- 一旦你有了 $Q^{*}$,最优策略就很简单了:总是选择Q值最高的动作:$\pi^{*}(s) = \arg\max_a Q^{*}(s, a)$。
|
||||
|
||||
- **动态规划**方法在已知转移概率和奖励(完整模型)时求解MDP。**策略评估**通过迭代应用贝尔曼方程直到收敛来计算给定策略的 $V^\pi$。**策略改进**利用价值函数并通过对最优动作贪心来构建更好的策略:$\pi'(s) = \arg\max_a \sum_{s'} P(s' \mid s, a)[R(s,a) + \gamma V^\pi(s')]$。
|
||||
|
||||
- **策略迭代**在评估和改进之间交替,直到策略停止变化。它保证收敛到最优策略。
|
||||
|
||||
- **价值迭代**将两个步骤合并为一个:重复应用贝尔曼最优方程直到 $V^{*}$ 收敛,然后提取策略。
|
||||
|
||||
$$V(s) \leftarrow \max_a \sum_{s'} P(s' \mid s, a) \left[ R(s, a) + \gamma \, V(s') \right]$$
|
||||
|
||||
- 动态规划需要知道 $P(s' \mid s, a)$,这通常不可行。在大多数真实问题中,智能体不知道环境的动态;它只能与环境交互。这就是**无模型**方法发挥作用的地方。
|
||||
|
||||
- **时序差分(TD)学习**在不了解模型的情况下从经验中学习。关键思想是**引导(bootstrapping)**:不等情节结束才计算实际回报 $G_t$,而是使用当前的价值函数对其进行估计:
|
||||
|
||||
$$V(s_t) \leftarrow V(s_t) + \alpha \left[ r_t + \gamma \, V(s_{t+1}) - V(s_t) \right]$$
|
||||
|
||||
- 括号中的项是**TD误差**:**TD目标**($r_t + \gamma V(s_{t+1})$)与当前估计 $V(s_t)$ 之间的差异。如果TD误差为正,说明该状态比预期好,我们增加其价值。如果为负,则减少其价值。
|
||||
|
||||

|
||||
|
||||
- TD学习在每一步之后(而不是完成整个情节后)进行更新,这使其比蒙特卡洛方法高效得多。它也适用于持续(非情节式)环境。
|
||||
|
||||
- **SARSA**(状态-动作-奖励-状态-动作)是将TD学习应用于Q值。智能体在状态 $s$ 下采取动作 $a$,观察奖励 $r$ 和下一状态 $s'$,然后根据其策略选择下一个动作 $a'$:
|
||||
|
||||
$$Q(s, a) \leftarrow Q(s, a) + \alpha \left[ r + \gamma \, Q(s', a') - Q(s, a) \right]$$
|
||||
|
||||
- SARSA是**在策略(on-policy)**:它使用智能体实际采取的动作进行更新,这包括了探索。这使得SARSA更为保守;它学习一个考虑自身探索噪声的策略。
|
||||
|
||||
- **Q学习**是最著名的RL算法。它类似于SARSA,但不同的是它使用最佳可能动作而非智能体实际采取的动作:
|
||||
|
||||
$$Q(s, a) \leftarrow Q(s, a) + \alpha \left[ r + \gamma \max_{a'} Q(s', a') - Q(s, a) \right]$$
|
||||
|
||||
- Q学习是**离策略(off-policy)**:它学习最优Q值,与正在执行的策略无关。智能体可以随机探索,同时仍然学习最优动作价值。这使得Q学习更具攻击性,通常收敛更快,但可能高估值。
|
||||
|
||||
- **探索 vs 利用**是基本困境:智能体应该利用已知信息(选择估计价值最高的动作)还是探索未知动作(可能发现更好的)?
|
||||
|
||||
- 最简单的策略是**ε-贪心**:以概率 $\epsilon$ 采取随机动作(探索);以概率 $1 - \epsilon$ 采取贪心动作(利用)。一种常见的时间表是从高 $\epsilon$(大量探索)开始,随时间衰减。
|
||||
|
||||
- 表格方法(在表中存储每个状态-动作对的价值)适用于小的离散状态空间。对于大或连续的状态空间,需要函数近似。**深度Q网络(DQN)** 使用神经网络来近似 $Q(s, a; \theta)$,其中 $\theta$ 是网络权重。
|
||||
|
||||
- DQN引入了两个关键的稳定技术。**经验回放**:不是从连续的转移中学习(高度相关),而是将转移存储在回放缓冲区中,并采样随机小批次进行训练。这打破了相关性并高效地重用数据。
|
||||
|
||||
- **目标网络**:使用一个单独的、缓慢更新的网络副本来计算TD目标。没有这个,每次更新网络时目标都会移动,造成"追自己尾巴"的不稳定性。目标网络定期更新(每 $N$ 步硬更新)或连续更新(软更新:$\theta^{-} \leftarrow \tau\theta + (1-\tau)\theta^{-}$)。
|
||||
|
||||
- DQN损失只是预测Q值与TD目标之间的均方误差:
|
||||
|
||||
$$\mathcal{L}(\theta) = \mathbb{E} \left[ \left( r + \gamma \max_{a'} Q(s', a'; \theta^{-}) - Q(s, a; \theta) \right)^2 \right]$$
|
||||
|
||||
- 到目前为止的所有方法都学习价值函数并从中推导策略。**策略梯度**方法采用不同方法:它们直接参数化策略 $\pi(a \mid s; \theta)$ 并通过梯度上升优化期望回报。
|
||||
|
||||
- **策略梯度定理**给出了期望回报相对于策略参数的梯度:
|
||||
|
||||
$$\nabla_\theta J(\theta) = \mathbb{E}_\pi \left[ \nabla_\theta \log \pi(a \mid s; \theta) \cdot G_t \right]$$
|
||||
|
||||
- 这说明:增加导致高回报的动作的概率,减少导致低回报的动作的概率。对数概率梯度给出了改变策略的方向,$G_t$ 则缩放改变的程度。
|
||||
|
||||
- **REINFORCE**是最简单的策略梯度算法。运行一个情节,为每一步计算回报 $G_t$,然后更新:
|
||||
|
||||
$$\theta \leftarrow \theta + \alpha \, \nabla_\theta \log \pi(a_t \mid s_t; \theta) \cdot G_t$$
|
||||
|
||||
- REINFORCE方差很高,因为 $G_t$ 是期望回报的噪声单样本估计。一个常见修复是减去一个**基线(baseline)**(通常是平均回报或学习到的价值函数)来降低方差而不引入偏差:
|
||||
|
||||
$$\theta \leftarrow \theta + \alpha \, \nabla_\theta \log \pi(a_t \mid s_t; \theta) \cdot (G_t - b)$$
|
||||
|
||||
- **演员-评论家(Actor-Critic)** 方法使用两个网络。**演员(actor)** 是策略 $\pi(a \mid s; \theta)$。**评论家(critic)** 是价值函数 $V(s; \phi)$,作为基线。优势 $A_t = r_t + \gamma V(s_{t+1}) - V(s_t)$ 替代了 $G_t - b$:
|
||||
|
||||
$$\theta \leftarrow \theta + \alpha \, \nabla_\theta \log \pi(a_t \mid s_t; \theta) \cdot A_t$$
|
||||
|
||||
- 评论家通过最小化TD误差来更新,与基于价值的方法相同。演员使用策略梯度更新,评论家的优势估计降低了方差。这是两全其美。
|
||||
|
||||

|
||||
|
||||
- **PPO**(近端策略优化)是实践中使用最广泛的策略梯度算法。它解决了一个关键问题:如果策略更新过大,性能可能灾难性地崩溃。
|
||||
|
||||
- PPO使用一个**裁剪的替代目标**。令 $r_t(\theta) = \frac{\pi(a_t | s_t; \theta)}{\pi(a_t | s_t; \theta_{\text{old}})}$ 为新旧策略之间的概率比。损失为:
|
||||
|
||||
$$\mathcal{L}^{\text{CLIP}}(\theta) = \mathbb{E} \left[ \min\!\left( r_t(\theta) A_t, \; \text{clip}(r_t(\theta), 1-\epsilon, 1+\epsilon) A_t \right) \right]$$
|
||||
|
||||
- 裁剪(通常 $\epsilon = 0.2$)防止比率远离1,使更新保持小而稳定。如果优势为正(动作好),比率上限为 $1 + \epsilon$。如果为负(动作差),比率下限为 $1 - \epsilon$。这比早期的信任区域方法(TRPO)更简单、更稳定。
|
||||
|
||||
- PPO被用于通过**RLHF**(基于人类反馈的强化学习)训练ChatGPT风格的模型。在RLHF中,一个奖励模型在人类偏好数据(人类更喜欢两个输出中的哪一个?)上训练,然后PPO优化语言模型策略以最大化这个学习到的奖励。
|
||||
|
||||
- **DPO**(直接偏好优化)通过完全消除奖励模型来简化RLHF。DPO不训练奖励模型然后运行RL,而是推导出一个闭式损失,直接从偏好数据优化策略:
|
||||
|
||||
$$\mathcal{L}_{\text{DPO}}(\theta) = -\mathbb{E} \left[ \log \sigma\!\left( \beta \log \frac{\pi_\theta(y_w \mid x)}{\pi_{\text{ref}}(y_w \mid x)} - \beta \log \frac{\pi_\theta(y_l \mid x)}{\pi_{\text{ref}}(y_l \mid x)} \right) \right]$$
|
||||
|
||||
- 这里 $y_w$ 是偏好的(胜出)回答,$y_l$ 是不被偏好的(失败)回答。DPO增加偏好输出的相对概率,并且比基于PPO的RLHF实现起来简单得多。
|
||||
|
||||
- RL算法中有两个重要区分。**在策略 vs 离策略**:在策略方法(SARSA, PPO)从当前策略生成的数据中学习;离策略方法(Q学习, DQN)可以从任何策略生成的数据中学习。离策略方法样本效率更高(它们重用旧数据),但可能不那么稳定。
|
||||
|
||||
- **基于模型 vs 无模型**:无模型方法(到目前为止讨论的所有方法)直接从经验中学习价值或策略。基于模型的方法学习环境的模型($P(s' \mid s, a)$ 和 $R(s, a)$)并用其进行规划(想象未来的轨迹而不实际采取动作)。基于模型的方法样本效率更高,但增加了学习精确模型的复杂性。
|
||||
|
||||
- 总结RL领域:
|
||||
|
||||
| 方法 | 类型 | 核心思想 | 优势 |
|
||||
|---|---|---|---|
|
||||
| 价值迭代 | DP, 基于模型 | 贝尔曼最优性 | 精确解(小MDP) |
|
||||
| SARSA | TD, 在策略 | 在策略学习Q | 保守、安全 |
|
||||
| Q学习 | TD, 离策略 | 学习Q*, 贪心目标 | 简单、有效 |
|
||||
| DQN | 深度, 离策略 | 神经Q + 回放 + 目标网络 | 扩展到高维状态 |
|
||||
| REINFORCE | 策略梯度 | log-概率 * 回报的梯度 | 简单的策略优化 |
|
||||
| 演员-评论家 | PG + 价值 | 演员 + 评论家降低方差 | 实用且灵活 |
|
||||
| PPO | PG, 裁剪 | 信任区域般的稳定性 | 行业标准 |
|
||||
| DPO | 直接偏好 | 跳过奖励模型 | 更简单的RLHF |
|
||||
|
||||
## 编程任务(使用CoLab或笔记本)
|
||||
|
||||
1. 为简单的网格世界实现价值迭代。计算最优价值函数并提取最优策略。将两者可视化为热力图和箭头图。
|
||||
```python
|
||||
import jax.numpy as jnp
|
||||
import matplotlib.pyplot as plt
|
||||
|
||||
# 4x4网格世界:目标在(3,3),每步奖励-1,目标处为0
|
||||
grid_size = 4
|
||||
gamma = 0.99
|
||||
goal = (3, 3)
|
||||
|
||||
# 动作:上、下、左、右
|
||||
actions = [(-1, 0), (1, 0), (0, -1), (0, 1)]
|
||||
action_names = ['up', 'down', 'left', 'right']
|
||||
action_arrows = ['\u2191', '\u2193', '\u2190', '\u2192']
|
||||
|
||||
def step(s, a):
|
||||
"""确定性转移。"""
|
||||
ns = (max(0, min(grid_size-1, s[0]+a[0])),
|
||||
max(0, min(grid_size-1, s[1]+a[1])))
|
||||
return ns
|
||||
|
||||
# 价值迭代
|
||||
V = jnp.zeros((grid_size, grid_size))
|
||||
for iteration in range(100):
|
||||
V_new = jnp.array(V)
|
||||
for i in range(grid_size):
|
||||
for j in range(grid_size):
|
||||
if (i, j) == goal:
|
||||
continue
|
||||
values = []
|
||||
for a in actions:
|
||||
ns = step((i, j), a)
|
||||
values.append(-1 + gamma * float(V[ns[0], ns[1]]))
|
||||
V_new = V_new.at[i, j].set(max(values))
|
||||
if jnp.max(jnp.abs(V_new - V)) < 1e-6:
|
||||
print(f"在{iteration+1}次迭代后收敛")
|
||||
break
|
||||
V = V_new
|
||||
|
||||
# 提取策略
|
||||
policy = [['' for _ in range(grid_size)] for _ in range(grid_size)]
|
||||
for i in range(grid_size):
|
||||
for j in range(grid_size):
|
||||
if (i, j) == goal:
|
||||
policy[i][j] = 'G'
|
||||
continue
|
||||
best_a = max(range(4), key=lambda a: -1 + gamma * float(V[step((i,j), actions[a])[0], step((i,j), actions[a])[1]]))
|
||||
policy[i][j] = action_arrows[best_a]
|
||||
|
||||
fig, axes = plt.subplots(1, 2, figsize=(10, 4))
|
||||
im = axes[0].imshow(V, cmap='YlOrRd_r')
|
||||
axes[0].set_title("最优价值函数")
|
||||
for i in range(grid_size):
|
||||
for j in range(grid_size):
|
||||
axes[0].text(j, i, f"{V[i,j]:.1f}", ha='center', va='center', fontsize=10)
|
||||
plt.colorbar(im, ax=axes[0])
|
||||
|
||||
axes[1].imshow(jnp.ones((grid_size, grid_size)), cmap='Greys', vmin=0, vmax=2)
|
||||
axes[1].set_title("最优策略")
|
||||
for i in range(grid_size):
|
||||
for j in range(grid_size):
|
||||
axes[1].text(j, i, policy[i][j], ha='center', va='center', fontsize=18)
|
||||
plt.tight_layout(); plt.show()
|
||||
```
|
||||
|
||||
2. 在简单的网格世界上实现表格Q学习。训练智能体,绘制学习曲线,显示学习到的Q值。
|
||||
```python
|
||||
import jax
|
||||
import jax.numpy as jnp
|
||||
import matplotlib.pyplot as plt
|
||||
|
||||
grid_size = 5
|
||||
goal = (4, 4)
|
||||
actions = [(-1,0), (1,0), (0,-1), (0,1)]
|
||||
|
||||
# Q表
|
||||
Q = {}
|
||||
for i in range(grid_size):
|
||||
for j in range(grid_size):
|
||||
Q[(i,j)] = [0.0] * 4
|
||||
|
||||
alpha = 0.1
|
||||
gamma = 0.95
|
||||
epsilon = 1.0
|
||||
epsilon_decay = 0.995
|
||||
min_epsilon = 0.01
|
||||
|
||||
def step(s, a_idx):
|
||||
a = actions[a_idx]
|
||||
ns = (max(0, min(grid_size-1, s[0]+a[0])),
|
||||
max(0, min(grid_size-1, s[1]+a[1])))
|
||||
r = 0.0 if ns == goal else -1.0
|
||||
done = ns == goal
|
||||
return ns, r, done
|
||||
|
||||
key = jax.random.PRNGKey(42)
|
||||
rewards_per_episode = []
|
||||
|
||||
for ep in range(500):
|
||||
s = (0, 0)
|
||||
total_reward = 0
|
||||
for _ in range(100):
|
||||
key, subkey = jax.random.split(key)
|
||||
if float(jax.random.uniform(subkey)) < epsilon:
|
||||
key, subkey = jax.random.split(key)
|
||||
a = int(jax.random.randint(subkey, (), 0, 4))
|
||||
else:
|
||||
a = max(range(4), key=lambda i: Q[s][i])
|
||||
|
||||
ns, r, done = step(s, a)
|
||||
total_reward += r
|
||||
# Q学习更新
|
||||
Q[s][a] += alpha * (r + gamma * max(Q[ns]) - Q[s][a])
|
||||
s = ns
|
||||
if done:
|
||||
break
|
||||
rewards_per_episode.append(total_reward)
|
||||
epsilon = max(min_epsilon, epsilon * epsilon_decay)
|
||||
|
||||
plt.figure(figsize=(8, 4))
|
||||
# 平滑曲线
|
||||
window = 20
|
||||
smoothed = [sum(rewards_per_episode[max(0,i-window):i+1])/min(i+1, window)
|
||||
for i in range(len(rewards_per_episode))]
|
||||
plt.plot(smoothed, color='#3498db', linewidth=1.5)
|
||||
plt.xlabel("Episode"); plt.ylabel("Total Reward (smoothed)")
|
||||
plt.title("Q-Learning on Gridworld")
|
||||
plt.grid(alpha=0.3); plt.show()
|
||||
|
||||
# 显示学到的策略
|
||||
arrow = ['\u2191', '\u2193', '\u2190', '\u2192']
|
||||
print("学到的策略:")
|
||||
for i in range(grid_size):
|
||||
row = ""
|
||||
for j in range(grid_size):
|
||||
if (i,j) == goal:
|
||||
row += " G "
|
||||
else:
|
||||
row += f" {arrow[max(range(4), key=lambda a: Q[(i,j)][a])]} "
|
||||
print(row)
|
||||
```
|
||||
|
||||
3. 在多臂老虎机问题上实现REINFORCE。展示策略如何随训练演变以偏向最佳臂。
|
||||
```python
|
||||
import jax
|
||||
import jax.numpy as jnp
|
||||
import matplotlib.pyplot as plt
|
||||
|
||||
# 5臂老虎机,不同期望奖励
|
||||
true_rewards = jnp.array([0.2, 0.5, 0.8, 0.3, 0.1])
|
||||
n_arms = len(true_rewards)
|
||||
|
||||
# 策略:在logits上的softmax
|
||||
logits = jnp.zeros(n_arms)
|
||||
lr = 0.1
|
||||
key = jax.random.PRNGKey(42)
|
||||
|
||||
policy_history = []
|
||||
reward_history = []
|
||||
|
||||
for step in range(2000):
|
||||
probs = jax.nn.softmax(logits)
|
||||
policy_history.append(probs)
|
||||
|
||||
# 采样动作
|
||||
key, subkey = jax.random.split(key)
|
||||
action = jax.random.choice(subkey, n_arms, p=probs)
|
||||
|
||||
# 获取奖励(伯努利分布)
|
||||
key, subkey = jax.random.split(key)
|
||||
reward = float(jax.random.uniform(subkey) < true_rewards[action])
|
||||
reward_history.append(reward)
|
||||
|
||||
# REINFORCE更新
|
||||
# grad log pi(a) = e_a - probs(对于softmax参数化)
|
||||
grad_log_pi = -probs.at[action].add(1.0) # one-hot(a) - probs
|
||||
logits = logits + lr * reward * grad_log_pi
|
||||
|
||||
policy_history = jnp.stack(policy_history)
|
||||
|
||||
fig, axes = plt.subplots(1, 2, figsize=(12, 4))
|
||||
colors = ['#3498db', '#e74c3c', '#27ae60', '#9b59b6', '#f39c12']
|
||||
for i in range(n_arms):
|
||||
axes[0].plot(policy_history[:, i], color=colors[i],
|
||||
label=f'臂{i} (真实={true_rewards[i]:.1f})', linewidth=1.5)
|
||||
axes[0].set_xlabel("步骤"); axes[0].set_ylabel("P(臂)")
|
||||
axes[0].set_title("策略演变 (REINFORCE)")
|
||||
axes[0].legend(fontsize=8); axes[0].grid(alpha=0.3)
|
||||
|
||||
# 平滑奖励
|
||||
window = 50
|
||||
smoothed = [sum(reward_history[max(0,i-window):i+1])/min(i+1,window)
|
||||
for i in range(len(reward_history))]
|
||||
axes[1].plot(smoothed, color='#27ae60', linewidth=1.5)
|
||||
axes[1].axhline(y=0.8, color='#e74c3c', linestyle='--', alpha=0.5, label='最佳臂')
|
||||
axes[1].set_xlabel("步骤"); axes[1].set_ylabel("平均奖励")
|
||||
axes[1].set_title("奖励随时间变化"); axes[1].legend()
|
||||
axes[1].grid(alpha=0.3)
|
||||
plt.tight_layout(); plt.show()
|
||||
```
|
||||
@@ -0,0 +1,264 @@
|
||||
# 分布式深度学习
|
||||
|
||||
*分布式训练将计算分散到多个GPU和机器上,以训练单个设备无法容纳或训练太慢的模型。本文件涵盖混合精度、数据并行、模型并行、流水线并行、ZeRO、FSDP、张量并行以及全规约等通信原语——这些对于大规模训练LLM至关重要。*
|
||||
|
||||
- 在单个GPU上训练大型神经网络最终会遇到瓶颈。模型可能无法放入内存,或者训练可能需要数月。分布式训练将工作分散到多个设备(GPU、TPU或整台机器)上,以更快地训练和训练更大的模型。本文件涵盖了实现这一目标的技术。
|
||||
|
||||
- 要理解为何分布式重要,从训练的**计算成本**开始。在一个包含 $d_{\text{in}}$ 个输入和 $d_{\text{out}}$ 个输出的密集层上,对一批 $B$ 个样本进行一次前向传播需要大约 $2 \cdot B \cdot d_{\text{in}} \cdot d_{\text{out}}$ 次FLOP(浮点运算):对输出矩阵的每个元素进行一次乘法和一次加法。反向传播的成本大约是前向传播的两倍(计算相对于输入和权重的梯度),因此一个密集层的一个训练步骤约为 $6 \cdot B \cdot d_{\text{in}} \cdot d_{\text{out}}$ 次FLOP。
|
||||
|
||||
- 对于隐藏维度为 $d$ 的Transformer层,自注意力块涉及四个投影(Q、K、V和输出),每个的成本为 $O(B \cdot n \cdot d^2)$ 次FLOP(其中 $n$ 是序列长度),加上注意力矩阵计算 $O(B \cdot n^2 \cdot d)$。前馈块有两个密集层,通常扩展到 $4d$ 再回来:$O(B \cdot n \cdot 8d^2)$。每层总计:大约 $O(B \cdot n \cdot 12d^2 + B \cdot n^2 \cdot d)$。乘以层数,你就会明白为什么训练GPT规模的模型需要数千个GPU小时。
|
||||
|
||||
- **内存墙**通常是更严格的约束。在训练期间,GPU内存必须同时容纳四样东西:
|
||||
|
||||

|
||||
|
||||
- **参数**:模型权重。一个70亿参数的模型在FP32中(每个参数4字节)仅权重就需要28 GB。
|
||||
- **梯度**:与参数大小相同。又是28 GB。
|
||||
- **优化器状态**:Adam维护两个额外的缓冲区(一阶和二阶矩估计),每个与参数大小相同。即使模型使用较低精度,这些也以FP32格式保存以确保数值稳定性。对于我们的7B模型,那就是 $2 \times 28 = 56$ GB。
|
||||
- **激活值**:在前向传播过程中保存下来供反向传播使用的中间值。大小取决于批量大小、序列长度和模型宽度。这通常是最主要的组成部分,并随批量大小线性增长。
|
||||
|
||||
- 对于使用FP32 Adam的7B模型:28(参数)+ 28(梯度)+ 56(优化器)= 112 GB,这还没算激活值。单个80 GB的A100 GPU无法容纳。这就是分布式策略至关重要的原因。
|
||||
|
||||
- **混合精度训练**是第一道防线。不是将所有内容存储在FP32(32位浮点)中,而是使用FP16或BF16(16位)进行前向和反向传播,同时将权重的FP32主副本保留给优化器更新。
|
||||
|
||||
- **FP16**具有高精度(10位尾数),但范围有限,可能导致上溢/下溢。损失缩放(在反向传播前将损失乘以一个大因子,然后将梯度除以相同因子)缓解了这个问题。
|
||||
|
||||
- **BF16**(脑浮点)具有与FP32相同的指数范围(8位指数),但精度较低(7位尾数)。它几乎从不溢出,很少需要损失缩放,因此使用更简单。BF16是现代Transformer训练的默认选择。
|
||||
|
||||
- 混合精度大致将激活值和梯度的内存减半(前向/反向传播期间的主要成本),同时将优化器状态保留在FP32中以确保数值稳定性。
|
||||
|
||||
- **数据并行**是最简单的分布式策略。你在 $N$ 个GPU上复制整个模型,将每个小批量分成 $N$ 个相等的块,并将一个块发送到每个GPU。每个GPU在其块上独立运行前向和反向传播。然后梯度在所有GPU上平均(使用全规约操作),每个GPU更新其本地模型副本。
|
||||
|
||||
- 从模型的角度来看,这相当于使用大了 $N$ 倍的小批量进行训练。如果每个GPU处理一个大小为 $B$ 的批次,则有效批量大小为 $N \cdot B$。
|
||||
|
||||

|
||||
|
||||
- 梯度平均可以同步或异步进行。**同步SGD**等待所有GPU完成后再进行平均,确保与使用更大批量的单GPU训练数学上等价。缺点是,最慢的GPU("掉队者")会拖慢所有人。
|
||||
|
||||
- **异步SGD**让每个GPU独立地更新一个共享的参数服务器,无需等待。这消除了掉队者问题,但引入了"陈旧梯度":一个GPU可能基于略微过时的参数计算梯度。陈旧梯度增加了噪声,可能减缓收敛。在实践中,带高效通信的同步SGD更受青睐。
|
||||
|
||||
- **梯度累积**是一种软件技巧,用于在有限硬件上模拟更大的批量大小。不必每个小批量做一次更新,而是运行多次前向/反向传播并累积梯度,然后做一次更新。这与更大批量得到相同的结果,而无需更多GPU内存用于激活值(一次只有一个小批量的激活值在内存中)。
|
||||
|
||||
- 当模型本身太大无法放入单个GPU时,需要**模型并行**。有两种主要形式。
|
||||
|
||||
- **张量并行**将单个层分割到多个GPU上。一个大的矩阵乘法 $Y = XW$ 可以按列分割:将 $W$ 分区为 $[W_1, W_2]$ 分布在两个GPU上,并行计算 $Y_1 = XW_1$ 和 $Y_2 = XW_2$,然后拼接。这适用于注意力投影和前馈层。它需要GPU之间快速通信(通常是节点内的NVLink),因为每层都必须组合部分结果。
|
||||
|
||||
- **流水线并行**将不同的层分配到不同的GPU上。GPU 0运行第1-4层,GPU 1运行第5-8层,依此类推。数据像流水线一样流经整个管道。朴素的方法有一个"流水线气泡":当GPU 0处理微批次1的前向传播时,GPU 1-3处于空闲状态。**微批处理**通过将小批量分割成更小的微批次来缓解这个问题,这些微批次按顺序流经流水线,使所有GPU大部分时间保持忙碌。
|
||||
|
||||
- **混合并行**结合了数据并行、张量并行和流水线并行。一个典型的大模型设置可能使用节点内的张量并行(8个GPU通过快速NVLink连接)、跨节点的流水线并行以及跨节点组的数据并行。这就是GPT-4和Llama等模型的训练方式。
|
||||
|
||||
- 分布式训练的效率在很大程度上取决于**通信**。关键操作是**全规约(all-reduce)**:给定 $N$ 个GPU上各有一个值,计算总和(或平均值)并将结果分发给所有GPU。
|
||||
|
||||
- 朴素的全规约将所有数据发送到一个GPU,求和,然后广播回来。通信量为 $O(N)$,并在根节点造成瓶颈。
|
||||
|
||||
- **环全规约(Ring all-reduce)** 要高效得多。将 $N$ 个GPU排列成一个环。每个GPU将其数据分割成 $N$ 块。在 $N - 1$ 步中,每个GPU向邻居发送一块,并从另一个邻居接收一块,累加部分和。再经过 $N - 1$ 步后,完整的总和传播到所有GPU。每个GPU的总数据传输量:数据大小的 $2(N-1)/N$ 倍,随着 $N$ 的增长趋近于 $2\times$。关键在于,这不随 $N$ 增加,使其带宽最优。
|
||||
|
||||

|
||||
|
||||
- **参数服务器**是一种替代架构,其中专用服务器节点保存模型参数。工作节点计算梯度并将其发送到服务器,服务器更新参数并将其发送回来。这更简单,但可能在服务器处造成通信瓶颈。
|
||||
|
||||
- **NCCL**(NVIDIA集合通信库)是GPU间通信的标准库。它提供了全规约、全收集、广播和其他集合操作的高效实现,自动为网络拓扑选择最佳算法。
|
||||
|
||||
- **缩放定律**描述了模型性能如何随计算量、数据量和模型大小而提升。原始的Kaplan等人(2020)缩放定律发现,损失随每个因素以幂律方式下降:
|
||||
|
||||
$$L(N) \propto N^{-\alpha_N}, \quad L(D) \propto D^{-\alpha_D}, \quad L(C) \propto C^{-\alpha_C}$$
|
||||
|
||||
- 其中 $N$ 是参数数量,$D$ 是数据集大小,$C$ 是计算预算。
|
||||
|
||||
- **Chinchilla缩放定律**(Hoffmann等人,2022)表明大多数模型训练不足:对于给定的计算预算,应该训练一个更小的模型,使用比以前认为的更多的数据。最优比例大约是每参数20个token。一个7B模型应该看到大约140B个token,而不是Llama 1在65B模型上使用的300B个token。这一发现将领域转向了"计算最优"训练。
|
||||
|
||||
- **混合专家(MoE)** 是一种在不按比例增加计算量的情况下扩展模型容量的架构。每个Transformer层不是使用一个前馈网络,而是有 $N$ 个"专家"网络(每个都是一个标准FFN)。一个**门控网络**(路由器)检查每个token并将其发送到top-$K$个专家(通常 $K = 1$ 或 $K = 2$)。
|
||||
|
||||

|
||||
|
||||
- 总参数量要大得多(因为有 $N$ 个专家),但每个token的FLOPs大致保持不变(因为每个token只有 $K$ 个专家激活)。例如,Mixtral 8x7B共有47B个参数,但每次前向传播只用大约13B,以较小模型的代价获得更大模型的性能。
|
||||
|
||||
- MoE带来了挑战。**负载均衡**:如果路由器将大多数token发送到同一个专家,其他专家就被浪费了。辅助损失鼓励均匀路由。**通信**:不同的专家可能位于不同的GPU上,因此路由token需要全对全通信,这很昂贵。
|
||||
|
||||
- **容错**在训练运行持续数周或数月、涉及数千个GPU时至关重要。如果单个GPU失效,你不想丢失所有进度。**检查点**定期将模型权重、优化器状态和训练状态(学习率、步数、数据位置)保存到磁盘。如果发生故障,你可以从最近的检查点重新开始。
|
||||
|
||||
- **梯度检查点**(也称为激活重计算)是一种内存优化,而非容错机制。在前向传播过程中,不是保存所有激活值供反向传播使用,而是只在某些检查点保存激活值。在反向传播过程中,从检查点重新计算缺失的激活值。这以计算换取内存:它使前向传播成本增加约33%,但可以将激活内存减少 $\sqrt{L}$ 倍(其中 $L$ 是层数)。
|
||||
|
||||
- 综合起来,训练前沿模型结合了所有这些技术:BF16混合精度、使用环全规约在数千个GPU上进行数据并行、节点内的张量并行、跨节点的流水线并行、减少内存的梯度检查点、提高参数效率的MoE,以及用于容错的定期检查点。系统工程与算法设计一样具有挑战性。
|
||||
|
||||
- 总结分布式训练工具包:
|
||||
|
||||
| 技术 | 作用 | 权衡 |
|
||||
|---|---|---|
|
||||
| 混合精度 (BF16) | 将激活值/梯度的内存减半 | 轻微数值差异 |
|
||||
| 数据并行 | 在GPU间扩展批量大小 | 梯度同步的通信开销 |
|
||||
| 张量并行 | 在GPU间分割层 | 需要快速互联 |
|
||||
| 流水线并行 | 在GPU间分割模型阶段 | 流水线气泡(计算浪费) |
|
||||
| 梯度累积 | 模拟大批量 | 更慢(多次前向/反向传播) |
|
||||
| 梯度检查点 | 减少激活内存 | 约多33%计算 |
|
||||
| 环全规约 | 高效的梯度平均 | 大模型受限于带宽 |
|
||||
| MoE | 更多容量,相同FLOPs | 负载均衡、路由复杂性 |
|
||||
| 缩放定律 | 指导计算分配 | 经验公式,未必在所有规模都成立 |
|
||||
|
||||
## 编程任务(使用CoLab或笔记本)
|
||||
|
||||
1. 计算Transformer层的FLOPs和内存需求。给定隐藏维度 $d$、序列长度 $n$、批量大小 $B$ 和层数,估计总训练成本。
|
||||
```python
|
||||
import jax.numpy as jnp
|
||||
|
||||
def transformer_layer_flops(d, n, B):
|
||||
"""一个Transformer层前向传播的近似FLOPs。"""
|
||||
# QKV投影:3 * (B * n * d * d) * 2(乘法-加法)
|
||||
qkv_flops = 3 * 2 * B * n * d * d
|
||||
# 注意力:(B * n * n * d) * 2 用于QK^T,(B * n * n * d) * 2 用于attn*V
|
||||
attn_flops = 2 * 2 * B * n * n * d
|
||||
# 输出投影:(B * n * d * d) * 2
|
||||
out_flops = 2 * B * n * d * d
|
||||
# FFN:两层,d->4d 和 4d->d:2 * (B * n * d * 4d) * 2
|
||||
ffn_flops = 2 * 2 * B * n * d * 4 * d
|
||||
return qkv_flops + attn_flops + out_flops + ffn_flops
|
||||
|
||||
def transformer_layer_memory(d, n, B, dtype_bytes=2):
|
||||
"""一个层的近似激活内存(字节)。"""
|
||||
# QKV:3 * B * n * d
|
||||
qkv_mem = 3 * B * n * d * dtype_bytes
|
||||
# 注意力权重:B * heads * n * n(近似 B * n * n * sizeof)
|
||||
attn_mem = B * n * n * dtype_bytes
|
||||
# FFN中间值:B * n * 4d
|
||||
ffn_mem = B * n * 4 * d * dtype_bytes
|
||||
return qkv_mem + attn_mem + ffn_mem
|
||||
|
||||
# 示例:GPT-2规模
|
||||
d, n, B, L = 1024, 1024, 8, 24
|
||||
fwd_flops = transformer_layer_flops(d, n, B)
|
||||
total_flops = 3 * L * fwd_flops # 前向+反向的3倍
|
||||
act_mem = L * transformer_layer_memory(d, n, B)
|
||||
param_count = L * (12 * d * d + 13 * d) # 近似
|
||||
|
||||
print(f"模型:d={d}, n={n}, B={B}, L={L}")
|
||||
print(f"参数:{param_count / 1e6:.0f}M")
|
||||
print(f"每步FLOPs:{total_flops / 1e12:.2f} TFLOPs")
|
||||
print(f"激活内存:{act_mem / 1e9:.2f} GB (BF16)")
|
||||
print(f"参数内存 (FP32):{param_count * 4 / 1e9:.2f} GB")
|
||||
print(f"Adam优化器内存:{param_count * 8 / 1e9:.2f} GB")
|
||||
print(f"总训练内存:{(param_count * 16 + act_mem) / 1e9:.2f} GB")
|
||||
```
|
||||
|
||||
2. 模拟数据并行训练。将数据集分割到多个"虚拟GPU"上,独立计算梯度,平均它们,并验证结果与单GPU训练匹配。
|
||||
```python
|
||||
import jax
|
||||
import jax.numpy as jnp
|
||||
|
||||
# 简单线性模型:y = wx + b
|
||||
key = jax.random.PRNGKey(0)
|
||||
X = jax.random.normal(key, (64, 4))
|
||||
w_true = jnp.array([1.0, -2.0, 3.0, 0.5])
|
||||
y = X @ w_true + 0.1 * jax.random.normal(key, (64,))
|
||||
|
||||
def loss_fn(w, X, y):
|
||||
return jnp.mean((X @ w - y) ** 2)
|
||||
|
||||
grad_fn = jax.grad(loss_fn)
|
||||
|
||||
# 单GPU:全批量梯度
|
||||
w = jnp.zeros(4)
|
||||
grad_single = grad_fn(w, X, y)
|
||||
|
||||
# 数据并行:分割到4个"GPU"上
|
||||
n_gpus = 4
|
||||
chunk_size = len(X) // n_gpus
|
||||
grads = []
|
||||
for i in range(n_gpus):
|
||||
X_chunk = X[i*chunk_size:(i+1)*chunk_size]
|
||||
y_chunk = y[i*chunk_size:(i+1)*chunk_size]
|
||||
grads.append(grad_fn(w, X_chunk, y_chunk))
|
||||
|
||||
# 全规约:平均梯度
|
||||
grad_parallel = jnp.mean(jnp.stack(grads), axis=0)
|
||||
|
||||
print("单GPU梯度:", grad_single)
|
||||
print("数据并行梯度(平均):", grad_parallel)
|
||||
print(f"匹配:{jnp.allclose(grad_single, grad_parallel, atol=1e-5)}")
|
||||
|
||||
# 训练两者并比较
|
||||
w_single, w_parallel = jnp.zeros(4), jnp.zeros(4)
|
||||
lr = 0.1
|
||||
for step in range(100):
|
||||
w_single = w_single - lr * grad_fn(w_single, X, y)
|
||||
|
||||
grads = [grad_fn(w_parallel, X[i*chunk_size:(i+1)*chunk_size],
|
||||
y[i*chunk_size:(i+1)*chunk_size]) for i in range(n_gpus)]
|
||||
avg_grad = jnp.mean(jnp.stack(grads), axis=0)
|
||||
w_parallel = w_parallel - lr * avg_grad
|
||||
|
||||
print(f"\n100步之后:")
|
||||
print(f"单GPU权重:{w_single}")
|
||||
print(f"数据并行权重:{w_parallel}")
|
||||
print(f"最大差异:{jnp.max(jnp.abs(w_single - w_parallel)):.2e}")
|
||||
```
|
||||
|
||||
3. 实现一个简单的混合专家层。创建一个门控网络,将token路由到top-K个专家并组合它们的输出。
|
||||
```python
|
||||
import jax
|
||||
import jax.numpy as jnp
|
||||
|
||||
def expert_fn(x, W1, b1, W2, b2):
|
||||
"""简单的2层FFN专家。"""
|
||||
h = jnp.maximum(0, x @ W1 + b1) # ReLU
|
||||
return h @ W2 + b2
|
||||
|
||||
def moe_layer(x, gate_W, experts_params, top_k=2):
|
||||
"""
|
||||
MoE前向传播。
|
||||
x: (batch, d_model)
|
||||
gate_W: (d_model, n_experts)
|
||||
experts_params: 每个专家的 (W1, b1, W2, b2) 列表
|
||||
"""
|
||||
n_experts = len(experts_params)
|
||||
|
||||
# 门控:计算路由分数
|
||||
gate_logits = x @ gate_W # (batch, n_experts)
|
||||
gate_probs = jax.nn.softmax(gate_logits, axis=-1)
|
||||
|
||||
# Top-K选择
|
||||
top_k_indices = jnp.argsort(-gate_probs, axis=-1)[:, :top_k]
|
||||
top_k_probs = jnp.take_along_axis(gate_probs, top_k_indices, axis=-1)
|
||||
# 重新归一化
|
||||
top_k_probs = top_k_probs / jnp.sum(top_k_probs, axis=-1, keepdims=True)
|
||||
|
||||
# 计算专家输出(简化:运行所有专家,稍后掩码)
|
||||
expert_outputs = jnp.stack([
|
||||
expert_fn(x, *experts_params[i]) for i in range(n_experts)
|
||||
], axis=1) # (batch, n_experts, d_model)
|
||||
|
||||
# 收集top-K专家输出并加权
|
||||
batch_idx = jnp.arange(x.shape[0])[:, None]
|
||||
selected_outputs = expert_outputs[batch_idx, top_k_indices] # (batch, top_k, d_model)
|
||||
output = jnp.sum(selected_outputs * top_k_probs[:, :, None], axis=1)
|
||||
|
||||
return output, gate_probs
|
||||
|
||||
# 设置
|
||||
key = jax.random.PRNGKey(42)
|
||||
batch, d_model, d_ff, n_experts = 8, 16, 32, 4
|
||||
|
||||
# 初始化专家
|
||||
experts_params = []
|
||||
for i in range(n_experts):
|
||||
k1, k2, key = jax.random.split(key, 3)[0], jax.random.split(key, 3)[1], jax.random.split(key, 3)[2]
|
||||
experts_params.append((
|
||||
jax.random.normal(k1, (d_model, d_ff)) * 0.1,
|
||||
jnp.zeros(d_ff),
|
||||
jax.random.normal(k2, (d_ff, d_model)) * 0.1,
|
||||
jnp.zeros(d_model),
|
||||
))
|
||||
|
||||
key, subkey = jax.random.split(key)
|
||||
gate_W = jax.random.normal(subkey, (d_model, n_experts)) * 0.1
|
||||
x = jax.random.normal(key, (batch, d_model))
|
||||
|
||||
output, gate_probs = moe_layer(x, gate_W, experts_params, top_k=2)
|
||||
|
||||
print(f"输入形状:{x.shape}")
|
||||
print(f"输出形状:{output.shape}")
|
||||
print(f"门控概率(第一个样本):{gate_probs[0]}")
|
||||
print(f"专家使用率(批量平均):")
|
||||
for i in range(n_experts):
|
||||
usage = jnp.mean(gate_probs[:, i])
|
||||
print(f" 专家 {i}: {usage:.3f}")
|
||||
```
|
||||
Reference in New Issue
Block a user