Ultralytics 图像分类推理深度解析:ClassificationPredictor 的变换拆分、预处理与后处理全流程指南

【免费下载链接】ultralytics Ultralytics YOLO26, YOLO11, YOLOv8 — object detection, instance segmentation, semantic segmentation, image classification, pose estimation, object tracking 【免费下载链接】ultralytics 项目地址: https://gitcode.com/GitHub_Trending/ul/ultralytics

导读

ClassificationPredictor 是 Ultralytics 图像分类(classify)任务在推理阶段的专用预测器,负责把一张或多张图像送入分类模型,并产出带概率分布(probs)的 Results 对象。它是 YOLO 系列分类模型(如 yolo26n-cls、yolo11n-cls、yolov8n-cls,甚至 torchvision 的 resnet18)运行 model.predict()、CLI yolo classify predict 时真正干活的引擎。阅读本文你将掌握:该类的职责边界、宿主端与设备端图像变换的智能拆分原理、preprocess/postprocess 的执行细节,以及如何在 Python 与 CLI 中正确调用并解读分类预测结果。

本文主体对应仓库 API 文档 models/yolo/classify/predict.md,核心实现见 ultralytics/models/yolo/classify/predict.py。

ClassificationPredictor 的定位:分类任务的专用预测器

在 Ultralytics 的架构里,一个 YOLO 模型按任务类型拆分为四件套:模型定义(Model)、训练器(Trainer)、验证器(Validator)和预测器(Predictor)。分类任务的预测器就是 ClassificationPredictor。

它定义在 predict.py,并在同目录的 init.py 中被导出,与 ClassificationTrainer、ClassificationValidator 并列。类文档明确指出它的职责:

A class extending the BasePredictor class for prediction based on a classification model. This predictor handles the specific requirements of classification models, including preprocessing images and postprocessing predictions to generate classification results.

也就是说,它继承通用的 BasePredictor(位于 ultralytics/engine/predictor.py),仅覆盖与分类任务强相关的三件事:图像变换(setup_source)、宿主端预缩放(pre_transform)、预处理(preprocess) 与 结果后处理(postprocess)。从源码结构看,类的主要公开方法与属性为:

成员类型作用
argsdict / SimpleNamespace预测配置参数(继承自 BasePredictor)
__init__(cfg, overrides, _callbacks)构造方法初始化并强制 task="classify"
setup_source(source)方法装配分类专用的图像变换流水线
pre_transform(im)方法在宿主端(CPU)完成 Resize + CenterCrop,保持 uint8
preprocess(img)方法将输入转换为模型兼容张量并归一化
postprocess(preds, img, orig_imgs)方法把原始 logits 封装为带 probs 的 Results

在模型任务分发表中的位置

ClassificationPredictor 并不是被直接硬编码调用的,而是通过任务分发表在运行时被选中。在 ultralytics/models/yolo/model.py 中,classify 键下显式登记了:

"classify": {
    "model": ClassificationModel,
    "trainer": yolo.classify.ClassificationTrainer,
    "validator": yolo.classify.ClassificationValidator,
    "predictor": yolo.classify.ClassificationPredictor,
},

因此当你执行 YOLO("yolo26n-cls.pt") 或 model=...-cls.pt 的推理时,框架会依据模型架构自动路由到本文的主角。分类模型的配置骨架可参考 ultralytics/cfg/models/26/yolo26-cls.yaml,其头部为 Classify 模块、默认 nc: 1000;仓库同目录下还提供了 v8、v11、v12 等多个系列的分类 YAML(如 yolov8-cls.yaml、yolo11-cls.yaml),以及基于 ResNet 骨干的分类配置 yolo11-cls-resnet18.yaml、yolov8-cls-resnet50.yaml 等。

构造与配置:任务被强制固定为 classify

ClassificationPredictor.__init__ 位于 predict.py,实现非常克制:

def __init__(self, cfg=DEFAULT_CFG, overrides=None, _callbacks: dict | None = None):
    super().__init__(cfg, overrides, _callbacks)
    self.args.task = "classify"

它把参数透传给父类 BasePredictor.__init__(位于 predictor.py),后者通过 get_cfg(cfg, overrides) 合并默认配置与用户覆盖项,并完成 save_dir 推导、conf 兜底(默认 0.25)、回调注册等通用初始化。随后唯一的关键动作是把任务强制写死为 "classify"——即使调用方传入了其他任务配置也不会生效,保证下游 Results 的组装走分类语义。

DEFAULT_CFG 来自 ultralytics/cfg/default.yaml。对分类预测而言,下面这些参数最常用到:

  • source:推理来源(图片路径 / 目录 / URL / 视频 / 摄像头编号等);
  • imgsz:输入尺寸,默认 640,分类推理常用 224(当模型自带预处理时会被覆盖,见下文);
  • device:cpu / 0 / mps 等运行设备;
  • save:是否保存推理结果图;save_txt 是否落盘标签文本;
  • verbose:是否打印日志;visualize:是否输出分类激活热力图;
  • augment:是否启用测试时增强;embed:是否返回指定层的特征嵌入。

也能推理 torchvision 模型

类文档的 Notes 中特别说明:torchvision 分类模型也可以直接传给 model 参数,例如 model="resnet18"。这让 ClassificationPredictor 成为一个通用分类推理入口——只要模型是标准的 (N, C, H, W) → (N, nc) 全连接式分类头即可复用同一套预处理与后处理逻辑。

setup_source:装配智能的分类变换流水线

setup_source(predict.py)是分类预测器区别于检测/分割预测器的核心,也是整个类中逻辑最精巧的部分。它先调用父类 BasePredictor.setup_source(predictor.py)完成 imgsz 校验与数据源加载,随后专门处理变换流水线:

import torchvision.transforms as T

super().setup_source(source)
transforms = getattr(self.model.model, "transforms", None)
size = getattr(transforms.transforms[0], "size", max(self.imgsz)) if transforms is not None else None
self.transforms = (
    transforms if size == max(self.imgsz) and self.model.format == "pt" else classify_transforms(self.imgsz)
)

决策逻辑可以概括为一条规则:

  • 若模型对象上携带了训练阶段保存的 transforms(PIL Compose),且其首个变换的 size 与用户指定的 imgsz 一致,且模型格式为 PyTorch(format == "pt"),则直接复用训练时校验过的预处理,保证推理数据分布与训练一致;
  • 否则回退到由 ultralytics/data/augment.py 的 classify_transforms(self.imgsz) 现场构造标准推理变换(默认最短边缩放到 224,再做中心裁剪)。

源码注释点明:YAML 临时构建的模型与旧版 checkpoint 常常缺失 transforms 属性,此时就会走回退分支,这是兼容性设计。

host 变换与 device 变换的拆分

接下来是最值得注意的部分——把 Compose 拆成两段(predict.py):

tfl = getattr(self.transforms, "transforms", ())
split = (
    type(self.transforms) is T.Compose
    and tuple(map(type, tfl)) == (T.Resize, T.CenterCrop, T.ToTensor, T.Normalize)
    and getattr(self.model, "channels", 3) == 3
)
self.host_transforms = T.Compose(tfl[:2]) if split else None
self.device_transform = tfl[-1] if split else None

拆分的前提非常苛刻:变换必须是 torchvision.transforms.Compose,且元素类型恰好是 (T.Resize, T.CenterCrop, T.ToTensor, T.Normalize) 这个标准四件套,同时模型输入通道为 3。满足时:

  • host_transforms = Resize + CenterCrop,在宿主端对 uint8 图像执行(CPU 友好、不吃显存);
  • device_transform = Normalize,交给设备端在张量上完成归一化。

这样做的工程收益很直接:几何变换(缩放、裁剪)天然适合 CPU + PIL 图像处理,而逐像素归一化在 GPU 张量上批量执行更高效,避免整幅 uint8 图像在 ToTensor 后于设备端重复搬运。若不满足拆分条件,则整条 Compose 保留在 self.transforms,在预处理里整体执行。

classify_transforms 的标准变换构成

回退分支使用的 classify_transforms(augment.py)组装了如下流水线:

tfl = [
    T.Resize(resize, interpolation=getattr(T.InterpolationMode, interpolation)),
    T.CenterCrop(size),
    T.ToTensor(),
    T.Normalize(mean=torch.tensor(mean), std=torch.tensor(std)),
]
return T.Compose(tfl)

其参数与语义为:

参数默认值说明
size224若为 int,表示最短边模式(保持宽高比的等比例缩放后再中心裁剪);若为 (h, w) 元组则直接精确缩放
mean / stdDEFAULT_MEAN / DEFAULT_STD归一化所用的逐通道统计量,对应 ImageNet 预训练约定
interpolation"BILINEAR"支持 NEAREST / BILINEAR / BICUBIC
crop_fractionNone已废弃参数,仅触发 deprecation 告警

注意 augment.py 中还有一个细节:正方形尺寸走 int 最短边模式以保留宽高比,非正方形才按精确 (h, w) 缩放。训练阶段的增强 classify_augmentations 则完全是另一套(含 RandomResizedCrop、翻转、HSV 抖动、randaugment、随机擦除等),用于 train.py,与推理无关。

pre_transform 与 preprocess:宿主端到设备端的两级处理

pre_transform:保持 uint8 的宿主端预处理

def pre_transform(self, im: list[np.ndarray]) -> list[np.ndarray]:
    return [np.array(self.host_transforms(Image.fromarray(x))) for x in im]

在拆分模式下,pre_transform 把每张 BGR 的 uint8 numpy 图转成 PIL 图像,套用 host_transforms(Resize + CenterCrop)后再转回 numpy。注释点明设计意图:"leaving uint8 BGR for the device-side conversion",即在这个阶段刻意不转浮点、不转 RGB,几何变换以低成本方式在 CPU 完成。

preprocess:两条执行路径

preprocess(predict.py)根据前面是否成功拆分,分两条路径执行:

def preprocess(self, img):
    if self.device_transform is None and not isinstance(img, torch.Tensor):
        img = torch.stack([self.transforms(Image.fromarray(cv2.cvtColor(x, cv2.COLOR_BGR2RGB))) for x in img], 0)
        img = img.to(self.model.device)
        return img.half() if self.model.fp16 else img.float()
    is_tensor = isinstance(img, torch.Tensor)
    img = super().preprocess(img)
    return img if is_tensor else self.device_transform(img)
  • 路径 A(未拆分):对每张图显式做 BGR→RGB 再走完整 Compose(Resize、CenterCrop、ToTensor、Normalize),torch.stack 成 batch,搬到模型设备并按 fp16 决策转半精度。
  • 路径 B(已拆分):先调用父类 BasePredictor.preprocess。父类实现(predictor.py)会把 uint8 numpy 栈转换为张量、完成 BHWC→BCHW 重排与 BGR→RGB 通道翻转、缩放到 0~1,并统一为 fp16/fp32。最后在设备端补上 device_transform(归一化)。

若 img 本身就是张量(例如调用方传入已经归一化好的 torch.Tensor),两条路径都会直接透传,避免重复归一化——这对应父类文档中 "already RGB and normalized to 0.0-1.0" 的输入约定。

postprocess:把 logits 变成带 probs 的 Results

推理完成后,postprocess(predict.py)把原始输出整理成结果对象:

def postprocess(self, preds, img, orig_imgs):
    if not isinstance(orig_imgs, list):  # 输入是 torch.Tensor
        orig_imgs = ops.convert_torch2numpy_batch(orig_imgs)[..., ::-1]
    preds = preds[0] if isinstance(preds, (list, tuple)) else preds
    return [
        Results(orig_img, path=img_path, names=self.model.names, probs=pred)
        for pred, orig_img, img_path in zip(preds, orig_imgs, self.batch[0])
    ]

三个关键动作:

  1. 原图还原:若原始输入是张量而非列表,则通过 ops.convert_torch2numpy_batch 转回 numpy 并做 [..., ::-1] 通道翻转(RGB→BGR),保证返回的 orig_img 是用户熟悉的 OpenCV BGR 图像;
  2. 预测解包:兼容模型可能返回 (preds, extra) 元组的情况,统一取出第一个元素作为 logits;
  3. 逐张封装:把每个样本的原始图、来源路径、类别名表 self.model.names 与预测概率 probs=pred 一起构造 Results 对象。

注意:分类结果里没有 boxes,核心载荷是 probs。它对应 results.py 的 Probs 类(继承 BaseTensor),提供便捷属性:top1(最高概率类别索引)、top1conf(最高置信度)、top5(前五类别索引列表)、top5conf(对应置信度)。因此,拿到 Results 后访问分类答案的标准姿势是:

for result in results:
    top1_id = result.probs.top1
    top1_conf = result.probs.top1conf
    top1_name = result.names[top1_id]

全流程串联:推理管线如何把它们组织起来

分类推理并不是孤立调用上述四个方法,而是由父类 BasePredictor.stream_inference(predictor.py)统一驱动,每个 batch 依次执行:

  1. setup_source → 装载数据源并(在分类场景)装配变换;
  2. preprocess(im0s) → 得到模型输入张量(并完成首次 warmup);
  3. inference(im) → 调用模型前向,得到 logits;
  4. postprocess(preds, im, im0s) → 产出 Results 列表;
  5. 依据 save / show 等参数调用 write_results,内部会 result.plot(...) 生成标注图(分类图会叠加 Top 类别标签与置信度文字)并落盘或显示。

整个过程被 threading.Lock 包裹以保证线程安全,三个关键环节(preprocess / inference / postprocess)分别计时写入 self.speed。另外,当 args.visualize=True 且模型是普通 PyTorch 基座模型时,inference 会改走 plotting.py 的 class_activation_map 输出分类激活热力图,用于可视化模型"看"图像的哪些区域。

实战:Python 与 CLI 两种调用方式

方式一:通过 YOLO 高层 API(推荐)

日常开发并不直接实例化 ClassificationPredictor,而是通过统一入口 YOLO 自动路由(任务分发表见前文 model.py):

from ultralytics import YOLO

model = YOLO("yolo26n-cls.pt")          # 或任意 -cls 权重,如 yolo11n-cls.pt / yolov8n-cls.pt
results = model.predict(source="https://ultralytics.com/images/bus.jpg", save=True)

for result in results:
    print(result.names[result.probs.top1], float(result.probs.top1conf))
    print(result.probs.top5, result.probs.top5conf)

方式二:直接使用 ClassificationPredictor

类文档自带的示例展示了底层直接驱动的写法(该用法适合需要深度定制回调、精细控制中间步骤的场景):

from ultralytics.utils import ASSETS
from ultralytics.models.yolo.classify import ClassificationPredictor

args = dict(model="yolo26n-cls.pt", source=ASSETS)
predictor = ClassificationPredictor(overrides=args)
predictor.predict_cli()

其中 predict_cli(继承自 predictor.py)会以生成器方式消费全部推理结果且不驻留内存,因此即便面对长视频或大目录也不会导致结果对象无限累积。若希望逐条拿到结果,应调用带 stream=True 的 __call__ 或直接迭代 stream_inference。

方式三:CLI

# 对图片/目录/视频推理,等价于 yolo classify predict
yolo predict model=yolo26n-cls.pt source='path/to/images' imgsz=224 save=True
yolo classify predict model=yolov8n-cls.pt source='https://ultralytics.com/images/bus.jpg'

CLI 内部会走 ClassificationPredictor.predict_cli 完成同样的流水线。对分类推理,imgsz 会被 setup_source 与模型自带变换做对齐校验:若 checkpoint 训练时保存了预处理且尺寸与 imgsz 一致则直接复用,否则以 classify_transforms(imgsz) 为准,因此把 imgsz 设为分类模型惯用的 224 是安全且推荐的做法。

兼容性提示:torchvision 权重(如 model="resnet18")同样可以在 YOLO(...) 与 ClassificationPredictor(...) 两条路径下直接推理,因为它们都落到同一个预测器上。若加载的权重缺失 transforms 属性(如 YAML 临时构建的模型),框架会自动回退到标准的 Resize + CenterCrop + 归一化流水线,无需手工干预。

小结:记住这三条主线

  1. 职责单一:ClassificationPredictor 只覆盖分类特有的变换装配与结果封装,其余加载、调度、保存、回调全部复用 BasePredictor,它通过 model.py 的任务分发表被自动选中;
  2. 两级流水线:setup_source 会优先复用 checkpoint 自带预处理,并把 (Resize, CenterCrop) | (ToTensor, Normalize) 拆成 host/device 两段以求效率;回退分支由 classify_transforms 兜底;
  3. 结果即 probs:postprocess 产出的每个 Results 携带 Probs,通过 top1 / top1conf / top5 / top5conf 即可直接读取分类答案,配合 names 映射即可得到人类可读的类别与置信度。

深入阅读时,建议结合类文档页 models/yolo/classify/predict.md 对照源码逐行学习,并参考同目录的 train.py 与 val.py 理解"训练增强 vs 推理变换"在设计上的分野。

【免费下载链接】ultralytics Ultralytics YOLO26, YOLO11, YOLOv8 — object detection, instance segmentation, semantic segmentation, image classification, pose estimation, object tracking 【免费下载链接】ultralytics 项目地址: https://gitcode.com/GitHub_Trending/ul/ultralytics

Logo

北京人形旗下天工造物具身智能开源社区,聚焦具身天工与慧思开物两大平台

更多推荐