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

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。

实现步骤

  1. 构造[B, S]的布尔型掩码:True表示对应位置需要被屏蔽(PyTorch会自动将这些位置的注意力权重置为负无穷,softmax后权重为0)。
  2. 将掩码传入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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.03 05:21:55