如何利用PyTorch Transform为输入输出图像对施加相同变换?
解决PyTorch中图像对的同步随机数据增强问题
在构建输入输出尺寸相同的图像增强模型时,核心需求是让随机变换的参数对输入、输出图像完全一致。以下是两种高效的解决方案:
方法一:使用TorchVision v2原生支持的成对变换
TorchVision 0.15及以上的v2版本,所有随机变换都原生支持对图像对(tuple/list形式)、字典甚至张量批量应用相同的随机参数,无需额外封装,是最简洁的方案。
代码示例
import torch from torchvision import transforms as v2_transforms # 定义包含随机变换的组合 paired_transform = v2_transforms.Compose([ v2_transforms.RandomHorizontalFlip(p=0.5), v2_transforms.RandomRotation(degrees=15), v2_transforms.ColorJitter(brightness=0.2, contrast=0.2), v2_transforms.ToTensor(), v2_transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) ]) # 假设input_img和target_img是PIL图像或张量格式的图像对 input_img, target_img = ... # 将图像对打包成tuple传入变换,自动同步参数 augmented_input, augmented_target = paired_transform((input_img, target_img))
原理:v2版本的变换在调用时会先一次性生成所有随机参数,再将相同参数应用到输入tuple中的每一个元素,确保输入输出的变换完全同步。
方法二:自定义成对变换类(兼容TorchVision v1.x)
如果仍在使用旧版TorchVision,可以通过自定义变换类,手动控制随机参数的生成与复用,实现同步变换。
代码示例
import random from torchvision import transforms, functional as F from PIL import Image # 自定义成对随机水平翻转 class PairedRandomHorizontalFlip: def __init__(self, p=0.5): self.p = p def __call__(self, img_pair): input_img, target_img = img_pair if random.random() < self.p: input_img = F.hflip(input_img) target_img = F.hflip(target_img) return input_img, target_img # 自定义成对随机旋转 class PairedRandomRotation: def __init__(self, degrees): self.degrees = degrees def __call__(self, img_pair): input_img, target_img = img_pair angle = random.uniform(-self.degrees, self.degrees) input_img = F.rotate(input_img, angle) target_img = F.rotate(target_img, angle) return input_img, target_img # 组合所有成对变换 paired_transform = transforms.Compose([ PairedRandomHorizontalFlip(p=0.5), PairedRandomRotation(degrees=15), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) ]) # 使用示例 input_img = Image.open("input.jpg") target_img = Image.open("target.jpg") aug_input, aug_target = paired_transform((input_img, target_img))
原理:每个自定义变换类在__call__方法中先生成随机参数(如翻转概率、旋转角度),再调用torchvision.transforms.functional中的无随机参数的变换函数,对输入和输出图像应用完全相同的变换操作。
内容的提问来源于stack exchange,提问作者user153245
相关产品推荐
相关产品推荐

