十年匠心定制 · 商业建站与技术教学双线并行 咨询热线:400-886-1026 service@lmnt.cn
ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

PyTorch Keypoint R-CNN自定义关键点检测实战:从标注到部署

PyTorch Keypoint R-CNN自定义关键点检测实战:从标注到部署 简介面向深度学习和计算机视觉开发者资源聚焦使用PyTorch框架中的Keypoint R-CNN训练自建数据集的关键点检测模型适合有基础目标检测知识、想掌握关键点检测完整流程的读者。压缩包共116个文件大小8.55MB包含训练用jpg图像、txt和json标注文件、Python训练脚本、ipynb交互式教程以及说明文档覆盖从数据准备、模型配置到训练评估和部署的主要环节。已有164人学习下载。资源提供了Keypoint R-CNN训练代码、标注转换脚本和目录结构清晰的笔记可帮助读者快速搭建自建数据集的关键点检测流程理解关键点坐标标注格式、损失计算和模型调参思路同时也可参考模型导出与部署的细节减少踩坑。1. 项目背景与核心思路做关键点检测这个方向很多人第一反应是直接上HRNet或者OpenPose这类专门设计的关键点模型但如果你面对的检测目标非常明确比如检测某个机械零件上的定位孔、一张卡片上的四个角点而且对部署成本和训练难度有要求PyTorch官方torchvision里那个Keypoint R-CNN往往是更省事的选项。这个模型最大的特点是它不需要你从零搭建网络也不需要自己写复杂的loss计算逻辑。它本质上是Faster R-CNN的扩展——在目标检测的基础上并行加了一个关键点预测分支。你只需要把自建数据集整理成它认识的格式然后调用现成的训练逻辑就能跑起来。我这次的项目就属于这种典型的自定义小目标场景需要检测一个矩形工件上的四个角点用来做后续的几何校正和尺寸测量。选择Keypoint R-CNN还有一个现实考虑它对硬件要求相对友好。相比目前动辄需要几十G显存的大模型Keypoint R-CNN用一块消费级显卡比如RTX 3060 12G就能完成训练和推理。这一点对于很多手里只有入门级GPU的开发者来说很重要。2. 数据集标注与格式转换2.1 标注工具选型与标注规范做关键点检测第一步就是给图片打点。关于标注工具我用的是labelmeGitHub上最主流的开源标注工具之一非常推荐。但有一个关键点需要提醒labelme默认标注完导出的是它自己的JSON格式而PyTorch的Keypoint R-CNN训练需要的是COCO格式的数据所以中间必须经历一次格式转换。标注的时候有几个细节会影响后续训练效果点的顺序要固定比如我标注矩形工件的四个角点约定顺序是左上、右上、右下、左下。如果第一张图先标左上再标右下第二张图先标右下再标左上模型会学得非常混乱。标注时必须保证同类目标的点顺序完全一致。尽量贴近边缘放大图片后精准点击不要凭感觉大致点一下。关键点标注的误差直接影响最终预测精度这一步偷懒后面训练出来的模型上限就很低。遮挡处理如果某个点在图片中被遮挡了我建议这张图要么弃用要么跳过这个目标不标。因为被遮挡的点本身信息不完整强行标注反而会给模型引入噪声。2.2 从labelme到COCO格式的转换逻辑COCO格式中关键点的存储形式是一个一维数组长度是num_keypoints * 3每组由[x, y, v]三个值组成。其中v代表可见性0表示该点未标注1表示已标注但被遮挡2表示已标注且可见。Keypoint R-CNN训练时模型只计算v1和v2的点的loss。而labelme的JSON里每个标注目标的points字段是一个二维数组形如[[x1, y1], [x2, y2], ...]。转换时需要注意COCO中的keypoints坐标是绝对像素坐标即原始图片的像素x、y值而不是归一化坐标。很多人在转换时容易在这里踩坑。另外COCO格式的categories字段必须包含keypoints和skeleton两个特殊字段categories [{ id: 1, name: workpiece, keypoints: [left_top, right_top, right_bottom, left_bottom], skeleton: [[0, 1], [1, 2], [2, 3], [3, 0]] }]skeleton字段定义了关键点之间的连接关系虽然Keypoint R-CNN训练时并未直接使用这个字段它主要被用于可视化但COCO格式校验时会检查这个字段是否存在建议保留。转换脚本的核心逻辑如下import json import os from PIL import Image def convert_labelme_to_coco(labelme_dir, img_dir, output_path): 将labelme标注的JSON文件转换为COCO格式 coco_output { images: [], annotations: [], categories: [{ id: 1, name: workpiece, keypoints: [left_top, right_top, right_bottom, left_bottom], skeleton: [[0, 1], [1, 2], [2, 3], [3, 0]] }] } annotation_id 1 image_id 1 for filename in os.listdir(labelme_dir): if not filename.endswith(.json): continue # 读labelme的标注文件 with open(os.path.join(labelme_dir, filename), r, encodingutf-8) as f: labelme_data json.load(f) # 获取图片信息 img_filename labelme_data[imagePath] img_path os.path.join(img_dir, img_filename) img Image.open(img_path) width, height img.size # 添加图片信息 coco_output[images].append({ id: image_id, file_name: img_filename, width: width, height: height }) # 处理每个标注目标 for shape in labelme_data[shapes]: if shape[label] ! workpiece: continue points shape[points] # 关键点坐标列表 # 计算边框关键点最小外接矩形 xs [p[0] for p in points] ys [p[1] for p in points] x_min, x_max min(xs), max(xs) y_min, y_max min(ys), max(ys) bbox [x_min, y_min, x_max - x_min, y_max - y_min] area bbox[2] * bbox[3] # 构建keypoints数组 [x1, y1, v1, x2, y2, v2, ...] keypoints [] for point in points: keypoints.extend([point[0], point[1], 2]) # v2表示可见 # 添加标注信息 coco_output[annotations].append({ id: annotation_id, image_id: image_id, category_id: 1, bbox: bbox, area: area, iscrowd: 0, keypoints: keypoints, num_keypoints: len(points) }) annotation_id 1 image_id 1 # 保存COCO格式的JSON with open(output_path, w, encodingutf-8) as f: json.dump(coco_output, f, ensure_asciiFalse) print(f转换完成共{image_id - 1}张图片{annotation_id - 1}个标注目标)这段脚本有几个需要根据自己项目调整的地方label字段要与你在labelme中实际标注的标签名称一致。keypoints的顺序要和标注时约定的一致。如果你有多个类别需要在categories里配置多个条目并在annotations中区分category_id。注意COCO格式的bbox坐标可以是浮点数但很多模型实现内部会做一些取整或clip操作建议确保bbox不超出图片边界否则训练时会报错或导致loss异常。3. 训练代码实现与参数调优3.1 数据集类与数据加载器PyTorch官方Torchvision中自带了torchvision.datasets.CocoDetection类可以直接读取COCO格式的JSON。但实际使用中我发现直接用这个类拿不到关键点标注信息需要做一层封装。所以我选择自己写数据集类这样对数据的控制更灵活也方便做数据增强。关键点检测的数据集类核心是实现__getitem__方法返回(image, target)对其中target是一个字典必须包含以下字段boxesTensor形状[N, 4]目标边界框labelsTensor形状[N]目标类别idkeypointsTensor形状[N, K, 3]关键点坐标及可见性image_idTensor图片idareaTensor目标面积iscrowdTensor是否为crowd目标import torch import torchvision from torch.utils.data import Dataset import json import os from PIL import Image import numpy as np class KeypointDataset(Dataset): def __init__(self, root_dir, annotation_file, transformsNone): self.root_dir root_dir self.transforms transforms with open(annotation_file, r, encodingutf-8) as f: self.coco_data json.load(f) # 按image_id索引图片和标注 self.images {img[id]: img for img in self.coco_data[images]} self.annotations {} for ann in self.coco_data[annotations]: img_id ann[image_id] if img_id not in self.annotations: self.annotations[img_id] [] self.annotations[img_id].append(ann) self.image_ids list(self.images.keys()) def __len__(self): return len(self.image_ids) def __getitem__(self, idx): img_id self.image_ids[idx] img_info self.images[img_id] # 加载图片 img_path os.path.join(self.root_dir, img_info[file_name]) image Image.open(img_path).convert(RGB) # 解析标注 anns self.annotations.get(img_id, []) boxes [] labels [] keypoints [] areas [] iscrowd [] for ann in anns: boxes.append(ann[bbox]) labels.append(ann[category_id]) keypoints.append(np.array(ann[keypoints]).reshape(-1, 3)) areas.append(ann[area]) iscrowd.append(ann[iscrowd]) boxes torch.as_tensor(boxes, dtypetorch.float32).reshape(-1, 4) labels torch.as_tensor(labels, dtypetorch.int64) keypoints torch.as_tensor(np.array(keypoints), dtypetorch.float32) areas torch.as_tensor(areas, dtypetorch.float32) iscrowd torch.as_tensor(iscrowd, dtypetorch.int64) image_id torch.tensor([img_id]) # 转换bbox格式从[x, y, w, h]转为[x1, y1, x2, y2] boxes[:, 2] boxes[:, 0] boxes[:, 2] boxes[:, 3] boxes[:, 1] boxes[:, 3] target { boxes: boxes, labels: labels, keypoints: keypoints, image_id: image_id, area: areas, iscrowd: iscrowd } if self.transforms is not None: image, target self.transforms(image, target) return image, target注意boxes从[x, y, w, h]到[x1, y1, x2, y2]的转换不能省。COCO格式存的是左上角坐标宽高而PyTorch检测模型接收的是左上角和右下角坐标。3.2 数据增强的正确姿势关键点检测的数据增强有一个特殊的地方对图像做翻转或旋转时关键点坐标必须同步变换。如果只对图像做增强而不更新关键点坐标模型的loss会飞掉。一个简单的水平翻转增强的实现如下import random class RandomHorizontalFlip: def __init__(self, flip_prob0.5): self.flip_prob flip_prob def __call__(self, image, target): if random.random() self.flip_prob: # 图像水平翻转 image image.transpose(Image.FLIP_LEFT_RIGHT) # 关键点x坐标需要镜像new_x width - old_x width image.size[0] boxes target[boxes].clone() boxes[:, [0, 2]] width - boxes[:, [2, 0]] keypoints target[keypoints].clone() keypoints[:, :, 0] width - keypoints[:, :, 0] # 注意关键点顺序也要反转因为我们的4个点有左右对应关系 # 假设点顺序是[左上, 右上, 右下, 左下]翻转后应该变成[右上, 左上, 左下, 右下] # 这里只是一个例子具体映射关系取决于你的点序定义 flipped_indices [1, 0, 3, 2] keypoints keypoints[:, flipped_indices, :] target[boxes] boxes target[keypoints] keypoints return image, target这个细节非常关键。我当时第一次做数据增强时只翻转了图像没动关键点结果训练了几个epoch后loss完全不下降后面检查才发现问题源头在这。3.3 模型创建与训练配置模型创建其实特别简单PyTorch提供了现成的构造方法。需要在torchvision.models.detection下导入keypointrcnn_resnet50_fpn并通过weights参数指定使用预训练权重。注意这个预训练权重是在COCO关键点数据集上训练的它本身已经学会了很多通用的特征我们后面要微调它来适配自己的数据。import torchvision # 加载预训练模型 model torchvision.models.detection.keypointrcnn_resnet50_fpn( weightstorchvision.models.detection.KeypointRCNN_ResNet50_FPN_Weights.COCO_V1 ) # 设置类别数背景 目标类别数背景固定为1 num_classes 2 # 1类目标 1个背景 in_features model.roi_heads.box_predictor.cls_score.in_features model.roi_heads.box_predictor torchvision.models.detection.faster_rcnn.FastRCNNPredictor(in_features, num_classes) # 设置关键点数 num_keypoints 4 in_features_kp model.roi_heads.keypoint_predictor.kps_score_lowres.in_features model.roi_heads.keypoint_predictor torchvision.models.detection.roi_heads.KeypointRCNNPredictor( in_features_kp, num_keypoints )有几点需要特别说明第一个num_classes 2是因为我们用了一个目标类别加一个背景类别。如果检测两种不同目标类别比如工件A和工件B这里就要设置为3。第二个num_keypoints是你定义的关键点数量。COCO预训练模型默认检测17个关键点人体姿态我们需要替换成自己的关键点数。如果显存有限可以考虑冻结backbone的部分层只训练roi_heads部分。但在我这个数据量下约2000张图即使全量微调RTX 3060也能跑得动。训练配置上我采用如下参数输入图片缩放到800x800以内保持宽高比初始学习率1e-4优化器用SGDmomentum0.9, weight_decay1e-4batch size 412G显存刚好训练20个epoch第14和第17个epoch把学习率衰减一半实际跑起来大概一个epoch在3分钟左右正好可以边跑边观察Tensorboard的loss曲线。4. 推理部署与常见问题排查4.1 推理代码与结果可视化训练完成后推理代码比训练代码简单很多。加载模型后设置model.eval()模式直接传入预处理后的图片即可。但这里有一个很容易被忽视的坑Keypoint R-CNN的输出结果是浮点坐标不是整数。如果你需要像素级的精确位置不要直接round取整后续最好配合其他算法进一步处理。标准的推理代码如下import torch import torchvision.transforms as T from PIL import Image # 加载训练好的模型 model torchvision.models.detection.keypointrcnn_resnet50_fpn( weightsNone, num_classes2, num_keypoints4 ) checkpoint torch.load(best_model.pth, map_locationcpu) model.load_state_dict(checkpoint[model_state_dict]) model.eval() # 推理单张图片 image Image.open(test.jpg).convert(RGB) transform T.Compose([T.ToTensor()]) image_tensor transform(image).unsqueeze(0) with torch.no_grad(): predictions model(image_tensor) # predictions[0]是一个字典包含以下键 # - boxes: [N, 4]检测框坐标 # - scores: [N]置信度 # - labels: [N]类别id # - keypoints: [N, K, 3]关键点坐标可见性 boxes predictions[0][boxes].numpy() scores predictions[0][scores].numpy() keypoints predictions[0][keypoints].numpy() # 过滤低置信度的框 threshold 0.7 keep scores threshold boxes boxes[keep] keypoints keypoints[keep] print(f检测到{len(boxes)}个目标)可视化时可以通过keypoints[:, :, 2]判断该关键点的预测置信度低于阈值比如0.5的点往往不可靠画图时用不同颜色区分即可。4.2 训练过程中的典型问题排查在实际操作中我遇到了几个问题分享出来供大家参考这些都属于常规操作中很容易踩到但很难搜到具体解决方案的类型问题一loss完全不动现象第一个epoch loss大概在1.7左右训练10个epoch后还在1.5附近徘徊几乎没有变化。排查思路检查num_keypoints设置是否正确。如果设置成了COCO的17但数据集的keypoints长度是4模型会在运行时丢出维度不匹配的错误或者默默学不到东西。检查target中的keypoints是否有正确的shape[N, K, 3]以及每个keypoint是否有可见性标记。很多自定义数据集转换代码会把v设为0这就等于告诉模型所有点都没有标注模型自然学不到任何东西。如果标记正确但loss还是不降试试用原始COCO图片做一个subset只保留前100张跑通全流程验证代码本身没有bug再放大数据集。问题二检测框位置总是偏移现象模型能检测到目标的框但框总是偏向某一侧没有紧紧包裹目标。这在数据量少且标注框很粗糙时经常出现。解决办法是检查数据标注的bbox是否太紧或太松。对于关键点检测来说bbox不需要特别精确因为它只是给RoIAlign做区域裁剪的参考。但是bbox过于不精确会让RoIAlign裁剪到的区域大量包含背景影响关键点特征提取。问题三某些关键点位置总是预测不准这类问题一般出在关键点语义不够明确的位置。比如我检测一个黑色工件上的黑孔关键点与背景对比度很低模型很容易把注意力放在纹理复杂的地方。一个缓解方案是加大对应目标的图像数量并弱化关键点附近的背景干扰。另一个方案是训练时用更小的学习率跑更多epoch让网络有更多时间收敛。如果还是不改善就要反思数据标注时点是否真的打在了正确位置。4.3 显存优化与推理速度的实测数据在RTX 3060 12G上batch size4训练时显存占用约6.8G稳稳的。如果显存更小比如8G建议把batch size降为2同时把最长边缩到640。实测batch size2和batch size4在最终精度上几乎没有差别差距在0.5%以内但对显存需求可下降40%。推理速度方面在一张1080P图片上不使用TensorRT加速的话RTX 3060的推理时间大概是210ms约4.7 FPS。如果对实时性有更高要求可以考虑把ResNet50 FPN的backbone换为MobileNetV3等轻量级网络但需要在torchvision层面做更多定制训练难度也会上升。我个人建议先跑通标准方案确认业务需求真的有高性能需求时再优化。5. 增强技巧与后续扩展方向5.1 扩大数据量的实用策略对于自建数据集很多人的痛点不是模型选型而是数据量不够。我这边能分享几个有效的方法多角度拍摄同样一个目标从不同角度、不同光照条件拍多张图。关键点检测模型对视角变化和光照变化比较敏感多角度拍摄能显著提升泛化能力。简单颜色抖动在数据增强中加入HueSaturationValue扰动随机改变色相、饱和度、明度成本极低但效果明显特别是对于目标颜色和背景颜色区分度较低的情况。合成数据如果你检测的目标是工业零件或固定结构物可以用3D建模软件批量渲染不同角度和光照下的图片然后在渲染图上标注关键点。这类合成数据的标注精度远高于人工标注用来做预训练非常合适。5.2 结合下游任务做后处理关键点检测往往不是最终目的它后续可能会接一个测量算法或校正逻辑。比如我做的这个矩形工件拿到四个角点后通过计算相邻点之间的距离就得到了工件的物理尺寸再结合相机标定参数可以做实际的毫米级测量。这里有一个经验关键点检测模型的输出是像素坐标在实际工程中直接使用会有亚像素级的偏差。如果对精度要求高可以在关键点周围取一个小窗口用OpenCV的cornerSubPix做亚像素精细化。实测下来可以把标准差从1.2像素降到0.4像素左右效果显著。5.3 模型训练到部署的完整链路建议最后给一个整体性建议项目不要停留在训练完模型、跑通推理就结束。我个人的经验是做到部署评估才算真正完成。具体来说至少应做以下三件事准备一个与训练集完全独立且分布不同的测试集比如换一台相机、换一个光照环境拍摄。在测试集上跑一次完备的评测指标比如OKSObject Keypoint SimilarityCOCO关键点任务标准指标不要只看Train loss。把模型导出为TorchScript这样后续可以方便地集成到C或Java部署环境避免被Python依赖绑死。拿TorchScript导出很简单# 导出为TorchScript model.eval() example torch.rand(1, 3, 800, 800) traced_script_module torch.jit.trace(model, example, strictFalse) traced_script_module.save(keypoint_rcnn_traced.pt)注意这里需要设置strictFalse因为模型内部的一些操作在trace时可能会触发非严格模式的警告但不影响实际功能。写在最后的一点心得体会跑完Keypoint R-CNN这个项目我最大的感受是这类官方预训练模型微调的方案确实能帮人把很大一部分精力从繁琐的模型结构设计中解放出来让你能专注在数据质量、业务逻辑这些真正决定项目成败的事情上。如果你的项目场景碰巧能用两阶段检测器的范式来覆盖Keypoint R-CNN无疑是一个可靠且低成本的起点。最后提醒一句标注数据时多花点心思把点打准了后面训练、调参、部署阶段能帮你省下几倍的时间。本文还有配套的精品资源点击获取
返回列表