087、YOLOv8改进实战:关键点检测头扩展,实现人体姿态与物体关键点联合检测

📅 2026/7/28 13:47:11 👁️ 阅读次数 📝 编程学习
087、YOLOv8改进实战:关键点检测头扩展,实现人体姿态与物体关键点联合检测

087、YOLOv8改进实战:关键点检测头扩展,实现人体姿态与物体关键点联合检测

从一次翻车现场说起

上个月接了个需求,要在工业质检场景里同时检测产品缺陷位置和关键装配点。甲方给的标注数据里既有bbox又有keypoints,我第一反应是跑两个模型——一个YOLOv8做检测,一个SimpleBaseline做姿态估计。结果部署的时候直接炸了:两套模型推理时间加起来快200ms,边缘盒子根本扛不住。

更坑的是,两个模型的输出坐标系还不一致,后处理要写一堆对齐逻辑。调试到凌晨三点,看着屏幕上歪歪扭扭的关键点连线,我意识到必须把检测头和关键点头合并成一个输出头。

为什么YOLOv8原生不支持关键点联合检测

翻源码的时候发现,YOLOv8的Detect模块输出维度是(4 + num_classes),每个anchor对应一个bbox和类别概率。关键点信息根本没地方塞。官方Pose模型倒是有关键点分支,但那是单独训练的,不能和检测任务共享特征。

核心问题在于:检测头和关键点头需要不同的特征分辨率。检测头喜欢大感受野来定位物体,关键点头需要高分辨率来精确定位点。强行共享一个输出层,要么检测不准,要么关键点飘得离谱。

动手改:在Detect模块里塞进关键点分支

我的做法是在YOLOv8的Detect类里新增一个并行的关键点预测分支,不破坏原有的检测逻辑。具体来说,在__init__方法里加一个self.kpt_branch

# 别这样写:直接把关键点堆到检测输出里# self.cv_kpt = nn.Conv2d(self.cv3[-1].out_channels, 17*3, 1) # 17个关键点,每个3维(x,y,visible)# 正确做法:保持检测头独立,新增并行分支self.kpt_branch=nn.Sequential(nn.Conv2d(self.cv3[-1].out_channels,128,3,padding=1),# 这里踩过坑,kernel=1感受野不够nn.BatchNorm2d(128),nn.SiLU(),nn.Conv2d(128,num_kpts*3,1)# 每个关键点输出x,y,visible)

注意num_kpts * 3这个设计。visible维度用来表示关键点是否可见,训练时如果某个点被遮挡,这个维度的loss要mask掉。一开始我没加visible,结果遮挡场景下关键点全往图像中心飘。

前向传播的坑

前向传播时,关键点分支和检测分支共享backbone和neck的特征。但这里有个细节:不同尺度的特征图对关键点检测的贡献不一样。

defforward(self,x):# x是三个尺度的特征图列表kpt_outs=[]fori,featinenumerate(x):# 只在大尺度特征图上做关键点预测ifi==0:# P3层,分辨率最高kpt_outs.append(self.kpt_branch(feat))else:# 小尺度特征图直接上采样后加进来kpt_outs.append(F.interpolate(self.kpt_branch(feat),size=kpt_outs[0].shape[2:],mode='bilinear'))# 融合多尺度关键点预测kpt_out=torch.stack(kpt_outs).mean(dim=0)

这里有个取舍:如果所有尺度都做关键点预测再融合,小目标的关键点精度会提升,但大目标的关键点反而会变模糊。我的经验是只保留P3和P4层,P5层直接丢掉——P5的感受野太大,关键点定位精度惨不忍睹。

Loss设计:别让关键点loss吃掉检测loss

联合训练的loss平衡是个大坑。一开始我把关键点loss的权重设成和检测loss一样,结果训练到一半发现模型只学关键点,bbox全乱飘。

# 踩坑代码:权重设置不合理loss=loss_det+loss_kpt# 关键点loss量级是检测loss的10倍# 正确做法:动态调整权重kpt_weight=0.25# 根据关键点数量调整,17个点用0.25,5个点用0.1loss=loss_det+kpt_weight*loss_kpt

关键点loss我用的是OKS(Object Keypoint Similarity)的变体,不是简单的MSE。OKS会根据关键点类型和物体尺度自动调整权重——比如眼睛这种小范围关键点,位置偏差的惩罚比手腕大得多。

defoks_loss(pred_kpts,gt_kpts,bbox_areas,sigmas):# sigmas是每个关键点的标准差,COCO数据集有预定义值# 这里踩过坑:bbox_areas要用sqrt,不然大物体loss太小d=(pred_kpts-gt_kpts).pow(2).sum(dim=-1)k=2*(bbox_areas.sqrt()*sigmas).pow(2)return(1-torch.exp(-d/k)).mean()

后处理:关键点怎么和检测框对齐

模型输出的是相对坐标,需要解码成绝对坐标。这里有个容易忽略的点:关键点的坐标应该基于检测框归一化,而不是基于图像。

defdecode_kpts(kpt_pred,bbox_pred,stride):# kpt_pred: [batch, anchors, num_kpts*3]# bbox_pred: [batch, anchors, 4]# 别这样写:直接乘stride# kpt_abs = kpt_pred * stride# 正确做法:先基于anchor中心点解码,再映射到bbox内部kpt_xy=kpt_pred[...,:2].sigmoid()# 归一化到[0,1]kpt_visible=kpt_pred[...,2:3].sigmoid()# 映射到bbox内部bbox_xy=bbox_pred[...,:2].sigmoid()*stride bbox_wh=bbox_pred[...,2:4].sigmoid()*stride kpt_abs_x=bbox_xy[...,0:1]+kpt_xy[...,0:1]*bbox_wh[...,0:1]kpt_abs_y=bbox_xy[...,1:2]+kpt_xy[...,1:2]*bbox_wh[...,1:2]returntorch.cat([kpt_abs_x,kpt_abs_y,kpt_visible],dim=-1)

这样设计的好处是:关键点天然和检测框绑定,不会出现关键点落在框外的情况。之前用图像坐标直接解码,经常出现关键点飞到框外几米远的情况。

训练技巧:数据增强要小心

关键点检测对数据增强特别敏感。随机裁剪和旋转会导致关键点位置和可见性发生变化。

# 自定义关键点增强classKeypointAugment:def__call__(self,img,bboxes,kpts):# 随机旋转angle=random.uniform(-30,30)img,bboxes,kpts=rotate(img,bboxes,kpts,angle)# 这里踩过坑:旋转后要重新计算可见性# 如果关键点旋转后超出图像边界,visible置0kpts[...,2]=(kpts[...,0]>=0)&(kpts[...,0]<img.shape[1])&\(kpts[...,1]>=0)&(kpts[...,1]<img.shape[0])# 随机遮挡:模拟关键点被遮挡ifrandom.random()<0.3:mask_h,mask_w=random.randint(20,60),random.randint(20,60)mask_x,mask_y=random.randint(0,img.shape[1]-mask_w),random.randint(0,img.shape[0]-mask_h)img[mask_y:mask_y+mask_h,mask_x:mask_x+mask_w]=0# 被遮挡区域内的关键点visible置0kpt_mask=(kpts[...,0]>=mask_x)&(kpts[...,0]<=mask_x+mask_w)&\(kpts[...,1]>=mask_y)&(kpts[...,1]<=mask_y+mask_h)kpts[kpt_mask,2]=0

部署踩坑:ONNX导出要改

导出ONNX时,关键点分支的dynamic shape会报错。因为不同batch size下,关键点数量是动态的。

# 导出时固定关键点数量classDetectWithKpts(nn.Module):defforward(self,x):det_out,kpt_out=self.detect(x)# 这里踩过坑:ONNX不支持动态reshape# 固定num_kpts,避免动态shapebatch,anchors,_=det_out.shape kpt_out=kpt_out.reshape(batch,anchors,-1,3)returndet_out,kpt_out

TensorRT部署时,关键点分支的精度会下降。我的解决方案是:在FP16推理时,关键点分支单独用FP32。虽然牺牲了一点速度,但关键点精度从0.72提升到0.81。

实际效果

在自建的工业数据集上,联合检测模型相比两个独立模型:

  • 推理速度:从180ms降到45ms(TensorRT FP16)
  • 关键点精度:OKS从0.68提升到0.74(共享特征让关键点学到更多上下文)
  • 检测精度:mAP基本持平,没有下降

最让我意外的是,联合训练后模型对遮挡场景的鲁棒性明显提升。因为检测分支和关键点分支互相监督——检测分支告诉关键点分支"这里有个物体",关键点分支反馈给检测分支"这个物体的关键点分布是这样的"。

个人经验

  1. 别贪心:关键点数量控制在17个以内,超过这个数loss平衡会变得极其困难。如果非要检测50个点,建议拆成多个关键点头。

  2. 数据质量比模型结构重要:关键点标注的噪声对精度影响巨大。我花了三周时间清洗标注数据,比改模型结构带来的提升大得多。

  3. 先跑通再优化:第一次实现时,先用最简单的MSE loss和固定权重跑通流程,再逐步替换成OKS loss和动态权重。一步到位容易debug到崩溃。

  4. 可视化debug:训练过程中实时可视化关键点预测结果,比看loss曲线有用十倍。我写了个回调函数,每100个epoch保存一次预测结果,一眼就能看出关键点是不是在乱飘。

  5. 边缘场景要单独处理:小目标(面积<32x32)的关键点检测效果很差,我的做法是单独训练一个小目标检测分支,和大模型做级联推理。

这个方案已经在三个工业场景落地,效果稳定。如果你也在做类似的需求,建议先从COCO关键点数据集开始验证,再迁移到自己的数据上。