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
相关产品推荐
相关产品推荐

