跳转到内容
新建笔记

torchvision:图像变换、视觉模型与框运算

torchvision 在视觉流程中的位置

跳转到“torchvision 在视觉流程中的位置”

torchvision 为 PyTorch 提供常用视觉数据集、图像变换、模型结构、图像读写和框运算。原笔记是一次 dir(torchvision) 的长列表,其中混有 torch、os、warnings 等导入对象、双下划线元信息和 _HAS_OPS 等内部状态。它们不是需要逐一学习的公开功能模块。

本文采用 torchvision 0.23.0 + PyTorch 2.8.0 CPU。所有示例使用自己生成的 PNG、随机初始化模型和手工框,不下载数据或预训练权重。

公开入口用途常见误解
torchvision.datasets内置视觉数据集与基类指定 download=True 可能触发下载;数据许可和划分仍需自行确认
torchvision.transforms.v2图像及其他视觉对象的组合变换转为浮点类型不必然缩放数值;图像、标签与边界框需要一致的几何变换
torchvision.models分类、检测、分割等模型和权重配置weights=None 是随机初始化,不含已训练能力
torchvision.io图像解码与编码RGB/BGR、位深、通道维度应显式核对
torchvision.opsNMS、框面积、IoU 等视觉算子有些算子依赖与 PyTorch 匹配的本地扩展
torchvision.utils拼图、画框与图像保存可视化输入范围和布局与模型输入约定可能不同

模块的稳定性要看相应文档标记;公开模块中也可能有 Beta 功能。后端查询函数与 _HAS_OPS 不是同一层接口,不能把内部布尔值当作通用安装诊断。torchvision 0.23 文档

图像:布局、类型、范围都要一致

跳转到“图像:布局、类型、范围都要一致”

本例先保存一张 8 位 RGB PNG,再通过 decode_image 读回。解码后是 C,H,W,不是绘图库常见的 H,W,C。ToImage 标记图像语义,不负责把 0–255 缩放到 0–1;ToDtype(..., scale=True) 才在这里同时处理 dtype 与范围转换。

from pathlib import Path
from tempfile import TemporaryDirectory
import torch
from torchvision.io import decode_image, ImageReadMode, write_png
from torchvision.transforms import v2
image = torch.zeros((3, 8, 12), dtype=torch.uint8)
image[0] = 255
image[1, :, 6:] = 128
with TemporaryDirectory() as directory:
path = Path(directory) / "rgb.png"
write_png(image, str(path))
decoded = decode_image(str(path), mode=ImageReadMode.RGB)
assert decoded.dtype == torch.uint8 and decoded.shape == (3, 8, 12)
assert torch.equal(decoded, image)
marked = v2.ToImage()(decoded)
assert marked.dtype == torch.uint8 and marked.max().item() == 255
transform = v2.Compose([
v2.ToImage(),
v2.Resize((16, 24), antialias=True),
v2.ToDtype(torch.float32, scale=True),
])
value = transform(decoded)
assert value.shape == (3, 16, 24) and value.dtype == torch.float32
assert 0 <= value.min().item() <= value.max().item() <= 1
normalized = v2.Normalize(mean=[0.5] * 3, std=[0.5] * 3)(value)
torch.testing.assert_close(normalized, value * 2 - 1)
display_rgb = value.permute(1, 2, 0).numpy()
assert display_rgb.shape == (16, 24, 3)
print("PNG roundtrip and image transforms passed")

这一检查针对自己生成的 8 位 PNG,不能推广为“所有图片解码都是 uint8”:例如 16 位 PNG 有不同 dtype 与范围。灰度图可为单通道,而 mode=RGB 明确要求转换为三通道。Normalize 对每个通道做减均值、除标准差,通常不把结果限制在 0–1。批量分类模型常用 N,C,H,W;多种 v2 变换可处理批量,但随机增强在批内如何共享参数仍应查具体算子的约定。图像解码、v2 变换、ToDtype

显示 RGB 图像通常需要 image.permute(1, 2, 0);灰度图只应去掉确定存在的单通道维。无条件 squeeze() 会同时去掉其他长度为 1 的维度,容易把批次或空间维一起丢掉。CSV 标签与 DataLoader 的完整例子见数据集创建。

模型结构和预训练权重要分别选择

跳转到“模型结构和预训练权重要分别选择”
import torch
from torchvision.models import resnet18, ResNet18_Weights
torch.manual_seed(23)
weights = ResNet18_Weights.IMAGENET1K_V1
preprocess = weights.transforms()
image = torch.randint(0, 256, (3, 256, 300), dtype=torch.uint8)
batch = preprocess(image).unsqueeze(0)
assert batch.shape == (1, 3, 224, 224)
model = resnet18(weights=None).eval()
with torch.inference_mode():
logits = model(batch)
assert logits.shape == (1, 1000) and torch.isfinite(logits).all()
print("random-weight ResNet18 output shape:", tuple(logits.shape))
print("preprocessing weight configuration:", weights.name)

访问权重枚举和 weights.transforms() 不会下载模型权重;将这个枚举传给 resnet18(weights=weights) 才会请求预训练参数,缓存中没有时需要下载。上例故意使用 weights=None,只检验接口和前向尺寸,随机 logits 不具有 ImageNet 分类意义。真正使用预训练模型时,应将权重与其预处理绑定,记录颜色空间、缩放/裁剪、归一化和类别表;不能只凭网络名称复制另一篇教程的均值方差。模型与预训练权重、ResNet18

NMS:按类别处理重叠框

跳转到“NMS:按类别处理重叠框”
import torch
from torchvision.ops import box_iou, nms, batched_nms
boxes = torch.tensor([[0., 0., 10., 10.],
[1., 1., 11., 11.],
[30., 30., 40., 40.]])
scores = torch.tensor([0.9, 0.8, 0.7])
categories = torch.tensor([0, 1, 0])
iou = box_iou(boxes[:1], boxes[1:2]).item()
assert abs(iou - 81 / 119) < 1e-6
assert nms(boxes, scores, 0.5).tolist() == [0, 2]
assert batched_nms(boxes, scores, categories, 0.5).tolist() == [0, 1, 2]
print("IoU:", iou)

这里使用 (x1,y1,x2,y2) 浮点框和互不相同的分数。普通 nms 不接收类别,重叠框 0、1 会互相竞争;batched_nms 按类别隔离,本例两框类别不同,所以都保留。它不是给图片批次自动做完整后处理:若把多张图的框一起输入,类别/分组键还需要体现图片归属。同分框的选择不要假定 CPU/GPU 完全一致;本例只验证 CPU 确定输入。NMS、batched_nms

若导入时报本地扩展或 torchvision::nms 不存在,先记录 torch、torchvision 版本和 CPU/CUDA 安装来源,核对官方配套版本。单独升级其中一个包、混用不同构建来源,可能使 Python 模块存在而本地算子不可用。本文实测配对为 torch 2.8.0+cpu 和 torchvision 0.23.0+cpu。

0.23 文档中的旧视频读写功能已标为弃用,并公告计划在 0.24 移除;不要把这里的图片读写示例推断为视频管线已验证。视频任务应另行核对目标版本与 TorchCodec 等专门方案。原有 get_image_backend / set_image_backend 也不意味着所有 IO 路径共用同一个 Python 图片后端。0.23 IO 与视频弃用说明