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

多分类任务中Sklearn交叉验证F1/精确率/召回率全为NaN求助

多分类任务cross_validate评分全为NaN的解决办法

问题根源

  1. 基础指标默认不支持多分类:sklearn里precision、recall、f1这些指标默认是给二分类用的,多分类场景下必须明确指定平均策略(比如macro、micro、weighted),不然不仅算不对,要是某折里有类别没被模型预测到,还会出现除以0的情况,直接返回NaN。
  2. roc_auc默认不兼容多分类:默认的roc_auc只处理二分类,多分类得指定用One-vs-Rest(OVR)还是One-vs-One(OVO)的计算逻辑,不然直接返回NaN。
  3. 单个指标失效牵连全部评分:只要scoring字典里有一个指标算不出来,cross_validate可能会把所有评分都设成NaN,这就是你只留accuracy和balanced_accuracy时正常,加其他指标就全挂的原因。

解决步骤

1. 直接用sklearn提供的多分类指标别名

sklearn已经给多分类场景预定义好了指标别名,直接用就行:

  • 精度(precision):precision_macro、precision_micro、precision_weighted
  • 召回率(recall):recall_macro、recall_micro、recall_weighted
  • F1值:f1_macro、f1_micro、f1_weighted
  • ROC-AUC:roc_auc_ovr(OVR策略)、roc_auc_ovo(OVO策略)

修改后的scoring字典示例:

scoring = {
    'accuracy': 'accuracy',
    'balanced_accuracy': 'balanced_accuracy',
    'precision_macro': 'precision_macro',
    'recall_weighted': 'recall_weighted',
    'f1_micro': 'f1_micro',
    'roc_auc_ovr': 'roc_auc_ovr'
}

2. 自定义评分器(更灵活)

如果需要自定义参数(比如指定average方式),可以用make_scorer来创建评分器,比如:

from sklearn.metrics import make_scorer, precision_score, roc_auc_score

scoring = {
    'accuracy': 'accuracy',
    'balanced_accuracy': 'balanced_accuracy',
    'precision_custom': make_scorer(precision_score, average='weighted'),
    'roc_auc_ovr_custom': make_scorer(roc_auc_score, multi_class='ovr', needs_proba=True)
}

注意:用roc_auc_score时,得确保模型支持predict_proba方法(决策树是支持的),并且要加needs_proba=True参数,因为ROC-AUC需要概率输出。

3. 检查数据分布

虽然用了RepeatedStratifiedKFold来保证每个折里都有所有类别,但如果某个类别的样本量特别少,还是可能出现某折里类别缺失的情况。可以先检查每个类别的样本数:

import numpy as np
print(np.bincount(y))

要是有样本数为1甚至0的类别,要么合并类别,要么想办法补充样本。

修改后的完整代码

from sklearn.model_selection import cross_validate, RepeatedStratifiedKFold
from sklearn.preprocessing import LabelEncoder
from sklearn.tree import DecisionTreeClassifier

# 适配多分类的scoring字典
scoring = {
    'accuracy': 'accuracy',
    'balanced_accuracy': 'balanced_accuracy',
    'precision_macro': 'precision_macro',
    'recall_weighted': 'recall_weighted',
    'f1_micro': 'f1_micro',
    'roc_auc_ovr': 'roc_auc_ovr'
}

def load_dataset(df):
    data = df.values
    X, y = data[:, :-1], data[:, -1]
    y = LabelEncoder().fit_transform(y)
    return X, y
 
def evaluate_model(X, y, model):
    cv = RepeatedStratifiedKFold(n_splits=10, n_repeats=3, random_state=1)
    scores = cross_validate(model, X, y, scoring=scoring, cv=cv, n_jobs=-1)
    return scores

model = DecisionTreeClassifier()

X, y = load_dataset(df2)
results_without_nlp = evaluate_model(X, y, model)

# 输出结果
print(results_without_nlp)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.12 06:45:31