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

自定义PyTorch函数:前向硬钳位+反向tanh梯度遇NaN问题求助

解决PyTorch自定义Clamp函数反向传播NaN问题

问题根源

你当前的代码在反向传播时,对[-1,1]区间内的输入保留了原torch.clamp的梯度(值为1),而对区间外的输入使用tanh的梯度,这会在x=-1和x=1处产生梯度不连续。这种不连续性会导致训练过程中的梯度震荡,进而引发NaN值。

最优实现方案

要实现前向传播与torch.clamp一致、反向传播使用类似tanh的平滑梯度,正确的做法是对所有输入值都使用tanh的梯度,而非仅区间外。这样梯度在整个定义域内都是连续平滑的,避免训练不稳定。

修正后的代码

import torch

class ClampWithTanhGrad(torch.autograd.Function):
    @staticmethod
    def forward(ctx, input):
        ctx.save_for_backward(input)
        # 前向传播完全等价于torch.clamp(input, -1, 1)
        return torch.clamp(input, -1, 1)
    
    @staticmethod
    def backward(ctx, grad_output):
        input, = ctx.saved_tensors
        # 计算tanh的梯度:1 - tanh(input)^2,等价于sech²(input)
        tanh_gradient = 1.0 - torch.tanh(input) ** 2.0
        # 将上游梯度与当前梯度相乘
        return grad_output * tanh_gradient

# 使用示例
clamp_with_tanh_grad = ClampWithTanhGrad.apply
x = torch.tensor([-2.0, -0.5, 0.0, 0.5, 2.0], requires_grad=True)
y = clamp_with_tanh_grad(x)
y.sum().backward()
print(x.grad)  # 输出平滑的梯度:[0.0707, 0.7865, 1.0, 0.7865, 0.0707]

可选变体(若需求为tanh(clamped_input)的梯度)

如果你实际需要的是经过clamp后的输入的tanh梯度(即梯度为1 - tanh(clamped_input)^2),可以修改反向传播代码如下:

@staticmethod
def backward(ctx, grad_output):
    input, = ctx.saved_tensors
    clamped_input = torch.clamp(input, -1, 1)
    tanh_gradient = 1.0 - torch.tanh(clamped_input) ** 2.0
    return grad_output * tanh_gradient

这种变体下,区间外的输入梯度会保持为1 - tanh(1)^2 ≈ 0.4199(而非趋近于0),适合需要在饱和区域保留一定梯度的场景。

内容的提问来源于stack exchange,提问作者Tob

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.24 21:37:53