拓冰建站拓冰建站
首页 / 资讯中心 / 正文

用Keypoint R-CNN训练自定义关键点检测:从标注到部署全流程

简介一份基于PyTorch Keypoint R-CNN训练自定义数据集关键点检测模型的完整工程资源适合具备一定深度学习基础、希望将关键点检测落地到自建数据集的开发者与研究者。资源围绕自建关键点数据的完整流程展开涵盖数据格式整理、标签转换、模型头部调整、训练超参数设置以及训练结果评估等环节能够帮助用户解决从数据准备到模型调优过程中常见的格式不统一、关键点数量适配等难题。压缩包共116个文件包括8个Python脚本、2个Jupyter Notebook、31个JSON标注文件、34个TXT文本文件以及39张示例图片另附说明文档和gitmodules配置整体仅8.55MB目录结构清晰便于按模块查阅和复用。目前已有164人学习下载。借助资源中的Convert_labels.ipynb与KeypointRCNN_training.ipynb等可执行文件用户可以快速掌握自建数据集上Keypoint R-CNN的构建、训练与评估流程并方便地迁移到自己的项目中减少重复搭建和排错时间。 大概两周前我在做一个工业零件装配角度检测的小需求需要在零件图像上定位几个关键点用来计算装配角度是否合格。一开始想省事直接调现成的MediaPipe结果发现它的人脸/人体关键点语义是固化的根本没法自定义成我需要的零件角点。换HRNet又嫌训练流程太重数据格式、评估脚本全要自己写试了两天就放弃了。最后兜兜转转落到了PyTorch自带的Keypoint R-CNN上——这个模型在torchvision里就有实现和预训练权重改一改输出层就能训自定义关键点。这篇就把我从标注、写Dataset、改造模型到训练调参的完整过程拆开讲给同样想用keypoint_rcnn训练自建关键点数据集的读者一条能直接走通的路。1. 为什么自建关键点检测我最终没选MediaPipe而选了Keypoint R-CNN1.1 Keypoint R-CNN的输出机制RoI之上的关键点热力图Keypoint R-CNN可以理解为Faster R-CNN的扩展。常规的目标检测只输出box和class而Keypoint R-CNN在Detection Head之外额外接了一个关键点分支。流程大致是这样图像进入ResNet50FPN的backbone提取多尺度特征RPN生成候选框接着RoIAlign把每个候选框对应的特征图抠出来喂给检测头做分类和框回归同时喂给关键点头输出K张热力图K就是你要检测的关键点数量。每张热力图上的峰值位置就是对应关键点的坐标。这个输出方式有个很实用的好处你不需要像回归坐标那样直接预测一个(x, y)值而是预测一张概率图再用soft-argmax或者峰值定位的方式把坐标解出来。峰值定位天然具备亚像素精度也比直接回归稳定得多。训练时使用的是基于高斯热图的MSE损失目标是把标注点周围的高斯区域“点亮”这一点在后面调参时会反复提到建议先记住。1.2 四种方案的选型对比为什么它适合中小型自定义数据集自建关键点数据集市面上能用的方案无外乎这几类我实际都试过方案优点缺点适合场景MediaPipe开箱即用推理快关键点稳定关键点语义固定无法自定义人脸、手势、人体姿态等标准任务HRNet姿态估计精度高社区资料多数据加载、训练pipeline、评估代码全要自己搭工程量大学术研究、追求SOTA指标CenterNet单阶段快关键点可自由定义需要自己写高斯核生成对标注质量敏感回归稳定性一般点稀疏、目标较小的场景Keypoint R-CNNtorchvision自带实现预训练权重丰富自定义关键点数量只要改参数速度和轻量化不如单阶段方案依赖GPU训练中小型自定义数据集、快速落地选Keypoint R-CNN的原因很简单它在一个相对完整的检测框架里把“自定义关键点”这件事的工程量降到了最低。预训练权重是COCO上的人体姿态模型backbone和RPN的特征提取能力已经很强了我们只需要把最后的关键点输出层改掉用少量标注数据微调就能收敛。对中小型项目来说这比从零训一个HRNet性价比高得多。2. 标注自建数据集COCO关键点格式与工具实操2.1 COCO格式里关键点的三种可见状态Keypoint R-CNN默认支持的标注格式是COCO Keypoint格式如果你自己写标注工具也需要转成这个格式。一份COCO标注JSON主要由三部分构成images、annotations、categories。其中关键点相关的核心字段是这样的{ images: [ {id: 1, file_name: 001.jpg, height: 480, width: 640} ], annotations: [ { id: 1, image_id: 1, category_id: 1, bbox: [100, 80, 200, 240], area: 48000, iscrowd: 0, keypoints: [120, 100, 2, 150, 140, 2, 180, 200, 1], num_keypoints: 3 } ], categories: [ { id: 1, name: part, keypoints: [point1, point2, point3], skeleton: [[1, 2], [2, 3]] } ] }keypoints数组是[x1, y1, v1, x2, y2, v2, ...]的平铺结构每三个数值代表一个关键点其中v是可见状态0表示该点未标注不在画面内或完全不可见1表示被遮挡但大致位置可推断2表示清晰可见。训练时Keypoint R-CNN会根据v0生成mask把这些关键点排除在损失函数之外。这里有个实操中容易踩的坑很多人把所有点都标成v2包括被遮挡的点。模型会强行去学习一个实际上不存在的关键点导致推理时在遮挡区域乱猜。我的建议是凡是不确定位置的点要么标v1要么干脆不标设为v0让模型学会“没有这个点”才是更合理的行为。2.2 标注工具与顺序统一关键点顺序比统一框更重要标注工具我用了LabelMe导出时选COCO格式即可。CVAT也可以但配置相对重一些单机做小数据集用LabelMe足够。LabelMe标注关键点时每个点会按你点击的顺序编号。这里最容易被忽略的一点是同一个类别下所有实例的关键点顺序必须保持一致。比如你对第一个零件先点左上角再点右下角第二个零件也必须先点左上角再点右下角。如果顺序混了模型会把这个点和那个点搞混loss直接崩。我的做法是标注前先把关键点列表写在纸上或做成一个Excel模板标注时严格按模板顺序逐一点击。另外bbox不一定要标得特别精准但必须把整个目标框住因为关键点头是在RoI的内部特征图上工作的。如果box太小RoIAlign截出的区域缺少上下文关键点定位精度会下降。2.3 关于多类别关键点的边界先别急着做多类如果你想让一个模型同时检测“人”的17个点和“猫”的9个点Keypoint R-CNN实现起来会比较麻烦。因为torchvision的关键点输出层是统一的num_keypoints通道不同类别有不同数量的关键点时你需要在head内部按类别做分支这偏离了原版实现改动成本很高。我自己的经验是自建项目最好一个数据集只做一种类别的关键点检测如果确实需要多类别优先拆成多个模型分别训练推理时再组合。省下的调试时间远超你省下的部署成本。3. 数据加载与增强关键点坐标同步是最大的坑3.1 Dataset返回的target字典到底要什么字段PyTorch自带的检测模型有一套固定的输入输出约定。我们写Dataset时__getitem__要返回(image, target)两个对象。image是FloatTensor[C, H, W]target是一个dict必须包含以下几个字段target { boxes: torch.as_tensor(boxes, dtypetorch.float32), # [N, 4] x1,y1,x2,y2 labels: torch.as_tensor(labels, dtypetorch.int64), # [N] image_id: torch.as_tensor([img_id], dtypetorch.int64), # [1] area: torch.as_tensor(area, dtypetorch.float32), # [N] iscrowd: torch.as_tensor([0] * N, dtypetorch.int64), # [N] keypoints: torch.as_tensor(kps, dtypetorch.float32), # [N, K, 2] keypoints_visible: torch.as_tensor(kps_v, dtypetorch.float32) # [N, K] }keypoints的最后一维是(x, y)而keypoints_visible则是0/1的可见掩码在构造时把COCO里v字段大于0的都映射为1v0映射为0。这个映射一定要在读取JSON的时候做不要在训练循环里临时处理否则会拖慢数据加载速度。转换时建议把所有坐标先统一成x1,y1,x2,y2的绝对坐标格式不要在Dataset里存归一化坐标因为模型内部会再做一次自己的坐标变换。3.2 albumentations做关键点增强的同步问题torchvision自带的transforms对检测任务的支持比较弱尤其不支持关键点同步变换。我自己用的是albumentations它对检测框和关键点的同步处理已经比较成熟import albumentations as A train_transform A.Compose([ A.HorizontalFlip(p0.5), A.RandomBrightnessContrast(p0.3), A.RandomSizedBBoxSafeCrop(width640, height640, p0.5), ], bbox_paramsA.BboxParams(formatpascal_voc, label_fields[labels]), keypoint_paramsA.KeypointParams(formatxy, remove_invisibleTrue) )注意两个细节。第一HorizontalFlip翻转后如果关键点存在“左右对称”的语义比如左眼和右眼必须手动交换它们的索引否则模型会学习到错误的左右对应关系。第二remove_invisibleTrue会把增强后跑出图像范围的关键点直接删掉但与此同时它也可能把该关键点的可见性改成False你需要把albumentations返回的keypoints和原始可见性对齐否则会出现索引错位。踩过一次大的我一开始用RandomSizedBBoxSafeCrop做裁剪增强但没有把增强后的关键点坐标同步映射到新的图像坐标系结果训练出来的模型在图像边缘区域的关键点全部偏了一大截。这类问题在代码层面很难肉眼发现最后是靠可视化增强后的图片和标注点才定位出来的。所以每次改完数据增强代码第一件事就是找几张图把增强结果画出来看一眼。3.3 输入尺寸与归一化对齐预训练权重torchvision的检测模型内部会有一个transform但它的默认行为只有ToTensor并不会做ImageNet归一化。而预训练权重是在ImageNet归一化分布下训练的所以我们需要在自定义Dataset里补上Normalize(mean[0.485,0.456,0.406], std[0.229,0.224,0.225])。输入尺寸上我的显卡是单张消费级GPU显存有限所以把长边resize到800、短边不小于600。如果目标特别小建议统一缩放到640×640甚至更大。这里要理解一个权衡Keypoint R-CNN的关键点头是在RoIAlign之后的小分辨率特征图上输出的目标越小、分辨率越低关键点定位越容易漂移。自建数据集里如果关键点本身只占几个像素那就尽量保持较大的输入尺寸哪怕牺牲一点batch size也值得。4. 模型改造与训练配置从预训练权重到自定义关键点4.1 预训练模型拉取与关键点头改造torchvision从0.13版本开始keypointrcnn_resnet50_fpn直接支持num_keypoints参数这是最省事的改造方式from torchvision.models.detection import keypointrcnn_resnet50_fpn from torchvision.models.detection.keypoint_rcnn import KeypointRCNN_ResNet50_FPN_Weights model keypointrcnn_resnet50_fpn( weightsKeypointRCNN_ResNet50_FPN_Weights.COCO_V1, num_classes2, # 1个类别 背景 num_keypoints6 # 你自己的关键点数量 )如果用的是老版本torchvision或者你想更细粒度地控制关键点分支就手动替换输出层from torchvision.models.detection import keypointrcnn_resnet50_fpn model keypointrcnn_resnet50_fpn(weightsKeypointRCNN_ResNet50_FPN_Weights.COCO_V1) in_channels model.roi_heads.keypoint_predictor.kp_score_lowres.in_channels model.roi_heads.keypoint_predictor.kp_score_lowres torch.nn.Conv2d( in_channels, num_keypoints, kernel_size1, stride1, padding0 )这两种方式的本质都是把Keypoint Head最后的输出通道从17改成你要的K。注意num_classes也要改因为你的数据集的语义不再是COCO的person即使你希望模型最终只检测一类目标也要设置num_classes2因为COCO的num_classes是包含背景的。4.2 优化器、学习率与loss权重的设置用预训练权重做迁移学习时学习率不能像从零训练那样设大。我试过几个组合最终稳定使用的是params [p for p in model.parameters() if p.requires_grad] optimizer torch.optim.SGD(params, lr2.5e-4, momentum0.9, weight_decay1e-4) lr_scheduler torch.optim.lr_scheduler.MultiStepLR(optimizer, milestones[8, 11], gamma0.1)如果你的数据量只有几百张建议前几个epoch冻结backbone和RPN只训练roi_heads。实现方式是把model.backbone和model.rpn的requires_grad设为False等检测头loss稳定之后再解冻所有层。原因很简单预训练backbone已经学会了通用特征数据量小时全量微调反而会把它带偏。关于loss权重torchvision内部把关键点loss的权重系数keypoint_loss_weight默认设为1.0。但实际训练中关键点loss是基于密集热力图的数值往往很大前期可能占整个loss的80%以上。如果发现总loss一直在降但关键点精度没有明显提升可以去torchvision/models/detection/roi_heads.py里把这个权重调小到0.5甚至0.2给分类和回归loss更多话语权。4.3 项目文件结构参考标题里的_keypoint_detection.zip暗示了这是一个完整项目包我建议在实际落地时也按下面这个结构组织代码这样别人拿到压缩包能快速跑起来keypoint_detection/ ├── data/ │ ├── annotations/ │ │ ├── instances_train.json │ │ └── instances_val.json │ ├── train/ │ └── val/ ├── datasets/ │ └── coco_keypoint_dataset.py ├── models/ │ └── keypoint_model.py ├── config.py ├── train.py ├── evaluate.py └── inference.pyconfig.py里集中放所有超参数包括学习率、batch size、关键点数量、类别名、数据路径等。不要把这些散落在各个脚本里否则换数据集时要改的东西太多很容易漏。5. 训练实测三个最影响精度的调优细节5.1 训练不收敛先检查这四件事如果你在训练中遇到loss不降、或者总在某个值附近震荡先别急着调模型架构按顺序排查这四个点第一标注JSON里的keypoints顺序和categories.keypoints定义是否一致。很多人改过类别名或关键点名称后忘了同步annotations里的数组顺序模型永远在学一个错乱的目标。第二是否有大量v0的关键点。如果一张图里半数关键点都没标注关键点head拿到的监督信号太稀疏自然学不准。要么补标要么干脆删除这些样本。第三学习率是否过大。预训练模型全量微调时lr2.5e-4起步是安全的如果用的AdamW我会降到1e-4。SGD配0.9的momentum在检测任务上效果很稳别一上来就换复杂优化器。第四数据加载的归一化是否正确。漏了Normalize的话输入分布和预训练权重不匹配loss初始值会异常偏高训练过程也会不稳。5.2 关键点偏移到物体外部完整的排查链路这类问题最容易让人崩溃因为训练loss看起来是正常的但推理时关键点就是偏到物体外面。我的排查顺序是这样的第一步可视化验证集ground truth。把标注点画在图像上确认标注本身没问题。第二步可视化增强后的图像和关键点排查数据增强是否破坏了坐标同步。我在这个环节吃过大亏所以现在对增强代码格外谨慎。第三步检查推理后处理。从模型输出的热力图取坐标时不要直接argmaxtorchvision内部用的heatmaps_to_keypoints会结合峰值附近的偏移量做亚像素细分直接取热力图argmax会带来1~2个像素的固定偏移自建小目标数据集上尤其明显。如果以上都没问题再考虑模型容量和输入尺寸把输入长边从800提到1000或者1200通常能缓解小目标偏移。5.3 评估与推理用PCK和soft-argmax拿准坐标COCO官方的OKS指标需要每个关键点配置sigma自建数据集通常没有这个先验所以我建议用PCKPercentage of Correct Keypoints来做评估。实现很简单def pck_metric(pred_kps, gt_kps, bbox, threshold0.1): bbox_diag (bbox[2] ** 2 bbox[3] ** 2) ** 0.5 dist torch.norm(pred_kps - gt_kps, dim-1) correct (dist threshold * bbox_diag).float().mean().item() return correct推理阶段的输出中predictions[0][keypoints]的形状通常是[N, K, 3]第三维是(x, y, score)。这个score就是热力图峰值可以当置信度用。建议在推理管线里加一个后处理先按score_thresh0.5过滤目标框再对框内关键点做一次可见性过滤把score过低的关键点置为不可见而不是强行输出一个坐标。最后分享一个实操中的小习惯每次跑全量训练前先拿20张图、1个epoch把整个pipeline跑通确认loss能正常下降、保存的checkpoint能正常加载、评估脚本能画出正确结果再启动正式训练。这个习惯帮我省下了大量反复调试的时间。如果你也准备在自建数据集上训练Keypoint R-CNN建议从这个流程开始一步步来整个过程并不复杂只是细节比较多。本文还有配套的精品资源点击获取
分享:

看完干货,该让你的企业上线了

免费需求沟通 · 48 小时内出具建站方案 · 河南本地可上门