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

如何在R/RStudio中从Keras模型提取训练模型并生成混淆矩阵?

在R的Keras中生成二分类混淆矩阵的方法

核心说明

history对象仅记录训练过程的指标变化(如损失、精度),不存储训练后的模型权重。你当前使用predict(mymodel, mytestdataset)是完全合理的——因为fit()方法会直接将训练得到的权重更新到mymodel对象中,无需从history里提取模型。

生成混淆矩阵的完整步骤

1. 将预测概率转换为类别标签

由于你的模型用sigmoid输出0-1的概率,需要设定阈值(通常用0.5)转为二分类标签:

# 以0.5为阈值,将概率转为0/1类别
pred_classes <- ifelse(predictions > 0.5, 1, 0)

2. 获取测试集的真实标签

根据你的测试集结构提取真实标签,比如:

# 假设测试集标签存储在mytestdataset$y中,根据实际情况调整
true_classes <- mytestdataset$y

3. 生成混淆矩阵

推荐使用caret包生成带详细指标的混淆矩阵,也可以用base R的table()快速生成:

方法一:用caret包(含精确率、召回率等指标)
# 安装并加载caret(首次使用需安装)
install.packages("caret")
library(caret)

# 生成混淆矩阵
cm <- confusionMatrix(as.factor(pred_classes), as.factor(true_classes))
# 打印完整结果(含真阳性、真阴性、精确率等)
print(cm)
方法二:用base R的table()
# 生成基础混淆矩阵表格
cm_table <- table(预测类别 = pred_classes, 真实类别 = true_classes)
print(cm_table)

额外:使用训练过程中保存的最优模型

你设置了callback_model_checkpoint保存最优模型到convnet.tf,如果想使用这个最优模型而非最后一轮训练的模型,可以加载它:

# 加载最优模型
best_model <- load_model_tf("convnet.tf")
# 用最优模型预测
predictions_best <- predict(best_model, mytestdataset)
# 后续步骤同前:转标签、生成混淆矩阵

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.25 11:07:04