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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.03 08:15:26