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

PyTorch 如何快速获取 Subset 子集的全部数据与标签

性能瓶颈根因

你的代码性能差核心来自两个问题:

  1. 列表推导式是单进程串行遍历Subset对象,每次取元素都会触发「磁盘读图片+执行transform」的全流程,没有并行优化
  2. 完全冗余的遍历逻辑: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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.06 23:18:04