torchvision数据集介绍

torchvision中有很多数据集,当我们写代码时指定相应的数据集指定一些参数,它就可以自行下载。

CIFAR-10数据集包含60000张32×32的彩色图片,一共10个类别,其中50000张训练图片,10000张测试图片。

torchvision数据集下载

导入torchvision包

import torchvision

下载CIFAR10数据集到文件夹下,root是存放数据集的相对路径

trian_set = torchvision.datasets.CIFAR10(root="./dataset_transform",train=True,download=True)

同样的下载CIFAR10测试集,其中测试集的train为False

test_set = torchvision.datasets.CIFAR10(root="./dataset_transform",train=False,download=True)

下载好后如下,有数据集的压缩文件和解压后的文件

查看CIFAR10数据集内容

输出数据集中的一项

print(test_set[0]) 

输出结果

第一部分是PIL图片的格式,第二部分是图片的target的编号,也就是图片的类型

所以,当我们获取数据集中的全部类型时

print(test_set.classes)

可以得到图片类型列表,说明上面的test_set[0]的target对应为cat

分别获得图片、target和对应的种类

img, target = test_set[0]
print(img)
print(target)
print(test_set.classes[target])

结果为

展示图片

img.show()

由于图片尺寸只有32*32,所以较为模糊

以下是完整代码

import torchvision

train_set = torchvision.datasets.CIFAR10(root="./dataset",train=True,download=True)
test_set = torchvision.datasets.CIFAR10(root="./dataset",train=False,download=True)

print(test_set[0])
print(test_set.classes)

img, target = test_set[0]
print(img)
print(target)
print(test_set.classes[target])

img.show()

Tensorboard查看内容

设置transform操作步骤,这里使用ToTensor操作即转为tensor类型

dataset_transform = torchvision.transforms.Compose([torchvision.transforms.ToTensor()])

将ToTensor应用到数据集中的每一张图片,每一张图片转为Tensor数据类型

train_set = torchvision.datasets.CIFAR10(root="./dataset",train=True,transform=dataset_transform,download=True)     
test_set = torchvision.datasets.CIFAR10(root="./dataset",train=False,transform=dataset_transform,download=True)   

生成p10的日志文件后,遍历测试集的10张图片

writer = SummaryWriter("p10")
for i in range(10):
    img,target = test_set[i]
    writer.add_image("test_set",img,i)

writer.close()

读取日志,在tensorboard中打开效果如下

完整代码

import torchvision
from torch.utils.tensorboard import SummaryWriter

dataset_transform = torchvision.transforms.Compose([torchvision.transforms.ToTensor()])     
train_set = torchvision.datasets.CIFAR10(root="./dataset",train=True,transform=dataset_transform,download=True)  
test_set = torchvision.datasets.CIFAR10(root="./dataset",train=False,transform=dataset_transform,download=True)   

writer = SummaryWriter("logs") 
for i in range(10):
    img, target = test_set[i]
    writer.add_image("test_set",img,i)
    print(img.size())

writer.close()

更多推荐