注意力机制深度解析:从原理到Transformer实战

1次阅读
没有评论

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

image.webp

传统序列建模的瓶颈

在处理文本、语音等序列数据时,传统 RNN/LSTM 面临两大核心痛点:

注意力机制深度解析:从原理到 Transformer 实战

  1. 梯度消失问题:当序列长度超过 50 步时,反向传播的梯度会指数级衰减,导致模型难以学习长距离依赖关系。数学上可表示为:
    $$\frac{\partial L}{\partial h_t} \approx \prod_{k=t}^{T} \frac{\partial h_{k+1}}{\partial h_k} \to 0 \quad (T\gg t)$$

  2. 顺序计算限制:RNN 的时序依赖性导致无法并行计算,处理长为 $n$ 的序列需要 $O(n)$ 时间步。这在处理万字长文或 DNA 序列时尤为致命

注意力机制的三要素

注意力机制通过 Query/Key/Value 分解实现动态权重分配,其核心计算流程为:

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

  • Query:当前需要计算的特征表示(如解码器当前词)
  • Key:待检索的特征集合(如编码器所有词)
  • Value:实际返回的特征信息(通常与 Key 维度相同)

常见变体对比

类型 计算公式 复杂度 适用场景
加性注意力 $v^T\tanh(W_q q + W_k k)$ $O(d^2)$ 低维空间
点积注意力 $q^Tk$ $O(d)$ 高维空间
缩放点积注意力 $q^Tk/\sqrt{d}$ $O(d)$ Transformer 默认

PyTorch 实现带 mask 的多头注意力

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

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

        # 线性变换层 (batch_size, seq_len, d_model) -> (batch_size, seq_len, d_model)
        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, q, k, v, mask=None):
        batch_size = q.size(0)

        # 线性变换 + 分头 (batch_size, seq_len, d_model) -> (batch_size, seq_len, n_heads, d_k)
        q = self.w_q(q).view(batch_size, -1, self.n_heads, self.d_k)
        k = self.w_k(k).view(batch_size, -1, self.n_heads, self.d_k)
        v = self.w_v(v).view(batch_size, -1, self.n_heads, self.d_k)

        # 转置为 (batch_size, n_heads, seq_len, d_k)
        q, k, v = q.transpose(1, 2), k.transpose(1, 2), v.transpose(1, 2)

        # 计算缩放点积注意力
        scores = torch.matmul(q, k.transpose(-2, -1)) / torch.sqrt(torch.tensor(self.d_k))

        # 应用 mask(如因果掩码)if mask is not None:
            scores = scores.masked_fill(mask == 0, -1e9)

        # 温度系数调节(训练稳定技巧)temperature = 1.0  # 可动态调整
        attn = F.softmax(scores * temperature, dim=-1)

        # 加权求和 + 合并头
        output = torch.matmul(attn, v)
        output = output.transpose(1, 2).contiguous().view(batch_size, -1, self.d_model)

        return self.w_o(output)

关键实现细节:
温度系数:通过调整 softmax 前的数值尺度控制注意力分布的尖锐程度
数值稳定性:对 masked 位置用 -1e9 替代负无穷,避免 NaN
分头计算:将 d_model 拆分为 n_heads 个 d_k 子空间,并行计算

Transformer 架构实战

在标准 Transformer 中,自注意力层通过以下方式增强模型能力:

  1. 编码器自注意力:建立输入序列全局依赖
  2. 解码器自注意力:结合因果掩码实现自回归生成
  3. 交叉注意力:连接编码器 - 解码器信息流

KV 缓存加速推理

在自回归生成时,可通过缓存历史 Key/Value 避免重复计算:

class DecoderLayer(nn.Module):
    def __init__(self, ...):
        self.self_attn = MultiHeadAttention()
        self.cross_attn = MultiHeadAttention()

    def forward(self, x, encoder_out, past_kv=None):
        # 自注意力(带因果掩码)self_attn_out = self.self_attn(
            q=x, k=x, v=x,
            mask=torch.tril(torch.ones(seq_len, seq_len))  # 下三角掩码
        )

        # 交叉注意力(使用编码器输出)cross_attn_out = self.cross_attn(
            q=self_attn_out,
            k=encoder_out,
            v=encoder_out
        )

        # 更新 KV 缓存
        new_kv = torch.cat([past_kv, current_kv], dim=1) if past_kv is not None else current_kv
        return cross_attn_out, new_kv

生产环境避坑指南

  1. 内存溢出(OOM)
  2. 解决方案:采用梯度检查点 (gradient checkpointing) 或分块注意力
  3. 示例:将序列分成 64-128 的块处理

  4. 训练不稳定

  5. 现象:损失出现 NaN/INF
  6. 对策:

    • 添加 LayerNorm
    • 使用 Xavier 初始化
    • 限制最大序列长度
  7. 长序列性能下降

  8. 优化方案:
    • 稀疏注意力(如 Longformer 的滑动窗口)
    • 线性注意力(Reformer 的 LSH 分桶)

IWSLT 实验对比

模型 BLEU-4 参数量 推理速度(词 / 秒)
LSTM+Attention 28.7 65M 120
Transformer(base) 32.1 65M 310
Transformer(big) 33.5 213M 190

动手挑战

尝试实现 局部窗口注意力
1. 修改注意力计算,使每个 token 只关注前后 $w$ 个邻居
2. 对比全局注意力的效果差异
3. 思考如何平衡局部与全局信息(提示:可参考 Swin Transformer)

参考文献

  1. Vaswani et al. Attention Is All You Need. NeurIPS 2017
  2. Dai et al. Transformer-XL. ACL 2019
  3. Kitaev et al. Reformer. ICLR 2020
正文完
 0
评论(没有评论)