首页 / PyTorch 入门教程 / TorchVision:数据集与预训练模型

PyTorch 入门教程

TorchVision:数据集与预训练模型

本教程共 60 篇 · 第 34 篇 · 更新于 2026-08-17 · 约 4 分钟阅读

PyTorchTorchVision数据集预训练模型ImageFolder计算机视觉

本节目标:认识 TorchVision 的组成,会用它加载经典数据集和预训练模型,并能把自己文件夹里的图片变成可用数据集。

从这一章开始,我们正式进入计算机视觉(Computer Vision,简称 CV)的地盘。写 CV 代码绕不开一个库:TorchVision。它是 PyTorch 官方配套的视觉工具箱,装 PyTorch 时一般一起装。

它到底装了什么

TorchVision 由几块组成,各有分工:

  • torchvision.datasets:内置几十个公开数据集,一行代码下载。
  • torchvision.models:现成的模型结构,附带预训练权重。
  • torchvision.transforms:图像变换与数据增强,第 12 章细讲过。
  • torchvision.io:读写图片、视频文件。
  • torchvision.ops:检测、分割用的专用算子。
  • torchvision.utils:画框、拼图、辅助函数。

新手先吃透 datasets 和 models 两块,其余用到再查。

Note

TorchVision 的版本号不跟 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 参数加载预训练权重。有了它,图像分类的准备工作从「折腾几天」变成「几行代码」。下一章,我们就把这些零件组装成一条完整的图像分类流水线。