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

数学原理
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 的核心组件,计算分为四步:
-
线性投影 :将输入 $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}$) -
缩放点积注意力 :
$$\text{Attention}(Q,K,V) = \text{softmax}(\frac{QK^T}{\sqrt{d_k}})V$$ -
多头拼接 :$h$ 个注意力头的结果拼接后线性变换
$$\text{MultiHead} = \text{Concat}(head_1,…,head_h)W_O$$ -
残差连接 :
$$\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 显存
避坑指南
- 注意力分数溢出 :
-
解决方法:确保除以 $\sqrt{d_k}$,使用混合精度训练时需监控 softmax 输入值范围
-
层归一化位置 :
- 原始 Transformer 在残差前做归一化,BERT 改为残差后(Post-LN)
-
错误实现会导致梯度消失
-
激活函数选择 :
- BERT 使用 GELU 而非 ReLU,误用会导致性能下降约 1 - 2 个点
$$\text{GELU}(x) = x\Phi(x)$$
总结与思考
这些计算公式直接影响模型捕获上下文信息的能力:
– 注意力机制决定 token 间交互方式
– 层归一化影响训练稳定性
– 前馈网络提供非线性变换能力
优化方向:
– 稀疏注意力降低计算复杂度
– 参数共享减少模型体积
– 量化感知训练提升推理速度
开放问题:
1. 如何设计更高效的注意力计算方式来处理长文档?
2. 预训练任务(MLM/NSP)与这些计算公式如何协同影响最终表现?
3. 在不同语言任务中,哪些公式参数最值得调整优化?
