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

如何让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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.13 08:20:11