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

如何从多分类kNN模型获取带相对值的混淆矩阵?

问题:多分类场景下如何在mlr3中获取带相对值的混淆矩阵?

我使用mlr3和mlr3learners包,基于mclust包的diabetes数据集构建了一个简单的kNN模型,通过glucose、insulin、sspg三个数值特征预测class类别,并借助mlr3measures包的指标和混淆矩阵评估模型性能、分析各类别的误分类情况。目前通过以下代码得到了绝对值混淆矩阵:

代码示例

# load packages
library(mlr3)
library(mlr3learners) # for classif.kknn
library(mlr3measures) # for confusion_matrix()
library(mclust) # for data(diabetes)

# load data
data(diabetes, package = "mclust")
diabetes <- as.data.table(diabetes)

# define task
diabetes_task <- as_task_classif(diabetes, 
                                 target = "class", 
                                 id = "diabetes")

# define ML algorithm
knn_model <- lrn('classif.kknn')

# partition data
splits <- partition(diabetes_task) 

# train model
knn_model$train(diabetes_task, 
                row_ids = splits$train)

# test model 
prediction <- knn_model$predict(diabetes_task, 
                                row_ids = splits$test)

# evaluate performance
prediction$confusion 

绝对值混淆矩阵

truth
 response   Chemical Normal Overt
   Chemical       10      2     0
   Normal          2     23     0
   Overt           0      0    11

我需要带相对值的混淆矩阵,但发现mlr3measures包的confusion_matrix()函数虽有relative = TRUE参数,但仅支持二分类场景。旧版mlr包实现起来较简单,请问多分类场景下如何简便获取带相对值的混淆矩阵?


解决方案

在mlr3的多分类场景中,直接用R基础包的prop.table()函数就能快速将绝对值混淆矩阵转换为相对值矩阵,无需额外依赖包,操作简便。

具体实现代码

在现有代码基础上添加以下代码即可:

# 按行计算相对值(每行代表预测类别,值为该类别下各真实类别的占比)
relative_confusion_row <- prop.table(prediction$confusion, margin = 1)
print(relative_confusion_row)

# 按列计算相对值(每列代表真实类别,值为该类别下各预测类别的占比)
relative_confusion_col <- prop.table(prediction$confusion, margin = 2)
print(relative_confusion_col)

# 全局相对值(所有样本中的占比)
relative_confusion_total <- prop.table(prediction$confusion)
print(relative_confusion_total)

参数说明

  • margin = 1:按行归一化,适合分析预测为某类的样本中真实类别的分布,比如能看到预测为Chemical的样本里,真实为Chemical的比例、误判为Normal的比例。
  • margin = 2:按列归一化,适合分析真实为某类的样本中预测类别的分布,比如能看到真实为Normal的样本里,被正确预测的比例、误判为Chemical的比例。
  • 不指定margin:计算全局占比,即每个单元格的数量占总测试样本数的比例。

示例输出(按行归一化)

以你的混淆矩阵为例,按行计算的相对值矩阵如下:

truth
 response   Chemical    Normal     Overt
   Chemical 0.8333333 0.1666667 0.0000000
   Normal   0.0800000 0.9200000 0.0000000
   Overt    0.0000000 0.0000000 1.0000000

这样就能清晰查看各类别的误分类比例,满足多分类场景的评估需求。


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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.08 03:33:11