基于梯度的机器学习 基于梯度的学习通过在损失曲面上顺着坡度一步步走,来优化模型参数。本文件涵盖线性回归、逻辑回归、softmax 分类、梯度下降的变体、正则化(L1/L2)以及偏差-方差权衡。 第 1 节里的经典方法用的是巧妙的启发式或解析解。本节讲的是靠"跟着梯度走"来学习的算法——在损失曲面上小步下山,直到找到好的参数。基于梯度的学习是从线性回归一直到最大神经网络背后共同的引擎。 线性回归(linear regression)是最简单的基于梯度的模型,而且它还有解析解,是绝佳的起点。
基于梯度的学习通过在损失曲面上顺着坡度一步步走,来优化模型参数。本文件涵盖线性回归、逻辑回归、softmax 分类、梯度下降的变体、正则化(L1/L2)以及偏差-方差权衡。
第 1 节里的经典方法用的是巧妙的启发式或解析解。本节讲的是靠"跟着梯度走"来学习的算法——在损失曲面上小步下山,直到找到好的参数。基于梯度的学习是从线性回归一直到最大神经网络背后共同的引擎。
**线性回归(linear regression)**是最简单的基于梯度的模型,而且它还有解析解,是绝佳的起点。模型就是一条直线(高维下是一个超平面):
用矩阵记号(来自第 2 章),把所有训练输入按行堆成矩阵 X,并通过追加一列 1 把偏置吸收进 w,就得到 \hat{y} = Xw。
目标是最小化均方误差(mean squared error, MSE),也就是预测值与真实值之差的平方的平均:
这里直接用到了第 2 章的矩阵求逆。表达式 X^T X 是一个 d \times d 的矩阵(d 是特征数),而 X^T y 是一个 d 维向量。正规方程一步就能给出精确的最优权重。
正规方程什么时候会失效?当 X^T X 奇异(不可逆)时,这发生在特征线性相关或者特征数多于样本数(d > n)的时候。这些情况下你需要正则化(后文介绍)或梯度下降。
**逻辑回归(logistic regression)**把线性模型改造用于二分类。我们不再预测连续值,而是想要一个 0 到 1 之间的概率。sigmoid 函数把任何实数压进这个范围:
sigmoid 有不错的性质:\sigma(0) = 0.5,z \to \infty 时 \sigma(z) \to 1,z \to -\infty 时 \sigma(z) \to 0,而且它的导数形式很优雅:\sigma'(z) = \sigma(z)(1 - \sigma(z))。
逻辑回归的损失函数是二元交叉熵(binary cross-entropy, BCE),它直接来自伯努利似然(第 5 章):
当真实标签为 1 时,只有第一项起作用,它会惩罚低的预测;当真实标签为 0 时,只有第二项起作用,它会惩罚高的预测。对数让"自信地犯错"的代价极其陡峭:真实标签为 1 时预测 0.01 的代价,远大于预测 0.4。
与线性回归的 MSE 不同,最小化 BCE 的权重没有解析解。我们需要一种迭代方法:梯度下降(gradient descent)。
梯度下降背后的直觉很简单:想象你在大雾中站在一片起伏的山地(损失曲面)上。你看不到全局最低点,但你能感觉到脚下的坡度。你往下走一步,再感受坡度,如此反复,最终会走到一个谷底。
梯度 \frac{\partial \mathcal{L}}{\partial w} 是一个指向最陡上升方向的向量。我们减去它,因为我们要往下走。这正是第 3 章的链式法则用在了损失函数上。
**批量梯度下降(batch gradient descent)**每一步都用整个训练集来算梯度。这给出的是精确梯度,但在 n 很大时代价高昂。
**随机梯度下降(stochastic gradient descent, SGD)**每步只用一个随机样本。梯度很嘈杂(它用单个样本来估计真实梯度),但每一步都极快。这种噪声其实还能帮助跳出浅浅的局部极小。
**小批量梯度下降(mini-batch gradient descent)**折中处理:每步用 B 个样本组成的一个小批量(通常是 32、64 或 256)。它在计算效率(对批量做向量化运算)和梯度质量之间取得了平衡。几乎所有深度学习用的都是小批量 SGD。
**反向传播(backpropagation)**是我们真正在神经网络这种参数众多的模型中计算梯度的方法。它就是把第 3 章的链式法则系统地应用在一个计算图上。
任何模型都可以表示成一个有向无环的计算图:输入流入,乘上权重,加在一起,穿过非线性函数,最终产生一个损失值。**前向传播(forward pass)**让数据从输入流向输出,算出输出(和损失)。
反向传播则让梯度逆向流动。从损失出发,你用链式法则在每个节点上算出损失相对于每个中间值的变化率。如果 L 依赖于 z,而 z 依赖于 w,那么:
每个节点只需要知道自己的局部导数和从上方流进来的梯度。这使反向传播模块化又高效:代价大约是前向传播的两倍(一遍前向、一遍反向)。
朴素的 SGD 有个毛病:它在曲率陡的方向上震荡,在平坦的方向上进展缓慢。**优化器(optimiser)**通过根据梯度历史调整步长来改进这一点。
**带动量的 SGD(SGD with momentum)**保留过去梯度的滑动平均(指数移动平均,来自第 4 章)。这能平滑震荡、加速沿一致方向的进展:
把它想象成一个滚下山坡的球:动量让它在一致方向上加速,并抑制左右的抖动。典型取值是 \beta = 0.9。
**Nesterov 加速梯度(Nesterov Accelerated Gradient, NAG)**是个巧妙的小改动:与其在当前位置算梯度,不如在"前瞻"位置 w - \eta \beta v_{t-1} 处算。这一修正步减少了冲过头:
问题在于:G_t 只增不减,所以有效学习率单调下降,最终小到什么都学不到。
RMSprop 用梯度的平方的指数移动平均替代累加和来修复这一点,这样近期的梯度比老旧的更重要:
默认超参数(\beta_1 = 0.9、\beta_2 = 0.999、\epsilon = 10^{-8})在很广的问题范围内都好用,这也是 Adam 成为大多数深度学习默认优化器的原因。
AdamW 把权重衰减与梯度更新解耦。标准的 L2 正则化和权重衰减对 SGD 是等价的,但对 Adam 并不等价。AdamW 把权重衰减直接作用在参数上,而不是把 \lambda w 加到梯度里。这带来了更好的泛化,如今是 Transformer 训练的标准做法:
除了 MSE 和 BCE,还有几种常用的损失函数。
平均绝对误差(Mean Absolute Error, MAE),也叫 L1 损失,取绝对差值的平均:\frac{1}{n}\sum|y_i - \hat{y}_i|。它对异常值比 MSE 更鲁棒,因为它不会把大误差平方放大。
Huber 损失结合了两者的优点:对小误差表现得像 MSE(平滑、易优化),对大误差表现得像 MAE(对异常值鲁棒)。它有一个阈值 \delta 控制过渡点。
**类别交叉熵(categorical cross-entropy, CCE)**把 BCE 推广到多类。如果 \hat{y}_k 是对类别 k 的预测概率,而真实类别是 c:
这就是正确类别的负对数概率。最小化交叉熵等价于最大化似然,这又连回第 5 章的信息论:交叉熵衡量的是,当你用预测分布而不是真实分布时,平均多花了多少比特。
**合页损失(hinge loss)**用于 SVM:\mathcal{L} = \max(0, 1 - y \cdot f(x))。它只惩罚落在间隔错误一侧或间隔之内的预测。一旦一个点被以足够置信度正确分类,损失就是零。
**正则化(regularisation)**通过给复杂模型加惩罚来防止过拟合。正则化后的损失是:
L2 正则化(Ridge、权重衰减)惩罚权重平方和:R(w) = \|w\|^2 = \sum w_i^2。它不鼓励任何单个权重变得过大,实际上是把所有权重都朝零收缩,但很少让它们恰好为零。
L1 正则化(Lasso)惩罚权重绝对值之和:R(w) = \|w\|_1 = \sum |w_i|。它鼓励稀疏性,把许多权重直接压到零,从而自动完成特征选择。
**弹性网(Elastic Net)**把两者结合:R(w) = \alpha \|w\|_1 + (1 - \alpha) \|w\|^2,调和了稀疏与收缩。
这里有一个漂亮的贝叶斯诠释(来自第 5 章)。L2 正则化等价于在权重上放一个高斯先验,再求 MAP 估计。L1 正则化对应于拉普拉斯先验。正则化强度 \lambda 控制了你有多信任先验相对于数据。
评估指标告诉你模型到底行不行。回归用 MSE 和 MAE 是标准做法。分类则更微妙。
**混淆矩阵(confusion matrix)**是二分类的一张四格表:
准确率(accuracy) = \frac{TP + TN}{TP + TN + FP + FN},在类别不平衡时会有误导。如果 99% 的邮件都不是垃圾邮件,一个永远预测"不是垃圾"的模型有 99% 的准确率,却毫无用处。
精确率(precision) = \frac{TP}{TP + FP} 回答:所有预测为正的里,有多少真的是正?高精确率意味着误报少。
召回率(recall)(敏感度) = \frac{TP}{TP + FN} 回答:所有真正为正的里,你抓住了多少?高召回率意味着漏报少。
F1 分数 = \frac{2 \cdot \text{precision} \cdot \text{recall}}{\text{precision} + \text{recall}} 是精确率和召回率的调和平均,兼顾两者。
ROC 曲线把分类阈值从 0 扫到 1,画出真阳性率(召回率)随假阳性率(\frac{FP}{FP + TN})的变化。完美的分类器会紧贴左上角。AUC(ROC 曲线下面积)用一个数概括表现:1.0 是完美,0.5 是瞎猜。
**交叉验证(cross-validation)**给出更可靠的泛化性能估计。在 k 折交叉验证中,你把数据分成 k 份,在 k-1 份上训练、在剩下那份上测试,然后轮转。所有 k 份测试性能的平均就是你的估计。这样所有数据都被同时用于训练和测试(只是从不同时),在数据稀缺时尤其有价值。
偏差-方差权衡(bias-variance tradeoff)(来自第 4 章)是机器学习里根本性的张力。一个模型的期望误差可以分解为:
偏差是错误假设带来的系统误差(比如用直线去拟合弯曲的数据)。方差是对训练数据波动的敏感度(比如用 20 次多项式去拟合噪声)。简单模型偏差高、方差低;复杂模型偏差低、方差高。甜点在总误差最小处。
**学习率调度(learning rate scheduling)**在训练过程中调整 \eta。常见策略:
**超参数调优(hyperparameter tuning)**是为学习率、批量大小、正则化强度等不被梯度下降学习的设定寻找好值的过程。常见做法:
**无调度学习(schedule-free learning)**干脆完全免去学习率调度的需要。它不按固定曲线衰减 \eta,而是维护两条序列:一条缓慢移动的迭代点平均 z_t(收敛到最优)和一条快速的探索性迭代点 y_t(梯度在它上面求值)。最终输出是那条平均序列,可以证明它的收敛速率追平事后看来最优的调度。这把"调度"整个从超参数里移除了——你只需设定基础学习率,剩下的交给优化器。SGD 和 Adam 的无调度变体都已显示出能追平甚至超过精心调过调度的对应版本。
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()
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()
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()
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}")