PyTorch学习笔记—— 数据集
在深度学习任务中,数据加载和处理是至关重要的一环。
自定义数据集
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),以下为其中两组图片。


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



所有评论(0)