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
相关产品推荐
相关产品推荐

