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

如何存储Torchvision/Albumentations中应用的精确变换参数?

如何保存torchvision与albumentations中图像增强的精确变换信息

torchvision 实现方式

torchvision的随机变换默认不会主动返回变换参数,但可以通过自定义包装类捕获每一次变换的具体参数:

  • 随机裁剪(RandomCrop)
    自定义包装类继承RandomCrop,在调用时记录裁剪的坐标参数:

    import torchvision.transforms as transforms
    
    class TrackableRandomCrop(transforms.RandomCrop):
        def __call__(self, img):
            # 获取并保存裁剪参数:(top, left, height, width)
            self.crop_params = self.get_params(img, self.size)
            return transforms.functional.crop(img, *self.crop_params)
    
    # 使用示例
    crop_transform = TrackableRandomCrop(size=(224, 224))
    aug_img = crop_transform(original_img)
    # 获取本次裁剪的精确参数
    crop_top, crop_left, crop_h, crop_w = crop_transform.crop_params
    
  • 亮度对比度调整(ColorJitter)
    同样通过包装类捕获亮度、对比度等随机因子:

    class TrackableColorJitter(transforms.ColorJitter):
        def __call__(self, img):
            # 获取并保存调整参数:(brightness_factor, contrast_factor, saturation_factor, hue_factor)
            self.jitter_params = self.get_params(self.brightness, self.contrast, self.saturation, self.hue)
            # 按顺序应用变换
            img = transforms.functional.adjust_hue(img, self.jitter_params[3])
            img = transforms.functional.adjust_saturation(img, self.jitter_params[2])
            img = transforms.functional.adjust_contrast(img, self.jitter_params[1])
            img = transforms.functional.adjust_brightness(img, self.jitter_params[0])
            return img
    
    # 使用示例
    jitter_transform = TrackableColorJitter(brightness=0.2, contrast=0.2)
    aug_img = jitter_transform(original_img)
    # 获取本次调整的精确参数
    brightness, contrast, saturation, hue = jitter_transform.jitter_params
    

    其他随机变换(如RandomHorizontalFlip)可以用类似逻辑,记录是否执行了翻转的布尔值即可。

albumentations 实现方式

albumentations原生支持返回变换的精确参数,无需额外包装,只需在定义变换时启用返回元数据的选项:

import albumentations as A
from albumentations.pytorch import ToTensorV2

# 定义变换时设置return_dict=True,启用元数据返回
transform = A.Compose([
    A.RandomCrop(height=224, width=224),
    A.RandomBrightnessContrast(brightness_limit=0.2, contrast_limit=0.2),
    ToTensorV2()
], return_dict=True)

# 应用变换,返回包含增强图像和变换元数据的字典
result = transform(image=original_img)
aug_img = result['image']
# 获取所有变换的精确参数,存储在'replay'字段中
transform_params = result['replay']

transform_params中包含了每一步变换的具体参数:比如RandomCrop的x_min、y_min、height、width,RandomBrightnessContrast的alpha(亮度因子)、beta(对比度因子)等。如果需要复用这些参数对其他图像执行完全相同的变换,可以调用transform.replay(transform_params, image=another_image)。


内容的提问来源于stack exchange,提问作者Sarthak Mehrotra

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.03 15:27:32