如何逐图控制或记录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
相关产品推荐
相关产品推荐

