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

自定义因果掩码下PyTorch MultiheadAttention不符合预期问题排查

自定义因果掩码在PyTorch nn.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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.15 11:50:54