如何基于XLNet或中文BERT实现Seq2SeqLM训练?
使用BERT/XLNet构建Seq2Seq模型的解决方案
你遇到的错误核心原因是:BERT(bert-base-chinese)和XLNet(xlnet-base-cased)属于单编码器架构模型,并非原生的Seq2Seq模型,因此无法通过AutoModelForSeq2SeqLM直接加载。要基于这类模型完成序列到序列任务(如粤译中),需要借助Transformers提供的EncoderDecoderModel框架,将编码器和解码器组合成完整的Seq2Seq模型。
修改步骤
1. 替换模型加载逻辑
删除原来的model = AutoModelForSeq2SeqLM.from_pretrained(checkpoint),根据你选择的预训练模型,使用以下代码构建Seq2Seq模型:
针对中文BERT(bert-base-chinese)
from transformers import EncoderDecoderModel, BertModel, BertConfig # 加载BERT作为编码器 encoder = BertModel.from_pretrained("bert-base-chinese") # 配置解码器:复用BERT配置,开启解码器模式与交叉注意力 decoder_config = BertConfig.from_pretrained("bert-base-chinese") decoder_config.is_decoder = True decoder_config.add_cross_attention = True decoder = BertModel.from_pretrained("bert-base-chinese", config=decoder_config) # 组合成EncoderDecoder模型 model = EncoderDecoderModel(encoder=encoder, decoder=decoder) # 配置必要的token参数(BERT默认pad token为[PAD],id=0) model.config.pad_token_id = tokenizer.pad_token_id # 设置解码器起始token(用[CLS])和结束token(用[SEP]) model.config.decoder_start_token_id = tokenizer.cls_token_id model.config.eos_token_id = tokenizer.sep_token_id
针对XLNet(xlnet-base-cased)
from transformers import EncoderDecoderModel, XLNetModel, XLNetConfig # 加载XLNet作为编码器 encoder = XLNetModel.from_pretrained("xlnet-base-cased") # 配置解码器:开启解码器模式与交叉注意力 decoder_config = XLNetConfig.from_pretrained("xlnet-base-cased") decoder_config.is_decoder = True decoder_config.add_cross_attention = True decoder = XLNetModel.from_pretrained("xlnet-base-cased", config=decoder_config) # 组合成EncoderDecoder模型 model = EncoderDecoderModel(encoder=encoder, decoder=decoder) # 配置XLNet的token参数 model.config.pad_token_id = tokenizer.pad_token_id model.config.decoder_start_token_id = tokenizer.cls_token_id # XLNet的<s>对应cls_token_id model.config.eos_token_id = tokenizer.sep_token_id # XLNet的</s>对应sep_token_id
2. 修正训练参数类名
你的代码中CustomSeq2SeqTrainingArguments属于笔误,需要替换为Transformers官方提供的Seq2SeqTrainingArguments:
from transformers import Seq2SeqTrainingArguments training_args = Seq2SeqTrainingArguments( output_dir="my-output-dir", evaluation_strategy="epoch", learning_rate=2e-5, per_device_train_batch_size=16, per_device_eval_batch_size=16, weight_decay=0.01, save_total_limit=3, num_train_epochs=2, predict_with_generate=True, remove_unused_columns=False, fp16=True, push_to_hub=False, # 暂不上传到Hub )
3. 其他注意事项
- 解码器必须开启
is_decoder=True和add_cross_attention=True,否则无法接收编码器的输出进行交叉注意力计算,无法完成Seq2Seq任务。 - 显式配置
pad_token_id、decoder_start_token_id、eos_token_id是必须的,否则训练或生成时会出现token不匹配的错误。 - 如果你想尝试混合架构(比如BERT做编码器,XLNet做解码器),只需替换对应部分的模型加载代码即可。
完整修改后训练流程
其余代码(数据集加载、预处理、评估函数、数据收集器等)无需修改,直接保留原逻辑即可。运行修改后的代码,就能基于BERT或XLNet完成Seq2Seq模型的训练。
内容的提问来源于stack exchange,提问作者Raptor
相关产品推荐
相关产品推荐

