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

如何利用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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.25 16:53:27