为何调用sklearn的compute_class_weight()时报参数数量不匹配错误?
问题成因
该报错由scikit-learn版本迭代引发的接口传参规则变更导致:
- scikit-learn 0.23及更早版本中,
compute_class_weight函数支持按位置依次传入class_weight、classes、y三个参数 - scikit-learn 0.24及后续版本中,函数仅第一个参数支持位置传入,剩余参数必须显式指定关键字才能传参,否则就会触发参数数量不匹配的类型错误
修复方案
修改原有代码,为后两个参数添加关键字声明即可正常运行:
from sklearn.utils import class_weight import numpy as np class_weights = class_weight.compute_class_weight( 'balanced', classes=np.unique(train_gen.classes), y=train_gen.classes) # 后续生成权重字典、传入模型训练的代码无需修改 train_class_weights = dict(enumerate(class_weights)) # model.fit_generator(..., class_weight=train_class_weights)
内容的提问来源于stack exchange,提问作者James__pxlwk
相关产品推荐
相关产品推荐

