YOLOv10 / YOLOv8 PosePredictor 源码解析:关键点姿态预测的推理链路与后处理实现
2026/9/15 17:15:14 网站建设 项目流程

YOLOv10 / YOLOv8 PosePredictor 源码解析:关键点姿态预测的推理链路与后处理实现

【免费下载链接】yolov10YOLOv10: Real-Time End-to-End Object Detection [NeurIPS 2024]项目地址: https://gitcode.com/GitHub_Trending/yo/yolov10

导读

本文围绕当前仓库 docs/en/reference/models/yolo/pose/predict.md 所引用的核心类PosePredictor展开,它是 Ultralytics YOLO 系列(包括本仓库的 YOLOv8 与 YOLOv10)在姿态估计任务(Pose Estimation)上的推理预测器。通过阅读本文,你将掌握PosePredictor的初始化机制、与检测预测器的继承关系、关键点后处理(NMS、坐标缩放、关键点恢复)的完整实现,以及如何通过 Python API 和 CLI 两种方式驱动姿态预测,并深入理解底层调用链与Results.keypoints的输出结构。

PosePredictor 在 YOLO 推理体系中的定位

在 Ultralytics YOLO 的架构中,姿态估计(Pose)与检测(Detect)、分割(Segment)、旋转框(OBB)、分类(Classify)并列为五大任务之一。PosePredictor是整个姿态推理链路的最终执行者,定义于 ultralytics/models/yolo/pose/predict.py。

它的类继承关系为:

BasePredictor(ultralytics/engine/predictor.py) ↑ 继承 DetectionPredictor(ultralytics/models/yolo/detect/predict.py) ↑ 继承 PosePredictor(ultralytics/models/yolo/pose/predict.py)

其中 BasePredictor 负责通用的推理框架:数据源加载、预处理、模型加载(AutoBackend)、流式推理调度、结果保存与可视化;DetectionPredictor 在其基础上实现了目标检测的 NMS 与检测框缩放;而PosePredictor则进一步在检测结果之上叠加关键点(keypoints)的后处理,从而输出带有人体/物体关键点坐标的姿态结果。

从任务注册机制看,ultralytics/models/yolo/model.py 中的task_mappose任务映射为:

"pose": { "model": PoseModel, "trainer": yolo.pose.PoseTrainer, "validator": yolo.pose.PoseValidator, "predictor": yolo.pose.PosePredictor, }

这意味着当用户加载yolov8n-pose.pt等姿态权重并调用model.predict()model(source=...)时,框架会自动选择PosePredictor完成推理。ultralytics/models/yolo/pose/init.py 将PoseTrainerPoseValidatorPosePredictor三个类统一对外导出。

初始化与 task 设置:mps 设备警告

PosePredictor的构造函数通过super().__init__(cfg, overrides, _callbacks)调用父类完成通用配置解析(get_cfg合并默认配置与覆盖项),随后执行两项姿态特有的初始化:

def __init__(self, cfg=DEFAULT_CFG, overrides=None, _callbacks=None): super().__init__(cfg, overrides, _callbacks) self.args.task = "pose" if isinstance(self.args.device, str) and self.args.device.lower() == "mps": LOGGER.warning( "WARNING ⚠️ Apple MPS known Pose bug. Recommend 'device=cpu' for Pose models. " )
  • self.args.task = "pose":将任务标识强制设为pose,后续的数据预处理(如pre_transform)、结果写入等逻辑都会依据该 task 选择姿态路径;
  • MPS 设备警告:当用户显式传入device=mps(Apple 芯片 GPU)时,会输出一条已知 Bug 的警告,并建议姿态模型改用device=cpu。这是仓库源码中明确记载的兼容性提示,属于姿态任务特有的注意事项,值得在部署到 macOS 时格外留意。

postprocess:姿态后处理的完整链路

PosePredictor的核心贡献集中在postprocess方法,它接收原始网络输出preds、经过 letterbox 预处理的张量img以及原始图像orig_imgs,返回一个Results对象列表。整个过程分为四步:

第一步:非极大值抑制(NMS)

preds = ops.non_max_suppression( preds, self.args.conf, self.args.iou, agnostic=self.args.agnostic_nms, max_det=self.args.max_det, classes=self.args.classes, nc=len(self.model.names), )

NMS 由 ultralytics/utils/ops.py 中的non_max_suppression函数实现,其默认阈值与 BasePredictor 保持一致:conf默认为 0.25、iou默认为 0.45、max_det默认为 300。与检测任务的关键区别是传入了nc=len(self.model.names)显式指定类别数,这是因为姿态输出张量的列布局为[x, y, w, h, class_conf, ...keypoints],NMS 需要精确切分类别置信度与关键点通道。

第二步:输入张量转 NumPy 批次

if not isinstance(orig_imgs, list): # input images are a torch.Tensor, not a list orig_imgs = ops.convert_torch2numpy_batch(orig_imgs)

当推理源为torch.Tensor(例如自定义预处理后直接喂入的张量)时,需先通过convert_torch2numpy_batch将其转换为 NumPy 批次,以保证后续按原始图像尺寸进行坐标还原。

第三步:检测框与关键点坐标缩放

pred[:, :4] = ops.scale_boxes(img.shape[2:], pred[:, :4], orig_img.shape).round() pred_kpts = pred[:, 6:].view(len(pred), *self.model.kpt_shape) if len(pred) else pred[:, 6:] pred_kpts = ops.scale_coords(img.shape[2:], pred_kpts, orig_img.shape)

这里体现了姿态后处理与检测后处理最核心的差异:

  • pred[:, :4]是检测框(xyxy),通过scale_boxes从 letterbox 后的尺寸映射回原图尺寸并四舍五入;
  • pred[:, 6:]是关键点通道。因为 NMS 输出布局为[x1, y1, x2, y2, score, cls, kpt_x1, kpt_y1, kpt_conf1, ...],所以从第 6 列开始才是关键点数据;
  • 关键点通过view(len(pred), *self.model.kpt_shape)重塑为(检测数, K, 2或3)的张量,kpt_shape来自模型配置(见 ultralytics/cfg/models/v8/yolov8-pose.yaml,COCO 姿态为kpt_shape: [17, 3],即 17 个关键点、每点[x, y, visible]三维);
  • 若某张图没有检测到目标(len(pred) == 0),则直接取pred[:, 6:]空张量,避免view因 0 长度报错;
  • 关键点坐标通过scale_coords做与原图等比例的反向缩放(含 padding 去除与越界 clip),与检测框共享同一套 letterbox 几何变换参数。

第四步:封装 Results 对象

results.append( Results(orig_img, path=img_path, names=self.model.names, boxes=pred[:, :6], keypoints=pred_kpts) )

每个结果被封装为 Results 对象,其中boxes=pred[:, :6](前 6 列:检测框 + 分数 + 类别)交给Boxes模块,keypoints=pred_kpts交给Keypoints模块。Keypoints类(ultralytics/engine/results.py)提供了xy(像素坐标)、xyn(归一化坐标)、conf(关键点置信度,仅在维度为 3 时可用)等属性,并会在conf < 0.5时将对应点的坐标置零以标记不可见点。这也是Results.keypoints能被results[0].plot()直接绘制出骨骼姿态图的数据基础。

与 DetectionPredictor 的对比:差异即姿态任务的精髓

对比 DetectionPredictor.postprocess 与PosePredictor.postprocess,可以清晰看到姿态任务的增量:

环节DetectionPredictorPosePredictor
NMS 类别数不显式传ncnc=len(self.model.names)
坐标缩放scale_boxes处理pred[:, :4]检测框缩放 +scale_coords处理关键点
关键点处理kpt_shape重塑 + 坐标反算 + 越界裁剪
Results 封装boxes=predboxes=pred[:, :6]+keypoints=pred_kpts

这个对比也解释了为什么姿态推理的底层调用链(预处理 → 推理 → 后处理 → 可视化)与检测完全一致,但结果里多出了keypoints张量维度——关键点分支在模型侧(PoseModel,见 ultralytics/nn/tasks.py)和损失侧(v8PoseLoss)均有独立实现。

使用姿势:Python API 与 CLI

官方文档给出的最小示例

PosePredictor类的 docstring 中给出了最直接的调用方式(与文档 docs/en/reference/models/yolo/pose/predict.md 完全一致):

from ultralytics.utils import ASSETS from ultralytics.models.yolo.pose import PosePredictor args = dict(model='yolov8n-pose.pt', source=ASSETS) predictor = PosePredictor(overrides=args) predictor.predict_cli()
  • model指定姿态权重(yolov8n-pose.pt或任意自定义.pt);
  • source支持图像、视频、目录、URL、摄像头等多类数据源;
  • predict_cli()来自 BasePredictor,会以生成器方式逐批消费数据流并打印/保存结果。

面向日常使用的推荐方式

更常见的做法是直接通过统一入口YOLO类加载权重(内部自动选择PosePredictor):

from ultralytics import YOLO model = YOLO('yolov8n-pose.pt') results = model('https://ultralytics.com/images/bus.jpg') # 单张推理 for r in results: boxes = r.boxes # Boxes:检测框与置信度 kpts = r.keypoints # Keypoints:姿态关键点 (N, 17, 3) print(kpts.xy.shape) # (N, 17, 2) 像素坐标 print(kpts.conf) # (N, 17) 关键点置信度(当 kpt 维度为 3 时)

CLI 方式等价:

yolo pose predict model=yolov8n-pose.pt source='https://ultralytics.com/images/bus.jpg'

更多预测参数(confiouimgszdevicesaveshowagnostic_nmsmax_det等)与完整用法可参考 docs/en/modes/predict.md;姿态任务的整体训练、验证、导出流程可参考 docs/en/tasks/pose.md。

后处理关键函数的源码级补充

为了让读者能精准调参,这里对postprocess依赖的两个核心工具函数做源码级说明(均位于 ultralytics/utils/ops.py):

  • scale_coords(img1_shape, coords, img0_shape, ...):当未显式传入ratio_pad时,按gain = min(h1/h0, w1/w0)计算等比缩放系数,并推导 letterbox 两侧 padding,然后依次执行「减去 padding → 除以 gain →clip_coords越界裁剪」。姿态关键点与检测框共用同一变换,因此不会出现骨骼与框错位。
  • non_max_suppression(prediction, conf_thres, iou_thres, ...):默认conf_thres=0.25iou_thres=0.45max_det=300max_wh=7680,支持classes白名单过滤与agnostic跨类 NMS。姿态推理时confioumax_detagnostic_nmsclasses五个参数均直接透传自预测配置。

总结

PosePredictor是 Ultralytics YOLO 姿态推理的最终落点:它复用了DetectionPredictor的检测后处理管线,仅在其上增加了关键点通道的重塑(依据模型配置中的kpt_shape)、关键点坐标的 letterbox 反向缩放与Results.keypoints的封装。理解 postprocess 实现 中「NMS → 张量转换 → 框/点双路缩放 → Results 封装」四步链路,即可在自定义推理脚本中准确获取关键点坐标,也能为修改姿态后处理逻辑(如自定义关键点过滤策略)提供准确的代码切入点。

【免费下载链接】yolov10YOLOv10: Real-Time End-to-End Object Detection [NeurIPS 2024]项目地址: https://gitcode.com/GitHub_Trending/yo/yolov10

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

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

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

立即咨询