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

Transformer多头自注意力的置换不变/等变性验证代码疑问

问题:Transformer多头自注意力的置换等变性验证错误排查

已知当函数f满足f(P(x)) = P(f(x))(其中P为置换操作)时,f具有等变性。为验证Transformer中Multi-Head Self Attention的置换不变性与等变性,编写了如下PyTorch代码:

import torch
import torch.nn as nn

multihead_attn = nn.MultiheadAttention(embed_dim=32, num_heads=4, batch_first=True)
x0 = torch.ones(11,32)
x1 = torch.ones(11,32)
for i in range(x0.size(0)):
    x0[i] *= i
    x1[i] *= (i+1) % x0.size(0)

x = torch.cat(
    (x0.unsqueeze(0), x1.unsqueeze(0))
    )

y0, y1 = multihead_attn(x,x,x)[0]
y0 = y0.squeeze(0)
y1 = y1.squeeze(0)

验证发现torch.equal(x0[1],x1[0])返回True,但torch.equal(y0,y1)返回False(看似不满足置换不变性),torch.equal(y0[1],y1[0])也返回False(看似不满足等变性)。


错误原因与修正方案

1. 混淆置换不变性与等变性的概念

纯自注意力(无位置编码)是置换等变的,而非置换不变的:

  • 置换不变性要求f(P(x)) = f(x),即输入置换后输出完全相同;
  • 等变性要求f(P(x)) = P(f(x)),即输入置换后输出也做对应置换。
    你用torch.equal(y0,y1)验证置换不变性是错误的,因为自注意力本身就不满足置换不变性。

2. 浮点数精度问题导致torch.equal返回False

torch.equal要求张量所有元素完全相等,但自注意力计算涉及softmax、矩阵乘法等浮点运算,即使理论上相等的张量,实际计算中会存在微小数值误差。应该使用torch.allclose来比较,允许设置合理的误差范围(默认容差已覆盖常见浮点误差)。

3. 等变性验证方式可更严谨

x1是x0循环左移一位的结果,对应的y1应该是y0循环左移一位的结果,直接验证整体对应关系更准确:

torch.allclose(y1, torch.roll(y0, shifts=-1, dims=0))

4. 固定随机种子提升复现性

添加随机种子固定代码,确保每次运行的模型参数一致,方便调试。


修正后的验证代码

import torch
import torch.nn as nn

# 固定随机种子,确保结果可复现
torch.manual_seed(42)
torch.cuda.manual_seed(42)

multihead_attn = nn.MultiheadAttention(embed_dim=32, num_heads=4, batch_first=True)
x0 = torch.ones(11,32)
x1 = torch.ones(11,32)
for i in range(x0.size(0)):
    x0[i] *= i
    x1[i] *= (i+1) % x0.size(0)

x = torch.cat((x0.unsqueeze(0), x1.unsqueeze(0)))

y0, y1 = multihead_attn(x,x,x)[0]
y0 = y0.squeeze(0)
y1 = y1.squeeze(0)

# 验证整体等变性
print(torch.allclose(y1, torch.roll(y0, shifts=-1, dims=0)))  # 输出True
# 验证单个位置对应关系
print(torch.allclose(y0[1], y1[0]))  # 输出True

内容的提问来源于stack exchange,提问作者fenaux

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.12 19:08:11