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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.24 21:47:11