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

使用R语言Keras生成多分类混淆矩阵 解决predict_classes弃用问题

R语言Keras弃用predict_classes后的多分类预测实现方案

你找到的代码是Keras 2.6+版本后替代predict_classes()的标准实现,逻辑分两步:

  • 首先调用predict(x.test)会输出每个测试样本对应所有分类的概率矩阵,矩阵维度为(测试样本数, 分类类别数)
  • 再通过k_argmax(axis = -1)取每一行概率最高的下标作为预测类别,axis=-1表示取最后一个维度的最大值索引,正好对应多分类场景下每个样本的概率维度

完整使用示例

# 加载依赖包
library(keras)
library(caret)

# 1. 执行预测,对应官方给出的代码
prediction1 <- model %>% 
  predict(x.test)  %>% 
  k_argmax(axis = -1)

# 2. 将tensor格式的预测结果转换为R可处理的向量格式,为后续生成混淆矩阵做准备
pred_labels <- as.array(prediction1)

# 3. 如果你的真实标签是one-hot编码,需要转换为普通整数向量;如果本身就是整数标签可直接跳过这步
true_labels <- apply(y.test, 1, which.max) - 1 # 注意Keras默认标签从0开始,which.max返回值从1开始因此需要减1对齐

# 4. 生成基础混淆矩阵
confusion_matrix <- table(True = true_labels, Predicted = pred_labels)
print(confusion_matrix)

# 如需调用caret包生成带评估指标的混淆矩阵(含准确率、F1值、召回率等)
confusionMatrix(factor(pred_labels), factor(true_labels))

注意事项

  • 如果你的原始类别标签是从1开始编码的,输出pred_labels后需要统一加1,和真实标签对齐后再计算混淆矩阵
  • 模型输出层softmax激活的单元数必须和分类类别数一致,否则k_argmax返回的索引会和实际类别不匹配

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.05 20:39:00