如何遍历DataLoader所有批次并正确保存全部图像?
问题分析与修复方案
你的代码仅保存第一个批次图像,核心问题有两个:
- 内层循环范围错误:用
test_loader.batch_size-1作为循环上限,既会漏掉当前批次最后一个样本,还会在DataLoader最后一批样本数小于设定batch_size时触发索引越界,导致后续批次循环直接中断。 - 文件名重复覆盖:每个批次的
i+1都从1开始,后一批次的文件会直接覆盖前一批次同名文件,视觉上表现为只保存了第一个批次。
修复后的代码
import os # 确保保存目录存在,避免因目录缺失导致保存失败 os.makedirs('/content/test_data', exist_ok=True) for batch_idx, (test_data, test_targets) in enumerate(test_loader): # 遍历当前批次所有样本,用len(test_data)适配最后一批的样本数量 for i in range(len(test_data)): # 提取单通道图像并转换为numpy数组 img = np.array(test_data[i, 0]) * 255 img = cv2.cvtColor(img, cv2.COLOR_GRAY2BGR) hsv = cv2.cvtColor(img, cv2.COLOR_BGR2HSV) low_black = np.array([0, 0, 0]) high_black = np.array([360, 255, 0]) mask = cv2.inRange(hsv, low_black, high_black) img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB) img[mask > 0] = random.choice(list(color_dict.values())) # 文件名加入batch_idx,避免不同批次文件重名覆盖 cv2.imwrite(f'/content/test_data/{test_targets[i].item()}_batch{batch_idx}_idx{i+1}.png', img)
关键修改说明
- 内层循环改用
len(test_data)遍历当前批次样本,适配任意批次的样本数量,包括不满batch_size的最后一批。 - 文件名新增
batch_idx标识,彻底解决跨批次文件重名覆盖问题。 - 增加
os.makedirs提前创建保存目录,避免因目录不存在导致的保存失败。
内容的提问来源于stack exchange,提问作者Sagnnik Biswas
相关产品推荐
相关产品推荐

