使用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
相关产品推荐
相关产品推荐

