BERT与Transformer架构深度解析:从原理到工程实践的关键差异

1次阅读
没有评论

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

image.webp

背景痛点

在实际 NLP 项目中,许多开发者容易混淆 BERT 和 Transformer 的适用场景,导致一些典型问题。最常见的是错误地使用 BERT 进行文本生成任务,结果发现模型根本无法输出合理结果。这是因为 BERT 本质上是一个双向编码器(Bidirectional Encoder),而标准的 Transformer 是编码器 - 解码器(Encoder-Decoder)结构,两者设计目标完全不同。

BERT 与 Transformer 架构深度解析:从原理到工程实践的关键差异

另一个常见问题是显存溢出(OOM)。有开发者尝试用 BERT 处理长文本(如整篇论文),结果遭遇显存不足。这是因为 BERT 的全连接自注意力机制(Full-connected Self-Attention)会消耗 O(n²)的内存,而标准 Transformer 的解码器可以通过掩码(Mask)实现自回归(Auto-regressive)生成,更适合长序列任务。

架构对比

结构差异

标准 Transformer 采用经典的编码器 - 解码器架构:

  • 编码器(Encoder):负责理解输入序列,通过多头自注意力(Multi-head Self-Attention)捕获全局依赖
  • 解码器(Decoder):在生成每个 token 时,只能看到前面的 token(通过注意力掩码实现)

而 BERT 则只保留了编码器部分,并做了两个关键改进:

  1. 双向上下文:通过掩码语言模型(Masked Language Model, MLM)训练目标,使模型能同时看到左右上下文
  2. 下一句预测(Next Sentence Prediction, NSP):增强模型理解句子间关系的能力

注意力机制差异

标准 Transformer 的解码器使用三角掩码(Triangular Mask)确保自回归属性,而 BERT 的 MLM 目标允许所有 token 相互关注:

# 标准 Transformer 解码器的注意力掩码(PyTorch 示例)# shape: [seq_len, seq_len]
mask = torch.triu(torch.ones(seq_len, seq_len), diagonal=1).bool()

# BERT 的注意力掩码(实际是 padding mask)# shape: [batch_size, seq_len]
mask = (input_ids != pad_token_id).long()

工程实现

注意力处理差异

在 HuggingFace 实现中,两者的注意力处理有显著不同。以下是关键代码对比:

# Transformer 解码器注意力实现(节选)def decoder_forward(self, hidden_states, attention_mask=None):
    # 自注意力层
    self_attn_output = self.self_attn(
        hidden_states,
        attention_mask=attention_mask,  # 三角掩码
    )

# BERT 注意力实现(节选)def encoder_forward(self, hidden_states, attention_mask=None):
    # 扩展 attention_mask 维度
    # 从 [batch_size, seq_len] 到[batch_size, 1, 1, seq_len]
    extended_attention_mask = attention_mask[:, None, None, :]

    self_attn_output = self.self_attn(
        hidden_states,
        attention_mask=extended_attention_mask,  # 只是 padding mask
    )

用 BERT 实现生成任务

虽然 BERT 不是为生成设计的,但可以通过以下方式改造:

  1. 添加因果掩码(Causal Mask)强制自回归属性
  2. 在顶层添加语言模型头(LM Head)
from transformers import BertLMHeadModel

model = BertLMHeadModel.from_pretrained('bert-base-uncased')

# 生成时添加三角掩码
def generate(self, input_ids, max_length):
    causal_mask = torch.triu(torch.ones(max_length, max_length), diagonal=1
    ).bool().to(input_ids.device)

    outputs = model(input_ids, attention_mask=causal_mask)
    return outputs

生产考量

显存占用对比

在 16GB 显存的 GPU 上(如 V100):

  • BERT-base:最大处理长度约 512 token
  • Transformer-base(解码器):可达 1024 token(因使用了 KV 缓存)

[CLS] token 的稳定性

虽然 BERT 官方推荐用[CLS] token 做分类,但实际上:

  • 在微调初期,[CLS]的表示质量较差
  • 更稳定的做法是使用平均池化(Mean Pooling):
# 更好的分类特征提取方式
pooled_output = last_hidden_state.mean(dim=1)  # [batch_size, hidden_size]

避坑指南

常见错误及解决方案

  1. 错误:直接使用 BERT 进行文本生成
  2. 解决:添加因果掩码并微调 LM Head

  3. 错误:用标准学习率微调 BERT

  4. 解决:使用线性缩放规则:lr = base_lr * batch_size / 256

  5. 错误:忽略序列长度限制

  6. 解决:长文本建议使用 Reformer 或 Longformer 变体

微调学习率设置

对于 batch size 调整,建议公式:

adjusted_lr = initial_lr * (current_batch_size / reference_batch_size)

其中 reference_batch_size 通常取 256。

延伸思考

如何改造 BERT 的 position embedding 来处理 10k 长度的法律文本?可以考虑:

  1. 使用相对位置编码(Relative Position Encoding)替代绝对位置编码
  2. 实现层次化位置编码(Hierarchical Position Encoding)
  3. 借鉴 Longformer 的稀疏注意力模式(Sparse Attention)

这些改造需要同时考虑模型效果和计算效率的平衡。

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