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

TransformerDecoder自定义掩码失效致输出NaN问题求助

问题:TransformerDecoder输出NaN,自定义掩码无法解决

问题描述

尝试仅基于首个token和掩码token预测掩码token,自定义create_causal_mask生成多头注意力掩码,但运行后返回NaN张量。已确认pad掩码与掩码token无交集,注意力掩码存在False值(非全屏蔽),严格遵循PyTorch官方文档中torch.nn.TransformerDecoder的掩码设置方式,仍无法解决问题。

问题代码

import torch
from torch import nn

torch.manual_seed(0)


class LookOnFirstDecoder(nn.Module):
    def __init__(self, depth, d_model, nhead, d_ff,
                 dropout, activation,
                sent_length, n_tokens, pad_idx
    ):
        super().__init__()
        """
        :param sent_length: max length of sentence
        :param n_tokens: number of tokens to use including mask and padding tokens
        :param pad_idx: index of padding to don't compute the gradient
        """
        self.d_model = d_model
        self.nhead = nhead
        self.n_tokens = n_tokens
        self.emb = nn.Embedding(
            num_embeddings=n_tokens,
            embedding_dim=d_model,
            padding_idx=pad_idx
        )

        self.pos_embed = nn.Parameter(
            torch.zeros(1, sent_length, d_model),
            requires_grad=True
        )
        torch.nn.init.normal_(self.pos_embed, std=.02)

        self.transformer = nn.TransformerDecoder(
            nn.TransformerDecoderLayer(
                d_model=d_model,
                nhead=nhead,
                dim_feedforward=d_ff,
                dropout=dropout,
                activation=activation,
                batch_first=True,
                norm_first=True,
            ),
            num_layers=depth,
        )

        self.fin_lin = nn.Linear(d_model, n_tokens)

    def create_causal_mask(self, mask):
        """
            The purpose is to create mask that allows all not first tokens
        look only on the first token and itself
        :param mask: (B, L)
        :return: (B * nhead, L)
        """

        mask[:, 0] = True  # to depend on first token
        b, l = mask.shape
        batch_causal_mask = ~torch.tril(mask.unsqueeze(-1) * mask.unsqueeze(-2))  # (B, L, L)
        # batch_causal_mask = torch.tril(torch.ones((b, l, l))).to("cuda") == 0

        # batch_causal_mask = torch.where(batch_causal_mask, 0, float('-inf'))
        print(f"Batch causal mask: \n{batch_causal_mask}")

        causal_mask = (
            batch_causal_mask.
            unsqueeze(1).  # (B, 1, L, L)
            expand(b, self.nhead, l, l).  # (B, nhead, L, L)
            reshape(b * self.nhead, l, l)  # (B * nhead, L, L)
        )

        return causal_mask

    def forward(self, tgt, memory, is_masked_mask, is_pad_mask):
        """
        :param tgt: (B, L)
        :param memory: (B, L1, D)
        :param is_masked_mask: (B, L) - True - mask token, False - not
        :param is_pad_mask: (B, L), True - pad token, False - not
        :return: tensor of shape (B, n_tokens)
        """
        b, l = tgt.shape
        tgt_tokens = self.emb(tgt) + self.pos_embed[:, :l].expand(b, l, self.d_model)

        tgt_tokens = self.transformer(
            tgt_tokens,
            memory,
            tgt_mask=self.create_causal_mask(is_masked_mask.clone()),
            tgt_is_causal=True,
            tgt_key_padding_mask=is_pad_mask
        )  # (B, L, D)

        fin_tokens = self.fin_lin(tgt_tokens[is_masked_mask])
        return fin_tokens


# my vocabulary
n_tokens = 10  # pad_idx - 9, mask_idx - 8
pad_idx = n_tokens - 1
mask_idx = n_tokens - 2

d_model = 4
nhead = 2
b, l = 3, 8

model = LookOnFirstDecoder(
    depth=2,
    d_model=4,
    nhead=2,
    d_ff=8,
    dropout=0.1,
    activation="gelu",
    sent_length=l,
    n_tokens=n_tokens,
    pad_idx=pad_idx
)

memory = torch.randn(b, l, d_model)

# so i create some random tokens, without padding and mask
in_tokens = torch.randint(0, mask_idx - 1, (b, l))

# mask and paddings add manually
in_tokens[0, 6:] = pad_idx
in_tokens[0, 5] = mask_idx

in_tokens[1, 7:] = pad_idx
in_tokens[1, 4] = mask_idx

in_tokens[2, 5:] = pad_idx
in_tokens[2, 0] = mask_idx

is_masked_mask = in_tokens == mask_idx
is_pad_mask = in_tokens == pad_idx

pred = model(in_tokens, memory, is_masked_mask, in_tokens == pad_idx)

print(f"In tokens: \n{in_tokens}")
print(f"Pad mask: \n{is_pad_mask}")
print(f"Masked mask: \n{is_masked_mask}")
print(f"Pred: \n{pred}")

依赖配置(requirements.txt)

torch == 2.1.1
torchvision == 0.16.1
xformers
albumentations==1.3.1

numpy == 1.26.2
scipy == 1.11.4
scikit-learn == 1.3.2
pandas == 2.1.4
matplotlib == 3.8.2
seaborn == 0.13.0

执行结果

Batch causal mask: 
tensor([[[False,  True,  True,  True,  True,  True,  True,  True],
         [ True,  True,  True,  True,  True,  True,  True,  True],
         [ True,  True,  True,  True,  True,  True,  True,  True],
         [ True,  True,  True,  True,  True,  True,  True,  True],
         [ True,  True,  True,  True,  True,  True,  True,  True],
         [False,  True,  True,  True,  True, False,  True,  True],
         [ True,  True,  True,  True,  True,  True,  True,  True],
         [ True,  True,  True,  True,  True,  True,  True,  True]],

        [[False,  True,  True,  True,  True,  True,  True,  True],
         [ True,  True,  True,  True,  True,  True,  True,  True],
         [ True,  True,  True,  True,  True,  True,  True,  True],
         [ True,  True,  True,  True,  True,  True,  True,  True],
         [False,  True,  True,  True, False,  True,  True,  True],
         [ True,  True,  True,  True,  True,  True,  True,  True],
         [ True,  True,  True,  True,  True,  True,  True,  True],
         [ True,  True,  True,  True,  True,  True,  True,  True]],

        [[False,  True,  True,  True,  True,  True,  True,  True],
         [ True,  True,  True,  True,  True,  True,  True,  True],
         [ True,  True,  True,  True,  True,  True,  True,  True],
         [ True,  True,  True,  True,  True,  True,  True,  True],
         [ True,  True,  True,  True,  True,  True,  True,  True],
         [ True,  True,  True,  True,  True,  True,  True,  True],
         [ True,  True,  True,  True,  True,  True,  True,  True],
         [ True,  True,  True,  True,  True,  True,  True,  True]]])
In tokens: 
tensor([[2, 1, 4, 0, 3, 8, 9, 9],
        [6, 4, 0, 6, 8, 0, 5, 9],
        [8, 2, 5, 2, 6, 9, 9, 9]])
Pad mask: 
tensor([[False, False, False, False, False, False,  True,  True],
        [False, False, False, False, False, False, False,  True],
        [False, False, False, False, False,  True,  True,  True]])
Masked mask: 
tensor([[False, False, False, False, False,  True, False, False],
        [False, False, False, False,  True, False, False, False],
        [ True, False, False, False, False, False, False, False]])
Pred: 
tensor([[nan, nan, nan, nan, nan, nan, nan, nan, nan, nan],
        [nan, nan, nan, nan, nan, nan, nan, nan, nan, nan],
        [nan, nan, nan, nan, nan, nan, nan, nan, nan, nan]],
       grad_fn=<AddmmBackward0>)

排查与解决

核心问题1:掩码逻辑与参数冲突

  1. 自动因果掩码与自定义掩码冲突:同时设置tgt_is_causal=True和自定义tgt_mask,会导致PyTorch自动生成的上三角因果掩码与你的自定义掩码叠加,部分位置可能被完全屏蔽,注意力分数归一化时出现0/0,最终产生NaN。
  2. 掩码格式不符合要求:PyTorch的TransformerDecoder中,tgt_mask若为bool类型,True表示该位置被屏蔽;你的代码中生成的掩码逻辑错误,部分位置的有效关注路径被完全切断。

核心问题2:首位掩码token的处理漏洞

第三组样本中,掩码token位于第0位,create_causal_mask中强制设置mask[:,0]=True,结合后续逻辑会导致该位置的注意力输入为空,触发NaN。

修复步骤

  1. 关闭自动因果掩码:将tgt_is_causal=True改为False,避免与自定义掩码冲突。
  2. 修正掩码生成逻辑:生成明确的允许关注路径,确保每个位置至少能关注到自身或首token,并将掩码转换为float类型(被屏蔽位置设为-inf,允许位置设为0)。
  3. 调整首位掩码token的处理逻辑:避免强制修改输入掩码的首位值,改为直接构建允许关注首token的规则。

修改后的关键代码:

def create_causal_mask(self, mask):
    """
        生成允许所有非首token仅关注首token和自身的掩码
        :param mask: (B, L) - True表示掩码token
        :return: (B * nhead, L, L)
    """
    b, l = mask.shape
    # 初始化基础掩码:默认所有位置都被屏蔽
    base_mask = torch.full((b, l, l), float('-inf'))
    # 允许每个位置关注首token
    base_mask[:, :, 0] = 0.0
    # 允许每个位置关注自身
    base_mask[:, torch.arange(l), torch.arange(l)] = 0.0
    
    print(f"Batch causal mask: \n{base_mask}")
    # 扩展为多头格式
    causal_mask = (
        base_mask.
        unsqueeze(1).
        expand(b, self.nhead, l, l).
        reshape(b * self.nhead, l, l)
    )
    return causal_mask

def forward(self, tgt, memory, is_masked_mask, is_pad_mask):
    b, l = tgt.shape
    tgt_tokens = self.emb(tgt) + self.pos_embed[:, :l].expand(b, l, self.d_model)

    tgt_tokens = self.transformer(
        tgt_tokens,
        memory,
        tgt_mask=self.create_causal_mask(is_masked_mask.clone()),
        tgt_is_causal=False,  # 关闭自动因果掩码
        tgt_key_padding_mask=is_pad_mask
    )

    fin_tokens = self.fin_lin(tgt_tokens[is_masked_mask])
    return fin_tokens

额外建议

  • 若不需要交叉注意力,可考虑改用TransformerEncoder,移除memory相关逻辑,减少复杂度。
  • 小尺寸d_model(如4)容易引发数值不稳定,建议适当增大d_model,或在归一化层中设置eps=1e-6提升稳定性。

内容的提问来源于stack exchange,提问作者First Name Second Name

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.30 20:38:10