PyTorch函数导数列表位置及不可微点处理方式查询
PyTorch相关技术问题解答
一、PyTorch函数导数定义的查询位置
- 官方文档的算子详情页:每个公开的PyTorch函数说明页,都会标注反向传播的计算逻辑,涉及不可微点的特殊处理也会在该部分明确说明
tools/autograd/derivatives.yaml配置文件:这是PyTorch源码中集中罗列所有原生算子导数定义的位置,所有自动生成的反向传播逻辑都基于该文件的规则生成,不需要逐份查找算子源码- 算子对应实现代码:
- Python层面实现的自定义算子,反向逻辑定义在对应
torch.autograd.Function子类的backward方法中,存放在torch/autograd目录下的对应文件 - C++层面实现的底层算子,反向逻辑与前向逻辑一同存放在
aten/src/ATen/native目录下的对应算子实现文件中
- Python层面实现的自定义算子,反向逻辑定义在对应
二、PyTorch对不可微点的处理逻辑
PyTorch对所有带不可微点的算子,默认遵循次梯度的通用规则处理,不同场景的具体逻辑如下:
- 激活类分段函数的不可微点:比如
ReLU在x=0处、torch.abs在x=0处,默认返回梯度为0 - 最值类多输入匹配的不可微点:比如
torch.max、torch.min在多个输入值相等的位置,会将梯度平均分配给所有值相等的输入张量 - 离散输出类函数的不可微点:比如
argmax、round、ceil、floor这类输出为离散值的函数,默认不可微,反向传播时梯度直接置为0,不会向上游传递梯度 - 定义域边界的不可微点:比如
torch.sqrt在x=0处、torch.log在x≤0处,会返回NaN梯度,若开启torch.autograd.detect_anomaly()会抛出异常提示,方便定位非法输入问题
若默认的不可微点处理逻辑不符合业务需求,可继承
torch.autograd.Function自定义对应算子的反向传播规则,覆盖默认的次梯度实现
内容的提问来源于stack exchange,提问作者Mo Elrefae
相关产品推荐
相关产品推荐

