TensorFlow中遍历张量并保存CNN预测图像的实现方案
解决方案:保存CNN预测结果对应的图像
要实现保存带有预测/真实类别和索引的图像,你需要在验证阶段获取预测标签的具体值,然后遍历每个batch的图像进行处理和保存。下面是具体的实现步骤和修改后的代码:
第一步:准备保存目录与依赖库
首先确保你导入了图像保存所需的库,并创建专门目录存放预测图像,避免文件混乱:
import os from skimage.io import imsave # 也可以用 PIL:from PIL import Image import numpy as np # 创建保存预测图像的目录,不存在则自动创建 pred_image_dir = "./predicted_images" os.makedirs(pred_image_dir, exist_ok=True)
第二步:修改验证循环,添加图像保存逻辑
在你的验证循环中,需要把y_pred_cls加入到sess.run的列表中(获取预测标签的numpy数组),同时还原预处理后的图像数据,最后遍历生成带信息的文件名并保存:
# Epoch completed, start validation print("{} Start validation".format(datetime.now().strftime('%Y-%m-%d %H:%M:%S'))) val_acc = 0. val_count = 0 cm_running_total = None global_idx = 0 # 全局索引,避免不同batch的图像编号重复 for _ in range(val_batches_per_epoch): batch_tx, batch_ty = val_preprocessor.next_batch(FLAGS.batch_size) # 新增y_pred_cls到run列表,获取预测标签的具体值 acc, loss, conf_m, y_pred = sess.run( [accuracy, cost, tf.confusion_matrix(y_true_cls, y_pred_cls, FLAGS.num_classes), y_pred_cls], feed_dict={x: batch_tx, y_true: batch_ty} ) if cm_running_total is None: cm_running_total = conf_m else: cm_running_total += conf_m val_acc += acc val_count += 1 # ---------- 新增:保存预测图像的核心代码 ---------- # 1. 还原图像数据:根据你的预处理逻辑调整 # 假设你预处理时把图像归一化到了[0,1],这里转成0-255的uint8格式 batch_images = (batch_tx * 255).astype(np.uint8) # 如果是归一化到[-1,1],则用下面的代码: # batch_images = ((batch_tx + 1) / 2 * 255).astype(np.uint8) # 2. 获取真实类别索引:如果batch_ty是one-hot编码,转成类别索引 y_true = np.argmax(batch_ty, axis=1) if len(batch_ty.shape) == 2 else batch_ty # 3. 遍历当前batch的每一张图像 for idx_in_batch in range(FLAGS.batch_size): # 生成文件名:包含全局索引、预测类别、真实类别 file_name = f"idx_{global_idx}_pred_{y_pred[idx_in_batch]}_true_{y_true[idx_in_batch]}.png" save_path = os.path.join(pred_image_dir, file_name) # 4. 保存图像(二选一即可) imsave(save_path, batch_images[idx_in_batch]) # 或者用PIL实现: # img = Image.fromarray(batch_images[idx_in_batch]) # img.save(save_path) global_idx += 1 # ---------- 保存图像代码结束 ---------- val_acc /= val_count s = tf.Summary(value=[ tf.Summary.Value(tag="validation_accuracy", simple_value=val_acc), tf.Summary.Value(tag="validation_loss", simple_value=loss) ]) val_writer.add_summary(s, epoch + 1)
关键细节说明
- 图像还原逻辑:必须和你预处理时的归一化方式对应,否则保存的图像会显示异常。
- 真实类别提取:如果
batch_ty是one-hot编码(形状为(batch_size, num_classes)),用np.argmax转成类别索引;如果已经是类别索引(形状为(batch_size,)),直接使用即可。 - 全局索引:用
global_idx累计所有验证图像的编号,避免不同batch的图像索引重复,方便后续查找特定图像。 - 目录创建:
os.makedirs(..., exist_ok=True)确保目录存在,不会因为目录未创建而报错。
内容的提问来源于stack exchange,提问作者A.M
相关产品推荐
相关产品推荐

