深入解读 PLBART:Transformers 中面向程序理解与生成的统一预训练模型
【免费下载链接】transformers🤗 Transformers: the model-definition framework for state-of-the-art machine learning models in text, vision, audio, and multimodal models, for both inference and training.项目地址: https://gitcode.com/GitHub_Trending/tra/transformers
PLBART(Program and Language BART)是一个基于 BART 架构的多语言 encoder-decoder 序列到序列模型,专门面向代码摘要(code summarization)、代码生成(code generation)与代码翻译(code translation)等编程语言(PL)与自然语言(NL)互转任务。本文以 docs/source/en/model_doc/plbart.md 为主线,结合本仓库中 configuration_plbart.py、tokenization_plbart.py、modeling_plbart.py 及对应测试,全面讲解 PLBART 的模型原理、语言 ID 特殊 token 的构造与编码格式、有监督微调与生成推理的完整用法,以及四类模型 Head 的实现细节。读完本文,你将能够在 Transformers 生态中用uclanlp/plbart-*系列 checkpoint 完成 Python/Java 等代码与英文之间的翻译、摘要、补全任务,并理解其输入输出格式为何与其他 BART 类模型不同。
一、模型背景与定位
PLBART 来自论文Unified Pre-training for Program Understanding and Generation(Wasi Uddin Ahmad 等,UCLA)。它是一个类 BART 的序列到序列模型,可用于:
- 代码摘要(code summarization):代码 → 自然语言描述;
- 代码生成(code generation):自然语言 → 代码;
- 代码翻译(code translation):一种编程语言 → 另一种编程语言,或代码 ↔ 英文。
预训练 checkpointplbart-base使用多语言去噪任务在 Java、Python 与英文文本上训练。模型通过去噪自编码(denoising autoencoding)学习大规模 Java/Python 函数及其关联 NL 文本;论文实验表明,除了生成类任务,PLBART 在程序修复(program repair)、克隆检测(clone detection)、漏洞代码检测等判别式任务上也展现出程序理解能力。PLBART 由社区成员 gchhablani 贡献至 Transformers,于 2022-02-18 合入仓库。
从实现上看,PLBART 完全复用 Transformers 的 BART 式分层结构(encoder + decoder + 共享词嵌入),但保留了一些与 fairseq 原版实现对齐的特殊设计,这正是它和其他 BART 模型最不一样的地方。模型代码位于 src/transformers/models/plbart/,由 modular_plbart.py 自动生成(文件头声明该文件不得手工编辑)。
二、区别于一般 BART 模型的两个核心设计
2.1 语言 ID token( )与输入输出格式约定
PLBART 是多语言模型,期望输入带有特殊语言 ID token(language id token,简称 ),以区分 Java、Python、英文等。与 MBart 等在序列最前面加语言 token 的做法不同,PLBART 的 采用"缀在末尾"的方式组织:
- 源端格式:
X [eos, src_lang_code],其中X为源文本; - 目标端格式:
[tgt_lang_code] X [eos],即解码器真正消费的序列以目标语言 ID 开头; bos(<s>)从不使用。
在 tokenizer 内部,这一约定通过special_tokens_pattern="prefix_suffix"与空的prefix_tokens实现:tokenization_plbart.py 中set_src_lang_special_tokens/set_tgt_lang_special_tokens都只设置suffix_tokens = [eos_token_id, cur_lang_code],前缀为空。也就是说 tokenizer 物理输出均为X [eos][<LID>];而模型侧通过特殊的shift_tokens_right(见下文 3.3 节)在右移时将末尾的 挪到序列首部,从而在解码器输入中得到[tgt_lang_code] X [eos]这一目标端形态,保证教师强制(teacher forcing)与自回归生成视角下解码总是以目标语言 ID 起头。文档同时提醒:微调时若只涉及单一语言,可以不附加语言 token(详见原论文)。
2.2 fairseq 词汇表对齐与 token 映射
PLBART tokenizer 基于 SentencePiece,模型文件为sentencepiece.bpe.model。为了与 fairseq 原版词汇表对齐,PLBartTokenizer特意处理了 "fairseq vocab 与 spm vocab 错位" 的问题:tokenization_plbart.py 的注释给出对齐示意——
- fairseq 前 4 个 token:
'<s>'=0, '<pad>'=1, '</s>'=2, '<unk>'=3; - spm 中
'<unk>'却占第 0 位。
因此 tokenizer 定义fairseq_tokens_to_ids = {"<s>": 0, "<pad>": 1, "</s>": 2, "<unk>": 3},并设置fairseq_offset = 1:SentencePiece 的真实子词从 spm ID 偏移后(spm_id + fairseq_offset)映射为 fairseq 侧 ID(_convert_token_to_id),语言代码等保留 token 也会从 added-token 表中移除,改为按 fairseq 语义统一管理,从而保证权重能够正确加载、vocab_size与模型 lm_head 尺寸一致。注意 tokenizer 需要一个额外的特殊<mask>偏移(base模式下vocab_size计算多+1),因为 PLBART 预训练时使用掩码去噪,推理端lang_code_to_id与id_to_lang_code两表则用来在生成时定位目标语言 ID。
三、PLBartTokenizer 用法详解
加载 tokenizer 与设置源/目标语言:
from transformers import PLBartTokenizer tokenizer = PLBartTokenizer.from_pretrained( "uclanlp/plbart-base", src_lang="en_XX", tgt_lang="python" )src_lang与tgt_lang均会经FAIRSEQ_LANGUAGE_CODES_MAP归一化为特殊格式(如python→__python__),支持的语言与对应代码在源码中显式列出:tokenization_plbart.py。
| language_codes 取值 | 可用语言 |
|---|---|
"base" | __java__、__python__、__en_XX__ |
"multi" | 在 base 之上增加__javascript__、__php__、__ruby__、__go__ |
- 传入
language_codes="base"(默认)使用 base 版词汇表,额外包含<mask>; - 传入
language_codes="multi"则额外注册四个语言代码,词汇表中不保留<mask>(多语言下游任务不做掩码),测试见 tests/models/plbart/test_tokenization_plbart.py。
编码约定:当把源文本作为第一个参数(或以关键字text)传入__call__时,tokenizer 编码源端格式;以text_target关键字传入目标文本时编码目标端格式。src_lang/tgt_lang还可用 setter 在运行时切换(见src_langproperty 与prepare_seq2seq_batch)。
构造PLBartTokenizer的其他参数:
vocab_file:SentencePiece 词表路径;bos_token="<s>"、eos_token="</s>"、sep_token="</s>"、cls_token="<s>"、unk_token="<unk>"、pad_token="<pad>"、mask_token="<mask>";language_codes:"base"或"multi";sp_model_kwargs:透传给SentencePieceProcessor.__init__,可用于启用 subword regularization(enable_sampling、nbest_size、alpha),实现细节与 BART 一致。
四、架构与配置:PLBartConfig 全参数
PLBartConfig的核心默认值定义在 configuration_plbart.py,对应uclanlp/plbart-basecheckpoint:
| 配置项 | 默认值 | 说明 |
|---|---|---|
vocab_size | 50005 | 词表大小(含 4 个基础 token、语言代码与 mask) |
max_position_embeddings | 1024 | 最大位置编码长度 |
d_model/hidden_size | 768 | 隐藏层维度 |
encoder_layers/decoder_layers | 6 / 6 | 编码器/解码器层数 |
encoder_ffn_dim/decoder_ffn_dim | 3072 / 3072 | FFN 中间维度 |
encoder_attention_heads/decoder_attention_heads | 12 / 12 | 注意力头数 |
activation_function | "gelu" | FFN 激活函数 |
dropout/attention_dropout | 0.1 / 0.1 | 全连接/注意力 dropout |
activation_dropout | 0.0 | FFN 激活后 dropout |
encoder_layerdrop/decoder_layerdrop | 0.0 / 0.0 | LayerDrop 比例 |
init_std | 0.02 | 权重初始化标准差 |
scale_embedding | True | 是否按 √d_model 缩放词嵌入 |
classifier_dropout | 0.0 | 分类头 dropout |
pad_token_id/bos_token_id/eos_token_id | 1 / 0 / 2 | 特殊 token ID |
forced_eos_token_id | 2 | 强制结束符 ID(生成必带</s>) |
is_encoder_decoder/is_decoder | True / False | seq2seq 架构标记 |
tie_word_embeddings | True | lm_head 与输入嵌入权重共享 |
use_cache | True | 生成时启用 KV 缓存 |
此外attribute_map将num_attention_heads → encoder_attention_heads、hidden_size → d_model等通用命名映射到 PLBART 命名,兼容其他 BART 类模型的加载约定。
直接以默认配置初始化模型:
from transformers import PLBartConfig, PLBartModel configuration = PLBartConfig() model = PLBartModel(configuration) # 随机权重模型内部实现要点
从 modeling_plbart.py 的结构看,PLBART 保留了 BART 家族特有的几个机制:
- 缩放词嵌入:
PLBartScaledWordEmbedding在nn.Embedding输出上乘以embed_scale;当config.scale_embedding=True时embed_scale = sqrt(d_model)(见 modeling_plbart.py 与 modeling_plbart.py),这一点需与 fairseq 训练保持一致; - 可学习位置嵌入:
PLBartLearnedPositionalEmbedding固定加offset=2再查表(modeling_plbart.py),保证 pad 占位与 BART 一致; - 注意力后端可插拔:
PLBartPreTrainedModel声明了_supports_flash_attn = True、_supports_sdpa = True、_supports_flex_attn = True(modeling_plbart.py),并在PLBartAttention.forward中通过ALL_ATTENTION_FUNCTIONS.get_interface(...)按config._attn_implementation分发到 eager / FlashAttention / SDPA / FlexAttention 实现——这也解释了官方模型页上的 FlashAttention 与 SDPA 徽章; - 梯度检查点:Encoder/Decoder 层继承
GradientCheckpointingLayer,_no_split_modules为["PLBartDecoderLayer", "PLBartEncoderLayer"],便于超大模型分布式切分。
五、监督训练:text-to-code 与 code-to-text
有监督微调直接把源文本与目标文本交给 tokenizer 即可(模型需预先加载到 device):
from transformers import PLBartTokenizer, PLBartForConditionalGeneration import torch model = PLBartForConditionalGeneration.from_pretrained( "uclanlp/plbart-base" ).to("cuda") tokenizer = PLBartTokenizer.from_pretrained( "uclanlp/plbart-base", src_lang="en_XX", tgt_lang="python" ) example_python_phrase = "def maximum(a,b,c):NEW_LINE_INDENTreturn max([a,b,c])" expected_translation_english = "Returns the maximum value of a b c." inputs = tokenizer( example_python_phrase, text_target=expected_translation_english, return_tensors="pt", ).to(model.device) model(**inputs)要点:
- 示例中 Python 代码里出现的
NEW_LINE_INDENT是预训练数据处理阶段使用的换行/缩进占位符,源码文本中直接原样包含它们即可; - 传入
text_target后 tokenizer 会切换到目标端模式,构建目标序列的 token 与labels; - 若需要计算 loss,可额外传
labels(shape(batch_size, seq_len),-100的位置被忽略)。当labels给出且未提供decoder_input_ids时,PLBartForConditionalGeneration.forward会用shift_tokens_right(labels, pad_token_id)自动构造解码器输入(见 modeling_plbart.py)。
3.x 附录:shift_tokens_right的特殊性
PLBART 没有像其他 BART 类模型那样使用统一decoder_start_token_id。其辅助函数 shift_tokens_right 的逻辑是:克隆标签序列 → 把-100替换成pad_token_id→ 取每行最后一个非 pad token(正是被缀在序列末尾的 )放到首位,其余 token 整体右移一位。因此解码器的第一个 token 天然就是目标语言 ID,这一点与文档所述目标格式[tgt_lang_code] X [eos]一致,也是推理时必须用decoder_start_token_id=语言ID的原因。
六、生成推理:Python → English 翻译实战
PLBartForConditionalGeneration继承GenerationMixin,可直接用generate解码。生成目标文本时,必须把decoder_start_token_id设为目标语言 ID;官方示例使用uclanlp/plbart-python-en_XX(在 Python 与英文上微调过的翻译模型):
from transformers import PLBartForConditionalGeneration, PLBartTokenizer tokenizer = PLBartTokenizer.from_pretrained( "uclanlp/plbart-python-en_XX", src_lang="python", tgt_lang="en_XX" ) example_python_phrase = "def maximum(a,b,c):NEW_LINE_INDENTreturn max([a,b,c])" inputs = tokenizer(example_python_phrase, return_tensors="pt").to(model.device) model = PLBartForConditionalGeneration.from_pretrained( "uclanlp/plbart-python-en_XX", device_map="auto" ) translated_tokens = model.generate( **inputs, decoder_start_token_id=tokenizer.lang_code_to_id["en_XX"] ) tokenizer.batch_decode(translated_tokens, skip_special_tokens=True)[0] # "Returns the maximum value of a b c."代码解读:
src_lang="python"、tgt_lang="en_XX"指定源/目标语言,tokenizer 据此把__python__/__en_XX__附加到对应序列;tokenizer.lang_code_to_id["en_XX"]返回英文语言 ID(该字典在 tokenizer 构造时按sp_model 大小 + fairseq_offset + 语言下标生成,如 docstring 所注对plbart-base大约对应 50003 等高位 ID),作为decoder_start_token_id使解码从英文 LID 起头;- 目标语言有单语言场景下也可不附加 LID,
model.generate(**inputs)同样可用(test_base_generate等测试正是用decoder_start_token_id=self.tokenizer.lang_code_to_id[src_lan]验证的)。
仓库对应测试集中在 tests/models/plbart/test_modeling_plbart.py(如test_java_cs_generate_one、test_java_cs_generate_batch、test_base_generate、test_sample_generate)与 tests/models/plbart/test_tokenization_plbart.py,可作为理解与验证格式约定的第一手资料。
七、四种模型 Head:从 seq2seq 到分类与纯解码
Transformers 仓库按用途提供了多套 PLBART 入口,均可通过from transformers import ...直接使用:
PLBartModel
无 Head 的裸 encoder-decoder 模型,输出Seq2SeqModelOutput(含last_hidden_state、past_key_values、decoder_attentions、cross_attentions与编码器侧各字段)。值得注意的是它的 forward 会在未提供decoder_input_ids时自动用shift_tokens_right从input_ids生成(modeling_plbart.py),因此可直接跑 denoising 预训练式的前向。
PLBartForConditionalGeneration
seq2seq 主模型:PLBartModel之上叠加lm_head(线性层,与输入嵌入共享权重)与final_logits_bias(nn.Buffer)。若labels未提供而decoder_input_ids未给,会默认按shift_tokens_right生成解码器输入以支持去噪;提供了labels则计算交叉熵损失并返回Seq2SeqLMOutput。若扩展词表,需调用resize_token_embeddings,它会同步_resize_final_logits_bias以维持偏置形状。
PLBartForSequenceClassification
用于程序理解类判别任务。在PLBartModel之上挂PLBartClassificationHead(Dense → tanh → dropout → 分类线性层),分类特征取自输入序列最后一个<eos>对应的隐状态(modeling_plbart.py),因此要求每个样本至少包含一个</s>token。labels支持回归(num_labels=1走 MSE)、单标签分类与多标签分类(自动推断problem_type),对应可复现论文中 clone detection、vulnerable code detection 这类"理解型"任务。
PLBartForCausalLM
纯解码器版语言模型。构造时会改写config.is_decoder=True、config.is_encoder_decoder=False,内部使用PLBartDecoderWrapper包装 decoder,并叠加与解码器词嵌入共享权重的lm_head(modeling_plbart.py)。它也可以在EncoderDecoderModel框架中作为解码器使用。例如从uclanlp/plbart-base直接加载后对一段文本做条件语言建模(docstring 内给出PLBartForCausalLM的完整示例)。
八、可迁移参考的任务指南
PLBART 覆盖的几类典型任务,在仓库中均有对应的端到端任务指南(与本模型正交、按AutoModel*通用接口组织),可直接迁移:
- 文本分类任务指南
- 因果语言建模任务指南
- 翻译任务指南
- 摘要任务指南
九、使用注意事项小结
- 语言 ID 必须正确:
src_lang/tgt_lang只接受java、python、en_XX(base),或扩展的javascript、php、ruby、go(multi),大小写与映射见FAIRSEQ_LANGUAGE_CODES_MAP; - 推理显式给解码起点:跨语言生成务必设置
decoder_start_token_id=tokenizer.lang_code_to_id[<tgt>];单语言微调场景可省略; bos不参与序列构建:不要手动拼接<s>,tokenizer 的prefix_suffix模式会自动只追加[eos][<LID>];- 数据集预处理占位符:代码文本中的换行/缩进占位符(如
NEW_LINE_INDENT)需要与训练语料一致地保留; - 注意力后端:默认支持 eager 与 SDPA,开启 FlashAttention 时需满足模型页面标注的能力约束与硬件前提;
- 以上 checkpoint、参数取值与调用方式均以本仓库实现为准;实际使用以
from_pretrained拉取到的 Hub 配置为准。
【免费下载链接】transformers🤗 Transformers: the model-definition framework for state-of-the-art machine learning models in text, vision, audio, and multimodal models, for both inference and training.项目地址: https://gitcode.com/GitHub_Trending/tra/transformers
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考