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

