如何用HuggingFace基于MLM等预训练任务微调facebook/mbart-large-cc25
HuggingFace框架下微调mBART-large-cc25做MLM任务实现步骤
1. 依赖安装
先安装所需的核心库:
pip install transformers datasets torch evaluate accelerate
2. 数据与基础组件准备
你的自有数据集只要包含单条文本列即可,支持csv、json、txt等常见格式,这里假设数据集文本列名为text。
首先加载模型和分词器,单语言训练可以指定对应语言代码(比如中文是zh_CN,多语言训练可省略该参数):
from transformers import AutoTokenizer, AutoModelForMaskedLM, DataCollatorForLanguageModeling model_name = "facebook/mbart-large-cc25" # 按你实际使用的语言调整src_lang和tgt_lang参数 tokenizer = AutoTokenizer.from_pretrained(model_name, src_lang="zh_CN", tgt_lang="zh_CN") model = AutoModelForMaskedLM.from_pretrained(model_name)
加载并预处理自有数据集:
from datasets import load_dataset # 替换成你自己的数据集路径和划分规则 dataset = load_dataset("csv", data_files="your_dataset.csv", split="train") def preprocess_function(examples): # mBART最大输入长度为1024,按需求调整截断、填充规则 return tokenizer(examples["text"], truncation=True, max_length=1024, padding="max_length") # 批量处理数据集,移除原文本列节省空间 tokenized_dataset = dataset.map(preprocess_function, batched=True, remove_columns=["text"])
准备MLM专用的数据收集器,自动完成输入掩码构造,mlm_probability是掩码比例,默认0.15和预训练时保持一致即可:
data_collator = DataCollatorForLanguageModeling( tokenizer=tokenizer, mlm=True, mlm_probability=0.15 )
3. 训练配置与启动
用Trainer API封装训练逻辑,不需要手动写训练循环:
from transformers import TrainingArguments, Trainer training_args = TrainingArguments( output_dir="./mbart-mlm-finetuned", # 模型日志、 checkpoint 保存路径 per_device_train_batch_size=2, # 按显存大小调整,24G显存可开到2~4 gradient_accumulation_steps=4, # 显存不够可以开梯度积累,等效扩大batch size learning_rate=2e-5, # 预训练模型微调常用学习率范围 num_train_epochs=3, weight_decay=0.01, logging_steps=10, save_strategy="epoch", fp16=True, # 安培及以上架构N卡可开混合精度加速 push_to_hub=False # 不需要上传到HuggingFace Hub就关闭 ) trainer = Trainer( model=model, args=training_args, train_dataset=tokenized_dataset, data_collator=data_collator, )
启动训练并保存最终模型:
trainer.train() # 训练完成后保存最终版本的模型和分词器 trainer.save_model("./mbart-mlm-finetuned-final") tokenizer.save_pretrained("./mbart-mlm-finetuned-final")
常见优化提示
- 显存不足的话可以在
TrainingArguments中添加gradient_checkpointing=True,能节省近50%显存,训练速度会略有下降 - 需要评估训练效果的话,可以把数据集拆分为训练集和验证集,在
Trainer中传入eval_dataset参数,训练过程会自动输出验证集困惑度 - 多语言数据集不需要指定语言代码,分词器可自动适配不同语言的输入
内容的提问来源于stack exchange,提问作者albero
相关产品推荐
相关产品推荐

