TorchVision 模型库与预训练权重完整指南:Multi-Weight API、推理预处理与模型检索实战
2026/9/21 15:42:11 网站建设 项目流程

TorchVision 模型库与预训练权重完整指南:Multi-Weight API、推理预处理与模型检索实战

【免费下载链接】visionDatasets, Transforms and Models specific to Computer Vision项目地址: https://gitcode.com/gh_mirrors/vi/vision

torchvision.models是 TorchVision 提供的模型子包,为图像分类、语义分割、目标检测、实例分割、人体关键点检测、视频分类与光流估计等任务提供了开箱即用的模型定义与预训练权重。本文以仓库文档 docs/source/models.rst 为主线,结合源码(如 torchvision/models/_api.py、torchvision/models/resnet.py)深度讲解 v0.13 引入的 Multi-Weight Support API、pretrained参数的弃用迁移、随权重内置的推理预处理变换,以及 v0.14 新增的模型/权重按名检索机制,并给出每一类任务可直接复制的推理代码。读完本文,你将能够用一行代码加载任意 TorchVision 预训练模型、正确复用官方推理预处理、按名称枚举与获取模型及其权重,并理解底层实现原理。

torchvision.models 子包能做什么

torchvision.models子包包含面向不同视觉任务的模型定义,覆盖以下任务类型:

任务模型家族(示例)对应源码模块
图像分类AlexNet、ResNet、VGG、DenseNet、EfficientNet、ConvNeXt、Swin Transformer、ViT 等torchvision/models
量化分类(INT8)GoogLeNet、Inception、MobileNetV2/V3、ResNet、ResNeXt、ShuffleNetV2torchvision/models/quantization
语义分割FCN、DeepLabV3、LRASPPtorchvision/models/segmentation
目标检测Faster R-CNN、SSD、SSDLite、RetinaNet、FCOStorchvision/models/detection
实例分割Mask R-CNNtorchvision/models/detection
人体关键点检测Keypoint R-CNNtorchvision/models/detection
视频分类R3D、MC3、R2plus1D、S3D、MViT、Swin3Dtorchvision/models/video
光流估计RAFT(large / small)torchvision/models/optical_flow

每种架构均提供"带预训练权重"与"不带权重"两种实例化方式。检测、实例分割与关键点检测模型以 TorchVision 内置的分类模型作为骨干网络初始化,且输入约定为Tensor[C, H, W]的列表(List[Tensor[C, H, W]]),这一点与分类/分割模型的单张批处理输入不同,使用时需注意区分。

预训练权重的加载机制

TorchVision 为每一种提供的架构都准备了预训练权重,权重加载基于 PyTorch 的torch.hub机制实现:实例化一个带权重的模型时,会先将权重文件下载到缓存目录,再加载进模型。

缓存目录可通过TORCH_HOME环境变量指定,默认遵循 PyTorch 的约定(~/.cache/torch/hub/checkpoints一类位置)。底层下载逻辑由torch.hub.load_state_dict_from_url承担,支持断点续传与哈希校验。在源码层面,权重的获取统一封装在WeightsEnum.get_state_dict()中(见 torchvision/models/_api.py),它内部调用load_state_dict_from_url(self.url, ...),模型构建函数则通过model.load_state_dict(weights.get_state_dict(progress=progress, check_hash=True))完成装载,例如 torchvision/models/resnet.py。

两个必须知道的使用注意点

  • 许可证责任:仓库提供的预训练模型可能带有源自训练数据集(如 ImageNet、COCO、Kinetics-400)的独立许可证与使用条款。是否拥有在自身场景下使用这些模型的权限,由使用者自行判断,官方不代为背书。
  • 序列化兼容性:使用旧版 PyTorch 创建的模型加载序列化state_dict是保证向后兼容的;但加载整个保存的模型(torch.save(model))或旧版序列化的ScriptModule,则不保证保持历史行为。因此官方推荐只做权重(state_dict)层面的迁移。

Multi-Weight Support API:一个模型多套权重

自 v0.13 起,TorchVision 为现有模型构建方法引入了新的 Multi-Weight Support API:每个模型构建函数都接受一个weights参数,其取值可以是该模型专属的权重枚举类(WeightsEnum子类)成员。以 ResNet-50 为例:

from torchvision.models import resnet50, ResNet50_Weights # 旧权重,ImageNet-1K 上 acc@1 为 76.130% resnet50(weights=ResNet50_Weights.IMAGENET1K_V1) # 新权重(新训练配方),acc@1 为 80.858% resnet50(weights=ResNet50_Weights.IMAGENET1K_V2) # 当前最佳可用权重(目前是 IMAGENET1K_V2 的别名) # 注意:DEFAULT 指向的权重可能随版本变化 resnet50(weights=ResNet50_Weights.DEFAULT) # 字符串同样受支持 resnet50(weights="IMAGENET1K_V2") # 不加载权重,随机初始化 resnet50(weights=None)

上面列出的 76.130% 与 80.858% 两个精度均来自源码中 ResNet50_Weights 的meta._metrics记录。从同一源码还可以看到,同一架构不同版本的权重甚至可能对应不同的推理预处理参数:IMAGENET1K_V1的 transform 是crop_size=224,而IMAGENET1K_V2crop_size=224, resize_size=232——这正是"预处理方式必须与权重匹配"的典型证据。

权重枚举与 Weights 数据结构

理解这套 API 的底层,需要看 torchvision/models/_api.py 中的两个核心类型:

  • Weights:一个@dataclass,聚合了权重的三个关键属性:
    • url:权重文件下载地址;
    • transforms:一个可调用对象构造器(而非已构造对象),用于构建该权重对应的推理预处理方法。采用构造器而非实例的原因在于预处理对象可能持有内存,延迟初始化更经济;
    • meta:与权重相关的元数据字典,既包含信息性属性(参数量num_params、FLOPs_ops、训练配方recipe、指标_metrics),也包含使用模型所必需的关键信息(如分类模型的categories类别列表)。
  • WeightsEnum:所有权重枚举的父类,继承自 PythonEnum,其枚举值必须是Weights类型。它对外暴露urltransformsmeta属性,并提供verify()用于统一解析字符串/枚举/None 三种形式的weights参数。

pretrained=True迁移到weights=

新 API 与旧 API 的调用一一对应,迁移非常直接:

from torchvision.models import resnet50, ResNet50_Weights # 加载预训练权重: resnet50(weights=ResNet50_Weights.IMAGENET1K_V1) resnet50(weights="IMAGENET1K_V1") resnet50(pretrained=True) # 已弃用 resnet50(True) # 已弃用 # 不加载权重: resnet50(weights=None) resnet50() resnet50(pretrained=False) # 已弃用 resnet50(False) # 已弃用

注意:pretrained参数目前处于弃用状态,使用它会触发警告,并将在 v0.15 中移除(当前仓库版本号为 0.30.0a0,见 version.txt,仍保留该弃用兼容层)。

在源码层面,这一兼容行为由 torchvision/models/_utils.py 中的handle_legacy_interface装饰器实现:它一方面通过kwonly_to_pos_or_kw恢复位置参数支持并发出弃用警告,另一方面把pretrained=True映射为对应权重的DEFAULT值、把pretrained=False映射为weights=None,同时给出迁移提示。因此旧代码在 v0.15 移除前依然可用,但新代码应一律采用weights=形式。

推理预处理:weights.transforms()是唯一正确入口

使用预训练模型前必须对输入做预处理(按正确分辨率与插值方式缩放、应用推理变换、重标定数值范围等)。这个问题没有统一标准答案——它取决于模型如何被训练,可能在不同模型家族、不同变体甚至不同权重版本之间都有差异。用错预处理会导致精度下降甚至输出错误。

所有预训练模型推理变换所需的完整信息都记录在其权重文档中。为了简化推理,TorchVision 将必要的预处理变换直接捆绑进每个权重,通过weight.transforms()访问:

# 初始化权重变换 weights = ResNet50_Weights.DEFAULT preprocess = weights.transforms() # 应用到输入图像 img_transformed = preprocess(img)

这些变换的具体行为定义在 torchvision/transforms/_presets.py,按任务分为五类:

变换类核心处理流程
ImageClassificationresize(默认 256,插值默认双线性、默认 antialias)→ center_crop(如 224)→ 转 float 并归一化(默认 mean=(0.485, 0.456, 0.406),std=(0.229, 0.224, 0.225))
VideoClassification逐帧 resize → center_crop → 归一化(默认 mean/std 为 Kinetics 统计量)→ 输出排列为(..., C, T, H, W)
SemanticSegmentation可选 resize → 归一化(与分类相同的 ImageNet 均值/方差)
ObjectDetection仅转 float 并重标定到[0.0, 1.0](不裁剪、不归一化,因检测模型内部自带归一化与尺寸处理)
OpticalFlow转 float,并将两帧分别归一化到[-1.0, 1.0](mean=std=0.5)

这印证了文档中的关键提示:必须使用与所选权重配套的transforms(),而不是自行拼装一套通用的 ImageNet 预处理——例如检测模型与光流模型的预处理就与分类模型截然不同。

训练/评估模式切换

部分模型包含 BatchNorm 等训练与推理行为不同的模块。使用前必须调用model.eval()切换到评估模式(训练时用model.train()):

# 初始化模型 weights = ResNet50_Weights.DEFAULT model = resnet50(weights=weights) # 切换为评估模式 model.eval()

按名称列出与检索模型/权重(v0.14+)

自 v0.14 起,TorchVision 提供了按名称列出与获取模型和权重的统一机制,四个公开函数为:get_modelget_model_weightsget_weightlist_models(均定义于 torchvision/models/_api.py)。

list_models:枚举可用模型

# 列出全部已注册模型 all_models = list_models() # 仅列出 torchvision.models 主模块下的模型 classification_models = list_models(module=torchvision.models)

list_models还支持include/exclude通配符过滤,过滤规则使用 Unix shell 风格通配符(fnmatch),多个过滤条件取并集后执行排除。

get_model:按名实例化

# 初始化模型 m1 = get_model("mobilenet_v3_large", weights=None) m2 = get_model("quantized_mobilenet_v3_large", weights="DEFAULT")

get_model(name, **config)会先通过get_model_builder(name)在注册表BUILTIN_MODELS中查找构建函数,再把**config原样透传给构建函数(name不区分大小写)。

get_weight / get_model_weights:获取权重

# 按全名获取权重枚举成员,如 "MobileNet_V3_Large_QuantizedWeights.DEFAULT" weights = get_weight("MobileNet_V3_Large_QuantizedWeights.DEFAULT") assert weights == MobileNet_V3_Large_QuantizedWeights.DEFAULT # 获取某个模型对应的权重枚举类 weights_enum = get_model_weights("quantized_mobilenet_v3_large") assert weights_enum == MobileNet_V3_Large_QuantizedWeights # 也可以直接传入构建函数 weights_enum2 = get_model_weights(torchvision.models.quantization.mobilenet_v3_large) assert weights_enum == weights_enum2

底层实现上,get_model_weights通过反射读取模型构建函数签名中weights参数的类型注解(WeightsEnum子类)来定位权重枚举(见_get_enum_from_fn);get_weight则按"枚举类.成员名"的格式在torchvision.models及其子模块中查找并返回对应的枚举成员。

通过 PyTorch Hub 使用模型

大多数预训练模型可以不安装 TorchVision、直接通过 PyTorch Hub 访问(本仓库的 hubconf.py 即负责注册这些入口):

import torch # 方式一:weights 参数直接传字符串 model = torch.hub.load("pytorch/vision", "resnet50", weights="IMAGENET1K_V2") # 方式二:先取权重枚举再传入 weights = torch.hub.load( "pytorch/vision", "get_weight", weights="ResNet50_Weights.IMAGENET1K_V2", ) model = torch.hub.load("pytorch/vision", "resnet50", weights=weights)

也可以枚举某个模型在 Hub 上可用的全部权重:

import torch weight_enum = torch.hub.load("pytorch/vision", "get_model_weights", name="resnet50") print([weight for weight in weight_enum])

例外情况torchvision.models.detection中的检测模型必须安装 TorchVision 才能使用,因为它们依赖自定义 C++ 算子(见 torchvision/csrc/ops 下的nmsroi_alignroi_poolps_roi_align等内核实现)。

各任务实战:完整推理示例

图像分类

分类模型的类别标签可以从weights.meta["categories"]中取得。以下示例读取仓库测试图片,完成"加载模型 → 预处理 → 前向 → 输出 Top-1 类别"的完整链路:

from torchvision.io import decode_image from torchvision.models import resnet50, ResNet50_Weights img = decode_image("test/assets/encode_jpeg/grace_hopper_517x606.jpg") # 第 1 步:使用最佳可用权重初始化模型 weights = ResNet50_Weights.DEFAULT model = resnet50(weights=weights) model.eval() # 第 2 步:初始化推理变换 preprocess = weights.transforms() # 第 3 步:应用推理预处理变换 batch = preprocess(img).unsqueeze(0) # 第 4 步:前向并打印预测类别 prediction = model(batch).squeeze(0).softmax(0) class_id = prediction.argmax().item() score = prediction[class_id].item() category_name = weights.meta["categories"][class_id] print(f"{category_name}: {100 * score:.1f}%")

量化分类(INT8)

以下架构提供 INT8 量化模型(带或不带预训练权重):GoogLeNet、Inception、MobileNetV2、MobileNetV3、ResNet、ResNeXt、ShuffleNetV2,对应源码位于 torchvision/models/quantization。量化模型的构建函数多一个quantize=True开关:

from torchvision.io import decode_image from torchvision.models.quantization import resnet50, ResNet50_QuantizedWeights img = decode_image("test/assets/encode_jpeg/grace_hopper_517x606.jpg") # 第 1 步:量化模型 + 最佳可用权重 weights = ResNet50_QuantizedWeights.DEFAULT model = resnet50(weights=weights, quantize=True) model.eval() # 第 2 步:初始化推理变换 preprocess = weights.transforms() # 第 3 步:应用推理预处理变换 batch = preprocess(img).unsqueeze(0) # 第 4 步:前向并打印预测类别 prediction = model(batch).squeeze(0).softmax(0) class_id = prediction.argmax().item() score = prediction[class_id].item() category_name = weights.meta["categories"][class_id] print(f"{category_name}: {100 * score}%")

量化权重同样提供独立的权重枚举(如ResNet50_QuantizedWeights),其精度表格在文档中以单作物(single crops)方式在 ImageNet-1K 上评测。

语义分割

分割模型输出字典{"out": ...}out为各像素类别得分,配合weights.meta["categories"]可以提取指定类别的 softmax 掩码:

from torchvision.io.image import decode_image from torchvision.models.segmentation import fcn_resnet50, FCN_ResNet50_Weights from torchvision.transforms.functional import to_pil_image img = decode_image("gallery/assets/dog1.jpg") # 第 1 步:初始化模型 weights = FCN_ResNet50_Weights.DEFAULT model = fcn_resnet50(weights=weights) model.eval() # 第 2 步:初始化推理变换 preprocess = weights.transforms() # 第 3 步:应用推理预处理变换 batch = preprocess(img).unsqueeze(0) # 第 4 步:前向并可视化预测 prediction = model(batch)["out"] normalized_masks = prediction.softmax(dim=1) class_to_idx = {cls: idx for (idx, cls) in enumerate(weights.meta["categories"])} mask = normalized_masks[0, class_to_idx["dog"]] to_pil_image(mask).show()

语义分割模型(FCN、DeepLabV3、LRASPP)的精度在 COCO val2017 中与 Pascal VOC 重叠的 20 个类别子集上评测。

目标检测

检测模型的输入是Tensor[C, H, W]列表(可不同尺寸),输出为prediction["boxes"]prediction["labels"]prediction["scores"]。构建时可额外传入推理阈值参数,例如box_score_thresh=0.9

from torchvision.io.image import decode_image from torchvision.models.detection import fasterrcnn_resnet50_fpn_v2, FasterRCNN_ResNet50_FPN_V2_Weights from torchvision.utils import draw_bounding_boxes from torchvision.transforms.functional import to_pil_image img = decode_image("test/assets/encode_jpeg/grace_hopper_517x606.jpg") # 第 1 步:初始化模型,并提高置信度阈值 weights = FasterRCNN_ResNet50_FPN_V2_Weights.DEFAULT model = fasterrcnn_resnet50_fpn_v2(weights=weights, box_score_thresh=0.9) model.eval() # 第 2 步:初始化推理变换 preprocess = weights.transforms() # 第 3 步:注意这里构建的是单元素列表 batch = [preprocess(img)] # 第 4 步:前向并可视化 prediction = model(batch)[0] labels = [weights.meta["categories"][i] for i in prediction["labels"]] box = draw_bounding_boxes(img, boxes=prediction["boxes"], labels=labels, colors="red", width=4, font_size=30) im = to_pil_image(box.detach()) im.show()

可用检测模型包括 Faster R-CNN、FCOS、RetinaNet、SSD、SSDLite(见 torchvision/models/detection),Box mAP 在 COCO val2017 上报告。

实例分割与关键点检测

  • 实例分割:Mask R-CNN(mask_rcnn.py),输出中除 boxes/labels/scores 外还包含masks;Box 与 Mask mAP 均在 COCO val2017 上报告。
  • 人体关键点检测:Keypoint R-CNN(keypoint_rcnn.py),输出包含keypointskeypoint_scores;关键点名称(17 个人体关键点)通过weights.meta["keypoint_names"]取得,而非categories。Box 与 Keypoint mAP 在 COCO val2017 上报告。

视频分类

视频模型的输入为(T, C, H, W)的视频帧张量,read_video配合output_format="TCHW"可直接得到该布局:

from torchvision.io.video import read_video from torchvision.models.video import r3d_18, R3D_18_Weights vid, _, _ = read_video("test/assets/videos/v_SoccerJuggling_g23_c01.avi", output_format="TCHW") vid = vid[:32] # 可选:截取前 32 帧缩短时长 # 第 1 步:初始化模型 weights = R3D_18_Weights.DEFAULT model = r3d_18(weights=weights) model.eval() # 第 2 步:初始化推理变换 preprocess = weights.transforms() # 第 3 步:应用推理预处理变换 batch = preprocess(vid).unsqueeze(0) # 第 4 步:前向并打印预测类别 prediction = model(batch).squeeze(0).softmax(0) label = prediction.argmax().item() score = prediction[label].item() category_name = weights.meta["categories"][label] print(f"{category_name}: {100 * score}%")

可用视频模型包括 R3D/MC3/R2plus1D(video/resnet.py)、S3D、MViT、Swin3D,精度在 Kinetics-400 上以 16 帧片段(clip length 16)单作物方式报告。

光流估计

光流模型位于 torchvision/models/optical_flow/raft.py,提供raft_largeraft_small两个构建函数。其输入为连续两帧图像,预处理(OpticalFlow预设)会将两帧归一化到[-1.0, 1.0]

兼容性与使用注意事项汇总

  1. 权重版本语义DEFAULT是当前最佳权重的别名,指向的目标可能随版本更新而改变,长期依赖请显式指定具体版本(如IMAGENET1K_V2)。
  2. pretrained弃用:v0.15 将移除该参数,新代码请统一使用weights=;移除前使用会收到弃用警告。
  3. 预处理必须与权重匹配:务必使用weights.transforms(),不要套用通用预处理;不同权重版本间的预处理参数可能不同(如 ResNet-50 V1 与 V2 的resize_size差异)。
  4. 推理前调用model.eval():含 BatchNorm 等模块的模型在训练与推理模式下行为不同。
  5. 检测模型依赖 C++ 算子torchvision.models.detection系列必须安装 TorchVision 本体,无法仅通过 PyTorch Hub 零依赖使用。
  6. 序列化边界state_dict加载保证向后兼容;整模型或旧版 ScriptModule 加载不保证历史行为。
  7. 模型输入约定差异:分类/分割/视频/光流接受张量输入,检测与实例/关键点模型接受List[Tensor[C, H, W]]

掌握以上内容后,你可以在 TorchVision 中自由组合"任务 → 架构 → 权重版本 → 配套预处理"四个维度,写出既正确又灵活的推理代码;需要进一步了解某个具体架构的细节时,可直接查阅 docs/source/models 下对应的架构文档(如 resnet.rst、faster_rcnn.rst),或深入 torchvision/models 对应源码阅读实现。

【免费下载链接】visionDatasets, Transforms and Models specific to Computer Vision项目地址: https://gitcode.com/gh_mirrors/vi/vision

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询