自定义因果掩码下PyTorch MultiheadAttention不符合预期问题排查
需求说明
想用PyTorch的nn.MultiheadAttention实现自注意力模块,目标是让每个token仅关注自身之前的token(不含自身),区别于默认允许关注自身的自回归因果掩码。
掩码实现与说明
自定义掩码生成函数:
def generate_causal_mask(seq_length): # 对角线及以上设为1,转换为bool后代表不可关注,实现仅关注之前的token mask = torch.triu(torch.full((seq_length, seq_length), 1, dtype=torch.float32), diagonal=0).bool() # 允许第一个token关注自身,避免NaN mask[0, 0] = False return mask
生成的掩码示例(seq_len=8):
tensor([[False, True, True, True, True, True, True, True], [False, True, True, True, True, True, True, True], [False, False, True, True, True, True, True, True], [False, False, False, True, True, True, True, True], [False, False, False, False, True, True, True, True], [False, False, False, False, False, True, True, True], [False, False, False, False, False, False, True, True], [False, False, False, False, False, False, False, True]])
其中True表示该位置不可被关注,仅第一个token允许关注自身。
复现代码
if __name__ == "__main__": embed_dim = 16 batch_size = 1 seq_len = 8 mha = nn.MultiheadAttention(embed_dim, num_heads=1, batch_first=True) x = torch.randn(batch_size, seq_len, embed_dim).requires_grad_(True) causal_mask = generate_causal_mask(seq_len) print(causal_mask) output, _ = mha(x, x, x, attn_mask=causal_mask) # 计算t=5位置输出对输入的梯度 t = 5 loss = output[:, t].sum().backward() print("Gradient of the token:") print(x.grad)
问题现象
打印t=5位置输入的梯度时,发现该位置输出仍依赖自身输入,与“仅关注之前token”的掩码设定不符。
疑问
此行为是nn.MultiheadAttention的bug,还是对attn_mask的理解有误?若为预期行为,如何正确实现“输出仅依赖之前token(不含自身)”的需求?
解答
问题根源:对attn_mask的作用理解偏差
attn_mask的作用仅在于限制注意力权重可以使用的键值对:被标记为True的位置会在注意力得分计算时被设为-inf,经softmax后权重趋近于0,即该位置的键值对不会被纳入输出的加权求和。但它不会改变查询向量的来源——当前token的查询向量q[t]仍然由输入x[t]经线性变换生成。
因此,即使掩码禁止关注自身,x[t]的变化会通过q[t]影响与所有允许的键(0到t-1)的注意力得分,进而改变注意力权重和最终输出,所以输出t必然依赖x[t],梯度不为0是预期行为,并非bug。
正确实现需求的两种方案
根据你的真实需求,可选择以下两种方案:
方案1:仅限制注意力权重不分配给自身(保留查询向量来自当前token)
如果你的需求只是“注意力权重仅分配给之前的token”,当前的掩码实现是正确的,梯度不为0属于正常现象,无需修改。此时输出t的计算逻辑是:用当前token的查询向量,对之前所有token的键值对做加权求和,符合“关注之前token”的语义(这里的“关注”指注意力权重的分配,而非完全不使用当前token的任何信息)。
方案2:让输出完全不依赖当前token的输入
如果你的目标是输出t完全不使用x[t]的任何信息,则需要调整查询向量的来源:
- 可以对输入做移位操作:将输入序列向右移位一位,使位置
t的查询向量来自x[t-1](位置0用占位符填充)。示例代码如下:
# 构造移位后的查询输入:x_shifted[t] = x[t-1],x_shifted[0] = 0 x_shifted = torch.cat([torch.zeros_like(x[:, :1, :]), x[:, :-1, :]], dim=1) # 使用移位后的x作为query,原x作为key和value output, _ = mha(x_shifted, x, x, attn_mask=causal_mask)
此时,位置t的查询向量来自x[t-1],注意力权重仅分配给0到t-1的token,输出t将完全不依赖x[t],梯度也会符合预期。
内容的提问来源于stack exchange,提问作者jackjack4468

