在深度学习任务中,数据加载和处理是至关重要的一环。

自定义数据集

        PyTorch 提供了强大的数据加载和处理工具,主要包括:

  • torch.utils.data.Dataset:数据集的抽象类,需要自定义并实现 __len__(数据集大小)和 __getitem__(按索引获取样本)

  • torch.utils.data.TensorDataset:基于张量的数据集,适合处理数据-标签对,直接支持批处理和迭代。

  • torch.utils.data.DataLoader:封装 Dataset 的迭代器,提供批处理、数据打乱、多线程加载等功能,便于数据输入模型训练。

  • torchvision.datasets.ImageFolder:从文件夹加载图像数据,每个子文件夹代表一个类别,适用于图像分类任务。

代码示例

import torchvision.transforms
from torch.utils.data import Dataset,DataLoader
import os
from PIL import Image
from torch.utils.tensorboard import SummaryWriter

class MyData(Dataset):
    def __init__(self,rootpath,labelpath):
        self.root_path = rootpath   # 图片数据根路径
        self.label_path = labelpath # 定义数据标签,图片数据已按标签分类
        self.data_path = os.path.join(self.root_path,self.label_path)  # 路径拼接,适用于不同系统
        self.img_path = os.listdir(self.data_path)  # 罗列出路径下的所有文件

    def __getitem__(self, idx):
        img_name = self.img_path[idx]
        img_item_path = os.path.join(self.root_path,self.label_path,img_name)
        img = Image.open(img_item_path)
        img_resize = img.resize((100,100))   # 要确保图片大小一致,以便后面图片拼接
        img_tensor = torchvision.transforms.ToTensor()(img_resize)  # 图片数据转换为tensor格式
        label = self.label_path
        return img_tensor,label  # 返回图片数据及标签

    def __len__(self):
        return len(self.img_path)  # 返回数据样本数量

rootpath = r'data\train'
ants_labelpath = 'ants'
bees_labelpath = 'bees'

ants_dataset = MyData(rootpath,ants_labelpath)
bees_dataset = MyData(rootpath,bees_labelpath)

train_dataset = ants_dataset + bees_dataset  # 拼接数据集

# 查看图片尺寸
print(ants_dataset[1][0].shape)
print(ants_dataset[2][0].shape)

# 数据加载,batch_size表示每次加载的图片数量,shuffle表示每次加载是否洗牌,drop_last表示剩余数据量不足batch_size时是否丢弃数据
ants_imgloader = DataLoader(ants_dataset,batch_size=64,shuffle=True,drop_last=False)

writer = SummaryWriter('logs_antsloader')
step = 0
for data in ants_imgloader:
    imgs,label = data   # 提取图片数据及标签
    writer.add_image('ants_img',imgs,step,dataformats='NCHW')   # dataformats需要根据实际情况设置
    step += 1
writer.close()

        运行程序后生成tensorboard日志文件,在终端输入tensorboard --logdir=logs_antsloader指令打开dataloader后的图片数据,如下图所示,每组图片64张(batch_size),最后一组不足64张,数据保留(drop_last为false)。

    

PyTorch 内置数据集

PyTorch 通过 torchvision.datasets 模块提供了许多常用的数据集,例如:

  • MNIST:手写数字图像数据集,用于图像分类任务。

  • CIFAR:包含 10 个类别、60000 张 32x32 的彩色图像数据集,用于图像分类任务。

  • COCO:通用物体检测、分割、关键点检测数据集,包含超过 330k 个图像和 2.5M 个目标实例的大规模数据集。

  • ImageNet:包含超过 1400 万张图像,用于图像分类和物体检测等任务。

  • STL-10:包含 100k 张 96x96 的彩色图像数据集,用于图像分类任务。

  • Cityscapes:包含 5000 张精细注释的城市街道场景图像,用于语义分割任务。

  • SQUAD:用于机器阅读理解任务的数据集。

以上数据集可以通过 torchvision.datasets 模块中的函数进行加载,也可以通过自定义的方式加载其他数据集。

代码示例

import torch
import torchvision
from torch import nn
from torch.utils.data import DataLoader
from torch.utils.tensorboard import SummaryWriter

# 数据transform处理
data_transform = torchvision.transforms.Compose([
                                                 torchvision.transforms.ToTensor()
                                                ])
# 下载pytorch官方CIFAR数据集
train_data = torchvision.datasets.CIFAR10('./data_CIFAR',train=True,transform=data_transform,download=True)
test_data = torchvision.datasets.CIFAR10('./data_CIFAR',train=False,transform=data_transform,download=True)

# train_data = torchvision.datasets.FashionMNIST('./data_fashionMNIST',train=True,transform=data_transform,download=True)
# test_data = torchvision.datasets.FashionMNIST('./data_fashionMNIST',train=False,transform=data_transform,download=True)

# 数据加载
test_loader = DataLoader(test_data,batch_size=64,drop_last=False)

writer = SummaryWriter('logs_loaderCIFAR')
step = 0
for data in test_loader:
    img,target = data
    writer.add_image('test_CIFAR10',img,step,dataformats='NCHW')
    step += 1
writer.close()

运行程序后生成tensorboard日志文件,在终端输入tensorboard --logdir=logs_loaderCIFAR指令打开dataloader后的图片数据,每组图片64张(batch_size),以下为其中两组图片。

Logo

魔乐社区(Modelers.cn) 是一个中立、公益的人工智能社区,提供人工智能工具、模型、数据的托管、展示与应用协同服务,为人工智能开发及爱好者搭建开放的学习交流平台。社区通过理事会方式运作,由全产业链共同建设、共同运营、共同享有,推动国产AI生态繁荣发展。

更多推荐