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

基于PyTorch优化ImageNet单类别批次采样效率

问题描述

需要在ImageNet数据集(1000类)上训练分类器,要求:

  • 每个batch包含64张同一类别的图像
  • 连续batch来自不同类别

基于已有思路实现了代码,但DS类中通过for循环遍历数据集构建类别索引列表的过程耗时过长,希望找到更高效的采样器构建方式。

优化方案

1. 直接利用ImageFolder的targets属性构建类别索引

ImageFolder数据集本身已存储所有样本的标签(train_dataset.targets),无需遍历整个数据集获取标签,直接对标签数组分组即可大幅缩短索引构建时间:

import torch
from torch.utils.data import Dataset

class DS(Dataset):
    def __init__(self, data, num_classes):
        super(DS, self).__init__()
        self.data = data
        self.data_len = len(data)
        
        # 直接通过targets数组快速构建类别索引
        targets = torch.tensor(data.targets)
        self.indices = []
        for cls in range(num_classes):
            # 获取当前类别所有样本的索引
            cls_indices = torch.where(targets == cls)[0].tolist()
            self.indices.append(cls_indices)

    def per_class_sample_indices(self):
        return self.indices

    def __getitem__(self, index):
        return self.data[index]

    def __len__(self):
        return self.data_len

2. 缓存类别索引文件

保留缓存逻辑,首次构建索引后保存为文件,后续训练直接加载,避免重复执行耗时的索引构建:

def main():
    file_path = "./cache"
    os.makedirs(file_path, exist_ok=True)
    file_name = 'per_class_sample_indices.pt'
    cache_path = os.path.join(file_path, file_name)
    
    if not os.path.exists(cache_path):
        print(f'缓存文件 {file_name} 不存在,开始构建索引...')
        ds = DS(train_dataset, num_classes)
        per_class_sample_indices = ds.per_class_sample_indices()
        torch.save(per_class_sample_indices, cache_path)
        print(f'索引已保存至 {cache_path}')
    else:
        per_class_sample_indices = torch.load(cache_path)
        print(f'已加载缓存的类别索引')
    
    # 初始化BatchSampler和DataLoader
    batch_sampler = BatchSampler(per_class_sample_indices, batch_size=args.batch_size)
    train_loader = torch.utils.data.DataLoader(
        train_dataset,
        num_workers=args.workers,
        pin_memory=True,
        batch_sampler=batch_sampler
    )
    
    # 验证采样逻辑(可选)
    labels = []
    for _, (_, _labels) in enumerate(train_loader):
        labels.append(torch.unique(_labels).item())
    print(f'遍历到的唯一类别数: {len(set(labels))}')

3. BatchSampler效率优化

使用yield逐个生成batch,避免一次性生成所有batch索引占用过多内存,同时优化类别选择逻辑,确保连续batch类别不同:

import random

class BatchSampler:
    def __init__(self, per_class_sample_indices, batch_size):
        self.per_class_sample_indices = per_class_sample_indices
        self.batch_size = batch_size
        self.class_list = list(range(len(per_class_sample_indices)))
        random.shuffle(self.class_list)
        
        # 预计算总样本数和batch数,避免重复计算
        self.total_samples = sum(len(indices) for indices in per_class_sample_indices)
        self.n_batches = self.total_samples // batch_size

    def __iter__(self):
        class_idx = 0
        for _ in range(self.n_batches):
            # 循环使用打乱后的类别列表,保证连续batch类别不同
            current_class = self.class_list[class_idx % len(self.class_list)]
            class_idx += 1
            
            cls_indices = self.per_class_sample_indices[current_class]
            # 随机采样batch_size个样本,样本不足时取全部
            batch = random.sample(cls_indices, self.batch_size) if len(cls_indices) >= self.batch_size else cls_indices.copy()
            yield batch

    def __len__(self):
        return self.n_batches

优化效果

  • 索引构建时间:从遍历百万级样本的数分钟,缩短至处理标签数组的数秒
  • 内存占用:yield方式生成batch,避免一次性存储所有batch索引,降低内存开销
  • 复用性:缓存机制确保仅首次训练需要构建索引,后续直接加载缓存文件

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.14 23:01:14