是否可以对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
相关产品推荐
相关产品推荐

