14.torchvision中的数据集使用
pytorch的官方文档https://docs.pytorch.org/vision/0.29/generated/torchvision.datasets.CIFAR10.html#torchvision.datasets.CIFAR10CIFAR10数据集数据集特点包含60,000张32×32像素的彩色图片分为10个类别每个类别6,000张图片训练集50,000张测试集10,000张类别包括飞机、汽车、鸟、猫、鹿、狗、青蛙、马、船、卡车CIFAR-10数据集下载使用关键参数root数据集存储路径必须设置trainTrue表示训练集False表示测试集transform对图像数据的预处理方法target_transform对标签数据的预处理方法downloadTrue时自动从网上下载数据集我们执行文件下载数据集但是速度会很慢如下下载技巧当下载速度慢时可以复制如上终端中的下载链接到迅雷等下载工具加速下载使用迅雷等下载工具获取数据集压缩包如cifar-10-python.tar.gz手动创建与代码中root参数同名的文件夹如./dataset将下载好的压缩包复制到目标文件夹校验机制程序运行时自动校验文件完整性若文件完整则跳过下载直接解压如下dataset文件夹放的是我们通过将下载链接通过迅雷工具下载后放置的这样下载的速度更快。然后执行文件的时候会校验我们预下载的这个文件的完整性。但是如果在终端执行文件进行数据集下载有的数据集下载的时候没有显示下载链接我们可以通过command点击数据集api,跳到这个数据集方法的实现处可以查看到对应的url如下我们将鼠标点击在如下代码的CIFAR10同时按下command键train_settorchvision.datasets.CIFAR10(root./dataset,trainTrue,transformdataset_transform,downloadTrue)我们就会跳转到方法的实现处我们可以看到url的下载链接开发建议保持downloadTrue便于自动校验推荐使用预下载自动校验的工作流程高效下载技巧通过迅雷加速下载复制源码中的URL到迅雷利用P2P/镜像资源。downloadTrue时优先检查本地缓存避免重复下载。数据集的数据结构每个样本返回元组(image, target)image是PIL.Image对象如下第一个样本的image是PIL.Image对象target是类别索引0-9如下第一个样本的索引是3可通过classes属性查看类别名称。如下打断点在变量面板处我们可以看到test_set是有classes属性的。print(test_set.classes)# 输出类别名称列表print(test_set.classes[target])# 输出具体类别名称样本的target便是类别列表的索引所以想知道样本是什么类别的图片我们就在数据集对象test_set.classes的类别列表里找到样本的target属性值这个索引对应的类别名称。如下我们知道第一个样本的类别是猫。我们可以调用show方法查看图片。importtorchvisionfromtorch.utils.tensorboardimportSummaryWriter train_settorchvision.datasets.CIFAR10(root./dataset,trainTrue,downloadTrue)test_settorchvision.datasets.CIFAR10(root./dataset,trainFalse,downloadTrue)print(test_set[0])print(test_set.classes)img,targettest_set[0]print(img)print(target)print(test_set.classes[target])img.show()将数据集中的数据转为tensor格式transform使用使用Compose组合多个transform操作将transform传入dataset的transform参数对数据集中的每张图片自动应用transformimporttorchvisionfromtorch.utils.tensorboardimportSummaryWriter dataset_transformtorchvision.transforms.Compose([torchvision.transforms.ToTensor()])train_settorchvision.datasets.CIFAR10(root./dataset,trainTrue,transformdataset_transform,downloadTrue)test_settorchvision.datasets.CIFAR10(root./dataset,trainFalse,transformdataset_transform,downloadTrue)writerSummaryWriter(p10)foriinrange(10):img,targettest_set[i]writer.add_image(test_set,img,i)writer.close()如下我们打印出来转换成tensor格式的图片数据集里的图片确实都转成了tensor格式tensorboard图像展示实现步骤创建SummaryWriter遍历数据集获取tensor格式图片使用add_image添加到tensorboard关闭writerimporttorchvisionfromtorch.utils.tensorboardimportSummaryWriter dataset_transformtorchvision.transforms.Compose([torchvision.transforms.ToTensor()])train_settorchvision.datasets.CIFAR10(root./dataset,trainTrue,transformdataset_transform,downloadTrue)test_settorchvision.datasets.CIFAR10(root./dataset,trainFalse,transformdataset_transform,downloadTrue)print(test_set[0])writerSummaryWriter(p10)foriinrange(10):img,targettest_set[i]writer.add_image(test_set,img,i)writer.close()tensor的格式的图片我们就可以用tensorboard去查看了。查看方法命令行运行tensorboard --logdir“p10”在浏览器打开localhost:6006查看图片COCO数据集关于如何查看torchvision各数据集的官方文档打开官网https://pytorch.org根据自己的环境中安装的torchvision的版本切换到对应版本的文档下通过在终端运行pip list,查看torchvision的版本如下是0.9.0版本的torchvision的coco数据集的文档参数配置root指定数据集保存的根目录路径annFile需要额外指定JSON标注文件的路径transform对PIL图像进行转换的函数如转换为Tensortarget_transform对标注目标进行转换的函数