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

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)

关键细节说明

  1. 图像还原逻辑:必须和你预处理时的归一化方式对应,否则保存的图像会显示异常。
  2. 真实类别提取:如果batch_ty是one-hot编码(形状为(batch_size, num_classes)),用np.argmax转成类别索引;如果已经是类别索引(形状为(batch_size,)),直接使用即可。
  3. 全局索引:用global_idx累计所有验证图像的编号,避免不同batch的图像索引重复,方便后续查找特定图像。
  4. 目录创建:os.makedirs(..., exist_ok=True)确保目录存在,不会因为目录未创建而报错。

内容的提问来源于stack exchange,提问作者A.M

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.29 09:02:54