PyTorch反向传播中max与absolute value操作的梯度下降实现机制咨询
PyTorch中max与绝对值操作的反向传播梯度计算
下面详细拆解这两种操作的反向传播逻辑:
1. Max(最大值)操作的反向传播
正向计算逻辑
max操作分为两种常见场景:
- 对单个张量全局取最大值:输出张量中的最大标量值
- 沿指定维度取最大值(如
torch.max(input, dim=d)):输出该维度上的最大值张量,同时返回最大值对应的索引张量——这个索引是反向传播的核心依据。
反向梯度计算
根据链式法则,max操作的梯度仅流向产生最大值的输入元素,其他元素的梯度为0:
- 全局取max:假设输入为张量
x,输出y=max(x),则y对x的梯度中,只有与y值相等的位置会分配梯度,其余为0。 - 沿维度取max:对于每个维度切片,只有该切片中最大值对应的索引位置会接收上游梯度,其余位置梯度为0。
PyTorch的具体实现
- 正向计算时,PyTorch会缓存最大值的索引信息;反向传播时,直接根据索引将上游梯度赋值到对应位置,其余位置填充0。
- 当存在多个元素等于最大值时(如
x=[2.,2.]),PyTorch会将上游梯度平均分配给所有最大值元素——这是子梯度的一种合理选择(因为max函数在多最大值点处不可导,子梯度可取任意和为上游梯度的非负权重)。
示例代码验证:
import torch x = torch.tensor([2., 2.], requires_grad=True) y = torch.max(x) y.backward() print(x.grad) # 输出: tensor([0.5000, 0.5000])
2. Absolute Value(绝对值)操作的反向传播
正向计算逻辑
绝对值操作的正向逻辑很直接:y = |x|,等价于y = x if x >= 0 else -x。
反向梯度计算
绝对值函数的导数在大部分点可定义:
- 当
x > 0时,导数为1; - 当
x < 0时,导数为-1; - 当
x = 0时,函数不可导,PyTorch采用子梯度策略,将该点的梯度设为0(这是子梯度的常见选择之一)。
PyTorch的具体实现
反向传播时,PyTorch会根据输入x的符号直接生成梯度:
- 遍历输入张量,
x > 0的位置梯度设为1,x < 0的位置设为-1,x = 0的位置设为0。
示例代码验证:
import torch x = torch.tensor([-1., 0., 1.], requires_grad=True) y = torch.abs(x) y.sum().backward() print(x.grad) # 输出: tensor([-1., 0., 1.])
内容的提问来源于stack exchange,提问作者Shin-chan Loves Elephants
相关产品推荐
相关产品推荐

