如何理解反向传播中的computation graph?求解梯度计算疑惑
计算图反向传播:偏导数与链式法则实例解析
先明确你例子中的计算逻辑:
- 前向传播:
y_pred = x1*w1 + x2*w2 + b = 7,真实标签y_true=4,损失采用L1绝对误差:Loss = |y_pred - y_true| = 3
你困惑的核心是把「损失值」和「梯度(偏导数)」搞混了,下面一步步拆解数学逻辑:
1. 损失函数的导数:梯度不是损失值
梯度描述的是损失函数随某个变量变化的速率,而非损失本身的大小。
对于L1损失Loss = |y_pred - y_true|:
- 当
y_pred > y_true时,损失函数的斜率为1(即dLoss/dy_pred = 1,因为y_pred每增加1,Loss增加1); - 当
y_pred < y_true时,斜率为-1; - 你这里
y_pred=7 > 4,所以dLoss/dy_pred = 1——这是反向传播的起点。
2. 用链式法则逐层计算梯度
反向传播的本质是链式法则:对于嵌套函数Loss(f(g(x))),dLoss/dx = dLoss/df * df/dg * dg/dx。我们从Loss开始,往前逐层推导每个变量的梯度:
例1:计算grad(Loss, x2)
变量关系链:Loss → y_pred → x2*w2 → x2
- 第一步:
dLoss/dy_pred = 1(已确定) - 第二步:
y_pred是x2*w2加上其他项,所以dy_pred/d(x2*w2) = 1(加法的导数是1) - 第三步:
d(x2*w2)/dx2 = w2(乘法偏导,w2视为常数),假设你的例子中w2=1,则这一项为1 - 链式相乘:
dLoss/dx2 = 1 * 1 * 1 = 1——这就是你看到的结果
例2:解释1*1*2的来源
假设你要计算grad(Loss, x1),且例子中w1=2:
- 变量关系链:
Loss → y_pred → x1*w1 → x1 - 链式相乘:
dLoss/dx1 = dLoss/dy_pred * dy_pred/d(x1*w1) * d(x1*w1)/dx1 = 1 * 1 * 2 = 2
这里的2就是权重w1的值,对应最后一步的偏导结果。
其他变量的梯度推导
grad(Loss, w2):dLoss/dw2 = dLoss/dy_pred * dy_pred/d(x2*w2) * d(x2*w2)/dw2 = 1 * 1 * x2 = 7(x2的值为7)grad(Loss, b):dLoss/db = dLoss/dy_pred * dy_pred/db = 1 * 1 =1(b的系数是1,偏导为1)
3. 核心误区纠正
你之前误以为grad(Loss, x2)是3,是把「损失的大小」和「损失随变量变化的速率」混淆了:
- Loss=3是当前的损失值,而梯度1的意思是:x2每增加1,Loss会增加1;x2每减少1,Loss会减少1。这正是梯度下降的依据——我们沿着梯度的反方向调整x2(或参数),就能让Loss降低。
内容的提问来源于stack exchange,提问作者Tyrone_Slothrop
相关产品推荐
相关产品推荐

