Sklearn1.0多标签随机森林传class_weight报ValueError问题
scikit-learn 1.0 多标签分类RandomForest传入class_weight列表报错排查
场景说明
- 任务为5标签多标签分类,每个标签取值为0或1,单样本可同时携带多个标签,各标签类别样本量分布不均
- 目标是通过
RandomForestClassifier的class_weight参数配置类别权重,缓解类别不平衡问题 - 运行环境为scikit-learn 1.0版本
参考的传参规则
根据官方文档描述,class_weight参数支持的传值类型包括字典、字典列表、"balanced"、None,默认值为None:
- 单输出场景权重格式为
{类别标签: 权重值},传None时所有类别权重默认取1 - 多输出(含多标签)场景可传入与标签列顺序一致的字典列表,需要为每一列的每个类别单独定义权重
4标签场景正确传参示例:
[{0: 1, 1: 1}, {0: 1, 1: 5}, {0: 1, 1: 1}, {0: 1, 1: 1}]
错误传参示例(仅指定单类权重):[{1:1}, {2:5}, {3:1}, {4:1}]
复现代码与报错
按照规则传入自定义权重列表的代码如下:
class_weights = [{0: 10, 1: 1}, {0: 6, 1: 1}, {0: 3, 1: 1}, {0: 1, 1: 1}, {0: 2, 1: 1}] forest = RandomForestClassifier(random_state=1, n_estimators=200, class_weight=class_weights)
运行后触发参数校验错误:
ValueError: class_weight must be dict, 'balanced', or None, got: [{0: 10, 1: 1}, {0: 6, 1: 1}, {0: 3, 1: 1}, {0: 1, 1: 1}, {0: 2, 1: 1}]
报错根因
文档中提到的“多输出场景支持传入字典列表格式class_weight”是scikit-learn更高版本才上线的特性,1.0版本的RandomForestClassifier参数校验逻辑仅认可字典、"balanced"、None三类合法输入,传入列表格式会直接判定为非法参数抛出错误,和配置的权重内容格式无关。
适配1.0版本的解决方案
- 方案1:直接传入
class_weight="balanced",模型会自动根据每个标签列的样本分布计算对应类别权重,计算规则为n_samples / (n_classes * np.bincount(y)),无需手动配置单列权重 - 方案2:需要自定义各标签权重时,使用
MultiOutputClassifier包装随机森林,为每个标签对应的基分类器单独配置class_weight,示例代码:
from sklearn.multioutput import MultiOutputClassifier from sklearn.ensemble import RandomForestClassifier # 按标签顺序初始化对应权重的基分类器 base_clfs = [ RandomForestClassifier(random_state=1, n_estimators=200, class_weight={0:10, 1:1}), RandomForestClassifier(random_state=1, n_estimators=200, class_weight={0:6, 1:1}), RandomForestClassifier(random_state=1, n_estimators=200, class_weight={0:3, 1:1}), RandomForestClassifier(random_state=1, n_estimators=200, class_weight={0:1, 1:1}), RandomForestClassifier(random_state=1, n_estimators=200, class_weight={0:2, 1:1}) ] # 包装为多输出分类器 model = MultiOutputClassifier(estimator=None, estimators=base_clfs)
- 方案3:将scikit-learn升级至1.2及以上版本,即可原生支持字典列表格式的class_weight传参,无需额外包装。
内容的提问来源于stack exchange,提问作者Energizer1
相关产品推荐
相关产品推荐

