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

Sklearn中NMF输入类别与转换类别映射匹配问题求助

解决NMF输出类别与真实类别自动映射的问题

不用手动试错,你可以通过**匈牙利算法(Kuhn-Munkres算法)**自动找到输出分量与真实类别之间的最优匹配,最大化分类准确率,具体步骤如下:

1. 生成NMF的预测类别

先训练NMF模型,对每个样本取权重最大的输出分量作为预测类别(如果你的真实类别是字符串,先把它编码为0、1、2这类整数,方便后续计算):

from sklearn.decomposition import NMF
import numpy as np

# 假设X是输入特征,y_true是编码后的真实类别(如[0,0,1,1,2,2])
nmf = NMF(n_components=3, random_state=42)
W = nmf.fit_transform(X)
y_pred = np.argmax(W, axis=1)  # 每个样本取权重最高的分量作为预测类别

2. 构建混淆矩阵并求解最优映射

利用混淆矩阵统计真实类别与预测类别的对应关系,再通过匈牙利算法找到能最大化准确率的映射:

from sklearn.metrics import confusion_matrix
from scipy.optimize import linear_sum_assignment

# 生成混淆矩阵,cm[i][j]代表真实类别i被预测为j的样本数
cm = confusion_matrix(y_true, y_pred)

# 把最大化准确率的问题转为最小化损失问题(用矩阵最大值减去每个元素)
cost_matrix = cm.max() - cm
# 用匈牙利算法求解最优分配
row_ind, col_ind = linear_sum_assignment(cost_matrix)

# 得到真实类别到预测类别的最优映射字典
class_mapping = {real_cls: pred_cls for real_cls, pred_cls in zip(row_ind, col_ind)}

3. 修正预测类别并计算准确率

用得到的映射关系修正原始的NMF预测类别,就能和真实类别对应上,进而计算准确率:

# 修正预测类别
y_pred_mapped = np.array([class_mapping[pred] for pred in y_pred])

# 计算最终准确率
accuracy = np.mean(y_true == y_pred_mapped)

补充说明

  • 这个方法完全自动化,不管类别数量多少都能高效运行,避免了手动匹配的低效
  • 如果真实类别是字符串(如"A"、"B"、"C"),可以先用LabelEncoder编码为整数,处理完后再转回字符串即可
  • 匈牙利算法的时间复杂度是O(n³),n为类别数,即使类别数达到几十也能快速完成计算

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.21 15:21:02