PyTorch中CNN图像数据增强:如何保留原数据与变换后数据
解决PyTorch(Monai)数据增强同时保留原图与变换图的问题
你当前用的Monai随机变换(RandRotate、RandFlip这类)是概率性触发的,每次只会输出原图或变换后的图,而非同时输出两者,所以数据集并没有真正扩充。要实现同时用原图和变换图训练,得主动让每个样本生成两个版本,下面给两种可行方案:
方案一:自定义双输出变换+数据集
写一个包装类,让变换同时返回原图和增强后的图,再通过自定义数据集把这两个版本拆成独立样本:
from monai.transforms import Compose, LoadImage, EnsureChannelFirst, ScaleIntensity, RandRotate, RandFlip from monai.data import Dataset, DataLoader import numpy as np # 包装变换,返回原图和增强图 class DualOutputTransform: def __init__(self, transform): self.transform = transform def __call__(self, data): original = data.copy() augmented = self.transform(data) return [original, augmented] # 定义必触发的增强变换(prob设为1.0确保每次都变换) base_aug_transforms = Compose([ LoadImage(image_only=True), EnsureChannelFirst(), ScaleIntensity(), RandRotate(range_x=np.pi / 12, prob=1.0, keep_size=True), RandFlip(spatial_axis=0, prob=1.0) ]) train_transforms = DualOutputTransform(base_aug_transforms) # 自定义数据集,展开每个样本的两个版本 class DualDataset(Dataset): def __init__(self, data, transform): super().__init__(data, transform) def __getitem__(self, idx): return self.transform(self.data[idx]) # 假设你的数据列表是data_list,每个元素是{"image": "图片路径"} data_list = [{"image": "img1.png"}, {"image": "img2.png"}, ...] train_dataset = DualDataset(data_list, train_transforms) # 自定义collate_fn,把批量里的样本对展开成单个样本 def collate_fn(batch): flattened = [] for pair in batch: flattened.extend(pair) from torch.utils.data._utils.collate import default_collate return default_collate(flattened) train_loader = DataLoader(train_dataset, batch_size=4, collate_fn=collate_fn)
方案二:复制数据集+合并
把原始数据列表复制一份,一份用仅做固定预处理的变换(保留原图),另一份用增强变换,最后合并两个数据集:
from monai.transforms import Compose, LoadImage, EnsureChannelFirst, ScaleIntensity, RandRotate, RandFlip from monai.data import Dataset from torch.utils.data import ConcatDataset, DataLoader # 原图预处理(无随机增强) original_transforms = Compose([ LoadImage(image_only=True), EnsureChannelFirst(), ScaleIntensity() ]) # 增强预处理(必触发变换) augmented_transforms = Compose([ LoadImage(image_only=True), EnsureChannelFirst(), ScaleIntensity(), RandRotate(range_x=np.pi / 12, prob=1.0, keep_size=True), RandFlip(spatial_axis=0, prob=1.0) ]) # 复制数据列表 data_list = [{"image": "img1.png"}, {"image": "img2.png"}, ...] original_data = data_list.copy() augmented_data = data_list.copy() # 创建两个数据集并合并 original_dataset = Dataset(original_data, original_transforms) augmented_dataset = Dataset(augmented_data, augmented_transforms) train_dataset = ConcatDataset([original_dataset, augmented_dataset]) train_loader = DataLoader(train_dataset, batch_size=4)
注意事项
- 如果需要多种增强方式(比如旋转、翻转、裁剪各一种),可以复制多份数据列表,分别应用不同的增强变换后再合并,进一步扩充数据集。
- 若保留
Rand系列变换的prob<1.0,会导致部分增强样本和原图重复,如需严格扩充数据,建议把增强变换的prob设为1.0。
内容的提问来源于stack exchange,提问作者Estel
相关产品推荐
相关产品推荐

