坐下 3 / 6
每个参数只问:我动一点,损失怎么变。
网络预测 0.35,答案是 1.00。里面五个参数。该改哪一个、朝哪边、改多少?反向传播给的是一份逐参数的账,不是一句「模型错了」。
一条链,五个数
最小网络:线性 → ReLU → 线性 → 半平方误差。一个样本。
z = w₁x₁ + w₂x₂ + b,h = ReLU(z),ŷ = v h + c,L = ½(ŷ − y)²。
取 x₁=1,x₂=2,y=1;参数 w₁=0.1,w₂=0.2,b=0,v=0.5,c=0.1。前向五个数:
z = 0.1×1 + 0.2×2 + 0 = 0.5。ReLU 原样放出。再 0.5×0.5 + 0.1 = 0.35。½(0.35−1)² = 0.21125。
退回去:上游 × 局部
每跨一个节点只做一件事。对 w₁ 这条路:∂L/∂w₁ = (∂L/∂ŷ)(∂ŷ/∂h)(∂h/∂z)(∂z/∂w₁)。半平方误差下 ∂L/∂ŷ = ŷ−y = −0.65。∂ŷ/∂h = v = 0.5,所以 ∂L/∂h = −0.325。z>0,ReLU 的局部导数是 1,∂L/∂z 仍是 −0.325。再乘 x₁=1,w₁ 的梯度是 −0.325。w₂ 同路上游,只是 x₂=2,梯度 −0.65。v 的梯度是 (ŷ−y)·h = −0.325,c 的梯度就是 ŷ−y = −0.65。五个参数的梯度按 w₁,w₂,b,v,c 排是 [−0.325, −0.65, −0.325, −0.325, −0.65]。梯度为负,减小损失要略增大该参数。
学习率 0.1,五个参数用这一轮的旧梯度同时减:w₁ 变成 0.1325,w₂ 变成 0.265。再前向,损失大约 0.108。步子小且梯度对,同一批数据的损失应当下降。对不上,先查符号和哪一步的局部导数。w₂ 的梯度恰好是 w₁ 的两倍,只因为这个样本里 x₂=2x₁,不是 w₂ 永远更重要。换一个 x,比例就变。
同一条串行路上,局部导数相乘。一个变量走两条路,贡献要相加。代码里是 x.grad += upstream * local。写成 *=,一条路是 0 时两条都没了。残差 h = x + relu(x) 对 x 的梯度应是 1 + relu'(x);漏掉那个 1,捷径等于没修。
五个参数必须用同一轮前向得到的旧梯度同时更新。先改 w₁ 再拿新 w₁ 去算 v 的步子,账会对不上。手算之后用中心差分抽查一个参数:[L(θ+ε)−L(θ−ε)]/(2ε),ε 从 1e-4 试起。相对误差到 1e-6 量级通常能信;符号反了先查损失有没有写反,差一个固定倍数先查 mean 和 sum。
损失有数,某个 .grad 却是 0
ReLU 在 z≤0 时局部导数是 0。把上面例子的 b 改成 −1,z = 0.1+0.4−1 = −0.5,h=0,ŷ=c=0.1,L=½(0.1−1)²=0.405,损失仍有限。c.grad=−0.9,反向跑过了;w₁.grad=0,这一跳被 ReLU 掐断。换几个样本、或把预激活往正处挪一点,梯度会回来。若所有 .grad 都是 None,更像没调用 backward(),或张量没 requires_grad。整数张量也不能回传。
backward 只填梯度,不改参数。训练循环里还要 zero_grad 再 step,下一页才展开。这一页先能指出:0 是路径断了,还是图根本没走回来。打印损失、h、每个 .grad 并排看,比单盯一个权重快得多。哪一项突然变 None,从那一项往输出端查调用;哪一项变成 0 而邻居还有数,查这一跳的局部导数。
点出卡在哪个节点
三次日志。loss 都有限。判断:死 ReLU、没 backward、还是分叉处乘错了。
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,另一个头明明还在用 h。
自己对过
第一则:c.grad 有数,反向跑过了,ReLU 把上游乘了 0。第二则:全是 None,没回传。第三则:分叉该加。