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

在R/RStudio中使用Keras:从TensorFlow数据集提取标签

解决TensorFlow测试集提取标签并生成混淆矩阵的方案

核心前提

加载测试集时必须设置shuffle=FALSE,确保数据集顺序与预测结果完全对应,这是标签和预测匹配的关键。


方法1:直接从TensorFlow Dataset提取真实标签

假设你用tf$keras$utils$image_dataset_from_directory加载测试集,代码示例如下:

library(tensorflow)
library(tfdatasets)

# 加载测试集(务必关闭shuffle)
test_ds <- tf$keras$utils$image_dataset_from_directory(
  "path/to/test_data",
  image_size = c(224L, 224L),
  batch_size = 32L,
  shuffle = FALSE
)

# 提取真实标签并转换为R向量
true_labels <- c()
iterator <- as_iterator(test_ds)

while (!is.null(batch <- iterator_get_next(iterator))) {
  # 每个batch的第二个元素是标签张量
  batch_labels <- as.array(batch[[2]])
  true_labels <- c(true_labels, batch_labels)
}

方法2:同时提取标签与图像文件路径

如果需要关联图像文件和预测结果,可以手动构建数据集并保留路径:

# 获取测试集所有文件路径并排序(保证与Dataset加载顺序一致)
test_files <- list.files("path/to/test_data", recursive = TRUE, full.names = TRUE)
test_files <- sort(test_files)

# 构建带路径的数据集
test_ds <- tensor_slices_dataset(test_files) %>%
  dataset_map(function(file_path) {
    # 读取并预处理图像(按需调整)
    img <- tf$io$read_file(file_path)
    img <- tf$image$decode_jpeg(img, channels = 3L)
    img <- tf$image$resize(img, c(224L, 224L))
    img <- tf$keras$applications$mobilenet_v2$preprocess_input(img)
    
    # 从路径提取标签(假设文件夹名为0/1,可根据实际调整映射逻辑)
    label <- tf$strings$split(file_path, "/") %>%
      tf$slice(start = -2L, size = 1L) %>%
      tf$strings$to_number(tf$int32)
    
    list(img, label, file_path)
  }) %>%
  dataset_batch(32L)

# 提取标签和文件路径
true_labels <- c()
file_paths <- c()
iterator <- as_iterator(test_ds)

while (!is.null(batch <- iterator_get_next(iterator))) {
  true_labels <- c(true_labels, as.array(batch[[2]]))
  file_paths <- c(file_paths, as.character(batch[[3]]))
}

生成混淆矩阵

将预测结果转换为类别标签,再与真实标签对比:

# 假设predictions是模型输出的二分类结果(softmax输出为(n,2),sigmoid为(n,1))
# 处理softmax输出:
pred_labels <- apply(predictions, 1, which.max) - 1  # 转换为0/1标签

# 处理sigmoid输出:
# pred_labels <- as.integer(predictions > 0.5)

# 生成混淆矩阵(用base R或caret包)
table(Predicted = pred_labels, Actual = true_labels)

# 或用caret包生成详细混淆矩阵
library(caret)
confusionMatrix(factor(pred_labels), factor(true_labels))

内容的提问来源于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:52:37