深入解析BERT多头注意力机制:从数学原理到PyTorch实现

1次阅读
没有评论

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

image.webp

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

在自然语言处理领域,传统的 RNN(循环神经网络)在处理长序列时存在明显的局限性。最突出的问题是:

深入解析 BERT 多头注意力机制:从数学原理到 PyTorch 实现

  • 梯度消失 / 爆炸 :随着序列长度的增加,RNN 难以有效捕捉远距离依赖关系
  • 顺序计算 :无法并行处理序列,训练速度受限于序列长度
  • 信息瓶颈 :最后一个隐状态需要压缩整个序列信息

而自注意力机制通过计算词与词之间的关联度,实现了:

  1. 直接建模任意距离的依赖关系
  2. 完全并行的序列计算
  3. 动态权重分配(不同位置关注不同重要性的上下文)

数学原理剖析

基础注意力计算

标准的缩放点积注意力公式为:

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

其中:
– $Q \in \mathbb{R}^{n\times d_k}$(查询矩阵)
– $K \in \mathbb{R}^{m\times d_k}$(键矩阵)
– $V \in \mathbb{R}^{m\times d_v}$(值矩阵)
– $\sqrt{d_k}$ 缩放因子防止内积过大导致 softmax 饱和

多头注意力扩展

将 Q、K、V 通过不同的线性投影拆分成 $h$ 个头:

$$
\begin{aligned}
head_i &= Attention(QW_i^Q, KW_i^K, VW_i^V) \
MultiHead(Q,K,V) &= Concat(head_1,…,head_h)W^O
\end{aligned}
$$

每个头的维度通常为 $d_{model}/h$,这样拼接后能保持总维度不变。

PyTorch 实现详解

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

class MultiHeadAttention(nn.Module):
    def __init__(self, d_model=512, num_heads=8):
        super().__init__()
        assert d_model % num_heads == 0, "d_model 必须能被 num_heads 整除"

        self.d_model = d_model
        self.num_heads = num_heads
        self.d_head = d_model // num_heads

        # 定义 QKV 的线性变换层
        self.wq = nn.Linear(d_model, d_model)  # [d_model, d_model]
        self.wk = nn.Linear(d_model, d_model)
        self.wv = nn.Linear(d_model, d_model)

        # 最终输出的线性层
        self.wo = nn.Linear(d_model, d_model)

    def forward(self, x, mask=None):
        """
        输入 x: [batch_size, seq_len, d_model]
        输出: [batch_size, seq_len, d_model]
        """
        batch_size, seq_len, _ = x.shape

        # 1. 计算 QKV [batch_size, seq_len, d_model]
        Q = self.wq(x)
        K = self.wk(x)
        V = self.wv(x)

        # 2. 拆分为多头 [batch_size, num_heads, seq_len, d_head]
        Q = Q.view(batch_size, -1, self.num_heads, self.d_head).transpose(1, 2)
        K = K.view(batch_size, -1, self.num_heads, self.d_head).transpose(1, 2)
        V = V.view(batch_size, -1, self.num_heads, self.d_head).transpose(1, 2)

        # 3. 计算缩放点积注意力 [batch_size, num_heads, seq_len, seq_len]
        attn_scores = torch.matmul(Q, K.transpose(-2, -1)) / torch.sqrt(torch.tensor(self.d_head))

        # 应用 mask(如需要)if mask is not None:
            attn_scores = attn_scores.masked_fill(mask == 0, -1e9)

        # softmax 归一化
        attn_weights = F.softmax(attn_scores, dim=-1)

        # 4. 计算注意力输出 [batch_size, num_heads, seq_len, d_head]
        attn_output = torch.matmul(attn_weights, V)

        # 5. 合并多头 [batch_size, seq_len, d_model]
        attn_output = attn_output.transpose(1, 2).contiguous()
        attn_output = attn_output.view(batch_size, seq_len, self.d_model)

        # 6. 最终线性变换
        output = self.wo(attn_output)
        return output

工程实践关键点

头数选择经验

头数 $h$ 与显存占用的关系近似为:

$$
显存 \approx batch_size \times seq_len^2 \times h \times \frac{d_{model}}{h}
$$

实践中建议:

  1. 基础模型(d_model=512):8 个头
  2. 大模型(d_model=1024):16 个头
  3. 头维度不应小于 64(保证每个头有足够表达能力)

常见陷阱规避

  1. 维度不匹配
  2. 确保 d_model % num_heads == 0
  3. 拼接前检查各头维度是否一致

  4. 掩码应用错误

  5. Decoder 需要严格的下三角掩码(因果注意力)
  6. Padding 掩码应在 softmax 前应用

  7. 梯度不稳定

  8. 使用缩放因子 $1/\sqrt{d_k}$
  9. 初始化时适当减小线性层权重

延伸思考方向

  1. 多头机制的普适性
  2. 在浅层网络(如 3 层 Transformer)中,多头是否仍优于单头?
  3. 不同层是否需要不同数量的注意力头?

  4. 注意力头可解释性

  5. 能否通过聚类等方法量化不同头捕获的语义特征?
  6. 特定头是否专门处理语法 / 语义等不同层面的信息?

实践建议

在真实业务场景中使用多头注意力时,建议:

  1. 先用小批量数据验证维度变换的正确性
  2. 使用 TensorBoard 可视化注意力权重分布
  3. 对长序列考虑内存优化的稀疏注意力实现

多头注意力机制是 Transformer 架构的核心创新,理解其实现细节能帮助我们更高效地调试模型,也能启发新的结构改进思路。

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