PyTorch中将数据加载函数改写为Dataset子类的优势解析
你现在用函数式写法实现数据加载确实能跑起来,但换成继承Dataset的类形式,好处真的挺多的,我给你一一拆解:
复用性拉满,模块化更清晰
函数式写法里,数据集路径、变换逻辑这些都硬编码在函数里了,要是后面要加载另一个数据集、换个变换组合,要么改原函数,要么重新写个新函数。而类的形式可以把通用逻辑封装起来,初始化时接收data_path、transform这类参数,后面加载不同数据集直接实例化就行,不用重复造轮子。比如你可以定义一个通用的自定义数据集类,后续不管是加载BSD还是其他数据集,只需要传不同的路径和变换就搞定。扩展性更强,自定义逻辑更方便
要是后面需要加一些特殊逻辑——比如加载时随机切换增强策略、读取额外的标签文件、处理非标准的数据集结构(不是ImageFolder那种按文件夹分类的)——类的形式就太香了。你只需要重写__getitem__和__len__这两个核心方法就行。比如要加载带额外标注的图片,在__init__里读取标注文件,__getitem__里返回图片+对应标注,比在函数里乱改逻辑清晰多了。贴合PyTorch生态,兼容性更好
PyTorch官方工具和大部分第三方库都是围绕Dataset类设计的,比如DataLoader天生就对Dataset实例有完美支持,还有ConcatDataset、Subset这类数据集组合工具,直接传你的Dataset实例就能用。函数式返回的DataLoader就没这么灵活,没法直接和这些工具对接。可读性和可维护性更高
类的形式把数据加载的各个环节拆成了不同方法:__init__负责初始化资源,__len__返回数据集长度,__getitem__负责获取单个样本。别人看代码的时候,一眼就能理清逻辑;后面要改某个环节,直接对应修改方法就行,不会牵一发而动全身,比一堆逻辑堆在函数里好维护太多。
给你把原来的函数改成类形式的示例参考:
from torch.utils.data import Dataset from torchvision import transforms, datasets class CustomBSDDataset(Dataset): def __init__(self, data_path, transform=None): self.data_path = data_path # 提供默认变换,也支持传入自定义变换 self.transform = transform or transforms.Compose([ transforms.Grayscale(num_output_channels=1), transforms.ToTensor() ]) self.image_dataset = datasets.ImageFolder(root=data_path, transform=self.transform) def __len__(self): # 返回数据集总长度 return len(self.image_dataset) def __getitem__(self, idx): # 返回指定索引的样本和标签 img, label = self.image_dataset[idx] return img, label # 使用方式也很简单 def load_dataset(size_batch, size): data_path = "/home/bledc/dataset/test_set/crops_BSD" dataset = CustomBSDDataset(data_path) train_loader = torch.utils.data.DataLoader( dataset, batch_size=size_batch, shuffle=True, num_workers=0, drop_last=True ) return train_loader
总的来说,类的形式更符合PyTorch的设计哲学,面对简单场景时函数写法够用,但稍微复杂一点的需求,类的灵活性、可维护性优势就会完全体现出来。
内容的提问来源于stack exchange,提问作者Bled Clement

