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

