PyTorch中nn.Transformer各类掩码失效问题求助
PyTorch Transformer层掩码不生效的原因及正确用法
常见导致掩码不生效的原因
- 掩码维度/类型错误:PyTorch对不同掩码的形状有严格要求:
src_mask/tgt_mask(自注意力掩码):需为[seq_len, seq_len]或[num_heads*batch_size, seq_len, seq_len],用于控制序列内部位置的注意力可见性src_key_padding_mask/tgt_key_padding_mask/memory_key_padding_mask(padding掩码):需为[batch_size, seq_len],标记padding token的位置(布尔类型中True表示该位置被mask;浮点类型需用-inf标记)
- 输入无有效掩码标记:如果输入序列没有padding,或设置的掩码全为
False(浮点型全为0),掩码不会产生任何影响 - 参数传递错误:调用层的forward方法时,误传参数名(比如把
src_key_padding_mask写成src_pad_mask)或传错参数位置,导致掩码未被正确应用
正确使用TransformerEncoderLayer掩码的示例
import torch import torch.nn as nn # 初始化EncoderLayer,注意设置batch_first=True匹配[batch, seq, dim]格式的输入 encoder_layer = nn.TransformerEncoderLayer( d_model=512, nhead=8, dim_feedforward=2048, batch_first=True ) # 构造输入:batch=2, seq_len=5, d_model=512 src = torch.randn(2, 5, 512) # 构造padding掩码:标记样本中的padding位置 src_key_padding_mask = torch.tensor([[False, False, False, True, True], [False, False, False, False, True]]) # 构造自注意力掩码:遮挡序列中位置i之后的所有位置(模拟类似Decoder的自回归效果) src_mask = torch.triu(torch.ones(5, 5) * float('-inf'), diagonal=1) # 不同掩码下的输出 output_no_mask = encoder_layer(src) output_with_pad_mask = encoder_layer(src, src_key_padding_mask=src_key_padding_mask) output_with_attn_mask = encoder_layer(src, src_mask=src_mask) # 验证输出差异 print("无掩码 vs 有padding掩码:", not torch.allclose(output_no_mask, output_with_pad_mask)) print("无掩码 vs 有自注意力掩码:", not torch.allclose(output_no_mask, output_with_attn_mask))
正确使用TransformerDecoderLayer掩码的示例
# 初始化DecoderLayer,同样设置batch_first=True decoder_layer = nn.TransformerDecoderLayer( d_model=512, nhead=8, dim_feedforward=2048, batch_first=True ) # 构造Decoder输入和Encoder输出(memory) tgt = torch.randn(2, 4, 512) memory = encoder_layer(src) # 构造各类掩码 tgt_key_padding_mask = torch.tensor([[False, False, True, True], [False, False, False, True]]) # Decoder输入的padding掩码 tgt_mask = torch.triu(torch.ones(4, 4) * float('-inf'), diagonal=1) # Decoder自回归掩码 memory_key_padding_mask = src_key_padding_mask # Encoder输出的padding掩码 # 不同掩码下的输出 output_no_mask = decoder_layer(tgt, memory) output_with_tgt_mask = decoder_layer(tgt, memory, tgt_mask=tgt_mask) output_with_pad_mask = decoder_layer(tgt, memory, tgt_key_padding_mask=tgt_key_padding_mask, memory_key_padding_mask=memory_key_padding_mask) # 验证输出差异 print("无掩码 vs 有自回归掩码:", not torch.allclose(output_no_mask, output_with_tgt_mask)) print("无掩码 vs 有padding掩码:", not torch.allclose(output_no_mask, output_with_pad_mask))
关键注意事项
- batch_first参数:若输入为
[batch_size, seq_len, d_model]格式,必须在初始化层时设置batch_first=True,否则掩码形状会与输入不匹配,导致失效 - 掩码类型一致性:布尔掩码和浮点掩码不能混用,确保掩码的标记方式符合PyTorch的要求(布尔
True/浮点-inf表示mask) - 掩码的作用场景:
*_mask主要用于控制序列内部的注意力可见性(如Decoder的自回归),*_key_padding_mask用于忽略padding token的注意力计算
内容的提问来源于stack exchange,提问作者haoxing
相关产品推荐
相关产品推荐

