You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何用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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.10.02 18:24:01