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

多标签分类中类别权重计算及样本权重采样的实现咨询

多标签分类任务中使用Scikit-Learn计算类别权重的指导

问题背景

拥有9973个训练样本的数据集,标签为独热编码格式,对应13个类别,训练标签形状为(9973, 13)。尝试以下代码时报错“参数过多”:

import numpy as np
from sklearn.utils.class_weight import compute_class_weight

y_integers = np.argmax(y, axis=1)
class_weights = compute_class_weight('balanced', np.unique(y_integers), y_integers)
d_class_weights = dict(enumerate(class_weights))

训练样本标签示例:

[0, 0, 0, 0, 0, 0, 1, 0, 0, 0, 0, 0, 0],
[0, 0, 1, 1, 0, 0, 0, 0, 0, 0, 0, 0, 0],
[0, 0, 0, 0, 0, 0, 0, 0, 0, 1, 0, 0, 0],
[0, 0, 1, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0],

疑问

  1. 如何在多标签分类问题中实现类别权重以解决数据不平衡问题?
  2. 目前代码已正常运行,但疑惑该方法是否适用于多标签场景,了解到可能需使用样本权重而非类别权重,如何实现样本权重采样?

解答

一、原代码的核心问题

你用np.argmax(y, axis=1)把独热编码转成单类别索引的做法完全不适用于多标签场景——它会直接丢弃样本中的多标签信息(比如第二个样本同时属于两个类别,转成索引后只保留一个)。而报错“参数过多”是因为Scikit-Learn新版本中compute_class_weight的参数顺序调整,需指定参数名调用,但即使修正参数,这个方法也不适合你的任务。

二、多标签场景的权重解决方案

多标签分类的不平衡需要同时考虑类别层面的不平衡和样本层面的不平衡,以下是具体实现方式:

1. 计算类别权重(针对每个类别)

为每个类别单独计算权重,反映该类别的样本占比:

import numpy as np

# 计算每个类别的样本出现次数(独热编码中1的数量)
class_counts = np.sum(y, axis=0)
total_samples = y.shape[0]
num_classes = y.shape[1]

# 计算balanced权重:总样本数/(类别数*该类别样本数)
class_weights = total_samples / (num_classes * class_counts)
# 转成字典格式,键为类别索引
class_weight_dict = dict(enumerate(class_weights))

如果使用MultiOutputClassifier,需要在子分类器中传入类别权重:

from sklearn.multioutput import MultiOutputClassifier
from sklearn.ensemble import RandomForestClassifier

clf = MultiOutputClassifier(
    RandomForestClassifier(class_weight=class_weight_dict),
    n_jobs=-1
)
clf.fit(X, y)

2. 计算样本权重(针对每个样本)

样本权重给每个样本分配权重,稀有类别占比高的样本权重更高,常见两种计算方式:

方式一:基于类别权重求和

将样本中所有正类别对应的权重相加:

# 每个样本的权重 = 样本标签中为1的类别权重之和
sample_weights = np.sum(y * class_weights, axis=1)

方式二:基于类别频率倒数求和

用类别频率的倒数计算,稀有类别对应的样本权重更高:

# 计算每个类别的频率
class_freq = class_counts / total_samples
# 样本权重 = 样本中所有正类别频率倒数之和
sample_weights = np.sum(y / class_freq, axis=1)

3. 使用样本权重训练模型

大部分Scikit-Learn模型的fit方法支持sample_weight参数,直接传入即可:

from sklearn.ensemble import RandomForestClassifier

clf = RandomForestClassifier()
clf.fit(X, y, sample_weight=sample_weights)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.23 03:27:38