第 2 课:自动微分 Autograd —— 让模型学会「自我修正」
第 2 课:自动微分 Autograd —— 让模型学会"自我修正"
代码位置:src/tensor.rs(grad / parents / 各运算的反向闭包)、src/autograd.rs(backward 拓扑排序) 演示入口:src/main.rs
1. 本课要搞懂的问题
- 什么是梯度?为什么"梯度下降"能让模型学习?
- 链式法则是怎么回事?为什么反向传播比正向求导更高效?
- 怎么用 Rust 把"求导"自动化,让代码自动算出每个参数的梯度?
2. 从"拟合直线"说起
上一课我们实现了矩阵乘法。现在想象一个最简单的"学习"任务:
已知一些点 (x, y),它们近似满足 y = 2x + 1,但参数 w、b 未知。怎么让程序自己找到 w=2、b=1?
思路:先随便给个初始值(比如 w=0, b=0),然后:
- 算出预测值和真实值的差距 → 损失 loss
- 问:如果我调大一点 w,loss 会变大还是变小? → 这个方向就叫梯度
- 沿着"让 loss 变小的方向"微调 w、b
- 重复直到 loss 足够小
第 2 步里的"损失随参数的变化率",数学上叫偏导数。所有参数偏导数的集合叫梯度。
3. 梯度下降公式
对参数 θ,更新规则只有一行:
θ = θ - 学习率 × ∂loss/∂θ
∂loss/∂θ > 0:θ 增大 loss 也增大 → 所以减去它,θ 变小∂loss/∂θ < 0:θ 增大 loss 减小 → 减去它,θ 变大- 学习率(learning rate)控制每次迈多大步子,太小太慢、太大震荡
我们采用的损失是 MSE(均方误差):
loss = Σ (pred - y)² # pred = x·w + b
4. 链式法则:复合函数求导的关键
问题来了:loss 是复合函数,loss = g(pred),而 pred = f(w)。
要求 ∂loss/∂w,不能直接求,需要一层层拆开:
链式法则:如果 z = g(f(x)),那么
dz/dx = (dg/df) × (df/dx)
类比"流水线":损失一路流经 pred 再到 w,每一段的"变化率"相乘就是总的"变化率"。
例:z = x*y + w,各梯度是:
∂z/∂x = y (对 x 求导时把 y 当常数)
∂z/∂y = x (对 y 求导时把 x 当常数)
∂z/∂w = 1 (加法求导是 1)
程序输出正是 dz/dx=3, dz/dy=2, dz/dw=1 ✓
5. 反向传播:把链式法则批量自动化
手工对每个参数求导太累,模型参数动辄百万个。怎么办?
核心思想:前向传播时,把每一步运算记录下来,形成一张计算图;反向传播时,从 loss 出发,把梯度沿图逐层往回传,每层的梯度都复用上一层的结果。
前向(边算边记): x ──┐
mul ──► t ──┐
y ──┘ add ──► z = 7
w ───────────┘
反向(梯度倒流): x ◄── grad=3
t ◄── grad=1 ──► w ◄── grad=1
y ◄── grad=2
关键好处:每个中间结果只算一次梯度,总计算量与前向相当,而不是对每个参数各算一次前向(那样要 O(参数个数) 倍计算)。这就是反向传播比"数值微分"快的原因。
6. Rust 实现:每个张量背一个"反向函数"
我们把张量升级成计算图节点:
pub struct Tensor {
data: Rc<RefCell<Vec<f32>>>, // 数值(共享可变)
shape: Vec<usize>,
grad: Rc<RefCell<Vec<f32>>>, // 梯度,初始全 0,反向时累加
requires_grad: bool, // 这个张量是否需要梯度(参数=true,数据=false)
parents: Vec<Tensor>, // 我是由哪些张量算出来的
backward: Option<BackwardFn>, // 我的反向函数:把梯度传给我的父母
}
用到的 Rust 核心特性:
Rc<T>(引用计数):计算图里一个节点可能被多个运算共享,Rc 让多份句柄指向同一份数据RefCell<T>(内部可变性):让"共享数据"同时可以"可变借用",这是 Rust 所有权模型下实现可变共享的标准手法- 每定义一个运算,就写两个东西:前向(算结果)+ 反向(算梯度传给输入)
7. 各运算的反向规则(背下这几条)
| 前向 | 反向(g 是输出的梯度) |
|---|---|
c = a + b |
∂a += g,∂b += g |
c = a - b |
∂a += g,∂b -= g |
c = a * b |
∂a += g*b,∂b += g*a |
c = a * s |
∂a += g*s |
c = a @ b |
∂a += g @ bᵀ,∂b += aᵀ @ g |
c = sum(a) |
∂a += g(每个元素都加同一个 g) |
以乘法为例看代码(注意累加 +=,因为一个张量可能被多条路径使用,梯度要累加):
pub fn mul(&self, other: &Tensor) -> Tensor {
// —— 前向 ——
let data = /* self * other 逐元素 */;
let mut result = Tensor::new(data, shape, self.requires_grad || other.requires_grad);
if result.requires_grad {
// 捕获需要的句柄
let rg = result.grad.clone(); // 读:自己的梯度
let sg = self.grad.clone(); // 写:给 self 累加梯度
let og = other.grad.clone();
let sd = self.data.clone(); // 读:self 的数值(计算 g*b 需要)
let od = other.data.clone();
result.backward = Some(Rc::new(move || {
let g = rg.borrow(); // 拿到"传到我这的梯度"
let sd_b = sd.borrow();
let od_b = od.borrow();
let mut sgm = sg.borrow_mut();
let mut ogm = og.borrow_mut();
for i in 0..g.len() {
sgm[i] += g[i] * od_b[i]; // ∂a = g * b
ogm[i] += g[i] * sd_b[i]; // ∂b = g * a
}
}));
}
result
}
8. 反向传播主流程:拓扑排序
图是有向无环图(DAG)。要保证"子节点先算,父节点后算",需要拓扑排序:
pub fn backward(&self) {
// 1. loss 自己的梯度 = 1(d loss / d loss = 1)
self.grad.borrow_mut()[0] = 1.0;
// 2. 迭代式 DFS 收集拓扑序:子在前,父在后(避免递归栈溢出,见第 4 课 6.2 节)
// 用显式栈模拟递归,三色标记法保证每个节点只入队一次
// 3. 逆序遍历,逐个调用 backward 闭包:梯度从子流向父
for t in order.iter().rev() { t.backward() }
}
用 Rc::as_ptr(&t.grad) 作为节点的唯一编号做去重(一个节点被多个运算共享时不能重复入队)。
9. 训练循环:四步舞
一次"学习"包含四步(也是所有深度学习框架训练循环的骨架):
// 1. 前向:算预测和损失
let loss = x.matmul(&w).sub(&y).mul(...).sum();
// 2. 反向:自动算每个参数的梯度
loss.backward();
// 3. 更新:w -= lr * 梯度(这里手动实现,展示底层)
w.set_data(vec![w.data()[0] - lr * w.grad()[0], ...]);
// 4. 清零梯度:否则下一轮会把旧梯度累加上去
w.zero_grad();
运行 cargo run 可以看到损失从 285 一路降到 0,w→2.00、b→1.00,模型"学会"了 y=2x+1。
10. 一个重要的细节:梯度累加
为什么用 += 而不是 =?因为一个张量可能被多条路径使用。
例:z = x*x。求梯度时 x 有两条路径(z 对第一个 x 的导数是第二个 x,对第二个 x 的导数是第一个 x),必须相加:∂x = g*x + g*x = 2gx。
代码里我们用 Rc::ptr_eq(&sg, &og) 判断两个父节点是不是同一个张量,是的话就合并累加。这也是测试 test_grad_accumulation_and_zero 验证的内容。
11. 运行与测试
cargo test # 8 个测试全部通过
cargo run # 演示:链式法则验证 + 线性回归训练
12. 动手练习
- 给张量加一个
pow(n)运算(c = a^n),推导并实现它的反向传播(提示:∂a = g * n * a^(n-1))。 - 思考:为什么 loss 必须是标量(0 维)才能调用
backward()?如果不是标量会怎样? - 把学习率从 0.001 改成 1.0,观察 loss 曲线,解释发生了什么(振荡/发散)。
- 这是本项目的"心脏"。试着不看代码,自己画出
loss = (x·w - y)²的计算图并手算∂loss/∂w。
13. 本课总结
- 梯度 = 损失对参数的"变化方向",梯度下降就是沿着反方向走
- 链式法则 + 计算图 = 反向传播,一次性高效算出所有梯度
- Rust 用
Rc<RefCell>+ 闭包实现了可自动求导的张量 - 亲手跑通了第一个"会学习"的程序:线性回归
- 下一课:补齐张量运算(广播、softmax 等),为神经网络打基础