PyTorch深度学习笔记(五)(torchvision数据集使用)
·
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()
更多推荐
所有评论(0)