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

Scikit-learn交叉验证自定义评分函数参数错误求解

多标签分类下适配cross_validate的balanced_accuracy评分函数实现

问题场景

针对14位二进制标签的不平衡多标签分类问题,使用Tree Parzen Estimation进行超参数优化时,尝试用balanced_accuracy_score做交叉验证,但该函数需要类别索引而非one-hot矩阵,因此编写了包装lambda函数:

_balanced_accuracy_score = lambda grnd, pred: balanced_accuracy_score(grnd.argmax(axis=1), pred.argmax(axis=1))

运行后报错:

<lambda>() takes 2 positional arguments but 3 were given

问题原因

cross_validate调用自定义评分函数时,可能会额外传入第三个参数sample_weight(比如启用权重相关设置时),而你的lambda函数仅定义了两个参数,导致参数不匹配。

解决方法

方法1:兼容第三个可选参数

修改函数,允许接收sample_weight参数(可选择是否使用):

# lambda简化写法
_balanced_accuracy_score = lambda grnd, pred, sample_weight=None: balanced_accuracy_score(grnd.argmax(axis=1), pred.argmax(axis=1))

# 普通函数写法(更易读,方便后续扩展)
def _balanced_accuracy_score(grnd, pred, sample_weight=None):
    y_true = grnd.argmax(axis=1)
    y_pred = pred.argmax(axis=1)
    # 如果需要支持样本权重,直接将参数传递给底层函数即可
    return balanced_accuracy_score(y_true, y_pred, sample_weight=sample_weight)

方法2:用make_scorer包装评分函数

借助sklearn的make_scorer工具,自动处理cross_validate的参数传递逻辑:

from sklearn.metrics import make_scorer, balanced_accuracy_score

def balanced_acc_multilabel(y_true, y_pred):
    # 将one-hot编码的标签矩阵转为类别索引
    true_labels = y_true.argmax(axis=1)
    pred_labels = y_pred.argmax(axis=1)
    return balanced_accuracy_score(true_labels, pred_labels)

# 生成可被cross_validate直接使用的评分器
scoring = make_scorer(balanced_acc_multilabel)

之后将scoring作为参数传入cross_validate即可。

重要提醒

注意argmax(axis=1)仅适用于单标签多分类场景(每个样本仅属于一个类别)。如果是真正的多标签分类(样本可同时属于多个类别),balanced_accuracy_score并不适用,需改用多标签专属指标,比如基于每个标签计算balanced accuracy后取均值,或使用label_ranking_average_precision_score等。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.08 13:25:17