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.ops | NMS、框面积、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 Pathfrom tempfile import TemporaryDirectoryimport torchfrom torchvision.io import decode_image, ImageReadMode, write_pngfrom torchvision.transforms import v2
image = torch.zeros((3, 8, 12), dtype=torch.uint8)image[0] = 255image[1, :, 6:] = 128with 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 torchfrom torchvision.models import resnet18, ResNet18_Weights
torch.manual_seed(23)weights = ResNet18_Weights.IMAGENET1K_V1preprocess = 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 torchfrom 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-6assert 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 与视频弃用说明