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

另一个常见问题是显存溢出(OOM)。有开发者尝试用 BERT 处理长文本(如整篇论文),结果遭遇显存不足。这是因为 BERT 的全连接自注意力机制(Full-connected Self-Attention)会消耗 O(n²)的内存,而标准 Transformer 的解码器可以通过掩码(Mask)实现自回归(Auto-regressive)生成,更适合长序列任务。
架构对比
结构差异
标准 Transformer 采用经典的编码器 - 解码器架构:
- 编码器(Encoder):负责理解输入序列,通过多头自注意力(Multi-head Self-Attention)捕获全局依赖
- 解码器(Decoder):在生成每个 token 时,只能看到前面的 token(通过注意力掩码实现)
而 BERT 则只保留了编码器部分,并做了两个关键改进:
- 双向上下文:通过掩码语言模型(Masked Language Model, MLM)训练目标,使模型能同时看到左右上下文
- 下一句预测(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 不是为生成设计的,但可以通过以下方式改造:
- 添加因果掩码(Causal Mask)强制自回归属性
- 在顶层添加语言模型头(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]
避坑指南
常见错误及解决方案
- 错误:直接使用 BERT 进行文本生成
-
解决:添加因果掩码并微调 LM Head
-
错误:用标准学习率微调 BERT
-
解决:使用线性缩放规则:
lr = base_lr * batch_size / 256 -
错误:忽略序列长度限制
- 解决:长文本建议使用 Reformer 或 Longformer 变体
微调学习率设置
对于 batch size 调整,建议公式:
adjusted_lr = initial_lr * (current_batch_size / reference_batch_size)
其中 reference_batch_size 通常取 256。
延伸思考
如何改造 BERT 的 position embedding 来处理 10k 长度的法律文本?可以考虑:
- 使用相对位置编码(Relative Position Encoding)替代绝对位置编码
- 实现层次化位置编码(Hierarchical Position Encoding)
- 借鉴 Longformer 的稀疏注意力模式(Sparse Attention)
这些改造需要同时考虑模型效果和计算效率的平衡。
