第 3 节 · 约 35 分钟
每个参数只问:我动一点,损失怎么变。
每一跳只做一件事:上游梯度乘局部导数。损失有限而某个 .grad 为 0,先指出卡在哪。
先讲
最小网络:线性 → ReLU → 线性 → 半平方误差。z = w₁x₁ + w₂x₂ + b,h = ReLU(z),ŷ = v h + c,L = ½(ŷ − y)²。半平方前面的 ½ 是为了让 ∂L/∂ŷ = ŷ−y。写成 (ŷ−y)²,所有梯度都会翻倍。取 x₁=1,x₂=2,y=1;w₁=0.1,w₂=0.2,b=0,v=0.5,c=0.1。前向五个数:z=0.5,h=0.5,ŷ=0.35,ŷ−y=−0.65,L=0.21125。
对 w₁ 这条路:∂L/∂w₁ = (∂L/∂ŷ)(∂ŷ/∂h)(∂h/∂z)(∂z/∂w₁)。同一条串行路上,局部导数相乘。一个变量走两条路,贡献要相加。代码里是 x.grad += upstream * local。写成 *=,一条路是 0 时两条都没了。残差 h = x + relu(x) 对 x 的梯度应是 1 + relu'(x)。分叉处上游分别是 g1=−0.325、g2=0,该写相加。五个参数必须用同一轮前向得到的旧梯度同时更新。学习率 0.1,w₁ 变成 0.1325,损失大约 0.108。先改 w₁ 再拿新值去算 v 的步子,账会对不上。
ReLU 在 z≤0 时局部导数是 0。把偏置改成 −1,z 变负,h=0,损失仍有限,c.grad 还有数,w₁.grad=0。这一跳被掐断。换几个样本、或把预激活往正处挪一点,梯度会回来。若所有 .grad 都是 None,更像没调用 backward(),或张量没 requires_grad。
打印损失、h、每个 .grad 并排看。哪一项突然变 None,从那一项往输出端查调用。哪一项变成 0 而邻居还有数,查这一跳的局部导数。中心差分抽查一个参数:[L(θ+ε)−L(θ−ε)]/(2ε),ε 从 1e-4 试起,抽查 w₁ 应落在 −0.325 附近。符号反了先查损失有没有写反。差一个固定倍数先查 mean 和 sum。
backward 只填梯度,不改参数。整数张量也不能回传。叶子没开 requires_grad,梯度也全是空。五个参数的梯度按 w₁,w₂,b,v,c 排是 [−0.325, −0.65, −0.325, −0.325, −0.65]。梯度为负,减小损失要略增大该参数。b 改成 −1 时 z=−0.5,L=0.405,c.grad=−0.9:反向跑过了,ReLU 把上游乘了 0。训练循环里还要 zero_grad 再 step,下一节才展开。这一节先能指出:0 是路径断了,还是图根本没走回来。
- z = 0.1×1 + 0.2×2 + 0 = 0.5
- h = ReLU(0.5) = 0.5
- ŷ = 0.5×0.5 + 0.1 = 0.35
- ŷ−y = −0.65,L = ½(−0.65)² = 0.21125
- 梯度 [w₁,w₂,b,v,c] = [−0.325, −0.65, −0.325, −0.325, −0.65]
你来做
三次日志。loss 都有限。点出卡在哪。
loss=0.405 z=−0.5 h=0 c.grad=−0.9 w1.grad=0.0
loss=0.21125 w1.grad=None v.grad=None c.grad=None
h 进两个头。你写 h.grad = g1 * g2。其中 g2=0,于是 w1.grad=0。
留下:0 是路径断了,还是图没走回来。下一节三拍循环和三条曲线。