如何优化torch.where与torch.sqrt实现以兼容PyTorch自动求导?
解决PyTorch Forward函数梯度NaN问题的优化实现
问题背景
原Forward函数实现如下,计算梯度时会出现NaN:
def forward(self, x): r = torch.sqrt(x[:, 0:1]**2 + x[:, 1:2]**2) value = torch.where(r <= 1, torch.sqrt(1.-r**2), -1) return value
问题根源:当r > 1时,1 - r²为负数,虽然torch.where不会取用torch.sqrt(1.-r**2)的结果,但PyTorch自动微分会计算整个张量的梯度,负数开根号会产生NaN,最终污染梯度。
现有规避方案:
def forward(self, x): i = x[:, 0:1]**2 + x[:, 1:2]**2 r = torch.sqrt(torch.relu(i)) alternative=torch.sqrt(torch.relu(1.-r**2)) value = torch.where(r <= 1, alternative, -1) return value
更优实现(最小改动原代码)
只需要针对1.-r**2添加ReLU保护,无需额外拆分计算步骤,改动极小:
def forward(self, x): r = torch.sqrt(x[:, 0:1]**2 + x[:, 1:2]**2) value = torch.where(r <= 1, torch.sqrt(torch.relu(1. - r**2)), -1) return value
说明
- 原代码中
x[:,0:1]**2 + x[:,1:2]**2本身是非负的,因此torch.sqrt不会产生NaN,无需额外添加torch.relu,这一步可以省略。 - 仅对
1. - r**2使用torch.relu,将负数钳位到0,避免开根号时产生NaN,同时完全保留原代码的结构逻辑。
内容的提问来源于stack exchange,提问作者BBB
相关产品推荐
相关产品推荐

