尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

深入解析 YOLOv10 图像分类推理:ClassificationPredictor 源码剖析与实战指南

深入解析 YOLOv10 图像分类推理:ClassificationPredictor 源码剖析与实战指南 深入解析 YOLOv10 图像分类推理ClassificationPredictor 源码剖析与实战指南【免费下载链接】yolov10YOLOv10: Real-Time End-to-End Object Detection [NeurIPS 2024]项目地址: https://gitcode.com/GitHub_Trending/yo/yolov10本文以 ultralytics/models/yolo/classify/predict.py 及其 API 参考文档 docs/en/reference/models/yolo/classify/predict.md 为线索系统讲解本仓库YOLOv10 Ultralytics 套件中图像分类预测组件ClassificationPredictor的实现原理、调用链与实战用法。读完本文你将掌握分类预测器与通用预测器的继承关系、预处理/后处理的逐行源码语义、Results/Probs结果对象的访问方式以及通过 Python 与 CLI 两种方式完成图像分类推理的完整流程。一、ClassificationPredictor分类任务专属的预测器在本仓库的 Ultralytics 架构中不同视觉任务检测、分割、姿态、OBB、分类各自拥有一套 trainer / validator / predictor 实现并通过task_map统一注册到 ultralytics/models/yolo/model.pyclassify: { model: ClassificationModel, trainer: yolo.classify.ClassificationTrainer, validator: yolo.classify.ClassificationValidator, predictor: yolo.classify.ClassificationPredictor, },ClassificationPredictor正是分类任务的推理执行者定义于 ultralytics/models/yolo/classify/predict.py并统一从 ultralytics/models/yolo/classify/init.py 导出from ultralytics.models.yolo.classify.predict import ClassificationPredictor from ultralytics.models.yolo.classify.train import ClassificationTrainer from ultralytics.models.yolo.classify.val import ClassificationValidator __all__ ClassificationPredictor, ClassificationTrainer, ClassificationValidator它的类图关系如下继承自 ultralytics/engine/predictor.py 中的BasePredictor复用了通用预测管线的全部骨架数据源加载、模型初始化、逐 batch 推理、结果保存与可视化。仅在__init__中把任务标记为classify并重写preprocess输入预处理与postprocess预测后处理两个钩子方法体现了典型的模板方法 子类定制设计。二、类的初始化任务标记与兼容性设计def __init__(self, cfgDEFAULT_CFG, overridesNone, _callbacksNone): Initializes ClassificationPredictor setting the task to classify. super().__init__(cfg, overrides, _callbacks) self.args.task classify self._legacy_transform_name ultralytics.yolo.data.augment.ToTensor初始化要点参数cfg默认为DEFAULT_CFG由 ultralytics/cfg/default.yaml 解析而来overrides是用户传入的参数覆盖字典_callbacks用于注入自定义回调。super().__init__会调用get_cfg合并配置并完成save_dir、默认置信度conf0.25见 ultralytics/engine/predictor.py等初始化。self.args.task classify是分类预测器的关键状态标记BasePredictor.setup_source会依据args.task classify决定是否加载分类专用的数据变换见下文第四节。self._legacy_transform_name记录旧版变换类的完整限定名用于在预处理时兼容旧权重/旧数据管线见第三节。与兄弟类的配套关系分类任务并非孤立存在其训练与验证由同目录下的 train.py 与 val.py 完成ClassificationTrainer.__init__会在无imgsz覆盖时默认置为224因为分类模型的标准输入尺寸是 224。ClassificationValidator.get_desc输出classes / top1_acc / top5_acc三列指标update_metrics中取min(nc, 5)个 top-k 预测用于评估 top-1 / top-5 准确率。三、preprocess把任意输入转换为模型可用的张量def preprocess(self, img): Converts input image to model-compatible data type. if not isinstance(img, torch.Tensor): is_legacy_transform any( self._legacy_transform_name in str(transform) for transform in self.transforms.transforms ) if is_legacy_transform: # to handle legacy transforms img torch.stack([self.transforms(im) for im in img], dim0) else: img torch.stack( [self.transforms(Image.fromarray(cv2.cvtColor(im, cv2.COLOR_BGR2RGB))) for im in img], dim0 ) img (img if isinstance(img, torch.Tensor) else torch.from_numpy(img)).to(self.model.device) return img.half() if self.model.fp16 else img.float() # uint8 to fp16/32逐行语义拆解张量直通如果输入已经是torch.Tensor例如用户直接传入 CUDA 张量跳过变换环节。新旧变换兼容通过检查self.transforms.transforms中是否含有_legacy_transform_nameultralytics.yolo.data.augment.ToTensor区分旧版输入为 numpy 数组直接套用变换与新版输入经cv2.cvtColor(im, cv2.COLOR_BGR2RGB)转成 RGB再用 PILImage.fromarray包装后套用 torchvision 变换。这一分支保证旧权重加载后推理不出错。类型与设备迁移统一to(self.model.device)并依据self.model.fp16选择半精度half()或单精度float()。注意这里不除以 255——归一化Normalize已由变换管线内部的 mean/std 完成这与检测任务在BasePredictor.preprocess中除以 255 的实现不同。与检测预测器对比BasePredictor.preprocessultralytics/engine/predictor.py执行 LetterBox 缩放、BGR→RGB、BHWC→BCHW 转置并除以 255而分类预测器使用中心裁剪CenterCrop而非 LetterBox分类任务不需要保边框坐标保持宽高比裁剪即可。四、数据变换管线classify_transforms 的组成ClassificationPredictor.preprocess中使用的self.transforms来自BasePredictor.setup_sourceultralytics/engine/predictor.pyself.transforms ( getattr( self.model.model, transforms, classify_transforms(self.imgsz[0], crop_fractionself.args.crop_fraction), ) if self.args.task classify else None )即优先使用模型自带的transforms属性ClassificationTrainer.get_dataloader会把验证集变换挂到模型上见 ultralytics/models/yolo/classify/train.py否则使用 ultralytics/data/augment.py 中的classify_transforms兜底def classify_transforms(size224, meanDEFAULT_MEAN, stdDEFAULT_STD, interpolationT.InterpolationMode.BILINEAR, crop_fraction1.0): ... scale_size math.floor(size / crop_fraction) tfl [T.Resize(scale_size[0], interpolationinterpolation)] tfl [T.CenterCrop(size)] tfl [T.ToTensor(), T.Normalize(mean..., std...)]关键参数size目标尺寸分类默认 224也可由imgsz覆盖。crop_fraction先按size / crop_fraction缩放再做中心裁剪。取 1.0 时等效直接缩放 中心裁剪取小于 1 的值可保留更多上下文信息常用于推理时提升精度。对应配置项见 ultralytics/cfg/default.yamlcrop_fraction: 1.0。训练阶段则使用带增强的classify_augmentationsultralytics/data/augment.py支持 hflip/vflip、auto_augmentrandaugment/autoaugment/augmix与erasing随机擦除等测试用例见 tests/test_python.py。五、postprocess封装为 Results 对象def postprocess(self, preds, img, orig_imgs): Post-processes predictions to return Results objects. if not isinstance(orig_imgs, list): # input images are a torch.Tensor, not a list orig_imgs ops.convert_torch2numpy_batch(orig_imgs) results [] for i, pred in enumerate(preds): orig_img orig_imgs[i] img_path self.batch[0][i] results.append(Results(orig_img, pathimg_path, namesself.model.names, probspred)) return results后处理逻辑若原始图像是 Tensor batch先用ops.convert_torch2numpy_batch转回 numpy便于可视化与保存。逐张构造 ultralytics/engine/results.py 中的Results对象把分类的原始概率向量pred挂到probs字段上同时记录原始图像、来源路径与类别名self.model.names。分类任务没有边界框因此Results.boxes为 None这与检测/分割任务的 postprocess 有本质区别。读取结果Probs 类的便捷属性分类预测结果通过Results.probs暴露类型为 ultralytics/engine/results.py 中的Probs继承BaseTensortop1得分最高类别的索引int(self.data.argmax())。top5得分前五类别的索引列表。top1conf/top5conf对应位置的置信度张量。这些属性带lru_cache多次访问不会重复计算。典型用法for r in results: print(r.probs.top1, r.probs.top1conf, r.probs.top5) # 索引与置信度 print(r.names[r.probs.top1]) # 类别名六、推理管线从数据源到输出的完整链路ClassificationPredictor复用的BasePredictor.stream_inferenceultralytics/engine/predictor.py是其核心执行循环流程为setup_model通过AutoBackend加载权重支持.pt、ONNX、TensorRT 等多种格式选择设备并fuseTrue融合层、eval()切换评估模式。setup_source解析输入源——支持图像、视频、目录、glob 通配符、URL 与摄像头/网络流见 ultralytics/engine/predictor.py 的模块注释。对分类任务同时安装classify_transforms。预热模型warmup后逐 batch 依次执行preprocess→inference→postprocess三个环节分别用ops.Profile计时最终在Results.speed中给出每张图的 preprocess / inference / postprocess 耗时。依据verbose、save、show、save_txt等参数决定打印日志、保存图像/视频或实时显示write_results。__call__提供streamTrue的生成器模式适合视频流/长序列避免结果在内存中无限累积streamFalse时一次性返回Results列表。CLI 入口predict_cli则直接消费生成器而不累积。七、实战Python 与 CLI 两种调用方式方式一通过 YOLO 高层 API推荐分类模型文件名带-cls后缀例如yolov8n-cls.pt本仓库支持 ImageNet 系列数据集配置见 ultralytics/cfg/datasets/ImageNet.yaml 与imagenet10等轻量数据集from ultralytics import YOLO model YOLO(yolov8n-cls.pt) # 官方预训练分类模型 # model YOLO(path/to/best.pt) # 自定义微调模型 results model(path/to/image.jpg) # 单张图片 # results model.predict(sourcedir/, saveTrue) # 批量目录 for r in results: print(r.probs.top1conf, model.names[r.probs.top1])CLI 等价命令yolo classify predict modelyolov8n-cls.pt sourcebus.jpg saveTrue任务关键字classify与modepredict等价source支持图片、目录、视频与网络流。方式二直接实例化 ClassificationPredictor源码级用法参照类文档给出的官方示例见 ultralytics/models/yolo/classify/predict.pyfrom ultralytics.utils import ASSETS from ultralytics.models.yolo.classify import ClassificationPredictor args dict(modelyolov8n-cls.pt, sourceASSETS) predictor ClassificationPredictor(overridesargs) predictor.predict_cli()sourceASSETS指向仓库内置示例素材目录ultralytics/assets/可直接用bus.jpg验证分类结果。overrides可覆盖任意默认配置例如dict(modelyolov8n-cls.pt, sourcebus.jpg, imgsz224, crop_fraction1.0, conf0.25, saveFalse, showFalse)。单元测试 tests/test_engine.py 中即用classify.ClassificationPredictor(overrides{imgsz: [64, 64]})加on_predict_start回调验证该路径可作为自行扩展回调的参考。兼容 torchvision 模型ClassificationPredictor类注释特别说明model参数同样接受 Torchvision 分类模型例如modelresnet18训练侧在 ultralytics/models/yolo/classify/train.py 的setup_model中通过torchvision.models.__dict__加载。这意味着你可以用同一套推理管线驱动第三方分类骨干网络。八、常用推理参数速查以下参数均在 ultralytics/cfg/default.yaml 中定义可用overrides或 CLIkeyvalue覆盖参数默认值分类推理中的含义imgsz640分类默认 224输入尺寸分类建议 224 或训练时所用尺寸crop_fraction1.0推理时先缩放再中心裁剪的比例小于 1 保留更多上下文conf0.25predict置信度阈值分类任务主要用于展示过滤save/showTrue / False是否保存结果图 / 弹窗显示verboseTrue是否逐 batch 打印推理日志含各环节耗时halfFalse是否使用 FP16 半精度推理GPU 上加速deviceNone推理设备如0、cpustreamFalse生成器模式长视频/流媒体建议开启九、小结ClassificationPredictor是本仓库图像分类能力的推理执行核心它继承BasePredictor复用整套推理管线仅通过标记任务 重写 pre/postprocess实现分类特化预处理采用 torchvision 风格变换缩放 中心裁剪 归一化并兼容旧版变换后处理将概率向量封装为Results.probs配合Probs.top1/top5等属性即可快速读取 Top-k 预测。配合同目录的ClassificationTrainer、ClassificationValidator与 docs/en/tasks/classify.md 中的训练、验证、导出流程即可在 YOLOv10 仓库中完成从数据准备到端侧部署的全链路图像分类开发。【免费下载链接】yolov10YOLOv10: Real-Time End-to-End Object Detection [NeurIPS 2024]项目地址: https://gitcode.com/GitHub_Trending/yo/yolov10创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表