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
相关产品推荐
相关产品推荐

