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:掩码逻辑与参数冲突
- 自动因果掩码与自定义掩码冲突:同时设置
tgt_is_causal=True和自定义tgt_mask,会导致PyTorch自动生成的上三角因果掩码与你的自定义掩码叠加,部分位置可能被完全屏蔽,注意力分数归一化时出现0/0,最终产生NaN。 - 掩码格式不符合要求:PyTorch的
TransformerDecoder中,tgt_mask若为bool类型,True表示该位置被屏蔽;你的代码中生成的掩码逻辑错误,部分位置的有效关注路径被完全切断。
核心问题2:首位掩码token的处理漏洞
第三组样本中,掩码token位于第0位,create_causal_mask中强制设置mask[:,0]=True,结合后续逻辑会导致该位置的注意力输入为空,触发NaN。
修复步骤
- 关闭自动因果掩码:将
tgt_is_causal=True改为False,避免与自定义掩码冲突。 - 修正掩码生成逻辑:生成明确的允许关注路径,确保每个位置至少能关注到自身或首token,并将掩码转换为float类型(被屏蔽位置设为
-inf,允许位置设为0)。 - 调整首位掩码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
相关产品推荐
相关产品推荐

