如何快速转换数据集格式以适配resnet18模型的训练输入要求
解决方案
最优方案:使用Dataset包装类(零前置转换开销)
不需要提前遍历全量数据集做转换,仅需要给原始数据集加一层轻量包装,在取样本时动态转换格式即可,前置准备时间小于1秒:
import torch class DatasetWrapper(torch.utils.data.Dataset): def __init__(self, original_dataset): self.original_dataset = original_dataset def __len__(self): return len(self.original_dataset) def __getitem__(self, idx): sample = self.original_dataset[idx] # 直接返回模型要求的(图像张量, 标签张量)元组格式 return sample['image'], sample['label']
使用方式:
# 包装原始数据集 wrapped_train_dataset = DatasetWrapper(trans_train_dataset) # 直接传入DataLoader用于训练即可,无需其他处理 train_loader = torch.utils.data.DataLoader(wrapped_train_dataset, batch_size=32, shuffle=True)
持久化存储方案
如果需要长期保存处理好的数据集,使用PyTorch原生的序列化工具,不要用csv/txt等文本格式:
适用场景:数据集可全部放入内存
# 仅当数据集能完整放入内存时使用 all_images = [] all_labels = [] for sample in trans_train_dataset: all_images.append(sample['image']) all_labels.append(sample['label']) # 拼接为批量张量,生成标准TensorDataset all_images = torch.stack(all_images) all_labels = torch.stack(all_labels) train_dataset = torch.utils.data.TensorDataset(all_images, all_labels) # 保存数据集,序列化原生张量格式,存储效率极高 torch.save(train_dataset, "train_dataset.pt") # 读取时直接加载即可,得到的就是符合输入要求的数据集 loaded_train_dataset = torch.load("train_dataset.pt")
适用场景:数据集过大无法放入内存
使用HDF5格式存储,借助h5py库实现磁盘级的高效随机读写,不需要把全量数据加载到内存。
注意事项
你贴出的ToTensor类代码存在缩进错误,__call__方法需要缩进为类的成员方法,否则运行会报错。
内容的提问来源于stack exchange,提问作者Tomasz Kaczmarski
相关产品推荐
相关产品推荐

