You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.28 10:20:17