MacBook Pro M1上PyTorch实现MultiHeadAttention测试异常求助
CS231n作业3 MultiHeadAttention测试异常排查
问题现象
我在完成CS231n作业3的Transformer_Captioning.ipynb时,测试MultiHeadAttention实现,要求相对误差小于1e-3,但运行后得到以下错误结果:
self_attn_output error: 0.449382070034207 masked_self_attn_output error: 1.0 attn_output error: 1.0
即使改用GitHub上能得到正确结果的代码,输出仍完全相同,怀疑遗漏了某些配置或步骤。
我的实现代码
MultiHeadAttention类
import torch import torch.nn as nn from torch.nn import functional as F import math class MultiHeadAttention(nn.Module): """ 实现《Attention Is All You Need》中简化版的掩码注意力机制 使用示例: attn = MultiHeadAttention(embed_dim, num_heads=2) # 自注意力 data = torch.randn(batch_size, sequence_length, embed_dim) self_attn_output = attn(query=data, key=data, value=data) # 双输入注意力 other_data = torch.randn(batch_size, sequence_length, embed_dim) attn_output = attn(query=data, key=other_data, value=other_data) """ def __init__(self, embed_dim, num_heads, dropout=0.1): super().__init__() assert embed_dim % num_heads == 0 # 初始化线性层,顺序固定以保证随机数生成一致 self.key = nn.Linear(embed_dim, embed_dim) self.query = nn.Linear(embed_dim, embed_dim) self.value = nn.Linear(embed_dim, embed_dim) self.proj = nn.Linear(embed_dim, embed_dim) self.attn_drop = nn.Dropout(dropout) self.n_head = num_heads self.emd_dim = embed_dim self.head_dim = self.emd_dim // self.n_head def forward(self, query, key, value, attn_mask=None): """ 计算掩码注意力输出,并行处理所有注意力头 形状说明: N=批量大小, S=源序列长度, T=目标序列长度, E=嵌入维度 输入: - query: 查询张量,形状(N, S, E) - key: 键张量,形状(N, T, E) - value: 值张量,形状(N, T, E) - attn_mask: 形状(S, T)的掩码数组,mask[i,j]==0表示源序列第i个token不应影响目标序列第j个token 返回: - output: 形状(N, S, E)的张量,根据注意力权重对value进行加权组合的结果 """ N, S, E = query.shape N, T, E = value.shape output = torch.empty((N, S, E)) H = self.n_head # 计算键、查询、值矩阵并拆分注意力头 K = self.key(key).view(N, T, H, E // H).moveaxis(1, 2) Q = self.query(query).view(N, S, H, E // H).moveaxis(1, 2) V = self.value(value).view(N, T, H, E // H).moveaxis(1, 2) # 计算注意力分数 Y = Q @ K.transpose(2, 3) / math.sqrt(self.head_dim) # 应用掩码 if attn_mask is not None: Y = Y.masked_fill(attn_mask == 0, float("-inf")) # 计算注意力权重并加权值矩阵,最后拼接头并投影 Y = self.attn_drop(F.softmax(Y, dim=-1)) @ V output = self.proj(Y.moveaxis(1, 2).reshape(N, S, E)) return output
测试代码
import torch import numpy as np if torch.backends.mps.is_available(): mps_device = torch.device("mps") x = torch.ones(1, device=mps_device) print(x) else: print("MPS device not found.") torch.manual_seed(231) # 选择唯一维度便于调试: N=1, H=2, T=3, E//H=4, E=8 batch_size = 1 sequence_length = 3 embed_dim = 8 attn = MultiHeadAttention(embed_dim, num_heads=2) # 自注意力测试 data = torch.randn(batch_size, sequence_length, embed_dim) self_attn_output = attn(query=data, key=data, value=data) # 掩码自注意力测试 mask = torch.randn(sequence_length, sequence_length) < 0.5 masked_self_attn_output = attn(query=data, key=data, value=data, attn_mask=mask) # 双输入注意力测试 other_data = torch.randn(batch_size, sequence_length, embed_dim) attn_output = attn(query=data, key=other_data, value=other_data) # 预期输出 expected_self_attn_output = np.asarray([[ [-0.2494, 0.1396, 0.4323, -0.2411, -0.1547, 0.2329, -0.1936, -0.1444], [-0.1997, 0.1746, 0.7377, -0.3549, -0.2657, 0.2693, -0.2541, -0.2476], [-0.0625, 0.1503, 0.7572, -0.3974, -0.1681, 0.2168, -0.2478, -0.3038]]]) expected_masked_self_attn_output = np.asarray([[ [-0.1347, 0.1934, 0.8628, -0.4903, -0.2614, 0.2798, -0.2586, -0.3019], [-0.1013, 0.3111, 0.5783, -0.3248, -0.3842, 0.1482, -0.3628, -0.1496], [-0.2071, 0.1669, 0.7097, -0.3152, -0.3136, 0.2520, -0.2774, -0.2208]]]) expected_attn_output = np.asarray([[ [-0.1980, 0.4083, 0.1968, -0.3477, 0.0321, 0.4258, -0.8972, -0.2744], [-0.1603, 0.4155, 0.2295, -0.3485, -0.0341, 0.3929, -0.8248, -0.2767], [-0.0908, 0.4113, 0.3017, -0.3539, -0.1020, 0.3784, -0.7189, -0.2912]]]) # 定义相对误差计算函数 def rel_error(x, y): return np.max(np.abs(x - y) / (np.maximum(1e-8, np.abs(x) + np.abs(y)))) print('self_attn_output error: ', rel_error(expected_self_attn_output, self_attn_output.detach().numpy())) print('masked_self_attn_output error: ', rel_error(expected_masked_self_attn_output, masked_self_attn_output.detach().numpy())) print('attn_output error: ', rel_error(expected_attn_output, attn_output.detach().numpy()))
环境配置
- MacBook Pro M1
- Python 3.8.12
- conda 4.11.0
- torch 2.1.0.dev20230415
- macOS Monterey Version 12.5
排查方向
- MPS设备兼容性问题:PyTorch开发版在MPS上的算子实现可能存在精度差异或bug,尝试强制使用CPU运行:
device = torch.device('cpu') attn = MultiHeadAttention(embed_dim, num_heads=2).to(device) data = torch.randn(batch_size, sequence_length, embed_dim).to(device) mask = mask.to(device) other_data = torch.randn(batch_size, sequence_length, embed_dim).to(device) - 随机种子全局一致性:补充NumPy种子设置,确保所有随机操作对齐:
np.random.seed(231) torch.manual_seed(231) - PyTorch版本切换:作业预期输出基于稳定版PyTorch生成,尝试切换到1.13或2.0稳定版。
- 掩码维度匹配:将掩码扩展到与注意力分数一致的维度,避免广播错误:
if attn_mask is not None: attn_mask = attn_mask.unsqueeze(0).unsqueeze(0) Y = Y.masked_fill(attn_mask == 0, float("-inf"))
内容的提问来源于stack exchange,提问作者ge1mina023
相关产品推荐
相关产品推荐

