下载链接:
本文中我们使用的是\(CIFAR\)数据集
具体网站:
torchvision中每个数据集的参数都是大同小异的,这里只介绍CIFAR10数据集
该数据集的数据格式为PIL格式
class torchvision.datasets.CIFAR10(root:str,train:bool=True,transform:Optional[Callable]=None,target_transform:Optional[Callable]=None,download:bool=False)
内置函数:
root(string):必须设置,输入数据集下载后存放在电脑中的路径
train(bool):True代表创建的一个训练集(train);False代表创建一个测试集(test)。
transform:对数据集中的数据进行变换
target_transform:对标签(target)数据进行变换
download(bool):True的时候会自动从网上下载这个数据集,False的时候则不会下载该数据集。
代码示例:
运行后直接下载数据集
需要注意的是,如果下载速度过慢,则可以在运行后,把弹出的网址单拎出来,放到迅雷等软件上进行下载
import torchvision
#设置训练集
#root:设置为相对路径,会在该.py文件下设置一个名为dataset的文件存放CIFAR10数据
#train: True,数据集为训练集
#download: 下载该数据集
train_set=torchvision.datasets.CIFAR10(root="./dataset",train=True,download=True)
#设置测试集;train=False
test_set=torchvision.datasets.CIFAR10(root="./dataset",train=False,download=True)
数据标签查看:
在运行上面的代码下载好数据集后,输入print(test_set[0)
,并使用一下pycharm的dubug功能,不难发现:
也就是说,数据标签有'airplane', 'automobile', 'bird', 'cat', 'deer', 'dog', 'frog', 'horse', 'ship', 'truck'十类,分别用整数0~9来表示
数据集包含的所有标签也可以用下面的代码打印出来:
print(test_set.classes)
#[Run] [airplane', 'automobile', 'bird', 'cat', 'deer', 'dog', 'frog', 'horse', 'ship', 'truck']
img,target=test_set[索引]
img,target=test_set[0]
print(img)
print(target,test_set.classes[target])
#[Run]
#<PIL.Image.Image image mode=RGB size=32x32 at 0x1DDF9FCD640>
#3 cat
img.show()
首先使用\(Compose\)去定义如何处理PIL图像数据
然后代入\(torchvision.datasets.CIFAR10\)中,处理里面的图像数据
#首先用Compose处理图像数据,可以先转为tensor格式,然后再裁剪等,这里只转tensor格式
import torchvision
dataset_transform=torchvision.transforms.Compose([
torchvision.transforms.ToTensor()
])
#定义transform=dataset_transform,使得图像数据类型转换为Compose中处理过后的
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)
from torch.utils.tensorboard import SummaryWriter
writer=SummaryWriter("p10")
for i in range(10): #显示test_set数据集中的前十张图片
img,target=test_set[i]
writer.add_image("test_set",img,i)
writer.close()