多头注意力机制入门指南:从理论到PyTorch实现

1次阅读
没有评论

共计 2863 个字符,预计需要花费 8 分钟才能阅读完成。

image.webp

引言:Transformer 架构中的注意力机制

Transformer 模型在自然语言处理领域取得了巨大成功,而其中的核心组件就是注意力机制。简单来说,注意力机制可以让模型在处理序列数据时,动态地关注不同位置的信息。想象一下你在阅读一篇文章时,大脑会自动聚焦于当前最相关的词语或句子——注意力机制的作用与此类似。

多头注意力机制入门指南:从理论到 PyTorch 实现

单头注意力回顾

在进入多头注意力之前,我们先简单回顾下单头注意力的基本原理。单头注意力通过三个关键向量进行计算:

  • Query(Q):表示当前要查询的内容
  • Key(K):表示待匹配的内容
  • Value(V):表示最终要提取的信息

计算过程可以概括为:
1. 计算 Q 和 K 的相似度
2. 通过 softmax 归一化得到注意力权重
3. 用权重对 V 进行加权求和

数学表达式为:

Attention(Q, K, V) = softmax(QK^T/√d_k)V

其中 d_k 是 Key 的维度,√d_k 用于缩放防止 softmax 值过大。

多头注意力的动机和优势

单头注意力有一个明显局限:它只能学习一种注意力模式。就像我们人类理解复杂问题时,会从多个角度思考一样,多头注意力通过并行计算多个注意力头,可以捕获不同子空间的信息。

多头注意力的主要优势包括:

  • 能够同时关注不同位置的不同信息模式
  • 提高了模型的表示能力
  • 允许信息在不同子空间中进行处理

数学原理详解

多头注意力的计算过程可以分为以下几个步骤:

  1. 线性投影:将 Q、K、V 分别投影 h 次(h 是头数)
  2. 缩放点积注意力:对每个头独立计算注意力
  3. 拼接输出:将所有头的输出拼接起来
  4. 最终投影:通过线性层得到最终输出

数学表达式为:

MultiHead(Q, K, V) = Concat(head_1, ..., head_h)W^O
where head_i = Attention(QW_i^Q, KW_i^K, VW_i^V)

PyTorch 实现

下面是一个完整的 PyTorch 实现,包含类型注解和详细注释:

import torch
import torch.nn as nn
import torch.nn.functional as F
from typing import Optional, Tuple

class MultiHeadAttention(nn.Module):
    def __init__(self, embed_dim: int, num_heads: int):
        super().__init__()
        assert embed_dim % num_heads == 0, "Embed dim must be divisible by num_heads"

        self.embed_dim = embed_dim
        self.num_heads = num_heads
        self.head_dim = embed_dim // num_heads

        # 定义 Q、K、V 和输出的投影矩阵
        self.q_proj = nn.Linear(embed_dim, embed_dim)
        self.k_proj = nn.Linear(embed_dim, embed_dim)
        self.v_proj = nn.Linear(embed_dim, embed_dim)
        self.out_proj = nn.Linear(embed_dim, embed_dim)

    def forward(
        self,
        query: torch.Tensor,
        key: torch.Tensor,
        value: torch.Tensor,
        mask: Optional[torch.Tensor] = None
    ) -> Tuple[torch.Tensor, torch.Tensor]:
        """
        前向传播
        Args:
            query: [batch_size, seq_len, embed_dim]
            key: [batch_size, seq_len, embed_dim]
            value: [batch_size, seq_len, embed_dim]
            mask: [batch_size, seq_len, seq_len] (可选)
        Returns:
            output: [batch_size, seq_len, embed_dim]
            attention_weights: [batch_size, num_heads, seq_len, seq_len]
        """
        batch_size = query.size(0)

        # 线性投影
        Q = self.q_proj(query)  # [batch_size, seq_len, embed_dim]
        K = self.k_proj(key)    # [batch_size, seq_len, embed_dim]
        V = self.v_proj(value)  # [batch_size, seq_len, embed_dim]

        # 重排形状以分离注意力头
        Q = Q.view(batch_size, -1, self.num_heads, self.head_dim).transpose(1, 2)
        K = K.view(batch_size, -1, self.num_heads, self.head_dim).transpose(1, 2)
        V = V.view(batch_size, -1, self.num_heads, self.head_dim).transpose(1, 2)

        # 计算缩放点积注意力
        attn_scores = torch.matmul(Q, K.transpose(-2, -1)) / (self.head_dim ** 0.5)

        # 应用 mask(如果有)
        if mask is not None:
            attn_scores = attn_scores.masked_fill(mask == 0, float('-inf'))

        attn_weights = F.softmax(attn_scores, dim=-1)
        output = torch.matmul(attn_weights, V)

        # 拼接所有头并做最终投影
        output = output.transpose(1, 2).contiguous().view(batch_size, -1, self.embed_dim)
        output = self.out_proj(output)

        return output, attn_weights

性能分析和优化建议

多头注意力的计算复杂度主要来自于矩阵乘法。对于长度为 n 的序列:

  • 时间复杂度:O(n^2 * d) 其中 d 是 embedding 维度
  • 空间复杂度:O(n^2 * h) 其中 h 是头数

优化建议:

  1. 对于长序列,考虑使用稀疏注意力或局部注意力
  2. 在 GPU 上,确保 batch_size 足够大以充分利用并行计算
  3. 适当选择头的数量 (通常 4 -16)

常见问题解答

Q: 为什么需要除以√d_k?
A: 防止点积结果过大导致 softmax 梯度消失。

Q: 如何选择头的数量?
A: 通常 embed_dim 能被头的数量整除。常见选择是 4 -16 个。

Q: 多头注意力层可以堆叠吗?
A: 可以,Transformer 就是堆叠多个多头注意力层。

课后练习

尝试实现一个改进版的多头注意力层,加入以下特性:

  1. 相对位置编码
  2. 残差连接和层归一化
  3. 注意力权重的 dropout

提示:可以参考 HuggingFace 的 Transformer 实现,但尽量自己动手实现核心部分。

正文完
 0
评论(没有评论)