微调BERT完成NLI任务时Hugging Face多分类评估报错咨询
多分类NLI任务的Hugging Face评估指标适配方案
你遇到的错误是因为f1、precision、recall这类指标默认采用average='binary'的计算方式,仅适用于二分类场景,而NLI是典型的三分类任务(对应蕴含、中立、矛盾三类),必须指定多分类兼容的平均策略。
解决方案:为指标指定多分类平均策略
不能直接通过字符串列表组合指标,需要为每个需要调整的指标单独加载并设置average参数,再进行组合。以下是具体实现:
示例代码(采用macro平均)
import evaluate # 组合多分类适配的指标:accuracy无需调整,f1/precision/recall指定macro平均 metric = evaluate.combine([ "accuracy", evaluate.load("f1", average="macro"), evaluate.load("precision", average="macro"), evaluate.load("recall", average="macro") ]) # 测试计算 metrics = metric.compute(predictions=[0,1,1,2], references=[0,2,1,0]) print(metrics)
可选的平均策略说明
根据任务需求选择合适的平均方式:
- macro:计算每个类别的指标值后取算术平均,平等对待每个类别,适合样本分布均衡的场景
- weighted:按每个类别的样本数量加权计算平均,适合样本不平衡的情况,避免小众类别被忽略
- micro:将所有类别的TP、FP、FN汇总后计算指标,更关注整体预测的正确性,与准确率在样本均衡时结果相近
- None:返回每个类别的单独指标值,适合需要分析单个类别表现的场景
带别名的指标组合写法
如果需要更清晰的指标命名,可以用字典形式组合:
metric = evaluate.combine({ "accuracy": evaluate.load("accuracy"), "f1_macro": evaluate.load("f1", average="macro"), "precision_weighted": evaluate.load("precision", average="weighted"), "recall_micro": evaluate.load("recall", average="micro") })
内容的提问来源于stack exchange,提问作者Ali ZareShahi
相关产品推荐
相关产品推荐

