torchvision中的数据使用

1、下载数据集

在pytorch官网中找到docs选择Domains,在该页面中有各种数据类型的数据集

在左边菜单栏中选择datasets

python 复制代码
import torchvision
train_set=torchvision.datasets.CIFAR10(root='/data',train=True,download=True)
test_set=torchvision.datasets.CIFAR10(root='./data',train=False,download=True)

2、Dataloader的使用

Dataloader参数介绍

  • dataset :加载的数据集,必须是 torch.utils.data.Dataset 的子类实例。

  • batch_size:每个批次的数据样本数,默认值为1。

  • shuffle:是否在每个周期开始时打乱数据,默认为 False。

  • sampler:定义从数据集中抽取样本的策略,如果指定,则忽略 shuffle 参数。

  • num_workers:用于数据加载的子进程数量,默认为0,表示数据将在主进程中加载。

  • collate_fn:如何将多个数据样本整合成一个批次,通常不需要指定。

  • pin_memory:如果为 True,会将数据放置到 GPU 上去,默认为 False。

  • drop_last:如果数据集大小不能被批次大小整除,是否丢弃最后一个不完整的批次,默认为 False。

python 复制代码
test_loader=DataLoader(dataset=test_set,batch_size=4,shuffle=True,num_workers=0,drop_last=False)
#获取一张图片的信息
img,target=test_set[0]
print(img.shape)
print(target)

writer=SummaryWriter("dataloader")
#taet_loader是一个迭代对象,用for循环进行迭代
step=0
for data in test_loader:
    imgs,targets=data
    # print(imgs.shape)
    # print(targets)
    writer.add_image("test_data",imgs,step,dataformats='NCHW')
    step+=1

writer.close()

添加轮次

python 复制代码
for epoch in range(2):
    step=0
    for data in test_loader:
        imgs,targets=data
        # print(imgs.shape)
        # print(targets)
        writer.add_image("Epoch:{}".format(epoch),imgs,step,dataformats='NCHW')
        step+=1
相关推荐
qyz_hr1 分钟前
央国企人力资源穿透式监管的关键领域、核心机制与数智化路径研究
大数据·人工智能
打工仔折腾 AI2 分钟前
网易UU远程实测:手机控电脑做Python开发和Agent调试的真实体验
人工智能·后端·python·智能手机·性能优化·电脑·ai agent 实战
2601_962381133 分钟前
做城市历史视频时,AI 自动生成变迁动效的实现路径拆解
人工智能·音视频
知几蜗牛3 分钟前
从Holo4理解GUI Agent的坐标离散化、反映射与误差
人工智能
haliu3 分钟前
【FHE】(十三):位级可复现性——4 线程与 8 线程为什么算出不同的结果
人工智能·嵌入式·c·fhe·推理引擎·c11·边缘推理·同态加密推理
Token掘金室5 分钟前
MCP 怎么接大模型?Model Context Protocol 接入教程
人工智能
谢亮_vipxieliang6 分钟前
容器日志收集与管理:从 stdout 规范到 ELK/Loki 落地
运维·网络·人工智能·elk·docker·容器
YOLO数据集集合8 分钟前
风力发电机检测数据集 | 风机检测 电缆塔识别 风电运维 无人机巡检 9122期
运维·人工智能·目标检测·目标跟踪·无人机·风力发电·电力巡检
像风一样自由20208 分钟前
42.VueReactNextjs如何为AI应用设计前端交互
前端·人工智能·大模型·交互·rag·智能体