如何让Torchvision ImageFolder同时加载原图与灰度图训练GAN上色模型
高效实现ImageFolder同时加载原始图与灰度变换图的方案
你现在训练图像上色GAN,需要用ImageFolder同时拿到原始图和对应的灰度图,还要兼顾大数据量下的效率对吧?核心痛点就是避免两次读取磁盘图像,毕竟IO是大数据训练的最大瓶颈之一,而且还要保证原始图和灰度图的变换完全对齐(比如随机裁剪、翻转的位置要一样)。
核心思路
不要让ImageFolder只输出一种图,而是自定义一个变换类,让它在一次图像加载后,同时生成经过相同几何变换的原始图和灰度图——这样既省了一次磁盘读取,又能保证两张图的空间对应关系,完美适配你的需求。
具体实现步骤
首先我们要写一个自定义的Transform类,把原代码里的随机变换逻辑拆出来,让原始图和灰度图共享所有随机变换参数,然后分别处理灰度转换和颜色抖动:
import torchvision.transforms as transforms from torchvision.transforms import functional as F import random import torchvision class OriginalAndGrayscaleTransform: def __init__(self, loadSize, fineSize): self.loadSize = loadSize self.fineSize = fineSize # 保留原代码的Resize选项 self.resize_options = [ transforms.Resize(loadSize, interpolation=1), transforms.Resize(loadSize, interpolation=2), transforms.Resize(loadSize, interpolation=3), transforms.Resize((loadSize, loadSize), interpolation=1), transforms.Resize((loadSize, loadSize), interpolation=2), transforms.Resize((loadSize, loadSize), interpolation=3) ] # 保留原代码的裁剪插值选项 self.crop_interpolations = [1, 2, 3] # 颜色抖动变换 self.color_jitter = transforms.ColorJitter(brightness=0.1, contrast=0.1) def __call__(self, img): # 先复制原始PIL图像,避免后续变换修改原对象 orig_img = img.copy() # 1. 随机选择Resize方式并应用到原始图 resize_transform = random.choice(self.resize_options) orig_img = resize_transform(orig_img) # 2. 生成随机裁剪参数,复用参数到原始图 crop_interp = random.choice(self.crop_interpolations) i, j, h, w = transforms.RandomResizedCrop.get_params( orig_img, scale=(0.8, 1.0), ratio=(0.9, 1.1) # 可根据需求调整scale/ratio,原代码默认是(0.08,1.0) ) orig_img = F.resized_crop(orig_img, i, j, h, w, (self.fineSize, self.fineSize), crop_interp) # 3. 随机水平翻转,同样复用参数 if random.random() > 0.5: orig_img = F.hflip(orig_img) # 4. 基于变换后的原始图生成灰度图,并应用颜色抖动 gray_img = F.rgb_to_grayscale(orig_img, num_output_channels=3) gray_img = self.color_jitter(gray_img) # 5. 转成张量返回 return F.to_tensor(orig_img), F.to_tensor(gray_img)
然后修改你的数据加载函数,把自定义变换传进去:
def load_data_bw(opt): datapath = '/content/gdrive/My Drive/faces/2003' # 使用我们自定义的变换类 transform = OriginalAndGrayscaleTransform(opt['loadSize'], opt['fineSize']) dataset = torchvision.datasets.ImageFolder(datapath, transform=transform) return dataset
使用方式
最后创建DataLoader并迭代,就能直接拿到迭代次数、原始图batch和灰度图batch了:
from torch.utils.data import DataLoader # 初始化数据集和加载器 training_dataset = load_data_bw(opt) training_data_loader = DataLoader(training_dataset, batch_size=opt['batchSize'], shuffle=True) # 迭代训练 for iteration, (orig_data, gray_data) in enumerate(training_data_loader, 1): # orig_data: 原始图像batch,shape [batch_size, 3, fineSize, fineSize] # gray_data: 灰度图像batch,shape [batch_size, 3, fineSize, fineSize] # 这里写你的GAN训练逻辑即可 pass
为什么这个方案高效?
- 一次IO读取:每张图像只从磁盘加载一次,所有变换都在内存中完成,彻底避免了重复读取的开销,大数据量下速度提升非常明显。
- 变换完全对齐:原始图和灰度图共享所有随机变换参数,保证了两者的空间位置完全对应,这对图像上色任务的训练至关重要。
- 兼容原有逻辑:完全保留了你原代码中的所有数据增强策略,包括不同插值方式的随机选择、颜色抖动等,不需要调整训练流程。
内容的提问来源于stack exchange,提问作者Riham Hazem
相关产品推荐
相关产品推荐

