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

1次阅读
没有评论

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

image.webp

背景痛点:序列建模的进化之路

在 Transformer 出现之前,循环神经网络 (RNN) 是处理序列数据的首选方案。但 RNN 存在明显的缺陷:

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

  • 梯度消失问题:随着序列长度增加,反向传播时梯度会指数级衰减,导致模型难以学习长距离依赖关系
  • 串行计算限制:必须按时间步顺序计算,无法充分利用 GPU 的并行计算能力

卷积神经网络 (CNN) 虽然可以通过并行计算缓解这个问题,但也有其局限性:

  • 局部感知野:需要堆叠多层卷积才能捕获全局信息
  • 固定模式识别:卷积核权重在不同位置共享,难以适应序列中的动态模式

数学原理:Scaled Dot-Product Attention

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

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

其中:
– $Q$ (Query/ 查询向量): 当前关注的 token 表示
– $K$ (Key/ 键向量): 用于匹配查询的上下文表示
– $V$ (Value/ 值向量): 实际提取的特征信息
– $d_k$: 键向量的维度

缩放因子 $\sqrt{d_k}$ 的作用
1. 防止点积结果过大导致 softmax 进入饱和区
2. 保持梯度稳定,有利于模型训练
3. 使注意力分布更加分散,增强模型表达能力

PyTorch 实现:模块化 MultiHeadAttention

import torch
import torch.nn as nn
import math

class MultiHeadAttention(nn.Module):
    """
    多头注意力机制实现
    Args:
        embed_dim: 输入 token 的维度
        num_heads: 注意力头数量
        dropout: 注意力权重 dropout 率
        bias: 是否在投影层使用偏置
        kdim: 键向量的维度(默认等于 embed_dim)
        vdim: 值向量的维度(默认等于 embed_dim)
    """
    def __init__(self, embed_dim: int, num_heads: int, 
                 dropout: float = 0.1, bias: bool = True,
                 kdim: int = None, vdim: int = None):
        super().__init__()
        self.embed_dim = embed_dim
        self.kdim = kdim if kdim is not None else embed_dim
        self.vdim = vdim if vdim is not None else embed_dim
        self.num_heads = num_heads
        self.head_dim = embed_dim // num_heads

        assert self.head_dim * num_heads == embed_dim, "embed_dim 必须能被 num_heads 整除"

        # 线性投影层
        self.q_proj = nn.Linear(embed_dim, embed_dim, bias=bias)
        self.k_proj = nn.Linear(self.kdim, embed_dim, bias=bias)
        self.v_proj = nn.Linear(self.vdim, embed_dim, bias=bias)

        # 输出投影层
        self.out_proj = nn.Linear(embed_dim, embed_dim, bias=bias)

        # Dropout 层
        self.dropout = nn.Dropout(dropout)

        # KV 缓存
        self.cache_k = None
        self.cache_v = None

    def forward(self, query: torch.Tensor, key: torch.Tensor, value: torch.Tensor,
                key_padding_mask: torch.Tensor = None, attn_mask: torch.Tensor = None,
                need_weights: bool = True, is_causal: bool = False):
        """
        前向传播计算
        Args:
            query: [batch_size, tgt_len, embed_dim]
            key: [batch_size, src_len, kdim]
            value: [batch_size, src_len, vdim]
            key_padding_mask: [batch_size, src_len]
            attn_mask: [tgt_len, src_len]
        """
        batch_size = query.size(0)
        tgt_len = query.size(1)
        src_len = key.size(1)

        # 线性投影
        q = self.q_proj(query)  # [B, T, E]
        k = self.k_proj(key)    # [B, S, E]
        v = self.v_proj(value)  # [B, S, E]

        # 重塑为多头形式
        q = q.view(batch_size, tgt_len, self.num_heads, self.head_dim).transpose(1, 2)  # [B, H, T, D]
        k = k.view(batch_size, src_len, self.num_heads, self.head_dim).transpose(1, 2)  # [B, H, S, D]
        v = v.view(batch_size, src_len, self.num_heads, self.head_dim).transpose(1, 2)  # [B, H, S, D]

        # 计算注意力分数
        attn_scores = torch.matmul(q, k.transpose(-2, -1)) / math.sqrt(self.head_dim)  # [B, H, T, S]

        # 应用注意力掩码
        if attn_mask is not None:
            attn_scores = attn_scores + attn_mask.unsqueeze(0).unsqueeze(0)

        if key_padding_mask is not None:
            attn_scores = attn_scores.masked_fill(key_padding_mask.unsqueeze(1).unsqueeze(2),
                float('-inf')
            )

        # 计算注意力权重
        attn_weights = torch.softmax(attn_scores, dim=-1)
        attn_weights = self.dropout(attn_weights)

        # 计算上下文向量
        context = torch.matmul(attn_weights, v)  # [B, H, T, D]
        context = context.transpose(1, 2).contiguous().view(batch_size, tgt_len, self.embed_dim)  # [B, T, E]

        # 输出投影
        output = self.out_proj(context)

        if need_weights:
            return output, attn_weights
        return output, None

    def update_kv_cache(self, key: torch.Tensor, value: torch.Tensor):
        """更新 KV 缓存"""
        if self.cache_k is None:
            self.cache_k = key
            self.cache_v = value
        else:
            self.cache_k = torch.cat([self.cache_k, key], dim=1)
            self.cache_v = torch.cat([self.cache_v, value], dim=1)
        return self.cache_k, self.cache_v

性能优化:Flash Attention 算法

Flash Attention 通过优化 GPU 内存访问模式显著提升计算效率:

  1. 分块计算(Tiling):将大型注意力矩阵分割为适合 SRAM 的小块
  2. 内存层次利用
  3. 全局内存 (HBM) 存储完整矩阵
  4. 共享内存 (SRAM) 缓存当前计算块
  5. 寄存器存储中间结果
  6. 重计算机制:在前向传播中丢弃中间结果,反向传播时重新计算
graph LR
    A[输入 Q,K,V 矩阵] --> B[分块加载到 SRAM]
    B --> C[计算局部注意力]
    C --> D[聚合全局结果]
    D --> E[写回 HBM]

避坑指南:实际部署经验

  1. 头维度分配策略
  2. 均等分割:每个头维度相同,实现简单
  3. 非对称分割:根据任务调整头维度,可能提升效果但增加实现复杂度

  4. 混合精度训练

  5. 使用 torch.cuda.amp 自动混合精度
  6. 注意 softmax 计算时的数值稳定性
  7. 建议对注意力分数进行最大减法(max subtraction)

  8. KV 缓存优化

  9. 增量解码时重用之前计算的 KV
  10. 注意内存管理,避免缓存无限增长

延伸思考:开放性问题

  1. 线性注意力 (Linear Attention) 在牺牲精确度的前提下获得 O(N)复杂度,这种权衡在不同应用场景中如何评估?
  2. 如何设计动态头机制,使模型能够根据输入内容自适应调整各头的关注范围?
  3. 在边缘设备部署时,有哪些量化压缩方法可以在保持模型性能的同时减少自注意力模块的计算开销?
正文完
 0
评论(没有评论)