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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.25 09:52:45