PyTorch TransformerEncoder掩码序列疑问:如何按批次/序列屏蔽输入?
在PyTorch TransformerEncoder中实现批次/序列级的特定元素屏蔽
首先明确TransformerEncoder的两个核心掩码参数的区别,这是解决问题的关键:
- src_mask:形状为
[S, S](或扩展为[num_heads*B, S, S]),用于控制序列内位置间的注意力关系,比如自回归任务中常用的上三角掩码,限制每个位置只能关注之前的位置。这个掩码是全局生效的,对批次内所有样本的序列规则一致,不适合按批次/样本单独屏蔽特定元素。 - src_key_padding_mask:形状为
[B, S],这才是你需要的批次级元素屏蔽掩码,用于指定每个样本的哪些位置需要被模型忽略(即不参与注意力计算)。
错误原因
你之前触发形状错误,是因为把[B, S]的掩码传给了src_mask参数,而非src_key_padding_mask。
实现步骤
- 构造
[B, S]的布尔型掩码:True表示对应位置需要被屏蔽(PyTorch会自动将这些位置的注意力权重置为负无穷,softmax后权重为0)。 - 将掩码传入
TransformerEncoder的src_key_padding_mask参数即可。
代码示例
import torch # 批次大小B=2,序列长度S=4,特征维度dim=512 B, S, dim = 2, 4, 512 # 构造批次级屏蔽掩码:按样本指定屏蔽位置 mask = torch.zeros(B, S, dtype=torch.bool) mask[0, 1:3] = True # 第1个样本屏蔽索引1、2的位置(第2、3个元素) mask[1, 0] = True # 第2个样本屏蔽索引0的位置(第1个元素) # 生成随机输入序列 src = torch.randn(B, S, dim) # 初始化TransformerEncoder encoder_layer = torch.nn.TransformerEncoderLayer(d_model=dim, nhead=8) encoder = torch.nn.TransformerEncoder(encoder_layer, num_layers=3) # 传入掩码,执行编码 output = encoder(src, src_key_padding_mask=mask)
扩展应用:实现“填补空缺”的抗噪训练
如果要让模型学习填补空缺,可在输入中随机替换部分token为掩码token(比如torch.zeros_like(src)或专门的mask embedding),同时用src_key_padding_mask屏蔽这些位置的注意力,迫使模型从其他有效位置的信息中重建被屏蔽的内容,以此提升抗噪能力和泛化性。
内容的提问来源于stack exchange,提问作者Ælex
相关产品推荐
相关产品推荐

