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

如何逐图控制或记录torchvision transform实际使用的参数?

解决方案

完全可以通过在Dataset类内部实现变换逻辑满足你的需求,这也是实现逐图参数控制、参数同步记录的最简便方案。具体实现步骤如下:

步骤1:实现可参数化的扰动函数

首先定义泊松噪声注入函数(torchvision无内置泊松噪声变换,需自行实现),高斯模糊可直接调用torchvision内置的GaussianBlur类:

import torch
from torch.utils.data import Dataset
from torchvision import transforms
from PIL import Image
import numpy as np

def add_poisson_noise(img, lam):
    # 输入img为归一化到[0,1]的张量,lam越大噪声强度越低
    noised_img = torch.poisson(img * lam) / lam
    return torch.clamp(noised_img, 0.0, 1.0)

步骤2:改写Dataset的__getitem__逻辑

在__getitem__方法中逐图采样扰动参数、应用变换,最后将图像、标签、实际使用的扰动参数一并返回:

class AugmentedImageDataset(Dataset):
    def __init__(self, img_paths, labels, 
                 blur_kernel_range=(3,9), 
                 sigma_range=(0.1, 5.0), 
                 poisson_lam_range=(10, 200)):
        self.img_paths = img_paths
        self.labels = labels
        # 存储扰动参数的采样范围
        self.blur_kernel_range = blur_kernel_range
        self.sigma_range = sigma_range
        self.poisson_lam_range = poisson_lam_range
        # 基础预处理:转张量、归一化到[0,1]
        self.base_transform = transforms.Compose([
            transforms.ToTensor(),
        ])
    
    def __len__(self):
        return len(self.img_paths)
    
    def __getitem__(self, idx):
        # 加载原始图像和标签
        img = Image.open(self.img_paths[idx]).convert("RGB")
        label = self.labels[idx]
        img = self.base_transform(img)
        
        # 采样高斯模糊参数并应用变换
        kernel_size = np.random.choice(range(self.blur_kernel_range[0], self.blur_kernel_range[1]+1, 2))
        sigma = np.random.uniform(*self.sigma_range)
        blur_aug = transforms.GaussianBlur(kernel_size=kernel_size, sigma=sigma)
        img = blur_aug(img)
        
        # 采样泊松噪声参数并应用变换
        poisson_lam = np.random.uniform(*self.poisson_lam_range)
        img = add_poisson_noise(img, poisson_lam)
        
        # 打包实际使用的扰动参数
        aug_params = {
            "gaussian_blur_kernel": kernel_size,
            "gaussian_blur_sigma": sigma,
            "poisson_lambda": poisson_lam
        }
        return img, label, aug_params

步骤3:正常调用Dataloader即可

无需向Dataloader传递任何transform参数,加载出的每一个batch会自动包含图像、标签、对应每张图的扰动参数:

from torch.utils.data import DataLoader

# 初始化数据集和加载器
dataset = AugmentedImageDataset(img_paths=your_img_path_list, labels=your_label_list)
dataloader = DataLoader(dataset, batch_size=32, shuffle=True)

# 遍历加载器
for imgs, labels, params in dataloader:
    # 后续训练/测试逻辑,params直接可用于结果统计
    pass

补充说明

  • 若需要实现逐图参数定制,只需要在__getitem__中根据图像idx、标签等信息调整参数采样逻辑即可,无需修改其他代码
  • 若不需要每次都同时应用两种扰动,可以新增概率参数控制扰动的触发概率,触发状态也可以一并存入aug_params用于后续统计

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.02 06:06:04