BERT模型与自注意力机制:从零开始理解Transformer核心原理

1次阅读
没有评论

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

image.webp

为什么需要 Transformer?

在自然语言处理(NLP)领域,传统的 RNN 和 LSTM 模型曾长期占据主导地位。但这些模型存在两个致命缺陷:

BERT 模型与自注意力机制:从零开始理解 Transformer 核心原理

  • 顺序计算限制:必须逐个处理序列中的词元,无法充分利用 GPU 并行计算能力
  • 长程依赖丢失:随着序列长度增加,早期时间步的信息会逐渐衰减(即便 LSTM 有所改善,但问题依然存在)

2017 年 Google 提出的 Transformer 架构彻底改变了这一局面。BERT(Bidirectional Encoder Representations from Transformers)作为其典型代表,通过完全基于注意力机制的设计,实现了:

  1. 真正的双向上下文建模
  2. 可并行计算的序列处理
  3. 对长距离依赖关系的直接捕捉

自注意力机制详解

核心思想:动态权重分配

自注意力(Self-Attention)的本质是让序列中的每个词元都能 ” 看到 ” 其他所有词元,并动态决定关注哪些部分。这个过程通过三个关键向量实现:

  • 查询向量(Query):当前词元想要查询什么信息
  • 键向量(Key):其他词元能提供什么信息
  • 值向量(Value):实际传递的信息内容

计算过程分步拆解

  1. 输入表示:假设输入序列有 n 个词元,每个词元的嵌入维度是 d。则输入矩阵 X ∈ ℝ^(n×d)

  2. 线性变换:通过可学习的权重矩阵生成 Q /K/V

    # PyTorch 实现(假设 d_model=512)W_Q = nn.Linear(512, 64)  # 通常使 dk=64
    W_K = nn.Linear(512, 64)
    W_V = nn.Linear(512, 64)
    
    Q = W_Q(X)  # (n, 64)
    K = W_K(X)  # (n, 64)
    V = W_V(X)  # (n, 64)

  3. 注意力分数计算(Scaled Dot-Product Attention):

    Attention(Q,K,V) = softmax(QK^T/√dk)V

    这里除以√dk(键向量维度)是为了防止点积结果过大导致 softmax 梯度消失

  4. 多头注意力:将这个过程并行执行 h 次(BERT-base 中 h =12),拼接后通过线性层融合

维度变化可视化

以 ” 我爱自然语言处理 ” 这句话为例(n=7,d=512,h=12):

  1. 输入 X: (7, 512)
  2. 单头 Q /K/V: (7, 64)
  3. 注意力矩阵: (7, 7) ← 每个单元格表示词元间关联度
  4. 多头拼接: (7, 768) ← 12 头×64 维
  5. 输出投影: (7, 512) ← 保持维度一致

关键实现细节

位置编码(Positional Encoding)

由于自注意力本身没有位置概念,必须显式添加位置信息。常用正弦 / 余弦函数生成:

def positional_encoding(max_len, d_model):
    position = torch.arange(max_len).unsqueeze(1)
    div_term = torch.exp(torch.arange(0, d_model, 2) * -(math.log(10000.0) / d_model))
    pe = torch.zeros(max_len, d_model)
    pe[:, 0::2] = torch.sin(position * div_term)
    pe[:, 1::2] = torch.cos(position * div_term)
    return pe

常见错误:直接使用可学习的位置嵌入(learned positional embeddings)时,如果训练数据长度不足,在推理长文本时会出现未初始化位置的问题

注意力掩码(Attention Mask)

两种主要应用场景:

  1. 填充掩码(Padding Mask):遮盖无效的 padding 位置(通常在 batch 处理时使用)
  2. 因果掩码(Causal Mask):防止解码器看到未来信息(BERT 不需要,但 GPT 类模型需要)
# 生成上三角矩阵(用于解码器)mask = torch.triu(torch.ones(seq_len, seq_len), diagonal=1).bool()
attention_scores.masked_fill_(mask, -1e9)

性能优化技巧

  1. 内存优化:当序列长度 >512 时,可以考虑:
  2. 使用稀疏注意力(如 Longformer 的滑动窗口模式)
  3. 分块计算(Reformer 的 LSH 注意力)

  4. 计算加速

  5. 使用 torch.baddbmm 替代矩阵乘法链
  6. 混合精度训练(AMP)

动手实验

推荐在 Colab 上运行这个简化版的多头注意力实现:

import torch
import torch.nn as nn
import math

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

        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)

    def forward(self, x, mask=None):
        batch_size, seq_len, _ = x.shape

        # 线性变换并分头 [batch, seq_len, h, d_k]
        Q = self.W_Q(x).view(batch_size, seq_len, self.h, self.d_k)
        K = self.W_K(x).view(batch_size, seq_len, self.h, self.d_k)
        V = self.W_V(x).view(batch_size, seq_len, self.h, self.d_k)

        # 转置为[batch, h, seq_len, d_k]
        Q, K, V = Q.transpose(1, 2), K.transpose(1, 2), V.transpose(1, 2)

        # 计算缩放点积注意力
        scores = torch.matmul(Q, K.transpose(-2, -1)) / math.sqrt(self.d_k)
        if mask is not None:
            scores = scores.masked_fill(mask == 0, -1e9)
        attention = torch.softmax(scores, dim=-1)

        # 输出拼接和投影
        output = torch.matmul(attention, V)
        output = output.transpose(1, 2).contiguous().view(batch_size, seq_len, -1)
        return self.W_O(output)

实践建议

  1. 调试时可以可视化注意力矩阵,观察模型关注的重点是否合理
  2. 开始时使用小头数(如 4 头),逐渐增加观察效果变化
  3. 在自定义任务中,尝试调整 d_k 维度(通常 64-256 之间)

通过理解这些基本原理,你就能更好地调整 BERT 模型适应自己的任务。下次我们将深入探讨 BERT 的预训练技巧和微调策略。

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