R Keras中如何用k_argmax获取可用于混淆矩阵的预测类别向量
问题场景
- 无TensorFlow、Keras相关使用经验,跟随RStudio Keras官方入门教程开展练习,使用代码如下:
library(keras) mnist <- dataset_mnist() mnist$train$x <- mnist$train$x/255 mnist$test$x <- mnist$test$x/255 model <- keras_model_sequential() %>% layer_flatten(input_shape = c(28, 28)) %>% layer_dense(units = 128, activation = "relu") %>% layer_dropout(0.2) %>% layer_dense(10, activation = "softmax") summary(model) model %>% compile( loss = "sparse_categorical_crossentropy", optimizer = "adam", metrics = "accuracy" ) # 注意:compile和fit会直接修改传入的model对象,和大多数R函数的复制修改逻辑不同 model %>% fit( x = mnist$train$x, y = mnist$train$y, epochs = 5, validation_split = 0.3, verbose = 2 ) predictions <- predict(model, mnist$test$x) head(predictions, 2) class_predictions <- predict(model, mnist$test$x) %>% k_argmax() class_predictions
- 现存问题:
predict_classes函数已被官方弃用,报错提示k_argmax()为替代方案,但无法将k_argmax()的输出转换为存储0-9预测数字的普通R向量,无法传入confusionMatrix函数构建混淆矩阵。
解决方法
k_argmax()默认返回TensorFlow张量对象,不属于R原生数据结构,无法直接被R的建模相关函数识别,加一步类型转换即可。
方法1:转换k_argmax()输出为R向量
直接在k_argmax()后接as.vector()即可拉取张量数值到R内存,生成普通整数向量。传入confusionMatrix前注意将预测值、真实值转为水平一致的因子,避免某类别预测缺失导致的水平不匹配报错,代码如下:
library(caret) # 生成0-9的普通整数预测向量 class_predictions <- predict(model, mnist$test$x) %>% k_argmax() %>% as.vector() # 构建混淆矩阵,手动指定因子水平保证匹配 confusionMatrix( data = factor(class_predictions, levels = 0:9), reference = factor(mnist$test$y, levels = 0:9) )
方法2:用base R函数直接从概率矩阵取预测类别
不依赖Keras的张量运算函数,直接对预测输出的概率矩阵用max.col()取每行最大值位置,注意R索引从1开始,结果减1即可得到0-9的标签:
# 先拿到测试集预测概率矩阵 pred_prob <- predict(model, mnist$test$x) # 取每行概率最大的位置,减1匹配0-9的标签规则 class_predictions <- max.col(pred_prob) - 1
两种方法得到的
class_predictions都是普通R整数向量,可以直接传入后续所有R原生函数做分析。
内容的提问来源于stack exchange,提问作者Hans
相关产品推荐
相关产品推荐

