AI的一切都是数据驱动的,数据集加载是深度学习的第一步,优质的AI离不开优质的数据集。
数据集是指训练模型的预筛选数据。
Pytorch提供了torch.utils.data.Dataset与torch.utils.data.DataLoader两个类,用于加载数据集。
数据集加载
torch.utils.data.Dataset
torch.utils.data.Dataset是一个抽象类,用于表示数据集。用户需要继承Dataset类,并实现__len__与__getitem__两个方法。
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可以将数据集分批次加载,以便于训练模型。
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模块中的MNIST、CIFAR10、CIFAR100等数据集。
Torchvision提供的数据集可查看Torchvision内置数据集
以GTSRB(德国交通标志识别基准 (GTSRB) )数据集为例,加载数据集的代码如下:
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数据集为例,其数据集的目录结构如下:
./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文件的内容示例如下:
Filename;Width;Height;Roi.X1;Roi.Y1;Roi.X2;Roi.Y2;ClassId
00000_00000.ppm;29;30;5;6;24;25;0在自定义数据集类中,我们不仅需要加载图片,还需要加载图片的标签信息。自定义数据集类的代码如下:
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.data与self.label中。在__getitem__方法中,我们加载了图片,并返回图片与标签信息。
__init__
__init__ 函数在实例化 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文件
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__ 函数用于返回数据集的长度,长度即数据集中的样本数量。
def __len__(self):
return len(self.data)__getitem__
__getitem__ 函数用于加载并返回数据集中的数据。在这个函数中,我们加载图片,并返回图片与标签信息。
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, labelDataLoader进阶
分批加载&重新排列
Dataset一次性获取所有数据,当数据量较大时,可能会导致内存不足。DataLoader可以将数据集分批次加载并重新排列,以便于训练模型。
关键参数可查看官方文档
具体来说,我们可以通过batch_size参数指定每个批次的大小,通过shuffle参数指定是否重新排列数据集。
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,获取每个批次的数据。
for data, label in train_dataloader:
print(data.shape, label.shape)