BERT自注意力机制深度解析:从数学原理到PyTorch实现

1次阅读
没有评论

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

image.webp

背景痛点:为什么需要自注意力机制

传统 RNN(循环神经网络)在处理长序列时面临两个核心问题:

BERT 自注意力机制深度解析:从数学原理到 PyTorch 实现

  1. 梯度消失 / 爆炸:随着序列长度增加,反向传播时梯度会指数级衰减或增长,导致模型难以训练。实验表明,当序列长度超过 50 时,LSTM 的准确率会下降约 30%
  2. 顺序计算限制:必须严格按时间步顺序计算,无法充分利用现代 GPU 的并行计算能力。在处理 1000 个 token 的序列时,RNN 的计算速度比 Transformer 慢约 200 倍

Transformer 架构通过自注意力机制 (Self-Attention) 解决了这些问题:

  • 任意两个 token 间可直接建立联系,最大路径长度仅为 O(1)
  • 计算过程天然适合并行化,理论 FLOPs 利用率可达 80% 以上

数学原理:缩放点积注意力

自注意力的核心计算公式如下:

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

其中各变量的物理意义:

  • $Q$ (Query): 查询向量,形状为 $[L_q, d_k]$
  • $K$ (Key): 键向量,形状为 $[L_k, d_k]$
  • $V$ (Value): 值向量,形状为 $[L_k, d_v]$
  • $\sqrt{d_k}$: 缩放因子,防止点积结果过大导致 softmax 梯度消失

关键数学推导步骤:

  1. 计算相似度矩阵:$S = QK^T$(形状 $[L_q, L_k]$)
  2. 缩放处理:$S’ = S/\sqrt{d_k}$
  3. 归一化:$A = \text{softmax}(S’)$(注意力权重)
  4. 加权求和:$O = AV$(形状 $[L_q, d_v]$)

PyTorch 实现详解

基础注意力模块

import torch
import torch.nn as nn
import torch.nn.functional as F
from einops import rearrange

class ScaledDotProductAttention(nn.Module):
    """
    输入形状:
        q: [batch_size, n_heads, L_q, d_k]
        k: [batch_size, n_heads, L_k, d_k] 
        v: [batch_size, n_heads, L_k, d_v]
        mask: [batch_size, L_q, L_k]
    """
    def __init__(self, dropout=0.1):
        super().__init__()
        self.dropout = nn.Dropout(dropout)

    def forward(self, q, k, v, mask=None):
        # 计算点积注意力分数 [batch_size, n_heads, L_q, L_k]
        scores = torch.matmul(q, k.transpose(-2, -1)) / (q.size(-1) ** 0.5)

        # 应用 mask(padding mask 或 sequence mask)if mask is not None:
            scores = scores.masked_fill(mask == 0, -1e9)

        # softmax 归一化
        attn = F.softmax(scores, dim=-1)
        attn = self.dropout(attn)

        # 加权求和 [batch_size, n_heads, L_q, d_v]
        output = torch.matmul(attn, v)
        return output

多头注意力整合

class MultiHeadAttention(nn.Module):
    def __init__(self, d_model=512, n_heads=8, dropout=0.1):
        assert d_model % n_heads == 0, "d_model must be divisible by n_heads"

        super().__init__()
        self.d_k = d_model // n_heads
        self.n_heads = n_heads

        # 线性变换层
        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)

        self.attention = ScaledDotProductAttention(dropout)
        self.dropout = nn.Dropout(dropout)
        self.layer_norm = nn.LayerNorm(d_model)

    def forward(self, q, k, v, mask=None):
        residual = q

        # 线性变换并分头 [batch_size, L, n_heads, d_k]
        q = rearrange(self.w_q(q), 'b l (h d) -> b h l d', h=self.n_heads)
        k = rearrange(self.w_k(k), 'b l (h d) -> b h l d', h=self.n_heads)
        v = rearrange(self.w_v(v), 'b l (h d) -> b h l d', h=self.n_heads)

        # 计算注意力
        output = self.attention(q, k, v, mask=mask)

        # 合并多头 [batch_size, L, d_model]
        output = rearrange(output, 'b h l d -> b l (h d)')
        output = self.w_o(output)

        # 残差连接 +LayerNorm
        output = self.dropout(output)
        output = self.layer_norm(output + residual)

        return output

显存优化技巧

  1. 梯度检查点
    from torch.utils.checkpoint import checkpoint
    output = checkpoint(self.attention, q, k, v, mask)
  2. 混合精度训练
    with torch.cuda.amp.autocast():
        output = model(inputs)
  3. 序列分块处理:当序列长度 >512 时,可采用分块计算注意力

常见错误与解决方案

  1. 忘记 LayerNorm
  2. 现象:模型训练不稳定,loss 震荡
  3. 修正:确保每个子层都有残差连接 +LayerNorm

  4. Mask 应用错误

  5. 现象:模型在验证集表现异常
  6. 检查:确保 padding mask 正确传递给每一层

  7. Dropout 设置不当

  8. 建议:attention dropout 通常设 0.1,hidden dropout 设 0.3

延伸思考

  1. 位置编码如何影响注意力机制的效果?能否用相对位置编码替代绝对位置编码?
  2. 当序列长度达到 1024 时,如何优化注意力矩阵的内存占用?
  3. 多头注意力的头数是否越多越好?如何设计实验验证最优头数?

完整实现可参考 Colab Notebook:点击访问

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