TorchVision:数据集与预训练模型
本教程共 60 篇 · 第 34 篇 · 更新于 2026-08-17 · 约 4 分钟阅读
本节目标:认识 TorchVision 的组成,会用它加载经典数据集和预训练模型,并能把自己文件夹里的图片变成可用数据集。
从这一章开始,我们正式进入计算机视觉(Computer Vision,简称 CV)的地盘。写 CV 代码绕不开一个库:TorchVision。它是 PyTorch 官方配套的视觉工具箱,装 PyTorch 时一般一起装。
它到底装了什么
TorchVision 由几块组成,各有分工:
torchvision.datasets:内置几十个公开数据集,一行代码下载。torchvision.models:现成的模型结构,附带预训练权重。torchvision.transforms:图像变换与数据增强,第 12 章细讲过。torchvision.io:读写图片、视频文件。torchvision.ops:检测、分割用的专用算子。torchvision.utils:画框、拼图、辅助函数。
新手先吃透 datasets 和 models 两块,其余用到再查。
NoteTorchVision 的版本号不跟 PyTorch 一一对应(PyTorch 2.x 配 TorchVision 0.x 系列),安装时配套装最新版即可,
pip install torchvision会自动处理版本匹配。验证版本:import torchvision; print(torchvision.__version__)。
datasets:一行下载数据集
自己找数据、写下载脚本很折腾。内置数据集把这一切包圆了。CIFAR-10、MNIST、FashionMNIST、ImageNet 等都有:
from torchvision import datasets, transforms
transform = transforms.Compose([
transforms.ToTensor(),
transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5)),
])
trainset = datasets.CIFAR10(
root='./data', # 存到哪
train=True, # True 训练集,False 测试集
download=True, # 没有就自动下载
transform=transform, # 读到数据后做的变换
)
print(len(trainset)) # 50000
root 是存储目录,download=True 会在缺失时自动下载,transform 把每张图先转张量再归一化。返回的数据集和普通 Dataset 用法一样:trainset[0] 是 (图片张量, 标签) 的元组,套个 DataLoader 就能批量训练。
不同数据集参数略有差异。比如 MNIST 是单通道图,归一化的 mean 和 std 各只有一个数:Normalize((0.1307,), (0.3081,))。用之前查一眼文档,记住「图是几通道」就不会写错。
除了分类数据集,torchvision.datasets 里还有检测和分割用的:COCO、VOC、Cityscapes 等。它们的用法不一样,读出来的不是「图 + 标签」,而是「图 + 标注字典」(边框、掩码等),第 37、38 章会用到,到时候再细讲。
ImageFolder:读你自己的图片
内置数据集玩腻了,你有自己的图片文件夹怎么读?前提是目录按「每个类别一个子文件夹」组织:
data/
猫/
a.jpg
b.png
狗/
c.jpg
d.jpg
ImageFolder 会把子文件夹名自动当成类别标签——猫是 0、狗是 1,按字母序排。这几乎成了约定俗成的数据集组织方式,网上教程的数据包十有八九长这样:
from torchvision import datasets, transforms
transform = transforms.Compose([
transforms.Resize((224, 224)),
transforms.ToTensor(),
])
dataset = datasets.ImageFolder(root='./data', transform=transform)
print(dataset.classes) # ['猫', '狗']
print(dataset[0][1]) # 第一张图的标签(0 或 1)
dataset.classes 是类别名列表,dataset.class_to_idx 是「类名 → 编号」的映射。图片尺寸不一没关系,Resize 统一到模型要的输入大小。
Tip数据增强别放在验证集上。训练集可以加 RandomHorizontalFlip、RandomCrop 这类随机变换,验证和测试集只做 Resize、ToTensor、Normalize,保证评估公平。这个原则第 12 章也强调过,记住就行。
models:预训练模型仓库
torchvision.models 里有几十个经典架构:ResNet、VGG、MobileNet、EfficientNet、Swin Transformer 等。带不带预训练权重,用 weights 参数控制:
from torchvision import models
# 带预训练权重(推荐写法)
model = models.resnet18(weights=models.ResNet18_Weights.IMAGENET1K_V1)
# 只要随机初始化的结构
model = models.resnet18(weights=None)
weights="DEFAULT" 等价于用该模型最新的默认权重,也常用。旧写法 pretrained=True 已弃用,会弹警告,别再用。
模型默认在 ImageNet(1000 类)上预训练,所以输出层是 1000 个节点。你的任务类别数不同,就要换掉最后一层。以 ResNet 为例:
import torch.nn as nn
model.fc = nn.Linear(model.fc.in_features, 10) # 改成 10 类
不同模型最后一层的名字不一样:ResNet 是 model.fc,VGG 和 AlexNet 是 model.classifier[6],MobileNetV2 是 model.classifier[1]。打印模型结构看一眼就能找到。换层之后怎么冻结、怎么微调,第 26 章有完整套路。
常见模型在 ImageNet 上的 Top-1 准确率大概如下:ResNet-18 约 70%,ResNet-50 约 80%,VGG-16 约 71%,MobileNetV2 约 72%,EfficientNet-B7 约 84%。参数越多不一定越准,但一定越吃显存。入门从 ResNet-18 或 MobileNetV2 起步就够。
拿预训练模型跑一次推理
模型加载完,直接对一张新图做预测。这里有个省心的细节:用 weights 枚举对象能顺便拿到配套的预处理方法,不用自己手写 Resize 和 Normalize 的参数:
from torchvision import models
from torchvision.io import read_image
weights = models.ResNet18_Weights.IMAGENET1K_V1
model = models.resnet18(weights=weights)
model.eval()
img = read_image("./dog.jpg") # 读成 (3, H, W) 张量
preprocess = weights.transforms() # 配套预处理:缩放 + 归一化
batch = preprocess(img).unsqueeze(0) # 加 batch 维
with torch.no_grad():
out = model(batch)
probs = torch.softmax(out, dim=1)
print(weights.meta["categories"][probs.argmax().item()])
weights.meta["categories"] 是 ImageNet 的 1000 个类别名列表,索引正好对上输出节点。别忘了 model.eval(),推理时 BatchNorm、Dropout 的行为和训练不一样。
预训练模型的底气
预训练模型为什么好用?因为它已经在 ImageNet(ILSVRC2012 训练集约 128 万张图、1000 类)上练过,学会了边缘、纹理、形状这些通用视觉特征。你把它拿来当起点,几百张自己的图就能调出一个不错的分类器,这就是迁移学习。具体怎么冻结、怎么微调,第 26 章有完整教程,这里先记住一句话:能用预训练就别从零训。
Tip不确定用哪个模型?入门选 ResNet-18:结构简单、权重小(约 45MB)、精度够用。移动端场景再看 MobileNetV2。
v2 变换与 tv_tensors
TorchVision 0.15 起推出了新的变换接口 torchvision.transforms.v2,旧的 transforms 还能用,但官方更推荐 v2。v2 最大的改进是支持目标检测和分割的标注一起变换,而且会用 tv_tensors 把数据包装起来:
from torchvision import tv_tensors
# tv_tensors.Image / BoundingBoxes / Mask 都是 Tensor 的子类
img = tv_tensors.Image(torch.rand(3, 224, 224))
print(type(img), img.shape) # <class 'torchvision.tv_tensors.Image'> torch.Size([3, 224, 224])
tv_tensors 本质是带「身份信息」的张量:变换知道它是图片还是标注框,就不会把坐标框当普通数值乱缩放。分类任务暂时用不上这些,但后面讲目标检测(第 37 章)会打交道,先混个脸熟。
还有个小工具值得记住:torchvision.utils.make_grid 把一批小图拼成一张大图,配合 matplotlib 展示样本、看看增强效果都很方便,第 31 章用过它。
小结
TorchVision 三件套:datasets 拿数据、transforms 处理数据、models 拿模型。内置数据集一键下载,ImageFolder 读自己的图,weights 参数加载预训练权重。有了它,图像分类的准备工作从「折腾几天」变成「几行代码」。下一章,我们就把这些零件组装成一条完整的图像分类流水线。