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

使用k_argmax()替代predict_classes()后如何获取CNN预测值并生成混淆矩阵

解决方法

1. 将Tensor对象转换为R可操作的向量

k_argmax()返回的是TensorFlow的tf.Tensor对象,无法直接用R常规函数处理,需先转换为R的整数向量:

# 转换为整数向量
cnn_pred_vec <- as.integer(cnn_pred)

# 查看前10个预测值验证转换结果
head(cnn_pred_vec, 10)

2. 生成混淆矩阵

假设真实标签存储在y_test中(若y_test是独热编码格式,需先执行k_argmax(y_test) %>% as.integer()转换为整数向量),可通过两种方式生成混淆矩阵:

方式一:用Base R的table()函数生成基础混淆矩阵

# 生成混淆矩阵
confusion_matrix <- table(实际标签 = y_test, 预测标签 = cnn_pred_vec)

# 打印矩阵
print(confusion_matrix)

方式二:用caret包生成带指标的详细混淆矩阵

如果需要包含准确率、召回率、F1值等指标,可使用caret包:

# 首次使用需安装包
install.packages("caret")
library(caret)

# 生成带评估指标的混淆矩阵(需将标签转为因子类型)
confusion_matrix <- confusionMatrix(factor(cnn_pred_vec), factor(y_test))

# 打印完整结果
print(confusion_matrix)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.23 05:24:30