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

如何优化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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.15 11:06:11