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

是否可以对facebook/bart-large-mnli模型进行重新训练?

基于自有数据集微调facebook/bart-large-mnli进行文本分类

如果你已经在用zero-shot分类处理任务,但想针对自有数据集和标签提升准确率,可以通过监督微调来适配facebook/bart-large-mnli模型,以下是完整实现流程:

一、数据准备

你的数据集需要是带标注的监督数据,每个样本包含待分类文本和对应的真实标签。比如可以用列表形式存储:

# 示例数据集,替换成你的自有数据
dataset = [
    {"text": "我计划去欧洲自驾游", "label": "travel"},
    {"text": "今天要做番茄炒蛋", "label": "cooking"},
    {"text": "报名了街舞培训班", "label": "dancing"},
    # 更多样本...
]

也可以从CSV/JSON文件加载,关键是明确文本字段和标签字段的对应关系。

二、加载模型与分词器

facebook/bart-large-mnli默认适配MNLI的3分类(蕴涵、矛盾、中性),我们需要修改num_labels为自定义标签数量,并重新初始化分类头:

from transformers import AutoModelForSequenceClassification, AutoTokenizer, Trainer, TrainingArguments
import torch
import numpy as np

# 加载预训练模型和分词器
model_name = "facebook/bart-large-mnli"
tokenizer = AutoTokenizer.from_pretrained(model_name)

# 映射自定义标签到id
label_list = ["travel", "cooking", "dancing"]  # 替换成你的标签集合
label2id = {label: idx for idx, label in enumerate(label_list)}
id2label = {idx: label for label, idx in label2id.items()}

# 加载模型,指定自定义标签数量
model = AutoModelForSequenceClassification.from_pretrained(
    model_name,
    num_labels=len(label_list),
    label2id=label2id,
    id2label=id2label
)

三、数据预处理

将文本转换为模型可接受的输入格式(input_ids、attention_mask),同时把标签转为id:

def preprocess_function(examples):
    # 对文本进行分词,截断/填充到模型最大长度
    tokenized_inputs = tokenizer(
        examples["text"],
        truncation=True,
        padding="max_length",
        max_length=512
    )
    # 转换标签为id
    tokenized_inputs["labels"] = [label2id[label] for label in examples["label"]]
    return tokenized_inputs

# 将示例数据集转换为Dataset格式(更适合transformers的Trainer)
from datasets import Dataset
train_dataset = Dataset.from_list(dataset).map(preprocess_function, batched=True)

四、配置训练参数

根据硬件资源调整训练基础参数:

training_args = TrainingArguments(
    output_dir="./bart-finetuned",  # 模型保存路径
    per_device_train_batch_size=4,  # 单GPU批次大小
    num_train_epochs=3,  # 训练轮数
    learning_rate=2e-5,  # BERT/BART类模型常用学习率范围2e-5~5e-5
    logging_dir="./logs",  # 日志保存路径
    logging_steps=10,
    save_total_limit=2,  # 最多保存2个模型
    fp16=True  # 支持混合精度训练的GPU可开启,加速训练
)

五、训练模型

使用Trainer类启动训练:

trainer = Trainer(
    model=model,
    args=training_args,
    train_dataset=train_dataset
)

# 开始训练
trainer.train()

# 保存微调后的模型和分词器
trainer.save_model("./bart-finetuned-final")
tokenizer.save_pretrained("./bart-finetuned-final")

六、使用微调后的模型推理

训练完成后,有两种方式进行推理:

方法1:用pipeline快速调用

from transformers import pipeline

classifier = pipeline(
    "text-classification",
    model="./bart-finetuned-final",
    tokenizer="./bart-finetuned-final"
)

# 测试样本
result = classifier("我准备去云南拍风景照")
print(result)
# 示例输出:[{'label': 'travel', 'score': 0.9876}]

方法2:手动推理(更灵活)

from transformers import AutoModelForSequenceClassification, AutoTokenizer

model = AutoModelForSequenceClassification.from_pretrained("./bart-finetuned-final")
tokenizer = AutoTokenizer.from_pretrained("./bart-finetuned-final")

text = "我准备去云南拍风景照"
inputs = tokenizer(text, return_tensors="pt", truncation=True, padding=True)

with torch.no_grad():
    outputs = model(**inputs)
    logits = outputs.logits
    predictions = torch.argmax(logits, dim=-1)

print(f"分类结果:{id2label[predictions.item()]}")
print(f"置信度:{torch.softmax(logits, dim=-1)[0][predictions.item()].item():.4f}")

注意事项

  • 数据集样本量建议至少几百条以上,否则微调效果可能不如zero-shot分类
  • 如果是多标签分类,需修改模型损失函数(如BCEWithLogitsLoss),并将标签转为one-hot编码
  • facebook/bart-large-mnli模型体积较大,训练时建议使用GPU,CPU训练速度极慢

内容的提问来源于stack exchange,提问作者Sunny

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.09 18:25:18