使用SimpleTransformers训练MBART时缺失decoder_start_token_id报错如何解决
问题解决方案
错误原因
该报错属于SimpleTransformers适配MBART模型时的参数缺失问题,和普通单语言BART不同,MBART作为多语言序列到序列模型,需要显式指定源/目标语言标识,同时你当前代码存在自定义令牌添加不全的问题,共同触发了参数缺失报错。
修复步骤
- 新增MBART必填的源语言、目标语言参数,你使用的
IlyaGusev/mbart_ru_sum_gazeta是俄文摘要模型,对应语言代码为ru_RU - 删除冗余的bos_token添加代码,该特殊令牌已在预训练模型中内置,手动添加不会生效反而可能引发冲突
- 补全解码器的自定义令牌添加逻辑,保证编码器、解码器词表大小一致
- 显式指定模型配置的decoder_start_token_id,避免框架自动读取失败
修改后可用代码
from simpletransformers.seq2seq import Seq2SeqModel, Seq2SeqArgs # Model Config model_args = Seq2SeqArgs() model_args.do_sample = True model_args.eval_batch_size = 4 model_args.evaluate_during_training = True model_args.evaluate_during_training_steps = 2500 model_args.evaluate_during_training_verbose = True model_args.fp16 = False model_args.learning_rate = 5e-5 model_args.max_length = 128 model_args.max_seq_length = 128 model_args.num_beams = 10 model_args.num_return_sequences = 3 model_args.num_train_epochs = 2 model_args.overwrite_output_dir = True model_args.reprocess_input_data = True model_args.save_eval_checkpoints = False model_args.save_steps = -1 model_args.top_k = 50 model_args.top_p = 0.95 model_args.train_batch_size = 4 model_args.use_multiprocessing = False # 新增MBART必填语言参数 model_args.src_lang = "ru_RU" model_args.tgt_lang = "ru_RU" model_ru = Seq2SeqModel( encoder_decoder_type="mbart", encoder_decoder_name="IlyaGusev/mbart_ru_sum_gazeta", args=model_args, use_cuda=True ) # 同时给编码器、解码器添加自定义令牌 new_tokens = ["token1", "token2"] model_ru.encoder_tokenizer.add_tokens(new_tokens) model_ru.decoder_tokenizer.add_tokens(new_tokens) # 显式指定decoder_start_token_id model_ru.model.config.decoder_start_token_id = model_ru.decoder_tokenizer.lang_code_to_id[model_args.tgt_lang] model_ru.model.resize_token_embeddings(len(model_ru.encoder_tokenizer)) model_ru.train_model(train, eval_data=dev)
如果修改后仍报错,可以额外手动补充pad_token_id配置:
model_ru.model.config.pad_token_id = model_ru.decoder_tokenizer.pad_token_id
内容的提问来源于stack exchange,提问作者LeOverflow
相关产品推荐
相关产品推荐

