前言

本节主要学习了Dataset数据加载代码实例

python中俩好用的函数

Python3.x相当于一个package,package里面有不同的区域,不同的区域有不同的工具。

Python语法有两大法宝:dir()、help() 函数。

dir() 

dir():打开,看见里面有多少分区、多少工具。

import torch
dir(torch)  # 查看torch包中有哪些区、有哪些工具

help()

help():说明书。

import torch
help(torch.cuda.is_available) # 查看 torch.cuda.is_available 的用法

Pytorch加载数据

Pytorch中加载数据需要Dataset、Dataloader。

 Dataset提供一种方式去获取每个数据及其对应的label,告诉我们总共有多少个数据。

 Dataloader为后面的网络提供不同的数据形式,它将一批一批数据进行一个打包。

常用数据集

常用的第一种数据形式,文件夹的名称是它的label。

常用的第二种形式,label为文本格式,文本名称为图片名称,文本中的内容为对应的label。

数据划分

所用到的蚂蚁蜜蜂数据集地址:https://pan.baidu.com/s/1jZoTmoFzaTLWh4lKBHVbEA 密码: 5suq

首先将train文件夹下的ants文件和bees文件分别改名为ants_image文件和bees_image文件,然后手动创建ants_label文件和bees_label文件作为标签存放文件,如下所示:

然后创建rename_dataset.py数据标签划分文件,进行标签划分

对于蚂蚁数据集部分

规定数据集路径和target路径文件名

root_dir = "dataset/train"
target_dir = "ants_image"

合并二者为图像路径,并获得路径下所有图像的地址

image_path = os.listdir(os.path.join(root_dir,target_dir))

label标签为target路径文件名经过划分后的第一项

label = target_dir.split("_")[0]

设定label输出路径文件名

out_dir = "ants_label"

对每个图像进行划分操作,去掉文件名后缀.jpg

for i in image_path:
    file_name = i.split(".jpg")[0]

在设定的输出地址中创建每个图像对应的label文件,并写入对应的标签

with open(os.path.join(root_dir,out_dir,"{}.txt".format(file_name)),"w") as f:
        f.write(label)

蜜蜂数据集同理如上

可以看到蚂蚁数据集和蜜蜂数据集的图片都生成了一一对应的标签文件:

完整代码:

import os

root_dir = "dataset/train"
target_dir = "ants_image"
image_path = os.listdir(os.path.join(root_dir,target_dir))
label = target_dir.split("_")[0]
out_dir = "ants_label"
for i in image_path:
    file_name = i.split(".jpg")[0]
    with open(os.path.join(root_dir,out_dir,"{}.txt".format(file_name)),"w") as f:
        f.write(label)

root_dir = "dataset/train"
target_dir = "bees_image"
image_path = os.listdir(os.path.join(root_dir,target_dir))
label = target_dir.split("_")[0]
out_dir = "bees_label"
for i in image_path:
    file_name = i.split(".jpg")[0]
    with open(os.path.join(root_dir,out_dir,"{}.txt".format(file_name)),"w") as f:
        f.write(label)

Dataset数据加载

创建一个继承Dataset的类MyData

class MyData(Dataset):

创建初始化类方法,当根据此类MyData创建一个事例对象时,会自动调用该函数

def __init__(self,root_dir,image_dir,label_dir):

该函数中为整个类提供全局变量,self相当于类中的全局变量

self.root_dir = root_dir
self.image_dir = image_dir
self.label_dir = label_dir

字符串拼接,根据是Windows或Lixus系统情况进行拼接

self.path = os.path.join(self.root_dir,self.image_dir)

获得路径下所有图片的地址

self.img_path = os.listdir(self.path)

完整方法代码

def __init__(self,root_dir,image_dir,label_dir):
    self.root_dir = root_dir
    self.image_dir = image_dir
    self.label_dir = label_dir
    self.path = os.path.join(self.root_dir,self.image_dir)
    self.img_path = os.listdir(self.path)

获取每一个图片的方法

def __getitem__(self, idx):

从图像地址的列表中读取对应位置的图像名称

img_name = self.img_path[idx]

接下来获取该图片的相对路径,将目录名、图像文件名和图像名合在一起

img_item_path = os.path.join(self.root_dir,self.image_dir,img_name)

根据相对路径读取该PIL类型图片

img = Image.open(img_item_path)

同时获取label的路径

label = self.label_dir

返回获取的图像和标签

return img,label

完整方法代码

def __getitem__(self, idx):
    img_name = self.img_path[idx]
    img_item_path = os.path.join(self.root_dir,self.image_dir,img_name)
    img = Image.open(img_item_path)
    label = self.label_dir
    return img,label

获取数据集长度的方法

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

设定路径名

root_dir = "dataset/train"
ants_label_dir = "ants_label"
bees_label_dir = "bees_label"
ants_image_dir = "ants_image"
bees_image_dir = "bees_image"

创建蚂蚁和蜜蜂的两个Dataset类的实例

ants_dataset = MyData(root_dir,ants_image_dir,ants_label_dir)
bees_dataset = MyData(root_dir,bees_image_dir,bees_label_dir)

输出数据集长度

print(len(ants_dataset))
print(len(bees_dataset))

两个数据集的集合为train_dataset

train_dataset = ants_dataset + bees_dataset

输出train_dataset的长度

print(len(train_dataset))

获取数据集第200张图片,打印其标签和详细信息,并展示图片

img,label = train_dataset[200]
print("label:",label)
print(getitem(train_dataset,200))
img.show()

完整代码

from operator import getitem
from torch.utils.data import Dataset
from PIL import Image
import os

class MyData(Dataset):
    def __init__(self,root_dir,image_dir,label_dir):
        self.root_dir = root_dir
        self.image_dir = image_dir
        self.label_dir = label_dir
        self.path = os.path.join(self.root_dir,self.image_dir)
        self.img_path = os.listdir(self.path)

    def __getitem__(self, idx):
        img_name = self.img_path[idx]
        img_item_path = os.path.join(self.root_dir,self.image_dir,img_name)
        img = Image.open(img_item_path)
        label = self.label_dir
        return img,label

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

root_dir = "dataset/train"
ants_label_dir = "ants_label"
bees_label_dir = "bees_label"
ants_image_dir = "ants_image"
bees_image_dir = "bees_image"
ants_dataset = MyData(root_dir,ants_image_dir,ants_label_dir)
bees_dataset = MyData(root_dir,bees_image_dir,bees_label_dir)

print(len(ants_dataset))
print(len(bees_dataset))
train_dataset = ants_dataset + bees_dataset   
print(len(train_dataset))

img,label = train_dataset[200]
print("label:",label)
print(getitem(train_dataset,200))
img.show()

分别输出蚂蚁和蜜蜂数据集长度和第200张图片的标签和详细信息并展示图片:

124
121
245
label: bees_label
(<PIL.JpegImagePlugin.JpegImageFile image mode=RGB size=441x500 at 0x1C1B0E39010>, 'bees_label')

Logo

码道开发者社区,聚焦华为云码道 CodeArts 代码智能体,沉淀 Agent、Skill、鸿蒙开发实战内容,供开发者查阅资料、交流技术、分享工程实践

更多推荐