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_branchnn.Sequential(nn.Conv2d(self.cv3[-1].out_channels,128,3,padding1),# 这里踩过坑kernel1感受野不够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):# 只在大尺度特征图上做关键点预测ifi0:# P3层分辨率最高kpt_outs.append(self.kpt_branch(feat))else:# 小尺度特征图直接上采样后加进来kpt_outs.append(F.interpolate(self.kpt_branch(feat),sizekpt_outs[0].shape[2:],modebilinear))# 融合多尺度关键点预测kpt_outtorch.stack(kpt_outs).mean(dim0)这里有个取舍如果所有尺度都做关键点预测再融合小目标的关键点精度会提升但大目标的关键点反而会变模糊。我的经验是只保留P3和P4层P5层直接丢掉——P5的感受野太大关键点定位精度惨不忍睹。Loss设计别让关键点loss吃掉检测loss联合训练的loss平衡是个大坑。一开始我把关键点loss的权重设成和检测loss一样结果训练到一半发现模型只学关键点bbox全乱飘。# 踩坑代码权重设置不合理lossloss_detloss_kpt# 关键点loss量级是检测loss的10倍# 正确做法动态调整权重kpt_weight0.25# 根据关键点数量调整17个点用0.255个点用0.1lossloss_detkpt_weight*loss_kpt关键点loss我用的是OKSObject 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)k2*(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_xykpt_pred[...,:2].sigmoid()# 归一化到[0,1]kpt_visiblekpt_pred[...,2:3].sigmoid()# 映射到bbox内部bbox_xybbox_pred[...,:2].sigmoid()*stride bbox_whbbox_pred[...,2:4].sigmoid()*stride kpt_abs_xbbox_xy[...,0:1]kpt_xy[...,0:1]*bbox_wh[...,0:1]kpt_abs_ybbox_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):# 随机旋转anglerandom.uniform(-30,30)img,bboxes,kptsrotate(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_wrandom.randint(20,60),random.randint(20,60)mask_x,mask_yrandom.randint(0,img.shape[1]-mask_w),random.randint(0,img.shape[0]-mask_h)img[mask_y:mask_ymask_h,mask_x:mask_xmask_w]0# 被遮挡区域内的关键点visible置0kpt_mask(kpts[...,0]mask_x)(kpts[...,0]mask_xmask_w)\(kpts[...,1]mask_y)(kpts[...,1]mask_ymask_h)kpts[kpt_mask,2]0部署踩坑ONNX导出要改导出ONNX时关键点分支的dynamic shape会报错。因为不同batch size下关键点数量是动态的。# 导出时固定关键点数量classDetectWithKpts(nn.Module):defforward(self,x):det_out,kpt_outself.detect(x)# 这里踩过坑ONNX不支持动态reshape# 固定num_kpts避免动态shapebatch,anchors,_det_out.shape kpt_outkpt_out.reshape(batch,anchors,-1,3)returndet_out,kpt_outTensorRT部署时关键点分支的精度会下降。我的解决方案是在FP16推理时关键点分支单独用FP32。虽然牺牲了一点速度但关键点精度从0.72提升到0.81。实际效果在自建的工业数据集上联合检测模型相比两个独立模型推理速度从180ms降到45msTensorRT FP16关键点精度OKS从0.68提升到0.74共享特征让关键点学到更多上下文检测精度mAP基本持平没有下降最让我意外的是联合训练后模型对遮挡场景的鲁棒性明显提升。因为检测分支和关键点分支互相监督——检测分支告诉关键点分支这里有个物体关键点分支反馈给检测分支这个物体的关键点分布是这样的。个人经验别贪心关键点数量控制在17个以内超过这个数loss平衡会变得极其困难。如果非要检测50个点建议拆成多个关键点头。数据质量比模型结构重要关键点标注的噪声对精度影响巨大。我花了三周时间清洗标注数据比改模型结构带来的提升大得多。先跑通再优化第一次实现时先用最简单的MSE loss和固定权重跑通流程再逐步替换成OKS loss和动态权重。一步到位容易debug到崩溃。可视化debug训练过程中实时可视化关键点预测结果比看loss曲线有用十倍。我写了个回调函数每100个epoch保存一次预测结果一眼就能看出关键点是不是在乱飘。边缘场景要单独处理小目标面积32x32的关键点检测效果很差我的做法是单独训练一个小目标检测分支和大模型做级联推理。这个方案已经在三个工业场景落地效果稳定。如果你也在做类似的需求建议先从COCO关键点数据集开始验证再迁移到自己的数据上。