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

视觉Transformer多头自注意力梯度爆炸(损失为NaN)问题排查

自定义多头自注意力模块导致NaN损失的问题排查

我自己实现的多头自注意力模块会让训练和验证损失变成NaN,移除该模块后训练恢复正常。我知道损失NaN通常和梯度爆炸有关,但找不到代码里的问题。对比PyTorch官方的nn.MultiheadSelfAttention结构类似,但用官方模块就不会出现梯度爆炸,自定义的却出问题且无报错信息,恳请帮忙排查代码问题:

class MultiHeadAttention(nn.Module):
  def __init__(self, in_dim, num_heads=8, dropout=0):
    super().__init__()
    self.num_heads = num_heads
    self.head_dim = in_dim // num_heads
    self.conv_q = nn.Conv2d(in_dim, in_dim, kernel_size=1)
    self.conv_k = nn.Conv2d(in_dim, in_dim, kernel_size=1)
    self.conv_v = nn.Conv2d(in_dim, in_dim, kernel_size=1)
    self.att_drop = nn.Dropout(dropout)
    self.proj = nn.Conv2d(in_dim, in_dim, kernel_size=1)
    self.proj_drop = nn.Dropout(dropout)

  def forward(self, x):

    b, _, h, w = x.shape
    
    q = self.conv_q(x)
    k = self.conv_k(x)
    v = self.conv_v(x)

    q = rearrange(q, "b (nh hd) h w -> b nh (h w) hd", nh=self.num_heads)
    k = rearrange(k, "b (nh hd) h w -> b nh (h w) hd", nh=self.num_heads)
    v = rearrange(v, "b (nh hd) h w -> b nh (h w) hd", nh=self.num_heads)

    att_score = q @ k.transpose(2, 3) ** (self.head_dim ** -0.5)
    att_score = F.softmax(att_score, dim=-1)
    att_score = self.att_drop(att_score)

    x = att_score @ v

    x = rearrange(x, 'b nh (h w) hd -> b (nh hd) h w', h=h, w=w)

    x = self.proj(x)
    x = self.proj_drop(x)

    return x

问题排查关键点

  • 缩放操作位置错误:你把注意力缩放因子(self.head_dim ** -0.5)用幂运算**作用在矩阵乘法结果上,这是核心错误。正确的做法是用乘法*把缩放因子乘到矩阵乘法的结果上,或者提前对q进行缩放。错误的幂运算会让原本正常的注意力分数直接爆炸成无穷大,后续softmax计算后就会出现NaN。修正后的代码:

    # 正确写法一:对矩阵乘法结果做缩放
    att_score = (q @ k.transpose(2, 3)) * (self.head_dim ** -0.5)
    # 或者正确写法二:提前缩放q
    q = q * (self.head_dim ** -0.5)
    att_score = q @ k.transpose(2, 3)
    
  • 数值稳定性验证:可以在forward函数中加入断言,快速定位哪一步出现NaN,比如:

    assert not torch.isnan(q).any(), "q contains NaN"
    assert not torch.isnan(k).any(), "k contains NaN"
    att_score = (q @ k.transpose(2, 3)) * (self.head_dim ** -0.5)
    assert not torch.isnan(att_score).any(), "att_score contains NaN before softmax"
    
  • 额外检查项:确认in_dim能被num_heads整除,否则self.head_dim会是整数除法后的非整数维度,后续维度拆分和计算会出现数值异常。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.18 02:50:12