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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.30 02:57:18