本文发表于 657 天前,其中的信息可能已经事过境迁

AI的一切都是数据驱动的,数据集加载是深度学习的第一步,优质的AI离不开优质的数据集。

数据集是指训练模型的预筛选数据。

Pytorch提供了torch.utils.data.Datasettorch.utils.data.DataLoader两个类,用于加载数据集。

数据集加载

  • torch.utils.data.Dataset

torch.utils.data.Dataset是一个抽象类,用于表示数据集。用户需要继承Dataset类,并实现__len____getitem__两个方法。

python
import torch
from torch.utils.data import Dataset

class MyDataset(Dataset):
    def __init__(self):
        self.data = torch.randn(100, 3, 32, 32)
        self.label = torch.randint(0, 10, (100,))

    def __len__(self):
        return len(self.data)

    def __getitem__(self, index):
        return self.data[index], self.label[index]
  • torch.utils.data.DataLoader

torch.utils.data.DataLoader是一个数据加载器,用于加载数据集。DataLoader可以将数据集分批次加载,以便于训练模型。

python
import torch
from torch.utils.data import DataLoader

dataset = MyDataset()
dataloader = DataLoader(dataset, batch_size=10, shuffle=True)

for data, label in dataloader:
    print(data.shape, label.shape)

使用内置数据集API

Pytorch还提供了一些常用的数据集,如torchvision.datasets模块中的MNISTCIFAR10CIFAR100等数据集。

Torchvision提供的数据集可查看Torchvision内置数据集

GTSRB(德国交通标志识别基准 (GTSRB) )数据集为例,加载数据集的代码如下:

python
import torch
from torch.utils.data import Dataset
from torchvision import datasets
from torchvision.transforms import ToTensor
import matplotlib.pyplot as plt


training_data = datasets.GTSRB(
    root="data", split="train", download=True, transform=ToTensor()
)

test_data = datasets.GTSRB(
    root="data", split="test", download=True, transform=ToTensor()
)

在以上代码中,datasets.GTSRB函数用于加载GTSRB数据集,root参数指定数据集的存储路径,split参数指定数据集的类型,download参数指定是否下载数据集,transform参数指定数据集的转换方式,在使用内置数据集时,需参考官方API,部分内置数据集的部分参数可能不同。

使用自定义数据集API

在实际应用中,我们使用的数据集可能不是内置数据集,而是自己收集的数据集。此时,我们需要自定义数据集类。

自定义数据集类需要继承torch.utils.data.Dataset类,需要实现__init____len____getitem__三个方法。

  • __init__方法用于初始化数据集,通常在此方法中加载数据集。

  • __len__方法用于返回数据集的长度。

  • __getitem__方法用于返回数据集中的数据。

以GTSRB数据集为例,其数据集的目录结构如下:

shell
./data/gtsrb
    GTSRB/
        Training/
            00000/
                00000_00000.ppm
                00000_00001.ppm
                ...
                GT-00000.csv
            00001/
                00001_00000.ppm
                00001_00001.ppm
                ...
            ...
        Final_Test/
            Images/
                00000.ppm
                00001.ppm
                ...
                GT-final_test.csv

在以上目录结构中,Training目录下包含了训练数据,Final_Test目录下包含了测试数据,每个类别的数据存储在一个单独的目录中,每个类别的数据包含了一个GT-*.csv文件,该文件包含了数据的标签信息。

csv文件的内容示例如下:

csv
Filename;Width;Height;Roi.X1;Roi.Y1;Roi.X2;Roi.Y2;ClassId
00000_00000.ppm;29;30;5;6;24;25;0

在自定义数据集类中,我们不仅需要加载图片,还需要加载图片的标签信息。自定义数据集类的代码如下:

python
import os
import pandas as pd
from torch.utils.data import Dataset
from torchvision.transforms import ToTensor
from PIL import Image


class GTSRB(Dataset):
    def __init__(self, root, split="train", transform=None):
        self.root = root
        self.split = split
        self.transform = transform
        self.data = []
        self.label = []

        if self.split == "train":
            data_dir = os.path.join(self.root, "GTSRB", "Training")
        elif self.split == "test":
            data_dir = os.path.join(self.root, "GTSRB", "Final_Test", "Images")

        for class_dir in os.listdir(data_dir):
            if not os.path.isdir(os.path.join(data_dir, class_dir)):
                continue

            csv_file = os.path.join(data_dir, class_dir, f"GT-{class_dir}.csv")
            if not os.path.exists(csv_file):
                continue

            df = pd.read_csv(csv_file, delimiter=";")
            for index, row in df.iterrows():
                filename = os.path.join(data_dir, class_dir, row["Filename"])
                label = row["ClassId"]
                self.data.append(filename)
                self.label.append(label)

    def __len__(self):
        return len(self.data)

    def __getitem__(self, index):
        image = Image.open(self.data[index])
        label = self.label[index]

        if self.transform:
            image = self.transform(image)
        else:
            image = ToTensor()(image)

        return image, label


# 随机挑选一张图片进行展示, 以及其对应的标签
import matplotlib.pyplot as plt
import numpy as np
import torch

dataset = GTSRB(root="./data/gtsrb", split="train")
imagePPM, label = dataset[np.random.randint(0, len(dataset))]
image = imagePPM.permute(1, 2, 0).numpy()
plt.imshow(image)
plt.title(f"Label: {label}")
plt.show()

在以上代码中,我们定义了一个GTSRB类,该类继承自torch.utils.data.Dataset类,实现了__init____len____getitem__三个方法。在__init__方法中,我们加载了GTSRB数据集,并将数据集的路径与标签信息保存在self.dataself.label中。在__getitem__方法中,我们加载了图片,并返回图片与标签信息。

  • __init__

__init__ 函数在实例化 Dataset 对象时运行一次。在这个函数中,我们初始化数据集的路径、数据集的类型、数据集的转换方式,并加载数据集。

python
def __init__(self, root, split="train", transform=None):
    self.root = root
    self.split = split
    self.transform = transform
    self.data = []
    self.label = []

    # 判断数据集的类型并读取对应的数据集
    if self.split == "train":
        data_dir = os.path.join(self.root, "GTSRB", "Training")
    elif self.split == "test":
        data_dir = os.path.join(self.root, "GTSRB", "Final_Test", "Images")

    # 遍历数据集的每个类别
    for class_dir in os.listdir(data_dir):
        # 判断是否为目录
        if not os.path.isdir(os.path.join(data_dir, class_dir)):
            continue

        # 读取类别对应的csv文件
        csv_file = os.path.join(data_dir, class_dir, f"GT-{class_dir}.csv")
        if not os.path.exists(csv_file):
            continue

        # 读取csv文件
        df = pd.read_csv(csv_file, delimiter=";")
        # 遍历csv文件的每一行
        for index, row in df.iterrows():
            # 读取图片的路径与标签信息
            filename = os.path.join(data_dir, class_dir, row["Filename"])
            label = row["ClassId"]
            self.data.append(filename)
            self.label.append(label)
  • __len__

__len__ 函数用于返回数据集的长度,长度即数据集中的样本数量。

python
def __len__(self):
    return len(self.data)
  • __getitem__

__getitem__ 函数用于加载并返回数据集中的数据。在这个函数中,我们加载图片,并返回图片与标签信息。

python
def __getitem__(self, index):
    # 读取图片
    image = Image.open(self.data[index])
    label = self.label[index]

    # 对图片进行转换,如果没有指定转换方式,则使用ToTensor转换
    if self.transform:
        image = self.transform(image)
    else:
        image = ToTensor()(image)

    return image, label

DataLoader进阶

分批加载&重新排列

Dataset一次性获取所有数据,当数据量较大时,可能会导致内存不足。DataLoader可以将数据集分批次加载并重新排列,以便于训练模型。

关键参数可查看官方文档

具体来说,我们可以通过batch_size参数指定每个批次的大小,通过shuffle参数指定是否重新排列数据集。

python
from torch.utils.data import DataLoader

train_dataloader = DataLoader(training_data, batch_size=64, shuffle=True)
test_dataloader = DataLoader(test_data, batch_size=64, shuffle=True)

遍历DataLoader

DataLoader是一个可迭代对象,我们可以通过for循环遍历DataLoader,获取每个批次的数据。

python
for data, label in train_dataloader:
    print(data.shape, label.shape)
评论 隐私政策