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

如何快速转换数据集格式以适配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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.06 11:57:02