深入解析BERT多头自注意力机制:从理论到PyTorch代码实现

1次阅读
没有评论

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

image.webp

为什么自注意力是 Transformer 的核心?

自注意力机制通过动态计算 token 间关联权重,彻底解决了 RNN 的长程依赖问题。它允许模型直接捕获任意位置的关系,为并行计算提供基础架构。正是这种特性让 Transformer 在捕捉复杂语义模式时展现出惊人效果。

深入解析 BERT 多头自注意力机制:从理论到 PyTorch 代码实现

数学原理拆解

缩放点积注意力公式

核心计算公式如下:
$$Attention(Q,K,V) = softmax(\frac{QK^T}{\sqrt{d_k}})V$$

  • $Q/K/V$ 分别表示查询 (Query)、键(Key)、值(Value) 矩阵
  • $d_k$ 是 key 向量的维度,缩放因子 $\sqrt{d_k}$ 用于防止点积结果过大导致 softmax 梯度消失
  • 计算过程可分为三步:
  • 计算 Q 与 K 的点积得到相似度分数
  • 缩放分数并做 softmax 归一化
  • 用注意力权重加权求和 V 矩阵

多头机制实现原理

多头注意力的关键在于:
$$MultiHead = Concat(head_1,…,head_h)W^O$$
其中每个头的计算为:
$$head_i = Attention(QW_i^Q,KW_i^K,VW_i^V)$$

  • 通过将 $Q/K/V$ 投影到 $h$ 个不同子空间(通常 $h=8$ 或 $12$)
  • 每个头学习不同的注意力模式(如局部 / 全局、语法 / 语义特征)
  • 最后拼接所有头输出并通过线性层 $W^O$ 融合

PyTorch 实战实现

import torch
import torch.nn as nn
import math

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

        # 线性变换层
        self.wq = nn.Linear(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)

        self.dropout = nn.Dropout(dropout)
        self.scale = 1 / math.sqrt(self.d_k)

    def forward(self, q, k, v, mask=None):
        # q/k/v shape: [batch, seq_len, d_model]
        batch_size = q.size(0)

        # 线性投影 + 分头 [batch, seq_len, num_heads, d_k]
        q = self.wq(q).view(batch_size, -1, self.num_heads, self.d_k)
        k = self.wk(k).view(batch_size, -1, self.num_heads, self.d_k)
        v = self.wv(v).view(batch_size, -1, self.num_heads, self.d_k)

        # 转置为 [batch, num_heads, seq_len, d_k]
        q, k, v = q.transpose(1, 2), k.transpose(1, 2), v.transpose(1, 2)

        # 计算注意力分数 [batch, num_heads, q_len, k_len]
        scores = torch.matmul(q, k.transpose(-2, -1)) * self.scale

        # 掩码处理(padding/sequence mask)if mask is not None:
            scores = scores.masked_fill(mask == 0, -1e9)

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

        # 加权求和 [batch, num_heads, seq_len, d_k]
        output = torch.matmul(attn, v)

        # 拼接所有头 [batch, seq_len, d_model]
        output = output.transpose(1, 2).contiguous()
        output = output.view(batch_size, -1, self.num_heads * self.d_k)

        return self.wo(output)

关键维度说明
– 输入输出保持 [batch, seq_len, d_model] 统一维度
– 分头后每个头的维度为d_k = d_model // num_heads
– 注意力分数矩阵形状为[batch, num_heads, q_len, k_len]

性能优化策略

计算复杂度分析

  • 时间复杂度:$O(n^2 \cdot d)$(n 为序列长度)
  • 空间复杂度:$O(n^2)$(需存储注意力矩阵)

优化方案
1. 滑动窗口注意力:限制每个 token 只关注局部邻域
2. 内存优化:
– 梯度检查点(gradient checkpointing)
– KV 缓存(解码时重复利用已计算的 K /V)

常见问题解决方案

梯度爆炸预防

  • 在残差连接前使用 LayerNorm(Post-LN 结构)
  • 初始化时缩小线性层权重范围

混合精度训练

with torch.cuda.amp.autocast():
    # 前向计算时自动转为 FP16
    output = attention_layer(q, k, v)

# 损失计算需保持 FP32
loss = loss_fn(output.float(), target)

思考题

  1. 注意力头差异化验证
  2. 可视化各头的注意力分布热力图
  3. 计算不同头注意力矩阵的相似度

  4. 长序列处理方法

  5. 位置编码外推(如 ALiBi)
  6. 动态稀疏注意力(如 Longformer 的局部 + 全局注意力)
正文完
 0
评论(没有评论)