pytorch5->datasets.CIFAR10简单例题
import torchvision
from torch.utils.tensorboard import SummaryWriter
import numpy as np
train_set = torchvision.datasets.CIFAR10(
root='./data',
train=True,
download=True
)
print(f'数据集大小:{len(train_set)}')
print(f'第一张图片类型:{type(train_set[0])}')
img, label = train_set[0]
print(f'图片尺寸:{img.size}')
print(f'图片类别:{label}')
print(f'类别名字:{train_set.classes[label]}')
writer = SummaryWriter('logs_dataset')
for i in range(10):
img, label = train_set[i]
img_np = np.array(img)
writer.add_image(f'cifar10/{train_set.classes[label]}', img_np, i, dataformats='HWC')
writer.close()
print('完成')
1.torchvision:PyTorch 的计算机视觉工具包,专门用来处理图像和视频数据的库。
2.train_set=torchvision.datasets.CIFAR10(
root =‘./data’,
train =True,
download =True
)
torchvision.datasets.CIFAR10:
torchvision,库。
datasets,模块。
CIFAR10,类,是一个在计算机视觉领域非常经典的彩色图像分类数据集,包含60,000张32x32像素的彩色图片
root =‘./data’,数据下载路径
train =True,(True,表示训练集,false,表示测试集)
训练集:50,000张图片,每个类别有5,000张,用于训练模型。
测试集:10,000张图片,每个类别有1,000张,用于评估模型效果。
训练集图片多,测试集图片少,不能反过来。如果训练少,测试多,它会有很多没见过的情况,你就测试不出来什么了
download =True,本地没有就下载
3.print(f’第一张图片类型:{type(train_set[0])}')
train_set[0],train_set对象正常不能[0],由于train_set[0]等价于train_set.getitem(0),getitem__魔法 方法内部实现了对象[索引]的功能
返回类型是元组。因为train_set[0]由.CIFAR10类实例化,CIFAR10中有__getitem 方法return img, label
4.img,label=train_set[0]
这里是解包,打包发生在return img, label 时
5.print(f’类别名字:{train_set.classes[label]}')
classes是列表,CIFAR10出生时自带,可以将label转化成name,比如label=1时,train_set.classes[label]是飞机.
AtomGit 是由开放原子开源基金会联合 CSDN 等生态伙伴共同推出的新一代开源与人工智能协作平台。平台坚持“开放、中立、公益”的理念,把代码托管、模型共享、数据集托管、智能体开发体验和算力服务整合在一起,为开发者提供从开发、训练到部署的一站式体验。
更多推荐



所有评论(0)