欢迎光临
我们一直在努力

torchvision数据集

一.基础使用(以 CIFAR10 / MNIST 为例)

1.导入依赖

import torch
from torchvision import datasets, transforms
from torch.utils.data import DataLoader

2.简单加载数据集

示例1:MNIST 手写数字

:::info
基本信息

  • 内容:手写阿拉伯数字 0~9,共10个类别
  • 图像规格:单通道灰度图,尺寸 28×28 像素
  • 数据量:
    • 训练集:60000 张
    • 测试集:10000 张
  • 特点:画面简单、背景干净、计算量小,适合新手入门、调试代码、测试模型
  • 只有黑白灰,每张图就一个数字,几乎无干扰,早期深度学习“Hello World”级数据集。
  • 归一化参数:mean=(0.1307) std=(0.3081)

:::

#1. 定义图像变换:转张量 + 归一化
transform = transforms.Compose([
transforms.ToTensor(), # PIL图片 -> 张量,像素值 [0,1]
transforms.Normalize((0.1307,), (0.3081,)) # MNIST专用均值、标准差
])
#2. 加载训练集、测试集
train_dataset = datasets.MNIST(
root="./data", # 数据存放路径
train=True, # True=训练集,False=测试集
download=True, # 自动下载,已有文件不会重复下载
transform=transform # 绑定预处理
)

示例2:CIFAR10 彩色图像

:::info
基本信息

  • 内容:10 类日常实物,类别固定:
    飞机、汽车、鸟类、猫、鹿、狗、青蛙、马、船、卡车
  • 图像规格:三通道彩色图(RGB),尺寸 32×32 像素
  • 数据量:
    • 训练集:50000 张
    • 测试集:10000 张
  • 特点:图像更小、色彩丰富、存在遮挡/模糊,难度比 MNIST 高,用来练习彩色图像分类、数据增强、CNN 网络设计
  • 归一化:mean=(0.4914,0.4822,0.4465) std=(0.2470,0.2435,0.2616)

:::

transform = transforms.Compose([
transforms.ToTensor(),
transforms.Normalize((0.4914, 0.4822, 0.4465), # 3通道均值
(0.2470, 0.2435, 0.2616)) # 3通道标准差
])
train_dataset = datasets.CIFAR10(
root="./data", #数据集存放在当前项目下 data 文件夹,没有文件夹会自动新建
train=True, #True=训练集(5万张图片) False=测试集(1万张图片)
download=True, #没有CIFAR10文件 ->自动联网下载数据集,已有文件就不再重复下载
transform=transform
)

二.核心参数说明(通用所有 torchvision 数据集)

  • root :数据集保存目录
  • train : True 训练集 / False 测试集
  • download : True 自动在线下载
  • transform :图像预处理流水线(最常用)
  • target_transform :对标签做变换(很少用)

三.搭配 DataLoader 批量取数据

Dataset 只负责读取单条数据,训练要用 DataLoader 分批、打乱、多线程加载:

from torchvision import datasets,transforms
from torch.utils.data import DataLoader
#训练集加载器
train_loader = DataLoader(
dataset=train_dataset,
batch_size=64, # 批次大小
shuffle=True, # 打乱数据(训练集开启,测试集关闭)
num_workers=0 # 多线程读取,Windows建议设为0
)
测试集加载器
test_loader = DataLoader(
dataset=test_dataset,
batch_size=64,
shuffle=False
)

遍历读取批次数据

#迭代取数据
for images, labels in train_loader:
# images: [batch_size, 通道数, H, W] 张量
# labels: [batch_size] 标签张量
print(images.shape, labels.shape)
break # 只看第一批

四.常用图像预处理 transforms

组合多种操作,放在 transforms.Compose 里:

train_transform = transforms.Compose([
transforms.RandomCrop(32, padding=4), # 随机裁剪
transforms.RandomHorizontalFlip(), # 随机水平翻转(数据增强)
transforms.ToTensor(),
transforms.Normalize(mean, std)
])
测试集不要随机增强
test_transform = transforms.Compose([
transforms.ToTensor(),
transforms.Normalize(mean, std)
])

五.加载自己的图片文件夹:ImageFolder

如果是自己的分类数据集(文件夹作为标签),用 ImageFolder ,用法和内置数据集一致。

代码使用

from torchvision.datasets import ImageFolder
transform = transforms.Compose([
transforms.Resize((64, 64)),
transforms.ToTensor()
])
#直接传入根目录
dataset = ImageFolder(root="./your_data", transform=transform)
dataloader = DataLoader(dataset, batch_size=8, shuffle=True)
查看类别对应关系
print(dataset.class_to_idx) # 文件夹名 -> 数字标签

六.完整可运行模板

import torch
import torch.nn as nn
from torchvision import datasets, transforms
from torch.utils.data import DataLoader
#1. 预处理
transform = transforms.Compose([
transforms.ToTensor(),
transforms.Normalize((0.1307,), (0.3081,))
])
#2. 加载数据集
train_ds = datasets.MNIST("./data", train=True, download=True, transform=transform)
train_loader = DataLoader(train_ds, batch_size=32, shuffle=True)
#3. 简单测试遍历
for img, label in train_loader:
print("图像形状:", img.shape)
print("标签:", label)
break

七.常见踩坑提醒

  • Windows 系统: DataLoader 的 num_workers 设为 0 ,否则报错。
  • 通道顺序:torchvision 输出是 [C, H, W] ,和 PIL/OpenCV [H,W,C] 相反。
  • 归一化:必须使用对应数据集的均值、标准差,不要全写 0.5。
  • 训练/测试集:训练集 shuffle=True ,测试集 shuffle=False 。
  • 赞(0)
    未经允许不得转载:171主机测评 » torchvision数据集
    分享到: 更多 (0)

    评论 抢沙发

    • 昵称 (必填)
    • 邮箱 (必填)
    • 网址