一.基础使用(以 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

