PyTorch图像增强是否增加本地数据集 如何保存增强数据均衡类别
增强图像未被统计的原因及去向
你当前使用的是PyTorch默认的在线实时增强逻辑,这种模式不会提前生成新的图像文件存入原始数据集:
- 你定义的
transforms变换只会在DataLoader迭代读取样本的瞬间,对当前加载到内存的原始图像当场执行随机变换,变换得到的张量会直接传入模型完成当前步的训练,用完就会被内存回收,不会写入本地磁盘,也不会往原始数据集的样本索引列表里新增条目。 - 训练时统计的图像数量本质是原始数据集的索引总数,每个epoch遍历的索引范围始终和原始数据集完全一致,只是同一张原始图像在不同epoch、不同step被读取时,会生成完全不同的随机增强版本,相当于每轮训练你用到的图像内容都有差异,但总计数永远等于原始样本数,自然看不到所谓“新增的增强图像”。
举个直白的例子:某类只有800张原始图,每个epoch训练时这800张图都会被随机旋转、裁剪、翻转生成800份新的变体用于训练,但这些变体不会被保存,下一个epoch又会基于原始图生成另一批全新的变体。
本地保存增强图像、实现类别均衡的实现方法
要把增强结果存为实体文件、把所有类别样本量对齐到4000张,需要用离线增强的方案,单独跑脚本生成图像存盘,不要直接在训练流程里靠实时transform实现,具体步骤如下:
- 前期准备
先统计每个类别的原始样本数,以样本最多的类别(即你提到的4000张的类)作为目标数量,计算每个类别需要补充生成的增强图数量,例如800张的类需要额外生成3200张。注意:不要直接在原始数据集目录下写入增强图,单独新建一个目录存放均衡后的数据集,避免损坏原始数据。
- 代码实现参考
先导入依赖库:
定义和训练时一致的增强策略:import os import random import shutil from PIL import Image from torchvision import transforms from tqdm import tqdm
配置路径、执行增强生成:aug_transform = transforms.Compose([ transforms.RandomRotation(30), transforms.RandomResizedCrop(140), transforms.RandomHorizontalFlip() ])# 配置原始数据集、均衡后数据集的根目录,按自己的实际路径修改 RAW_DATA_ROOT = "./raw_dataset" BALANCED_DATA_ROOT = "./balanced_dataset" os.makedirs(BALANCED_DATA_ROOT, exist_ok=True) # 读取所有类别文件夹 class_list = [d for d in os.listdir(RAW_DATA_ROOT) if os.path.isdir(os.path.join(RAW_DATA_ROOT, d))] # 统计每个类的原始样本量 class_img_map = {} for cls in class_list: cls_path = os.path.join(RAW_DATA_ROOT, cls) img_list = [f for f in os.listdir(cls_path) if f.lower().endswith(('.png', '.jpg', '.jpeg', '.bmp'))] class_img_map[cls] = img_list max_sample_num = max([len(v) for v in class_img_map.values()]) # 逐类处理 for cls in class_list: target_cls_dir = os.path.join(BALANCED_DATA_ROOT, cls) os.makedirs(target_cls_dir, exist_ok=True) raw_cls_dir = os.path.join(RAW_DATA_ROOT, cls) raw_imgs = class_img_map[cls] # 先把该类所有原始图复制到均衡数据集目录 for img_name in raw_imgs: src_path = os.path.join(raw_cls_dir, img_name) dst_path = os.path.join(target_cls_dir, img_name) if not os.path.exists(dst_path): shutil.copy2(src_path, dst_path) # 计算需要补充的增强图数量 aug_need = max_sample_num - len(raw_imgs) if aug_need <= 0: continue # 循环生成增强图并保存 for aug_idx in tqdm(range(aug_need), desc=f"处理类别 {cls}"): # 随机选一张原始图做增强 picked_img_name = random.choice(raw_imgs) picked_img = Image.open(os.path.join(raw_cls_dir, picked_img_name)).convert('RGB') aug_res = aug_transform(picked_img) # 命名加aug前缀避免和原始图重名 save_name = f"aug_{aug_idx}_{picked_img_name}" aug_res.save(os.path.join(target_cls_dir, save_name)) - 补充说明
脚本运行完成后,BALANCED_DATA_ROOT目录下每个类别的样本数都会对齐到4000张,后续训练直接加载这个目录下的数据集即可,此时统计到的样本数就是均衡后的总数。
如果不想占用额外磁盘存增强图,也可以直接在训练时使用WeightedRandomSampler设置采样权重,提高少样本类别的被采样概率,配合在线增强也能解决类别不均衡问题,不需要额外生成实体文件。
内容的提问来源于stack exchange,提问作者Manjunath D
相关产品推荐
相关产品推荐

