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

微调HuggingFace零样本模型遇评估问题:准确率低+F1指标报错

问题描述

我正在基于HuggingFace微调facebook/bart-large-mnli模型,使用的训练参数与代码如下:

training_args = TrainingArguments(
    output_dir=model_directory,      # 输出目录
    num_train_epochs=30,              # 训练总轮次
    per_device_train_batch_size=1,  # 单设备训练批次大小 - 原计划16,超过会触发内存不足
    per_device_eval_batch_size=2,   # 单设备评估批次大小 - 原计划64,超过会触发内存不足
    warmup_steps=500,                 # 学习率预热步数
    weight_decay=0.01,               # 权重衰减强度
)

model = BartForSequenceClassification.from_pretrained("facebook/bart-large-mnli")

trainer = Trainer(
    model=model,                          # 待训练的🤗 Transformers模型实例
    args=training_args,                   # 上述训练参数
    compute_metrics=compute_metrics,      # 指标计算函数
    train_dataset=train_dataset,          # 训练数据集
    eval_dataset=test_dataset             # 测试数据集
)

# 启动训练
trainer.train()

初始使用的compute_metrics函数:

import numpy as np
from datasets import Dataset, load_metric
from transformers import EvalPrediction

def compute_metrics(p: EvalPrediction):
  metric_acc = load_metric("accuracy")
  preds = p.predictions[0] if isinstance(p.predictions, tuple) else p.predictions
  preds = np.argmax(preds, axis=1)
  result = {}
  result["accuracy"] = metric_acc.compute(predictions=preds, references=p.label_ids)["accuracy"]
  return result

但无论调整训练/测试数据量、训练轮次,调用trainer.evaluate()得到的准确率始终固定为0.5。我有两个问题:

  1. 如何提升模型的准确率?
  2. 如何添加F1等其他评估指标?

我尝试修改compute_metrics添加F1指标,代码如下:

def compute_metrics(p: EvalPrediction):
  load_accuracy = load_metric("accuracy")
  load_f1 = load_metric("f1")
  preds = p.predictions[0] if isinstance(p.predictions, tuple) else p.predictions
  preds = np.argmax(preds, axis=1)
  result = {}
  result["accuracy"] = load_accuracy.compute(predictions=preds, references=p.label_ids)["accuracy"]
  result["f1"] = load_f1.compute(predictions=preds, references=p.label_ids)["f1"]
  return result

但调用trainer.evaluate()时触发错误:

ValueError: pos_label=1 is not a valid label. It should be one of [0, 2]

补充信息:

  • 使用的分词器:
from transformers import BartTokenizerFast
tokenizer = BartTokenizerFast.from_pretrained('facebook/bart-large-mnli')
  • 数据集创建与转换采用标准文本分类适配逻辑

解决方案

一、解决F1指标计算的报错问题

报错核心原因是load_metric("f1")默认适配二分类场景,默认正类标签为1,但你的数据集标签范围是[0,2](属于三分类场景),需要明确指定average参数适配多分类计算:

修改后的compute_metrics函数:

import numpy as np
from datasets import load_metric
from transformers import EvalPrediction

def compute_metrics(p: EvalPrediction):
    metric_acc = load_metric("accuracy")
    metric_f1 = load_metric("f1")
    
    preds = p.predictions[0] if isinstance(p.predictions, tuple) else p.predictions
    preds = np.argmax(preds, axis=1)
    
    accuracy = metric_acc.compute(predictions=preds, references=p.label_ids)["accuracy"]
    # 根据任务需求选择average参数:
    # - "macro":计算各类别F1的算术平均值
    # - "weighted":按类别样本量加权计算F1平均值
    # - "micro":计算全局精确率、召回率后得到的F1
    f1_score = metric_f1.compute(predictions=preds, references=p.label_ids, average="macro")["f1"]
    
    return {"accuracy": accuracy, "f1": f1_score}

如果你的任务实际是二分类(仅标签用了0和2),也可以手动指定pos_label=2,但更推荐根据实际分类类型选择average参数。

二、提升模型准确率的常见优化方向

1. 检查数据集处理逻辑

  • 确认标签映射正确性:facebook/bart-large-mnli默认是三分类(ENTAILMENT=0, NEUTRAL=1, CONTRADICTION=2),如果你的任务是二分类,需确保模型分类头输出维度与任务匹配,或调整数据集标签到对应范围。
  • 验证输入格式:MNLI是文本对任务,输入需包含premise和hypothesis;如果是单文本分类任务,需调整输入格式(比如将单文本作为premise,hypothesis设为固定分类描述),或修改模型分类头适配单文本输入。
  • 排查数据问题:检查训练/测试集是否存在数据泄露、标签错误,或样本分布极度不均衡(比如两类样本各占50%,随机预测也能得到0.5准确率)。

2. 调整训练参数

  • 指定学习率:默认学习率可能不适配任务,建议添加learning_rate=2e-5或5e-5(大模型微调常用范围)到TrainingArguments中。
  • 模拟大批次训练:内存不足时,可设置gradient_accumulation_steps=8(等价于批次大小为8),提升训练稳定性。
  • 添加早停机制:设置evaluation_strategy="epoch"+early_stopping_patience=3,每轮评估后根据验证集性能提前停止,避免过拟合。
  • 调整权重衰减:尝试将weight_decay改为0.001或0.1,观察模型性能变化。

3. 模型适配与初始化

  • 尝试部分参数微调:比如仅训练分类头,或使用peft库做LoRA微调,减少内存占用的同时避免破坏预训练模型的通用特征。
  • 检查分类头初始化:如果任务与MNLI差异较大,可重新初始化模型分类头,避免预训练权重干扰。

4. 其他优化

  • 添加梯度裁剪:在TrainingArguments中设置max_grad_norm=1.0,防止梯度爆炸。
  • 更换学习率调度器:设置lr_scheduler_type="cosine",让学习率在训练后期逐步下降,提升收敛效果。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.21 18:07:02