Ultralytics 图像分类推理深度解析:ClassificationPredictor 的变换拆分、预处理与后处理全流程指南
Ultralytics 图像分类推理深度解析:ClassificationPredictor 的变换拆分、预处理与后处理全流程指南
导读
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)。从源码结构看,类的主要公开方法与属性为:
| 成员 | 类型 | 作用 |
|---|---|---|
args | dict / 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(PILCompose),且其首个变换的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)
其参数与语义为:
| 参数 | 默认值 | 说明 |
|---|---|---|
size | 224 | 若为 int,表示最短边模式(保持宽高比的等比例缩放后再中心裁剪);若为 (h, w) 元组则直接精确缩放 |
mean / std | DEFAULT_MEAN / DEFAULT_STD | 归一化所用的逐通道统计量,对应 ImageNet 预训练约定 |
interpolation | "BILINEAR" | 支持 NEAREST / BILINEAR / BICUBIC |
crop_fraction | None | 已废弃参数,仅触发 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])
]
三个关键动作:
- 原图还原:若原始输入是张量而非列表,则通过
ops.convert_torch2numpy_batch转回 numpy 并做[..., ::-1]通道翻转(RGB→BGR),保证返回的orig_img是用户熟悉的 OpenCV BGR 图像; - 预测解包:兼容模型可能返回
(preds, extra)元组的情况,统一取出第一个元素作为 logits; - 逐张封装:把每个样本的原始图、来源路径、类别名表
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 依次执行:
setup_source→ 装载数据源并(在分类场景)装配变换;preprocess(im0s)→ 得到模型输入张量(并完成首次warmup);inference(im)→ 调用模型前向,得到 logits;postprocess(preds, im, im0s)→ 产出Results列表;- 依据
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 + 归一化流水线,无需手工干预。
小结:记住这三条主线
- 职责单一:
ClassificationPredictor只覆盖分类特有的变换装配与结果封装,其余加载、调度、保存、回调全部复用 BasePredictor,它通过 model.py 的任务分发表被自动选中; - 两级流水线:
setup_source会优先复用 checkpoint 自带预处理,并把(Resize, CenterCrop) | (ToTensor, Normalize)拆成 host/device 两段以求效率;回退分支由classify_transforms兜底; - 结果即 probs:
postprocess产出的每个 Results 携带Probs,通过top1/top1conf/top5/top5conf即可直接读取分类答案,配合names映射即可得到人类可读的类别与置信度。
深入阅读时,建议结合类文档页 models/yolo/classify/predict.md 对照源码逐行学习,并参考同目录的 train.py 与 val.py 理解"训练增强 vs 推理变换"在设计上的分野。
更多推荐
所有评论(0)