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

PyTorch自注意力实现正确性验证及优化建议咨询

自注意力实现问题排查与优化

存在的核心问题

你的代码能运行但不符合自注意力的数学定义,主要问题如下:

  • 转置操作错误:queries.T会对整个张量转置,若输入是(batch_size, seq_len, embedding_dim)的标准格式,这会把batch和序列维度打乱。正确做法是仅对最后两个维度转置,保证batch维度不变。
  • 缺少缩放因子:自注意力需要将注意力分数除以√embedding_dim,否则当embedding维度较大时,softmax后的权重会过于集中,导致梯度消失或训练不稳定。
  • Softmax维度未指定:默认softmax会在dim=0计算,这会错误地跨batch或序列维度归一化,必须指定dim=-1在序列维度上计算注意力权重。
  • 矩阵乘法顺序错误:根据自注意力公式Attention(Q,K,V) = softmax(QK^T/√d_k)V,应该用注意力分数乘以values,而非values乘以分数,你的代码维度匹配仅在极端场景(如单样本单序列)下巧合成立。

修正后的代码

import torch
import torch.nn as nn

class SelfAttention(nn.Module):
    def __init__(self, embedding_dim):
        super(SelfAttention, self).__init__()
        self.embedding_dim = embedding_dim
        
        self.keys = nn.Linear(embedding_dim, embedding_dim)
        self.queries = nn.Linear(embedding_dim, embedding_dim)
        self.values = nn.Linear(embedding_dim, embedding_dim)

    def forward(self, x, mask=None):
        batch_size, seq_len, _ = x.shape
        
        keys = self.keys(x)  # shape: (batch_size, seq_len, embedding_dim)
        queries = self.queries(x)  # shape: (batch_size, seq_len, embedding_dim)
        values = self.values(x)  # shape: (batch_size, seq_len, embedding_dim)
        
        # 计算注意力分数:Q @ K^T / √d
        scores_prime = torch.matmul(queries, keys.transpose(-2, -1)) / torch.sqrt(torch.tensor(self.embedding_dim, dtype=torch.float32))
        
        # 应用掩码(可选)
        if mask is not None:
            scores_prime = scores_prime.masked_fill(mask == 0, -1e9)
        
        # 在序列维度上计算softmax
        scores = nn.functional.softmax(scores_prime, dim=-1)
        
        # 计算上下文向量:注意力权重 @ V
        context_vectors = torch.matmul(scores, values)
        
        return context_vectors

优化方向

  • 支持注意力掩码:上述代码已预留mask参数,可添加padding mask(屏蔽填充token)或sequence mask(屏蔽未来token,用于Decoder),适配实际NLP任务。
  • 实现多头注意力:将embedding_dim拆分为多个头,每个头学习不同的注意力模式,提升模型的表达能力,这是Transformer的核心组件。
  • 使用内置高效函数:PyTorch 2.0及以上版本提供torch.nn.functional.scaled_dot_product_attention,内置优化(如FlashAttention),可替代手动实现的分数计算与softmax,大幅提升训练效率。
  • 模块化拆分:将线性层合并为一个大的线性层(一次计算Q、K、V),减少代码冗余,例如用nn.Linear(embedding_dim, 3*embedding_dim),然后拆分出Q、K、V。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.07 09:05:33