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

优化PyTorch版Mix-up数据增强实现的性能求助

问题描述

我实现了mix-up图像增强的代码,但运行速度极慢。像按0.5权重缩放图像再逐像素求和这类操作似乎难以避免,天生就慢。这个方案要用到强化学习场景里,得处理6400万张图像,所以必须大幅提速。

注:以下是原作者的实现代码,和我的代码逻辑完全一致,速度应该同样慢。

import torch
import utils
import os
import torch.nn.functional as F
import torchvision.transforms as TF
import torchvision.datasets as datasets

dataloader = None
data_iter = None

def _load_data(
    sub_path: str, batch_size: int = 256, image_size: int = 84, num_workers: int = 16
):
    global data_iter, dataloader
    for data_dir in utils.load_config("datasets"):
        if os.path.exists(data_dir):
            fp = os.path.join(data_dir, sub_path)
            if not os.path.exists(fp):
                print(f"Warning: path {fp} does not exist, falling back to {data_dir}")
            dataloader = torch.utils.data.DataLoader(
                datasets.ImageFolder(
                    fp,
                    TF.Compose(
                        [
                            TF.RandomResizedCrop(image_size),
                            TF.RandomHorizontalFlip(),
                            TF.ToTensor(),
                        ]
                    ),
                ),
                batch_size=batch_size,
                shuffle=True,
                num_workers=num_workers,
                pin_memory=True,
            )
            data_iter = iter(dataloader)
            break
    if data_iter is None:
        raise FileNotFoundError(
            "failed to find image data at any of the specified paths"
        )
    print("Loaded dataset from", data_dir)


def _load_places(batch_size=256, image_size=84, num_workers=16, use_val=False):
    partition = "val" if use_val else "train"
    sub_path = os.path.join("places365_standard", partition)
    print(f"Loading {partition} partition of places365_standard...")
    _load_data(
        sub_path=sub_path,
        batch_size=batch_size,
        image_size=image_size,
        num_workers=num_workers,
    )


def _load_coco(batch_size=256, image_size=84, num_workers=16, use_val=False):
    sub_path = "COCO"
    print(f"Loading COCO 2017 Val...")
    _load_data(
        sub_path=sub_path,
        batch_size=batch_size,
        image_size=image_size,
        num_workers=num_workers,
    )

def _get_data_batch(batch_size):
    global data_iter
    try:
        imgs, _ = next(data_iter)
        if imgs.size(0) < batch_size:
            data_iter = iter(dataloader)
            imgs, _ = next(data_iter)
    except StopIteration:
        data_iter = iter(dataloader)
        imgs, _ = next(data_iter)
    return imgs.cuda()

def load_dataloader(batch_size, image_size, dataset="coco"):
    if dataset == "places365_standard":
        if dataloader is None:
            _load_places(batch_size=batch_size, image_size=image_size)
    elif dataset == "coco":
        if dataloader is None:
            _load_coco(batch_size=batch_size, image_size=image_size)
    else:
        raise NotImplementedError(
            f'overlay has not been implemented for dataset "{dataset}"'
        )

def random_mixup(x, dataset="coco"):
    """Randomly overlay an image from Places or COCO"""
    global data_iter
    alpha = 0.5

    load_dataloader(batch_size=x.size(0), image_size=x.size(-1), dataset=dataset)

    imgs = _get_data_batch(batch_size=x.size(0)).repeat(1, x.size(1) // 3, 1, 1)

    return ((1 - alpha) * (x / 255.0) + (alpha) * imgs) * 255.0

优化方案

1. 砍掉冗余的数值转换

当前代码里的((1 - alpha) * (x / 255.0) + (alpha) * imgs) * 255.0做了两次无意义的浮点转换(除以255再乘回去),直接简化计算逻辑:

  • 如果输入x是0-255的整数张量,先把imgs(ToTensor输出的0-1浮点)转成0-255再混合:(1-alpha)*x + alpha*(imgs*255.0)
  • 或者全程用0-1浮点计算,后续模型需要0-255时再统一转换,避免来回折腾。

修改后的混合代码:

# 全程用0-1浮点计算,省去来回转换
x_norm = x.float() / 255.0
mixed = (1 - alpha) * x_norm + alpha * imgs
# 若后续需要0-255,仅需一次转换:mixed * 255.0

2. 用内存高效的操作替代repeat

代码里的imgs.repeat(1, x.size(1) // 3, 1, 1)会复制内存,换成expand(仅扩展张量视图,不复制数据)能省内存和时间:

imgs = _get_data_batch(batch_size=x.size(0)).expand(-1, x.size(1), -1, -1)

3. 数据加载是提速核心(6400万张图的IO开销远大于计算)

  • 拉满num_workers:根据CPU核心数设置,比如32或64(别超过CPU逻辑核心数的80%)
  • 开启persistent_workers:DataLoader初始化时加persistent_workers=True,避免每个epoch重建worker进程
  • 换高效数据格式:用LMDB、WebDataset替代ImageFolder,大幅降低磁盘IO延迟;内存够的话直接把数据集缓存到内存
  • 调大batch_size:用更大的batch(比如1024),榨干GPU的批量计算能力
  • GPU端做数据增强:把RandomResizedCrop、RandomHorizontalFlip换成Kornia的GPU版操作,减少CPU-GPU数据传输

修改后的DataLoader示例:

dataloader = torch.utils.data.DataLoader(
    datasets.ImageFolder(
        fp,
        TF.Compose(
            [
                TF.RandomResizedCrop(image_size),
                TF.RandomHorizontalFlip(),
                TF.ToTensor(),
            ]
        ),
    ),
    batch_size=1024,  # 调大batch
    shuffle=True,
    num_workers=32,  # 拉满worker数
    pin_memory=True,
    persistent_workers=True,  # 保留worker进程
    prefetch_factor=2,  # 预取2个batch
)

4. 干掉全局变量与重复初始化

当前用全局的dataloader和data_iter,不仅容易出问题,每次调用random_mixup都检查加载也浪费时间:

  • 提前初始化好所有需要的dataloader,存在字典里直接调用
  • 预取多个batch到内存,避免每次_get_data_batch的异常处理开销

5. 用JIT编译融合计算操作

PyTorch的JIT编译能自动融合多个操作,减少GPU kernel调用次数,直接给混合逻辑加装饰器:

@torch.jit.script
def fast_mix(x: torch.Tensor, imgs: torch.Tensor, alpha: float = 0.5):
    x_norm = x.float() / 255.0
    return (1 - alpha) * x_norm + alpha * imgs

6. 硬件级优化(可选)

  • 部署阶段用TensorRT编译混合逻辑,进一步提升速度
  • 多GPU场景用DistributedDataLoader分散数据加载压力

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.18 21:20:43