PyTorch 如何快速获取 Subset 子集的全部数据与标签
性能瓶颈根因
你的代码性能差核心来自两个问题:
- 列表推导式是单进程串行遍历
Subset对象,每次取元素都会触发「磁盘读图片+执行transform」的全流程,没有并行优化 - 完全冗余的遍历逻辑:
ImageFolder初始化阶段已经扫描全目录生成了samples(图像路径+标签对)和targets(所有标签列表)两个属性,无需遍历数据集就能拿到所有标签信息,无需走Subset的迭代逻辑
最优实现方案
1. 毫秒级获取标签(无需加载图像)
直接从原数据集的targets属性按索引提取,完全没有IO开销:
import torch from sklearn.model_selection import train_test_split from torchvision import datasets as dset dataset = dset.ImageFolder("/path_to_folder", transform = transform) train_idx, test_idx = train_test_split(list(range(len(dataset))), test_size=0.2, stratify=dataset.targets) # 直接批量取标签,比迭代快几个数量级 train_labels = torch.tensor(dataset.targets)[train_idx] test_labels = torch.tensor(dataset.targets)[test_idx]
2. 高效加载所有图像到内存
如果需要把全量图像加载为张量/NumPy数组,用多进程批量加载:
from torch.utils.data import DataLoader, Subset train_set = Subset(dataset, train_idx) # num_workers设置为你的CPU物理核心数,batch_size尽量拉满,shuffle设为false避免额外开销 train_loader = DataLoader( train_set, batch_size=512, num_workers=8, pin_memory=True, shuffle=False ) # 批量拼接所有数据,比循环append效率更高 train_data_batches = [] train_label_batches = [] for data, label in train_loader: train_data_batches.append(data) train_label_batches.append(label) train_data = torch.cat(train_data_batches, dim=0) train_labels = torch.cat(train_label_batches, dim=0) # 如需转NumPy直接调用.numpy()即可 # train_data_np = train_data.numpy() # train_labels_np = train_labels.numpy()
长期优化建议
如果该数据集会反复使用,建议提前把所有图像预处理后序列化存储为HDF5或者LMDB格式,下次加载时直接读取单一大文件,避免大量小文件的随机IO开销,加载速度可以提升10倍以上。
内容的提问来源于stack exchange,提问作者Angelus
相关产品推荐
相关产品推荐

