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

PyTorch中数据增强效果验证及增强后数据集数量查看方法

验证数据增强操作是否生效

方法1:可视化对比原图与增强结果

直接通过可视化能最直观验证所有增强操作是否生效,包括几何变换(翻转、旋转、透视等)和颜色变换的效果,同时还能确认输入图与掩码的同步变换是否正确(这对Unet任务至关重要)。

示例代码:

from PIL import Image
import torchvision.transforms as transforms

# 初始化你的数据集
dataset = ProcessTrainDataset(x_paths, y_paths)  # x_paths、y_paths是原始数据的路径列表

# 获取第一张图的原图
raw_x = Image.open(dataset.x[0])
raw_y = Image.open(dataset.y[0]).convert("L")

# 获取增强后的图(因为有随机操作,多次调用dataset[0]会得到不同结果)
aug_x, aug_y = dataset[0]

# 将张量转回PIL格式用于显示/保存
to_pil = transforms.ToPILImage()
aug_x_pil = to_pil(aug_x)
aug_y_pil = to_pil(aug_y)

# 对比显示(也可以保存到本地查看)
raw_x.show(title="原始输入图")
aug_x_pil.show(title="增强后输入图")
raw_y.show(title="原始掩码")
aug_y_pil.show(title="增强后掩码")

多次运行获取增强结果,能看到随机翻转、旋转、透视变形、亮度对比度变化等效果,同时确认掩码和输入图的变换完全同步。

方法2:检查张量数值变化

通过对比原图与增强后张量的统计特征或特定像素值,验证变换是否生效:

# 把原图转成张量
raw_x_tensor = transforms.ToTensor()(raw_x)

# 对比均值、标准差(颜色变换会改变这些值)
print("原图均值:", raw_x_tensor.mean().item())
print("增强后均值:", aug_x.mean().item())
print("原图标准差:", raw_x_tensor.std().item())
print("增强后标准差:", aug_x.std().item())

# 对比特定位置像素(验证几何变换)
print("原图左上角像素值:", raw_x_tensor[:, 0, 0].tolist())
print("增强后对应位置像素值:", aug_x[:, 0, 0].tolist())
查看增强后的数据集数量

首先要注意你当前代码的小问题:self.x_augmented和self.y_augmented没有在__init__中初始化,运行时会触发AttributeError,需要先在__init__里添加self.x_augmented = []和self.y_augmented = []。

情况1:动态生成增强样本(常用方式)

如果你的代码是每次调用__getitem__时实时生成增强结果(不存储所有增强样本),那么数据集的实际有效长度就是原始数据集的长度,即len(dataset)返回的len(self.x)。这种方式下,每个epoch都会生成全新的增强样本,不需要额外存储,节省内存。

情况2:预生成并存储所有增强样本

如果需要预先生成固定数量的增强样本并存储,可修改代码实现批量扩充,此时直接通过len(dataset)查看增强后的总数量:

class ProcessTrainDataset(Dataset):
    def __init__(self, x, y, aug_times=2):
        self.x = x
        self.y = y
        self.aug_times = aug_times  # 每张原始图生成aug_times个增强样本
        self.x_augmented = []
        self.y_augmented = []

        self.pre_process = transforms.Compose([transforms.ToTensor()])
        self.transform_data = transforms.Compose([transforms.ColorJitter(brightness=0.2, contrast=0.2)])
        self.transform_all = transforms.Compose([
            transforms.RandomVerticalFlip(),
            transforms.RandomHorizontalFlip(),
            transforms.RandomRotation(10),
            transforms.RandomPerspective(distortion_scale=0.2, p=0.5),
            transforms.RandomAffine(degrees=0, translate=(0.2,0.2), scale=(0.9,1.1)),
        ])

        # 预先生成所有增强样本
        for idx in range(len(self.x)):
            for _ in range(self.aug_times):
                img_x = Image.open(self.x[idx])
                img_y = Image.open(self.y[idx]).convert("L")
                
                img_x = self.pre_process(img_x)
                img_y = self.pre_process(img_y)
                img_all = torch.cat([img_x, img_y])
                img_all = self.transform_all(img_all)
                img_x, img_y = img_all[:-1, ...], img_all[-1:,...]
                img_x = self.transform_data(img_x)
                
                self.x_augmented.append(img_x)
                self.y_augmented.append(img_y)

    def __len__(self):
        return len(self.x_augmented)

    def __getitem__(self, idx):
        return self.x_augmented[idx], self.y_augmented[idx]

此时len(dataset)会返回原始数据集数量 × aug_times,即增强后的总样本数。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.24 19:05:33