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

