深入解析Transformer中的自注意力机制:从数学原理到实现细节

1次阅读
没有评论

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

image.webp

传统序列模型的局限性

在 Transformer 出现之前,RNN 和 LSTM 是处理序列数据的主流方法。但它们存在两个致命缺陷:

深入解析 Transformer 中的自注意力机制:从数学原理到实现细节

  1. 顺序计算依赖 :必须逐个处理序列元素,无法并行计算。对于长度为 n 的序列,时间复杂度高达 O(n)
  2. 长距离依赖衰减 :信息通过隐藏状态逐步传递,超过 20 步后梯度消失严重(参考 [Hochreiter 1991])

自注意力机制原理

QKV 计算过程

给定输入矩阵 $X \in \mathbb{R}^{n \times d_{model}}$,通过三个权重矩阵得到查询 (Query)、键 (Key)、值 (Value):

$$
Q = XW^Q, \quad W^Q \in \mathbb{R}^{d_{model} \times d_k}
$$
$$
K = XW^K, \quad W^K \in \mathbb{R}^{d_{model} \times d_k}
$$
$$
V = XW^V, \quad W^V \in \mathbb{R}^{d_{model} \times d_v}
$$

缩放点积注意力

计算注意力权重的完整公式:

$$
Attention(Q,K,V) = softmax(\frac{QK^T}{\sqrt{d_k}})V
$$

缩放因子 $\sqrt{d_k}$ 防止点积结果过大导致 softmax 梯度消失

多头注意力

将 QKV 拆分成 h 个头并行计算:

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

PyTorch 实现

单头注意力

import torch
import torch.nn as nn
import torch.nn.functional as F

class ScaledDotProductAttention(nn.Module):
    def __init__(self, d_k):
        super().__init__()
        self.d_k = d_k

    def forward(self, Q, K, V, mask=None):
        # Q,K,V shape: (batch, seq_len, d_k)
        scores = torch.matmul(Q, K.transpose(-2, -1)) / (self.d_k ** 0.5)

        if mask is not None:
            scores = scores.masked_fill(mask == 0, -1e9)

        attn = F.softmax(scores, dim=-1)
        output = torch.matmul(attn, V)
        return output

完整多头注意力

class MultiHeadAttention(nn.Module):
    def __init__(self, d_model=512, h=8):
        super().__init__()
        assert d_model % h == 0
        self.d_k = d_model // h
        self.h = h

        self.W_Q = nn.Linear(d_model, d_model)
        self.W_K = nn.Linear(d_model, d_model)
        self.W_V = nn.Linear(d_model, d_model)
        self.W_O = nn.Linear(d_model, d_model)

    def forward(self, x, mask=None):
        batch_size = x.size(0)

        # 线性投影
        Q = self.W_Q(x)  # (batch, seq, d_model)
        K = self.W_K(x)
        V = self.W_V(x)

        # 分头处理
        Q = Q.view(batch_size, -1, self.h, self.d_k).transpose(1,2)
        K = K.view(batch_size, -1, self.h, self.d_k).transpose(1,2)
        V = V.view(batch_size, -1, self.h, self.d_k).transpose(1,2)

        # 计算注意力
        scores = torch.matmul(Q, K.transpose(-2, -1)) / (self.d_k ** 0.5)
        if mask is not None:
            mask = mask.unsqueeze(1)  # 广播到所有头
            scores = scores.masked_fill(mask == 0, -1e9)

        attn = F.softmax(scores, dim=-1)
        context = torch.matmul(attn, V)

        # 合并多头输出
        context = context.transpose(1,2).contiguous()
        context = context.view(batch_size, -1, self.h * self.d_k)

        return self.W_O(context)

实践经验

超参选择策略

  • 头数 h :通常 4 -16 之间,建议先用 8
  • 模型维度 d_model:一般 512/768/1024,需能被 h 整除

内存优化技巧

  1. 梯度检查点:torch.utils.checkpoint
  2. 序列分块:将长序列拆分为多个子序列
  3. 混合精度训练:torch.cuda.amp

梯度问题预防

  • 初始化:使用 Xavier/Glorot 初始化
  • 层归一化:每个子层后添加 LayerNorm
  • 学习率预热:前 4000 步线性增加学习率

性能对比

模型 时间复杂度 空间复杂度 并行性
RNN O(n) O(1)
LSTM O(n) O(1)
Transformer O(n²) O(n²) 完全

实际测试中(n=512),Transformer 比 LSTM 快 3 - 5 倍

延伸思考

  1. 如何改造自注意力机制使其适用于图像数据?
  2. 当序列长度达到 10 万时,有哪些优化方法?
  3. 为什么说自注意力机制本质是一种图神经网络?
正文完
 0
评论(没有评论)