深入浅出自注意力机制:从原理到PyTorch实战

1次阅读
没有评论

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

image.webp

1. 核心概念:Q/K/ V 矩阵与缩放点积注意力

自注意力机制的核心是通过三个矩阵(Query、Key、Value)建立序列元素间的关联。给定输入序列 $X \in \mathbb{R}^{n \times d}$(n 为序列长度,d 为特征维度),计算过程如下:

深入浅出自注意力机制:从原理到 PyTorch 实战

  1. 线性投影
    $$Q = XW_Q, \quad K = XW_K, \quad V = XW_V$$
    其中 $W_Q, W_K, W_V \in \mathbb{R}^{d \times d_k}$ 为可训练参数矩阵

  2. 注意力分数
    $$\text{Attention}(Q, K, V) = \text{softmax}\left(\frac{QK^T}{\sqrt{d_k}}\right)V$$
    缩放因子 $\sqrt{d_k}$ 用于防止点积结果过大导致 Softmax 梯度消失

2. PyTorch 完整实现

import torch
import torch.nn as nn
import math

class SelfAttention(nn.Module):
    def __init__(self, embed_size, heads):
        super().__init__()
        self.embed_size = embed_size
        self.heads = heads
        self.head_dim = embed_size // heads

        assert self.head_dim * heads == embed_size, "Embed size needs division by heads"

        # 线性投影层
        self.values = nn.Linear(self.head_dim, self.head_dim, bias=False)
        self.keys = nn.Linear(self.head_dim, self.head_dim, bias=False)
        self.queries = nn.Linear(self.head_dim, self.head_dim, bias=False)
        self.fc_out = nn.Linear(heads * self.head_dim, embed_size)

    def forward(self, x, mask=None):
        N = x.shape[0]
        seq_len = x.shape[1]

        # 分割多头 (batch, seq_len, heads, head_dim)
        x = x.view(N, seq_len, self.heads, self.head_dim)

        queries = self.queries(x)
        keys = self.keys(x)
        values = self.values(x)

        # 计算缩放点积注意力
        energy = torch.einsum("nqhd,nkhd->nhqk", [queries, keys])
        if mask is not None:
            energy = energy.masked_fill(mask == 0, float("-1e20"))

        attention = torch.softmax(energy / (self.embed_size ** (1/2)), dim=3)
        out = torch.einsum("nhql,nlhd->nqhd", [attention, values])

        # 合并多头输出
        out = out.reshape(N, seq_len, -1)
        return self.fc_out(out)

# 位置编码示例
class PositionalEncoding(nn.Module):
    def __init__(self, d_model, max_len=100):
        super().__init__()
        pe = torch.zeros(max_len, d_model)
        position = torch.arange(0, max_len, dtype=torch.float).unsqueeze(1)
        div_term = torch.exp(torch.arange(0, d_model, 2).float() * (-math.log(10000.0) / d_model))
        pe[:, 0::2] = torch.sin(position * div_term)
        pe[:, 1::2] = torch.cos(position * div_term)
        self.register_buffer('pe', pe.unsqueeze(0))

    def forward(self, x):
        return x + self.pe[:, :x.size(1)]

3. 性能优化策略

  1. 计算复杂度分析
  2. QK^T 乘法:$O(n^2d)$
  3. Softmax 计算:$O(n^2)$
  4. 与 V 相乘:$O(n^2d)$

  5. 优化技巧

  6. 使用 torch.einsum 替代逐元素计算
  7. 混合精度训练(AMP)减少显存占用
  8. 当序列长度 >512 时考虑使用 Flash Attention

4. 常见错误与解决方案

  • 错误 1:未缩放 Attention 分数
  • 现象:训练初期出现 NaN 损失
  • 修复:确保除以 $\sqrt{d_k}$

  • 错误 2:忽略 padding 影响

  • 现象:模型对填充位置过度关注
  • 修复:添加 mask 矩阵masked_fill(-1e20)

  • 错误 3:多头维度分配不当

  • 现象:head_dim 非整数导致维度错误
  • 修复:添加assert embed_size % heads == 0

5. 扩展思考与实践

  1. Attention 可视化

    # 获取 attention 权重
    attention = model.get_attention(inputs)
    plt.imshow(attention[0].detach().numpy())  # 可视化第一个头

  2. 长序列优化方案

  3. 使用稀疏 Attention(如 Longformer)
  4. 采用分块计算(Reformer 的 LSH Attention)
  5. 尝试线性 Attention 变体(Performer)

实践建议

从简单的序列分类任务开始(如 IMDB 影评分类),逐步尝试以下改进:
1. 对比单头与多头注意力的效果差异
2. 添加 / 移除位置编码观察性能变化
3. 在自定义数据集上可视化 Attention 热力图

代码仓库推荐参考 HuggingFace 的 transformers 库实现,其中包含了工业级的 Attention 优化方案。

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