共计 2182 个字符,预计需要花费 6 分钟才能阅读完成。
背景痛点:为什么 BART 预训练让人头疼?
刚接触 BART(Bidirectional and Auto-Regressive Transformers)时,我发现三个高频问题:

- 数据构造复杂 :需要同时处理编码器(Encoder) 的破坏文本和解码器 (Decoder) 的原始文本重建
- 显存黑洞:Seq2Seq 结构相比 BERT 消耗更多显存,batch_size 稍大就 OOM(Out Of Memory)
- 超参敏感:dropout、learning rate 等参数对重建任务的影响比想象中更大
技术对比:BART 凭什么特别?
和 BERT、GPT 相比,BART 的预训练目标很独特:
- BERT:纯编码器,使用掩码语言模型(Masked Language Model)
- GPT:纯解码器,自回归 (Auto-Regressive) 生成
-
BART:编码器 - 解码器联合训练,通过以下方式破坏输入文本:
-
随机替换 token(类似 BERT)
- 删除片段
- 打乱句子顺序
- 然后让解码器还原原始文本
数学表达其损失函数:
$$
\mathcal{L}{BART} = -\sum, z)
$$}^T \log P(x_t | x_{<t
其中 $z$ 是编码器处理的损坏文本。
核心实现:PyTorch 实战代码
1. 数据预处理
from transformers import BartTokenizer
tokenizer = BartTokenizer.from_pretrained('facebook/bart-base')
def corrupt_text(text):
"""随机执行文本破坏操作"""
tokens = text.split()
# 这里简化为随机 mask 15% 的 token
return ''.join(['[MASK]' if random.random() < 0.15 else t
for t in tokens
])
2. 动态填充 DataCollator
from dataclasses import dataclass
from torch.nn.utils.rnn import pad_sequence
@dataclass
class BartDataCollator:
def __call__(self, batch):
input_ids = [torch.tensor(x['input_ids']) for x in batch]
labels = [torch.tensor(x['labels']) for x in batch]
# 动态填充到 batch 内最大长度
input_ids = pad_sequence(input_ids, batch_first=True, padding_value=tokenizer.pad_token_id)
labels = pad_sequence(labels, batch_first=True, padding_value=-100)
return {'input_ids': input_ids, 'labels': labels}
3. 梯度累积 + 混合精度训练
from torch.cuda.amp import GradScaler, autocast
scaler = GradScaler()
accum_steps = 4 # 累积 4 个 batch 的梯度
for epoch in range(epochs):
model.train()
for i, batch in enumerate(train_loader):
with autocast():
outputs = model(**batch)
loss = outputs.loss / accum_steps
scaler.scale(loss).backward()
if (i+1) % accum_steps == 0:
scaler.step(optimizer)
scaler.update()
optimizer.zero_grad()
避坑指南:血泪经验总结
- 标签错位问题
- 现象:模型输出乱码
-
解决:确保 decoder_input_ids 比 labels 向右偏移一位
-
梯度爆炸
- 现象:loss 突然变成 nan
-
解决:添加
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) -
显存不足
- 现象:CUDA out of memory
- 解决:
- 使用梯度累积
- 尝试
model.gradient_checkpointing_enable()
性能验证:batch_size 的影响
| batch_size | 吞吐量(tokens/sec) | GPU 显存占用(GB) |
|---|---|---|
| 8 | 1200 | 10.2 |
| 16 | 2100 | 18.7 |
| 32 | 3800 | OOM |
测试环境:NVIDIA V100 32GB, CNN/DailyMail 数据集
延伸思考:尝试不同破坏策略
可以修改 corrupt_text() 函数实现:
-
Span Masking:连续 mask 整段文本
# 示例:随机选择 2 个 span,每个 span 长度 3 - 5 个 token spans = [(start, start+length) for ...] -
Sentence Permutation:打乱句子顺序
建议记录不同策略下的验证集 loss 变化,找到最适合你任务的破坏方式。
写在最后
通过这次 BART 预训练实践,我最大的体会是:理解模型设计初衷比盲目调参更重要。BART 的编码器 - 解码器结构让它特别适合文本生成类任务,但相应地也需要更精细的数据处理。建议新手先用小批量数据跑通全流程,再逐步扩大规模,这样调试效率会高很多。
正文完
