BERT多头注意力机制入门指南:从理论到PyTorch实现

1次阅读
没有评论

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

image.webp

为什么需要注意力机制

在传统的 RNN 处理序列数据时,存在两个主要问题:一是长距离依赖难以捕捉,二是无法并行计算。Transformer 架构通过自注意力机制完美解决了这两个痛点。自注意力层能够直接计算序列中任意两个位置的关系权重,不受距离限制,且所有位置的注意力权重可以并行计算。

BERT 多头注意力机制入门指南:从理论到 PyTorch 实现

BERT 作为 Transformer 的典型代表,其核心正是多头注意力机制。这种机制可以理解为让模型同时从多个不同的角度(即多个头)来关注输入序列的不同部分,从而捕获更丰富的语义信息。

数学原理拆解

自注意力机制的核心计算涉及三个关键矩阵:查询矩阵 Q(Query)、键矩阵 K(Key)和值矩阵 V(Value)。它们的计算过程如下:

  1. 线性变换
    $$Q = XW_Q, \quad K = XW_K, \quad V = XW_V$$
    其中 $X \in \mathbb{R}^{n\times d_{model}}$ 是输入序列,$W_Q,W_K,W_V \in \mathbb{R}^{d_{model}\times d_k}$ 是可学习参数矩阵

  2. 注意力得分
    $$Attention(Q,K,V) = softmax(\frac{QK^T}{\sqrt{d_k}})V$$
    这里除以 $\sqrt{d_k}$ 是为了防止点积结果过大导致 softmax 梯度消失

  3. 多头扩展
    将 Q、K、V 分别拆分为 h 个头(BERT-base 是 12 个头),每个头独立计算注意力:
    $$MultiHead = Concat(head_1,…,head_h)W^O$$
    $$where \ head_i = Attention(QW_i^Q, KW_i^K, VW_i^V)$$

单头 vs 多头效率分析

  • 单头注意力
    时间复杂度 $O(n^2 \cdot d)$,其中 n 是序列长度,d 是特征维度

  • 多头注意力
    虽然看起来计算量增加了 h 倍,但因为每个头的维度降为 $d/h$,实际总复杂度仍为 $O(n^2 \cdot d)$

关键优势在于:
1. 参数矩阵被分解到多个头,更容易训练
2. 不同头可以关注不同位置的模式(如语法、指代等)
3. 在 GPU 上可以完美并行计算

PyTorch 实现详解

1. 基础模块定义

import torch
import torch.nn as nn
import math

class MultiHeadAttention(nn.Module):
    def __init__(self, d_model=768, n_heads=12, dropout=0.1):
        super().__init__()
        assert d_model % n_heads == 0
        self.d_k = d_model // n_heads
        self.n_heads = n_heads

        # 线性变换矩阵
        self.wq = nn.Linear(d_model, d_model)  # (768,768)
        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)

2. 多头拆分与注意力计算

def split_heads(self, x):
    # x 形状: (batch, seq_len, d_model)
    batch_size = x.size(0)
    return x.view(batch_size, -1, self.n_heads, self.d_k).transpose(1, 2)
    # 输出形状: (batch, n_heads, seq_len, d_k)

def forward(self, q, k, v, mask=None):
    # 1. 线性变换并分头
    q = self.split_heads(self.wq(q))  # (batch, heads, q_len, d_k)
    k = self.split_heads(self.wk(k))
    v = self.split_heads(self.wv(v))

    # 2. 计算缩放点积注意力
    attn = torch.matmul(q, k.transpose(-2, -1)) * self.scale
    if mask is not None:
        attn = attn.masked_fill(mask == 0, -1e10)
    attn = self.dropout(torch.softmax(attn, dim=-1))

    # 3. 加权求和并合并多头
    output = torch.matmul(attn, v)  # (batch, heads, q_len, d_k)
    output = output.transpose(1, 2).contiguous() \
             .view(output.size(0), -1, self.d_model)
    return self.wo(output)

3. 残差连接与层归一化

class TransformerBlock(nn.Module):
    def __init__(self, d_model, n_heads, dropout=0.1):
        super().__init__()
        self.attention = MultiHeadAttention(d_model, n_heads)
        self.norm1 = nn.LayerNorm(d_model)
        self.norm2 = nn.LayerNorm(d_model)
        self.ffn = nn.Sequential(nn.Linear(d_model, 4*d_model),
            nn.GELU(),
            nn.Linear(4*d_model, d_model)
        )
        self.dropout = nn.Dropout(dropout)

    def forward(self, x, mask=None):
        # 注意力子层
        attn_out = self.attention(x, x, x, mask)
        x = self.norm1(x + self.dropout(attn_out))

        # 前馈子层
        ffn_out = self.ffn(x)
        return self.norm2(x + self.dropout(ffn_out))

实战避坑指南

梯度消失问题

  1. 初始化技巧
  2. 使用 Xavier 初始化注意力层的权重
  3. 值矩阵 V 的初始方差应设为 $1/\sqrt{d_k}$

  4. 学习率预热

    optimizer = AdamW(model.parameters(), lr=5e-5)
    scheduler = get_linear_schedule_with_warmup(
        optimizer, 
        num_warmup_steps=1000, 
        num_training_steps=total_steps
    )

注意力掩码应用

  • Padding 掩码 :处理变长序列时屏蔽无效位置

    # seq_len=512, valid_len= 实际长度
    mask = torch.arange(seq_len)[None, :] < valid_len[:, None]  

  • 因果掩码 :防止解码器看到未来信息

    mask = torch.triu(torch.ones(seq_len, seq_len), diagonal=1).bool()

显存优化

  1. 梯度检查点

    from torch.utils.checkpoint import checkpoint
    def custom_forward(*inputs):
        x, mask = inputs
        return transformer_block(x, mask)
    
    output = checkpoint(custom_forward, x, mask)

  2. 混合精度训练

    scaler = GradScaler()
    with autocast():
        outputs = model(inputs)
        loss = criterion(outputs, labels)
    scaler.scale(loss).backward()
    scaler.step(optimizer)
    scaler.update()

可视化与分析

我准备了一个 Colab notebook,可以直观观察不同注意力头的聚焦模式: 查看 Notebook

示例可视化代码:

import seaborn as sns
import matplotlib.pyplot as plt

def plot_attention(attention_weights, layer_idx, head_idx):
    plt.figure(figsize=(10,8))
    sns.heatmap(attention_weights[layer_idx][head_idx].cpu().detach())
    plt.title(f"Layer {layer_idx} Head {head_idx}")
    plt.show()

思考与延伸

最后留给读者一个思考题:如何设计实验验证不同注意力头确实捕获了不同的语义信息?这里给出两个思路方向:

  1. 模式分析 :对同一层的不同头计算注意力权重的互信息,如果差异大则说明分工明确
  2. 消融实验 :冻结某些头的参数,观察模型在不同 NLP 任务上的表现变化

多头注意力机制就像团队合作,每个成员(注意力头)各司其职,有的关注局部语法,有的把握全局语义。希望通过本文的讲解,你能真正理解并驾驭这一强大工具。

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