BART生成实战
BART介绍
BART(Bidirectional and Auto-Regressive Transformer)是 Facebook AI Research 于 2019 年提出的序列到序列预训练模型,核心是结合双向 Transformer 编码器与自回归 Transformer 解码器,通过去噪自编码预训练,同时具备文本理解与生成能力,适配多种 NLP 任务。以下从核心架构、预训练机制、下游微调、性能与应用、局限与变体展开详细介绍。
一、核心架构
BART 基于 Transformer 的 Encoder-Decoder 架构,关键设计如下:
| 组件 | 核心特性 | 作用 |
|---|---|---|
| 双向编码器 | 同 BERT,全量上下文注意力,GeLU 激活,参数正态初始化(N (0,0.02)) | 编码受损文本的全局语义,输出上下文表征 |
| 自回归解码器 | 同 GPT,单向注意力(仅左向),与编码器交互(跨注意力) | 自回归重构原始文本,生成连贯输出 |
| 层数配置 | 基础版:编码器 / 解码器各 6 层;大型版:各 12 层 | 平衡性能与效率 |
二、预训练机制(去噪自编码)
核心逻辑是 “破坏输入→重建原始文本”,通过多种噪声策略迫使模型学习文本结构与语义,噪声类型包括:
- Token 掩码:随机替换 token 为 [MASK](类似 BERT)。
- 文本删除:随机删除部分 token,模型需推断缺失内容。
- 句子重组:打乱输入句子顺序。
- 文本填充:连续 token 段替换为单个 [MASK](区别于 BERT 的独立掩码)。
- 文档旋转:随机选 token 作为文档起始,旋转整段文本。
预训练目标是最小化重建文本与原始文本的交叉熵,编码器输出作为解码器的 “记忆” 参与每一层计算。
三、下游任务微调
BART 通过灵活调整适配不同任务,常见微调方式如下:
| 任务类型 | 微调方式 | 示例 |
|---|---|---|
| 文本生成 | 直接用 Encoder-Decoder 生成目标文本 | 摘要生成、机器翻译、对话生成 |
| 文本分类 | 编码器输出接分类头,或解码器末尾 token 接分类器 | 情感分析、新闻分类 |
| 问答系统 | 编码器处理问题 + 上下文,解码器生成答案 | SQuAD 等抽取 / 生成式问答 |
| 词元分类 | 编码器输出接 token 级分类头 | 命名实体识别(NER)、词性标注 |
四、性能与应用场景
- 核心优势
- 兼顾理解与生成:双向编码保障语义理解,自回归解码保障生成连贯。
- 泛化能力强:多样预训练噪声适配多任务,在 GLUE、SQuAD、CNN/Daily Mail 等数据集上性能优异(如 ROUGE-1 达 42.9486)。
- 典型应用
- 文本摘要:BART-large-cnn 在新闻摘要任务上表现突出。
- 机器翻译:编码器处理源语言,解码器生成目标语言。
- 对话系统:生成自然回复,适配开放域 / 任务型对话。
- 文本纠错:重建受损文本的能力直接用于语法 / 拼写纠错。
五、局限性与主流变体
- 局限性
- 长文本处理效率低:Transformer 注意力复杂度 O (n²),长文本推理慢、显存开销大。
- 生成多样性不足:自回归生成易重复,依赖解码策略(如 beam search、top-k 采样)优化。
- 领域适配成本高:通用预训练在专业领域(如医疗、法律)需额外微调。
- 主流变体
- BART-large-cnn:在 CNN/Daily Mail 上微调,摘要任务标杆模型。
- mBART:多语言版,支持 100 + 语言翻译与生成。
- BART-base:轻量版,适合资源受限场景。
六、与 BERT、GPT 的核心差异
| 模型 | 架构 | 核心能力 | 预训练目标 | 典型任务 |
|---|---|---|---|---|
| BERT | 仅双向编码器 | 文本理解 | 掩码预测 + 下一句预测 | 分类、NER、问答 |
| GPT | 仅自回归解码器 | 文本生成 | 自回归语言建模 | 文本续写、对话 |
| BART | Encoder-Decoder | 理解 + 生成 | 去噪自编码 | 摘要、翻译、生成 + 理解混合任务 |
总结
BART 以 Encoder-Decoder 架构融合双向编码与自回归生成,通过去噪预训练实现理解与生成的统一,是 NLP 领域的 “多面手”,尤其在文本生成类任务中优势显著,其变体与微调策略持续推动下游应用落地。
BART生成实战
首先读取数据,分割数据集,将训练集分为训练集(90%)和验证集(10%)
process_data.py
import pandas as pd #处理表格数据
pre_train_file= "data/train.csv"
train_df = pd.read_csv(pre_train_file,header=None,names=["id","input","tgt"]) #读入数据
print(train_df.head())
#下面两句:将"data/train.csv"里的数据分割成训练集(90%)和验证集(10%)
train_data = train_df.sample(frac=0.9, random_state=0, axis=0) #从train_df采样0.9(90%)的比例作为训练数据
val_data = train_df[~train_df.index.isin(train_data.index)] #干啥的, 过来用
train_data.to_csv("data/pro_train_data.csv", index=False,header=False) #提取出来的训练集保存到"data/pro_train_data.csv"
val_data.to_csv("data/pro_val_data.csv", index=False,header=False) #提取出来的验证集保存到"data/pro_val_data.csv"
处理词表:

这里使用第三种方法:重新制作词表
pro_vocab.py
import sys
import torch
from collections import Counter
from transformers import BertTokenizer
from transformers import BartConfig
from transformers import BartForConditionalGeneration
from model_utils.config import parse_args
args = parse_args() #设置 ,字典, 属性类 config {}
def load_data(path):
with open(path, 'r', encoding='utf-8') as f:
lines = f.readlines()
datas = []
for line in lines:
line = line.strip().split(",")
if len(line) == 3:
# 训练集
text, target = line[1].split(" "), line[2].split(" ")
datas.append(text + target)
else:
text = line[1].split(" ")
datas.append(text)
return datas
train_data = load_data('./data/train.csv')
token2count = Counter() #计数工具 哈希表
for i in train_data:
token2count.update(i) #不需要知道原理
tail = []
ct = 0
for k, v in token2count.items():
if v >= ct:
tail.append(k)
tail.sort()
vocab = tail
vocab.insert(0,"[PAD]")
vocab.insert(100,"[UNK]")
vocab.insert(101,"[CLS]")
vocab.insert(102,"[SEP]")
vocab.insert(103,"[MASK]")
vocab.insert(104,"[EOS]")
# tokenizer = BertTokenizer.from_pretrained(args.pre_model_path)
# vocabs = tokenizer.get_vocab() #获取模型词表
# new_vocabs = list(vocabs.keys())
# print(len(vocabs))
# count = 0
# for v in vocab: #mn复杂度
# if v not in vocabs:
# count += 1
# new_vocabs.append(v)
# print(len(new_vocabs))
new_vocabs = vocab #重新制作词表
with open(args.pre_model_path+'/vocab.txt', 'w', encoding='utf-8') as f:
for v in new_vocabs:
f.write(f"{v}\n") #保存为新的词表
model = BartForConditionalGeneration.from_pretrained(args.pre_model_path) #Bart模型
model.resize_token_embeddings(len(new_vocabs)) #因为词表增加了,所以也要改变模型
state_dict = model.state_dict()
torch.save(state_dict, args.pre_model_path+'/pytorch_model.bin')
bartconfig = BartConfig.from_pretrained(args.pre_model_path)
bartconfig.vocab_size = len(new_vocabs)
bartconfig.save_pretrained(args.pre_model_path) #保存新的模型
进行自监督预训练:

MLM 预训练
掩码语言模型(Masked Language Model,MLM)是 BERT 提出的双向上下文预训练范式,核心是随机掩码文本中 15% 的 Token,让双向 Transformer 编码器通过左右全量上下文预测原始 Token,仅对掩码位计算交叉熵损失,核心目标是让模型学习文本的语义、句法关联与长距离依赖,是 NLP 文本理解类任务的基础预训练方式,为模型输出高质量双向上下文表征提供支撑。
MLM 在 BART 中的作用
BART 以去噪自编码为核心预训练目标,MLM 是其核心去噪策略之一,且针对 BART 的 Encoder-Decoder 架构做了简化(直接替换为 [MASK],移除 BERT 的随机替换 / 保留规则,因 BART 预训练为 Seq2Seq 文本还原,无预训练 - 微调一致性问题),核心作用有三:
- 为 BART 的双向编码器奠定细粒度语义理解能力,是编码器捕捉 Token 级语义关联的核心手段;
- 为自回归解码器提供精准生成引导,让解码器通过跨注意力对齐编码器表征,还原掩码位置的原始 Token,保证生成的语义准确性;
- 让 BART 兼具文本理解能力,使其不仅适配生成类任务,还能高效微调至分类、NER 等理解类任务,成为通用 NLP 模型。
此外,BART 的 MLM 可单独或与句子重组、文本填充等去噪策略叠加,与其他策略形成Token 级 + 结构级的互补,让模型学习从词元到篇章的全层级文本特征。

BERT 标准 MLM 的核心实现
# 核心循环:遍历 随机数列表 + 原始token_id列表,一一配对处理每个token
for r, i in zip(rands, text_ids): # r=当前位置的0~1随机数,i=当前位置的原始token_id
if r < 0.15 * 0.8:
# 情况1:概率占比 12% (0.15*0.8) → 【掩码替换】
input_ids.append(self.tk.mask_token_id) # 输入:替换成MASK标记
output_ids.append(i) # 标签:原始token_id → 让模型预测「这个位置原本是什么」
elif r < 0.15 * 0.9:
# 情况2:概率占比 1% (0.15*0.9 - 0.15*0.8) → 【原样保留】
input_ids.append(i) # 输入:还是原始token_id不变
output_ids.append(i) # 标签:原始token_id → 让模型「自己预测自己」
elif r < 0.15:
# 情况3:概率占比 2% (0.15 - 0.15*0.9) → 【随机替换】
input_ids.append(np.random.randint(self.spNum,self.tkNum)) # 输入:随机选一个非特殊符号的token_id
output_ids.append(i) # 标签:原始token_id → 让模型「根据随机词预测原本的词」
else:
# 情况4:概率占比 85% (1 - 0.15) → 【完全保留,不参与预测】
input_ids.append(i) # 输入:原始token_id不变
output_ids.append(-100) # 标签:-100 → 损失函数忽略该位置,模型不需要预测
一、先明确:这段代码的核心目标
实现 BERT 的经典 MLM 训练策略:对输入的文本 token 序列 (text_ids) 做「随机掩码 + 三类不同处理」,最终生成两个序列:
input_ids:模型的输入序列(被掩码 / 替换 / 保留后的 token)output_ids:模型的监督标签序列(用来计算 loss 的标准答案)
训练逻辑:让模型根据
input_ids去预测output_ids里非-100的 token,学懂文本的语义关联
二、BERT 官方【15% 黄金掩码规则】完整拆解
这段代码完美复现了 Google 在 BERT 论文中公布的标准 MLM 规则,这个15%的分配比例是论文实验最优值,所有预训练代码都遵循这个规则,必须记住:
✔️ 整体规则:只对文本中 15% 的 token 进行「处理 + 预测」,剩下 85% 的 token 完全不动、不预测
✔️ 15% 的被处理 token,内部再做「三级拆分」,比例固定:
- 12% 的 token(占整体,15% × 80%) → 替换为
[MASK]标记 - 1% 的 token(占整体,15% × 10%) → 保持原始 token 不变
- 2% 的 token(占整体,15% × 20%) → 替换为 词表中随机的一个普通 token(非特殊符号)
✔️ 三者加总 = 15%,剩余 85% 完全保留 + 标签置为 - 100(不预测)
BART 预训练流程
你之前看的是 BERT 的 MLM 掩码逻辑,BART 是Facebook 在 BERT、GPT 之后的融合型模型,它的预训练是 **【损坏文本 + 文本复原】的自编码任务 **,整体流程非常简洁,所有核心点我用大白话讲透,和你之前的 MLM 代码能完全对应上,好理解✅
核心结论先记:BART 的预训练核心 = 一个「文本降噪 + 完形填空」的过程,模型学会把「被随机破坏的乱序文本」恢复成「通顺的原始文本」。
一、BART 的模型基础
BART 是标准的【Encoder-Decoder 编码器 - 解码器】结构(Transformer 完整结构),和 BERT(只有 Encoder)、GPT(只有 Decoder)都不同:
- Encoder 编码器:输入「被随机破坏的文本」,负责提取文本的特征表示;
- Decoder 解码器:在 Encoder 的特征基础上,自回归生成「完整的原始文本」;
补充:BART 的预训练目标,就是把这个「编码器输入、解码器输出」的任务做对,学会文本修复。
二、BART 预训练 完整核心流程
✅ 步骤 1:准备「干净的原始文本」
拿到无任何修改的正常文本句子 / 段落 → 做常规的 tokenize 分词 → 得到 原始 token 序列(记为 src),比如:我/爱/中/国/的/山/河
这一步和 BERT 的预处理完全一样,没任何区别。
✅ 步骤 2:对原始文本做【随机文本破坏】(核心!BART 的核心预处理,替代了 BERT 的 15% 掩码)
这是 BART 最关键的一步,也是和 BERT 最大的区别:BERT 只是「局部掩码 token」,BART 是「花式破坏整段文本」BART 会对原始文本src执行 1 种或多种随机的「文本损坏操作」,生成一个 被污染、被打乱的损坏文本序列(记为 tgt_input),损坏操作是随机选的,非常灵活,官方标配的破坏方式有这些(都是常用的,随机触发):
- ✔️ Token 掩码(最常用,和你写的 BERT MLM 完全一致):随机选 15% 的 token,按你代码里的
80%MASK+10%保留+10%随机替换规则处理; - ✔️ 句子打乱:如果文本是多句话,随机打乱句子的顺序(比如 句 1 + 句 2 + 句 3 → 句 2 + 句 1 + 句 3);
- ✔️ 文本片段删除:随机删掉文本中的连续一段 token(比如 我 / 爱 /【删掉】/ 山 / 河);
- ✔️ 文本片段替换:随机把一段连续 token 替换成一个
[MASK]; - ✔️ 文本长度截断 / 填充:随机截断文本或填充无意义 token。
✅ 核心目的:让输入的文本「有残缺、有错误、有乱序」,制造一个带噪声的文本,给模型出难题。
✅ 步骤 3:执行「降噪自编码」预训练任务(BART 的核心目标)
把两步得到的两个序列,喂给 BART 的 Encoder-Decoder,执行训练,这一步是 BART 预训练的唯一目标:
- 把 被破坏的文本序列 → 输入到 BART 的编码器 (Encoder) 中;
- 要求 BART 的解码器 (Decoder) → 以「自回归」的方式,逐字逐句生成出完整的原始文本序列;
- 损失计算:用交叉熵损失,计算「解码器生成的文本」和「原始干净文本」的差异,反向传播更新模型参数。
✅ 一句话说清 BART 预训练目标:给模型看一段被弄脏的话,让模型把它复原成原来通顺的原话。
三、BART 与 BERT 预训练的核心区别
你刚看完 BERT 的 MLM 掩码代码,这里对比一下,2 个核心差异,1 个关键相同点,瞬间分清两者,不混淆:
✅ 相同点
BART 里也包含了 BERT 的 15% MLM 掩码规则(就是你写的那段r<0.15*0.8的逻辑),MLM 是 BART 众多文本破坏方式中的核心一种。
✅ 核心不同点 1:训练任务不同(本质区别)
- BERT:单向填空任务 → 输入带
[MASK]的文本,模型只需要预测[MASK]的位置是什么 token 即可,输出是「单个 token 的预测值」; - BART:完整生成任务 → 输入损坏文本,模型需要「从头到尾完整生成一整段原始文本」,输出是「整段文本序列」,难度远大于 BERT。
✅ 核心不同点 2:模型结构 + 预测方式不同
- BERT:只有 Encoder,并行预测所有掩码位置,一次性输出所有预测结果;
- BART:Encoder+Decoder,自回归预测,生成文本时只能「看前面生成的内容,预测下一个 token」,和人类写字一样,从左到右逐个写。
四、BART 预训练的核心优势(为什么要这么设计)
BART 把 BERT 的「理解能力」和 GPT 的「生成能力」融合了,这也是它能成为目前最强的通用预训练模型之一的原因:
- ✔️ 学会「文本理解」:因为要修复损坏文本,必须先看懂文本的语义、语法、逻辑关系(继承 BERT 的优点);
- ✔️ 学会「文本生成」:因为要完整生成原始文本,掌握了连续文本的创作能力(继承 GPT 的优点);
- ✔️ 泛化能力极强:BART 的预训练任务足够难,学好后,下游所有 NLP 任务(分类、翻译、摘要、问答、续写、改错)都能适配,效果都很好。
✅ 终极极简总结
- BART 预训练 = 文本损坏 + 文本复原,核心是「降噪自编码」任务;
- 输入是「随机破坏的文本」,输出要求是「原始干净文本」,模型用 Encoder-Decoder 完成这个任务;
- BART 兼容 BERT 的 MLM 掩码逻辑,同时学会理解 + 生成,是比 BERT 更全能的模型。
CiderD_scorer
✅ 1. 核心定位
CiderD_scorer 是 CIDEr-D(Consensus-based Image Description Evaluation with Diversity)指标的官方评分器,由 Vedantam 等人在 CIDEr 基础上改进,专门用于图像描述 / 文本生成任务中,评估模型生成文本与参考文本的语义一致性,解决了原始 CIDEr 的 “gaming” 漏洞(如重复关键词刷分、句长异常等),是目前图像描述领域的核心评估工具。
✅ 2. 核心用途
最主要用在 图像生成标题 / 图文生成 / 摘要生成 任务中,判断模型写的句子是不是「人话」、和标准答案贴合度有多高,是该领域的主流评估指标(和 BLEU、ROUGE、METEOR 并列核心指标)。
✅ 3. 核心原理(一句话懂)
基于 1~4 元词组 (n-gram) 做计算,给不同词组做 TF-IDF 权重(稀有词组权重更高),计算「生成文本」与「参考文本」的余弦相似度,同时解决了原始 CIDEr 的漏洞:加了重复词惩罚(防止模型无脑重复关键词刷分)+句长差异惩罚(防止生成过长 / 过短句),评估结果更贴合人类的主观打分。
✅ 4. 关键特点
- 🍎 优点:比 BLEU 更看重语义匹配(BLEU 只看表面词重合),对「意思对、用词不同」的生成文本更友好,评价更合理;
- ⚡ 分数范围:一般 0~10 分,越高越好(满分代表生成文本和参考文本完全一致);
- 📌 CIDEr vs CIDEr-D:CIDEr-D 是 CIDEr 的优化版,D=Diversity(去冗余 / 多样性),修复了原始 CIDEr 被「重复词刷分」的 bug,现在所有场景都用 CIDEr-D,没人用原始 CIDEr。
✅ 5. 补充关联(和 BART 相关)
BART 如果做文本生成 / 图文生成 / 摘要任务,训练完模型后,一定会用 CIDEr-D_scorer 做最终效果评估,和 BLEU、ROUGE 一起构成完整的评估体系。
极简总结
- CIDEr-D_scorer → 文本生成任务的语义匹配评分器,图像描述任务必用;
- 核心优势:防刷分、语义匹配更准,评分贴合人类判断;
- 分数越高 → 模型生成的文本越贴合参考标准答案。
可能会问的问题:
一、这个生成任务用的什么模型、为什么用这个模型
Bart模型
二、用的评价指标是什么、为什么用这个评价指标
CiderD_scorer,选 CiderD_scorer,本质是选图像描述场景下 “最准、最严谨、最通用”的评估方式 —— 既解决了传统指标的评价偏差,又适配任务需求,还是领域共识,能让模型效果的评估结果可信、可对比、贴合人类判断。
三、输入是怎么处理的
改变词表
AtomGit 是由开放原子开源基金会联合 CSDN 等生态伙伴共同推出的新一代开源与人工智能协作平台。平台坚持“开放、中立、公益”的理念,把代码托管、模型共享、数据集托管、智能体开发体验和算力服务整合在一起,为开发者提供从开发、训练到部署的一站式体验。
更多推荐



所有评论(0)