BERT预训练模型计算公式详解:从数学原理到代码实现

1次阅读
没有评论

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

image.webp

背景介绍

BERT(Bidirectional Encoder Representations from Transformers)是自然语言处理领域的重要里程碑,通过预训练 - 微调范式显著提升了各类 NLP 任务的表现。其核心在于 Transformer 架构和掩码语言建模(MLM)任务,其中预训练阶段的计算公式直接决定了模型捕获语义信息的能力。理解这些公式对调试模型、优化性能至关重要。

BERT 预训练模型计算公式详解:从数学原理到代码实现

数学原理

1. 输入嵌入层

BERT 的输入是词嵌入(Token Embeddings)、段嵌入(Segment Embeddings)和位置嵌入(Position Embeddings)的总和:

$$\text{Embedding} = E_w + E_s + E_p$$

其中:
– $E_w \in \mathbb{R}^{d_{model}}$ 是词 ID 映射得到的嵌入
– $E_s$ 区分句子 A /B(单句任务时为全 0)
– $E_p$ 使用固定位置编码公式:

$$PE_{(pos,2i)} = \sin(pos/10000^{2i/d_{model}}})$$
$$PE_{(pos,2i+1)} = \cos(pos/10000^{2i/d_{model}}})$$

2. 多头注意力机制

这是 BERT 的核心组件,计算分为四步:

  1. 线性投影 :将输入 $X$ 分别映射为 Q /K/V
    $$Q = XW_Q, K = XW_K, V = XW_V$$
    ($W_Q,W_K,W_V \in \mathbb{R}^{d_{model} \times d_k}$)

  2. 缩放点积注意力
    $$\text{Attention}(Q,K,V) = \text{softmax}(\frac{QK^T}{\sqrt{d_k}})V$$

  3. 多头拼接 :$h$ 个注意力头的结果拼接后线性变换
    $$\text{MultiHead} = \text{Concat}(head_1,…,head_h)W_O$$

  4. 残差连接
    $$\text{Output} = \text{LayerNorm}(X + \text{Dropout}(\text{MultiHead}))$$

3. 前馈网络层

包含两个全连接层和 GELU 激活:
$$\text{FFN}(x) = \text{GELU}(xW_1 + b_1)W_2 + b_2$$
其中中间维度通常扩大 4 倍(如 BERT-base 中 $d_{ff}=3072$)

4. Layer Normalization

对特征维度进行标准化:
$$\mu = \frac{1}{d}\sum_{i=1}^d x_i$$
$$\sigma = \sqrt{\frac{1}{d}\sum_{i=1}^d (x_i-\mu)^2}$$
$$\text{LayerNorm}(x) = \gamma \cdot \frac{x-\mu}{\sigma + \epsilon} + \beta$$

代码实现(PyTorch)

import torch
import torch.nn as nn
import math

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

        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)

        self.dropout = nn.Dropout(dropout)
        self.layer_norm = nn.LayerNorm(d_model)

    def forward(self, x, mask=None):
        residual = x
        batch_size = x.size(0)

        # 1. 线性投影
        Q = self.W_Q(x).view(batch_size, -1, self.num_heads, self.d_k).transpose(1,2)
        K = self.W_K(x).view(batch_size, -1, self.num_heads, self.d_k).transpose(1,2)
        V = self.W_V(x).view(batch_size, -1, self.num_heads, self.d_k).transpose(1,2)

        # 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)
        attn = torch.softmax(scores, dim=-1)
        attn = self.dropout(attn)
        context = torch.matmul(attn, V)

        # 3. 多头拼接
        context = context.transpose(1,2).contiguous().view(batch_size, -1, self.num_heads * self.d_k)
        output = self.W_O(context)

        # 4. 残差连接
        return self.layer_norm(residual + self.dropout(output))

性能考量

计算复杂度

  • 注意力机制:$O(n^2 \cdot d)$(n 为序列长度)
  • 前馈网络:$O(n \cdot d^2)$
  • 实际训练中,长序列处理是主要瓶颈

内存占用

  • 主要消耗在注意力矩阵:$batch \times heads \times seq_len^2$
  • BERT-base 处理 512 长度序列时,单样本约需 1GB 显存

避坑指南

  1. 注意力分数溢出
  2. 解决方法:确保除以 $\sqrt{d_k}$,使用混合精度训练时需监控 softmax 输入值范围

  3. 层归一化位置

  4. 原始 Transformer 在残差前做归一化,BERT 改为残差后(Post-LN)
  5. 错误实现会导致梯度消失

  6. 激活函数选择

  7. BERT 使用 GELU 而非 ReLU,误用会导致性能下降约 1 - 2 个点
    $$\text{GELU}(x) = x\Phi(x)$$

总结与思考

这些计算公式直接影响模型捕获上下文信息的能力:
– 注意力机制决定 token 间交互方式
– 层归一化影响训练稳定性
– 前馈网络提供非线性变换能力

优化方向:
– 稀疏注意力降低计算复杂度
– 参数共享减少模型体积
– 量化感知训练提升推理速度

开放问题:
1. 如何设计更高效的注意力计算方式来处理长文档?
2. 预训练任务(MLM/NSP)与这些计算公式如何协同影响最终表现?
3. 在不同语言任务中,哪些公式参数最值得调整优化?

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