ARTICLE DETAIL

资讯详情

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

PyTorch自建关键点数据集训练Keypoint R-CNN指南

PyTorch自建关键点数据集训练Keypoint R-CNN指南 简介基于PyTorch关键点R-CNN训练自建数据集的完整工程资料包面向有深度学习基础、做姿态估计或人脸关键点任务的开发者。压缩包共116个文件约8.55MB含8个Python脚本、2个Jupyter Notebook、31个JSON标注、34个TXT说明及39张JPG图片其中模型训练和标签转换两个Notebook可还原从标注整理到训练评估的全流程。当前已有166人浏览学习。资源覆盖数据准备、预处理、模型配置、训练评估与部署等环节并提供从标注转换到训练推理的完整实现思路适合用来快速搭建自己的关键点检测模型并扩展改进。1. 自建关键点数据集为什么先看keypoint_rcnn做关键点检测的pytorch实战第一反应往往是MediaPipe那种开箱即用的方案或者自己搭一个回归网络直接吐坐标。但换成自建数据集、目标物体又是工业零件或动物这类非标准类别时这两条路都不太稳MediaPipe的拓扑固定自研小网络又容易在遮挡、多目标、小目标场景下翻车。torchvision里的keypoint_rcnn属于两阶段检测加关键点热力图回归COCO预训练权重可以直接作为底座换掉分类头、换掉关键点头就能训自己的类别。这篇笔记把搭建自建数据集的标注结构、训练脚本、参数设置和坑位一次说透读完能直接照着一套流程跑通训练和推理。整个方案的关键不是模型本身而是数据格式和坐标对齐。keypoint_rcnn对输入标注有严格要求bbox、keypoints、可见性标记这三样东西错了模型再强也训不出来。下面从模型结构开始把适配逻辑讲清楚。2. keypoint_rcnn的结构与自建数据集的适配点2.1 检测头加热力图头输出到底长什么样keypoint_rcnn本质上是Mask R-CNN去掉mask分支换成keypoint分支检测部分沿用Faster R-CNN的RPN加ROIAlignkeypoint分支输出的是K个通道的低分辨率热力图默认是56×56。训练时对每个标注点生成高斯响应图推理时在这个热力图里找峰值再映射回原图坐标整个过程不需要自己解码坐标模型直接吐出影像坐标。关键点分支输出的形状是(batch_rois, K, 56, 56)K就是你要预测的关键点数量。torchvision的keypointrcnn_resnet50_fpn预训练版本K17对应COCO人体关键点。用在自建数据集上第一个要解决的问题就是你的关键点数量不等于17或者你不需要17个点比如只测4个角点那输出层必须重建。import torchvision from torchvision.models.detection import keypoint_rcnn # 自建数据集K6目标类别只有1类 model keypoint_rcnn.keypointrcnn_resnet50_fpn( pretrainedFalse, num_classes2, # 背景 1个目标类别 num_keypoints6, box_detections_per_img10 ) # 加载COCO预训练权重跳过形状不匹配的层 checkpoint torchvision.models.detection.keypoint_rcnn.keypointrcnn_resnet50_fpn( pretrainedTrue ).state_dict() # 手动删除关键点头和分类头的权重 for k in list(checkpoint.keys()): if keypoint_predictor in k or box_predictor in k: checkpoint.pop(k) model.load_state_dict(checkpoint, strictFalse) print(model.roi_heads.keypoint_predictor.kps_score_lowres)这段代码的关键在于strictFalse它会自动忽略形状不匹配的层。删除box_predictor和keypoint_predictor之后再加载是因为这两个输出层的输出维度与自建数据集不一致保留原始权重会直接报错或产生错误映射。num_classes2是背景加一个目标类别如果你的数据集有多个类别这里要相应加一。这里有个容易踩的细节pretrainedTrue会自动下载权重但下载失败时会回退到随机初始化。首次运行时建议先单独跑一次权重下载确认文件完整再进训练流程。2.2 为什么自建数据集优先选它而不是自研回归网络关键点检测有两条技术路线。一条是heatmap回归也就是keypoint_rcnn走的路线另一条是直接用全连接层回归坐标值常见于轻量级网络。后者实现简单模型小显存占用低但在遮挡、截断、多人或密集目标场景下精度下降很快。坐标回归本质是让网络在全局特征上直接预测数值目标的细微位移很难被激活函数表达出来。heatmap路线的优势在于每个关键点都有空间位置信息即使目标部分被遮挡只要周围上下文还在热力图仍然能给出响应。而且56×56的分辨率对关键点定位来说足够配合ROIAlign的框内特征提取小目标也能被单独放大处理。这就是为什么同类任务里标注质量没问题的情况下keypoint_rcnn通常比自研回归网络好调收敛更稳OCR评估指标也更好看。torchvision实现的损失函数里keypoint_loss用的是基于热力图的二值交叉熵loss配合高斯核生成的soft target。也就是说标注点周围不是非0即1而是以真值点为中心生成一个高斯分布网络学会的是预测这个分布而不是硬猜一个像素位置。这带来一个实际好处即使标注位置有几像素偏差训练也不会剧烈震荡。代价是显存占用比回归网络高不少。batch_size和输入尺寸都要配套调整后面训练章节会给出具体建议。2.3 模型结构里需要改动的三个位置自建数据集要改三处。第一是roi_heads.box_predictor它决定检测头的类别数第二是roi_heads.keypoint_predictor它决定关键点输出的通道数第三是RPN的anchor生成器大部分自建目标尺寸和COCO人体差异较大默认anchor可能完全覆盖不到。anchor这块很多人忽略。默认anchor sizes是(32, 64, 128, 256, 512)如果你的目标是一个小零件在800×800的输入下可能只有几十像素宽最小anchor是32像素勉强能覆盖但召回会很低。更小尺度的目标比如螺丝、芯片引脚需要把sizes改成(8, 16, 32, 64, 128)。from torchvision.models.detection.rpn import AnchorGenerator anchor_generator AnchorGenerator( sizes((8, 16, 32, 64, 128),) * 5, # 5个FPN层每层都要配 aspect_ratios((0.5, 1.0, 2.0),) * 5 ) model.rpn.anchor_generator anchor_generatorFPN有5层输出每层特征图对应的anchor感受野不同sizes必须是5个tuple。低层特征图负责小目标高层负责大目标这样设置后小目标也能产生足够多的正样本。如果目标分布比较均匀可以只把第一个tuple改成(8, 16, 32, 64, 128)其余保持默认训练时观察RPN的loss曲线再决定是否全改。3. 把原始图片整理成COCO标注训练前的数据准备3.1 COCO关键点标注的JSON结构字段一个都不能少keypoint_rcnn的训练接口直接吃COCO格式的target字典所以数据准备的核心是生成一个符合COCO规范的关键点标注JSON。这个JSON包含三个顶层数组。{ images: [ { id: 0, file_name: img_001.jpg, height: 1080, width: 1920 } ], annotations: [ { id: 0, image_id: 0, category_id: 1, keypoints: [412, 318, 2, 500, 402, 2, 610, 390, 2, 545, 505, 2], num_keypoints: 4, bbox: [390, 280, 260, 270], area: 70200, iscrowd: 0 } ], categories: [ { id: 1, name: circuit_board, keypoints: [corner_1, corner_2, corner_3, corner_4], skeleton: [[1, 2], [2, 3], [3, 4], [4, 1]] } ] }核心字段里keypoints是长度为3K的平铺数组每三个数代表一个关键点前两个是x、y坐标第三个是可见性v。v2表示可见且已标注v1表示被遮挡但位置可推断v0表示不在图内不需要预测。注意如果一张图里某个目标的关键点全不可见那这个目标的annotation可以直接剔除不参与训练。bbox有几个容易被忽视的要求。bbox不一定要严格框住整个目标但建议框住所有可见关键点之后外扩10到20像素。原因在于keypoint head只在ROI内部提取特征如果bbox把关键点切在边缘ROIAlign的采样会让关键点特征被截断。目标整体出现的场景用目标检测框做bbox也可以但如果bbox来自另一个检测模型要检查框和关键点是否对齐。area字段是bbox宽乘高COCO评估时会按area分组计算AP。iscrowd固定写0出现遮挡重叠时不要用iscrowd标记直接让两个目标各自标注。categories里的keypoints数组是字符串列表它的顺序就是关键点索引顺序skeleton定义了关键点之间的连线只用于可视化和OKS评估不影响训练。3.2 从标注工具导出到COCO一个可抄的转换脚本常见标注工具导出格式五花八门LabelMe导出的是JSON数组每张图一个文件。转换时最常见的坑是坐标坐标系不一致LabelMe的坐标是相对原图的像素坐标而部分工具会用归一化坐标转换前看清楚一个点。import json import os import glob from PIL import Image def labelme_to_coco(labelme_dir, image_dir, output_json, kp_names): images [] annotations [] ann_id 0 for img_id, label_path in enumerate(glob.glob(os.path.join(labelme_dir, *.json))): with open(label_path, r, encodingutf-8) as f: label json.load(f) img_path os.path.join(image_dir, label[imagePath]) img_w, img_h Image.open(img_path).size images.append({ id: img_id, file_name: label[imagePath], height: img_h, width: img_w }) # 取第一个目标的标注自建数据集一般单目标 for shape in label[shapes]: points shape[points] # [[x1, y1], [x2, y2], ...] if len(points) ! len(kp_names): continue xs [p[0] for p in points] ys [p[1] for p in points] x_min, y_min min(xs), min(ys) x_max, y_max max(xs), max(ys) pad 10 x_min max(0, x_min - pad) y_min max(0, y_min - pad) x_max min(img_w, x_max pad) y_max min(img_h, y_max pad) bbox [x_min, y_min, x_max - x_min, y_max - y_min] keypoints [] for x, y in points: keypoints.extend([float(x), float(y), 2]) annotations.append({ id: ann_id, image_id: img_id, category_id: 1, keypoints: keypoints, num_keypoints: len(points), bbox: bbox, area: bbox[2] * bbox[3], iscrowd: 0 }) ann_id 1 categories [{ id: 1, name: object, keypoints: kp_names, skeleton: [[i, i 1] for i in range(1, len(kp_names))] }] with open(output_json, w, encodingutf-8) as f: json.dump({images: images, annotations: annotations, categories: categories}, f) print(fconverted {len(images)} images, {len(annotations)} annotations) # 使用示例 labelme_to_coco( labelme_dir./labelme_jsons, image_dir./images, output_json./annotations/train.json, kp_names[corner_1, corner_2, corner_3, corner_4] )这个脚本按单目标场景处理每个labelme文件只取第一个shape列表。如果你的数据里有多个目标、多张图需要维护目标到图像的映射关系。pad10是经验值小目标建议打大一点大目标10像素影响不大。3.3 transform预处理缩放和翻转时关键点必须同步变换训练时输入图片会被缩放到短边800像素以内这个操作发生在dataloader里。如果直接使用torchvision.transforms里的Resize只会处理图像不会处理坐标。torchvision的v2版本内置了带box和keypoints的变换但手动实现更透明方便排查坐标问题。import torch import random from torchvision import transforms as T from PIL import Image def scale_coords(keypoints, scale_x, scale_y): kps keypoints.clone().float() kps[:, 0] * scale_x kps[:, 1] * scale_y return kps def flip_coords(keypoints, img_width): kps keypoints.clone().float() kps[:, 0] img_width - kps[:, 0] return kps class KeypointTransform: def __init__(self, min_size800, max_size1333): self.min_size min_size self.max_size max_size def __call__(self, image, target): orig_w, orig_h image.size # 等比缩放短边对齐min_size scale min(self.min_size / min(orig_w, orig_h), self.max_size / max(orig_w, orig_h)) new_w int(orig_w * scale) new_h int(orig_h * scale) image image.resize((new_w, new_h), Image.BILINEAR) scale_x new_w / orig_w scale_y new_h / orig_h # 同步缩放关键点坐标和bbox if keypoints in target: target[keypoints] scale_coords(target[keypoints], scale_x, scale_y) if boxes in target: target[boxes] target[boxes] * torch.tensor( [scale_x, scale_y, scale_x, scale_y], dtypetorch.float32 ) # 随机水平翻转关键点x坐标同步翻转 if random.random() 0.5: image T.functional.hflip(image) target[keypoints] flip_coords(target[keypoints], new_w) target[boxes][:, [0, 2]] new_w - target[boxes][:, [2, 0]] # 转成tensor并用ImageNet参数归一化 image T.functional.to_tensor(image) image T.functional.normalize( image, mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225] ) return image, target这里最容易出错的是翻转后的bbox处理。target[boxes][:, [0, 2]] new_w - target[boxes][:, [2, 0]]这一行如果直接用x w - x会得到翻转后的坐标但要注意顺序左边坐标变成右边坐标两者要交换位置。上面的写法是先取原box的[x0, x1]翻转后变成[w - x1, w - x0]所以赋值时要反着写。min_size800是模型预训练时的输入规格如果你改成min_size1024模型精度可能略升但显存消耗大增。自建数据集里目标尺寸偏小建议从800开始训练时观察小目标的召回率再决定是否升到1024。归一化参数必须跟预训练一致否则加载的权重相当于换了一个输入分布loss曲线会异常。4. 用pytorch训练keypoint_rcnn最小可复现的训练脚本与超参数4.1 训练主循环从dataloader到loss回传把标注JSON转成torchvision的target字典这是训练前最后一公里。dataloader的collate_fn要特殊处理因为每张图的target长度不一致torch默认的batch函数会报错。import os import json import torch from torch.utils.data import Dataset, DataLoader from PIL import Image from keypoint_transform import KeypointTransform class KeypointDataset(Dataset): def __init__(self, coco_json, image_dir, transformsNone): with open(coco_json, r) as f: self.coco json.load(f) self.image_dir image_dir self.transforms transforms self.img_id_to_file {img[id]: img[file_name] for img in self.coco[images]} self.ann_by_image {} for ann in self.coco[annotations]: self.ann_by_image.setdefault(ann[image_id], []).append(ann) self.valid_image_ids list(self.ann_by_image.keys()) def __len__(self): return len(self.valid_image_ids) def __getitem__(self, idx): img_id self.valid_image_ids[idx] img_path os.path.join(self.image_dir, self.img_id_to_file[img_id]) image Image.open(img_path).convert(RGB) anns self.ann_by_image[img_id] boxes [] keypoints [] for ann in anns: boxes.append(ann[bbox]) kps ann[keypoints] # 已经是3K平铺格式 keypoints.append(kps) if len(boxes) 1: boxes_t torch.tensor(boxes, dtypetorch.float32) kps_t torch.tensor(keypoints, dtypetorch.float32).view(1, -1, 3) else: boxes_t torch.tensor(boxes, dtypetorch.float32) kps_t torch.tensor(keypoints, dtypetorch.float32).view(len(boxes), -1, 3) target { boxes: boxes_t, labels: torch.ones((boxes_t.shape[0],), dtypetorch.int64), keypoints: kps_t, } if self.transforms: image, target self.transforms(image, target) return image, target def collate_fn(batch): images, targets zip(*batch) images list(images) targets list(targets) return images, targets train_dataset KeypointDataset( coco_json./annotations/train.json, image_dir./images, transformsKeypointTransform() ) train_loader DataLoader( train_dataset, batch_size2, shuffleTrue, num_workers4, collate_fncollate_fn, pin_memoryTrue )训练循环本身很短关键是loss的聚合方式。torchvision的模型前向在训练模式下返回一个dict五个loss加起来反传。import torch.optim as optim from torch.optim.lr_scheduler import StepLR model model.to(device) optimizer optim.SGD( model.parameters(), lr0.005, momentum0.9, weight_decay0.0005 ) scheduler StepLR(optimizer, step_size3, gamma0.1) num_epochs 20 for epoch in range(num_epochs): model.train() total_loss 0.0 for images, targets in train_loader: images [img.to(device) for img in images] targets [{k: v.to(device) for k, v in t.items()} for t in targets] loss_dict model(images, targets) losses sum(loss for loss in loss_dict.values()) optimizer.zero_grad() losses.backward() optimizer.step() total_loss losses.item() scheduler.step() avg_loss total_loss / len(train_loader) print(fepoch {epoch 1}/{num_epochs}, loss{avg_loss:.4f}) print(f objectness{loss_dict[objectness].item():.4f}, frpn_box_reg{loss_dict[loss_rpn_box_reg].item():.4f}, fbox_cls{loss_dict[loss_box_cls].item():.4f}, fbox_reg{loss_dict[loss_box_reg].item():.4f}, fkeypoint{loss_dict[loss_keypoint].item():.4f})StepLR每3个epoch把学习率缩小10倍20个epoch训练量下这个节奏合适。loss_dict里有5个键objectness、loss_rpn_box_reg、loss_box_cls、loss_box_reg、loss_keypoint。前三个对应RPN和检测头的分类与回归最后一个是关键点热力图的损失。观察这个dict能定位问题后面排查章节会详细说。4.2 必调的六个超参数与loss曲线解读第一个是lr。官方预训练微调用0.005起步自建数据集从零训练加载了预训练权重但数据分布差异大建议降到0.001。学习率过高的典型特征是前几个epoch loss反而往上走降一个数量级后曲线立刻平滑。第二个是batch_size。keypoint_rcnn是两阶段模型显存占用是同等分类模型的3到5倍。2张图一个batch几乎是入门卡12GB显存的极限。batch_size太小导致BN统计不稳但keypoint_rcnn用的是同步批量归一化每个batch只有2张图时BN近似失效所以不要额外加BN层默认结构就好。第三个是step_size。学习率衰减节奏要和epoch数匹配20个epoch的规模下每3个epoch衰减一次合理。训练数据量很大几千张时改成每5个epoch一次。第四个是box_detections_per_img。这个参数控制每张图最多输出多少个检测框推理和训练都会用到。多目标场景保持默认100单目标场景改成10反而能减少误检。第五个是min_size。输入尺寸决定了小目标的特征分辨率800是下限1024会好一些但显存占用涨一倍。目标像素宽度小于32时优先加大min_size而不是改anchor。第六个是num_workers。这个参数不影响精度但影响训练速度Windows下建议设为0或2Linux可以设4到8。设太高的典型问题是数据加载进程内存溢出现象是训练到某个epoch时突然卡死。loss曲线的解读比绝对值更重要。正常收敛的过程是objectness从2.0左右先降rpn_box_reg跟着降box_cls在1.0到2.0之间波动keypoint loss一开始很大COCO热力图要重新学后面稳步下降。如果objectness先降到0.1以下而box_reg还在高位说明RPN已经能找到目标但框回归没收敛检查bbox标注是否越界。如果keypoint loss卡在某个值反复横跳常见原因是不可见点的v值标错导致loss计算混乱。4.3 保存checkpoint与中断恢复训练中断是常态20个epoch的工程跑到第15个epoch时断电、显存被别的任务挤掉都意味着前功尽弃。保存checkpoint时要把关键信息都带上。def save_checkpoint(model, optimizer, scheduler, epoch, path): torch.save({ epoch: epoch, model_state: model.state_dict(), optimizer_state: optimizer.state_dict(), scheduler_state: scheduler.state_dict(), }, path) def load_checkpoint(model, optimizer, scheduler, path): ckpt torch.load(path, map_locationcpu) model.load_state_dict(ckpt[model_state]) optimizer.load_state_dict(ckpt[optimizer_state]) scheduler.load_state_dict(ckpt[scheduler_state]) return ckpt[epoch] 1恢复训练时epoch数要从保存的下一个epoch开始否则StepLR会重复衰减学习率。特别注意加载checkpoint时如果模型还没重建keypoint_predictor必须先重建再load否则strictTrue会报错strictFalse又会丢失关键点头的权重。5. 避坑自建关键点数据集在训练中的5个高频翻车点5.1 现象num_classes或num_keypoints与预训练权重冲突load时报错或loss爆炸你按网上的教程用pretrainedTrue加载权重然后直接把model.roi_heads.keypoint_predictor换掉train开始后第一个epoch loss就到了几十甚至NaN。原因是pretrainedTrue的权重里keypoint_predictor是17通道输出你重建后这个head的权重是随机初始化的大学习率下随机初始化的输出层会把loss拉爆。而且如果直接model.load_state_dict(pretrained_dict, strictTrue)keypoint_predictor的shape不匹配直接抛异常。解决方法是先加载预训练权重再重建head或者重建后用strictFalse加载预训练权重。重建后的关键点头学习率要单独调小可以在优化器里针对keypoint_predictor参数设一个较小的学习率比如0.0005其余层保持0.005。这样随机初始化的new head不会在前期剧烈震荡。5.2 现象训练正常、loss下降但推理时关键点坐标全部偏到图外一个常见情况是训练时loss很漂亮地降到0.2附近你兴冲冲跑推理发现box位置基本准但关键点像乱洒的芝麻落在目标外的随机位置。这种现象九成是坐标没有对齐。训练时输入图像被缩放标注的keypoints也同步缩放了这部分没问题。问题出在推理时的预处理。推理时你用cv2.resize把图片缩放后直接丢进模型模型输出的是缩放后图片上的关键点坐标你没有把它还原到原始分辨率。输出坐标是在800×1200的输入图上计算的直接叠加到1920×1080的原图上当然全歪。写推理脚本时把缩放系数留好。target[keypoints]是归一化到800尺寸的坐标推理时得到pred_keypoints要乘回原来的scale_x和scale_y。# 推理时的坐标还原 scale_x orig_w / input_w scale_y orig_h / input_h pred_keypoints[:, :, 0] * scale_x pred_keypoints[:, :, 1] * scale_y5.3 现象box_loss收敛很快keypoint_loss死活不掉训练到第8个epochobjectness和box_reg都已经掉到0.5以下keypoint loss还在3.0以上纹丝不动。原因大概率是热力图的soft target生成问题。torchvision在训练时会把关键点坐标映射到56×56的热力图空间然后以每个真值点为中心画高斯核。如果目标在原图里只有40×40像素缩放到56×56热力图后几个关键点可能全部落在一个像素附近高斯核相互重叠网络根本分不清哪个点是哪个。遇到这个问题首先检查标注的bbox是否过小。给bbox增加外扩范围让每个关键点之间的像素距离至少大于3个热力图像素。其次把min_size从800升到1024让目标在特征图上的分辨率更大。最后一种办法是减少关键点的数量如果汽配零件上有8个点但其中有几个几乎重叠在图像上合并掉冗余点比硬训更高效。5.4 现象pytorch环境搭建正常训练时显存OOM或CUDA设备不支持这类问题和pytorch安装密切相关。用anaconda配置pytorch环境时如果直接从默认channel装CPU版本训练时报错CUDA not available还是小事装了GPU版本但CUDA版本和驱动不匹配报的是CUDA error: no kernel image is available for execution on the device这种报错更隐蔽。常见原因是用conda install pytorch而不是conda install pytorch pytorch-cudaxx.x -c pytorch -c nvidia。用conda装pytorch虽然也能装上但cudatoolkit版本往往偏旧和新驱动、新显卡不一定兼容。Windows下还经常出现torch.cuda.is_available()返回True但实际计算时kernel不兼容的情况。解决方法是先跑一行确认设备可用性再开始训练。另外OOM问题把batch_size降为1关闭pin_memory减少num_workers这是最直接的后悔药。5.5 现象验证集OKS/AP指标不错可视化检出的关键点却是错的训练完用COCO API评估AP有0.8很漂亮但可视化时发现有的图片关键点标在背景上或者一个目标的点串到旁边目标上。问题出在评估指标和实际坐标的差异上。COCO的OKS评估对关键点位置有一个容忍半径半径取决于bbox面积。bbox面积大OKS容忍范围就大坐标偏了几个像素完全不扣分。所以AP高不代表像素级定位准。可视化时发现坐标偏移到背景说明关键点头的定位精度其实不够只是OKS指标没暴露出来。这种场景下需要分辨两个可能一是多个目标靠得很近ROI的特征混在了一起解决方法是降低box_detections_per_img或者提高NMS阈值二是热力图本身存在偏移比如目标边缘纹理干扰。前者是参数问题后者需要检查训练集的标注是否在关键点定义上有歧义。建议按像素误差dist sqrt((x_pred - x_gt)^2 (y_pred - y_gt)^2)除以目标尺寸单独计算一个定位误差指标比AP更直观。6. 导出与部署阶段把训练好的模型转成ONNX并跑通推理6.1 用torch.onnx.export导出关键点检测模型训练结束后把权重转成ONNX脱离pytorch环境做部署这是工程落地的常见一步。keypoint_rcnn的导出要注意输出是变长列表需要固定输入尺寸来减少导出时shape推导的错误。import torch from torchvision.models.detection import keypoint_rcnn model keypoint_rcnn.keypointrcnn_resnet50_fpn( pretrainedFalse, num_classes2, num_keypoints6, box_detections_per_img10 ) ckpt torch.load(./checkpoints/best.pth, map_locationcpu) model.load_state_dict(ckpt[model_state]) model.eval() # 固定输入尺寸导出前一定确保eval模式 dummy torch.randn(1, 3, 800, 1280) torch.onnx.export( model, dummy, ./keypoint_rcnn.onnx, opset_version17, input_names[images], output_names[boxes, scores, labels, keypoints, keypoints_scores], dynamic_axes{images: {0: batch_size}} )导出后先用onnxruntime做一次输出形状检查确认keypoints输出是(batch, N, 6, 2)和(batch, N, 6)。torchvision模型在推理时输出的boxes是绝对坐标不是归一化坐标这点和很多检测模型的导出习惯不同部署时别再做一次归一化还原。6.2 部署推理的预处理对齐和后处理细节ONNX推理时的预处理和训练时保持一致那段代码不能省ImageNet归一化的mean和std写错任何一个数模型精度都会明显下降。推理后处理里除了坐标还原还要做两个操作按scores过滤低置信度box以及按keypoints_scores过滤低置信度的关键点。import onnxruntime as ort import numpy as np sess ort.InferenceSession(./keypoint_rcnn.onnx) input_tensor np.random.randn(1, 3, 800, 1280).astype(np.float32) # 实际替换为预处理后的图 outputs sess.run(None, {images: input_tensor}) boxes, scores, labels, keypoints, kp_scores outputs keep scores[0] 0.5 valid_boxes boxes[0][keep] valid_kps keypoints[0][keep] valid_kp_scores kp_scores[0][keep] # 按关键点置信度过滤低于0.3的点置为不可见 valid_kps[valid_kp_scores 0.3] 0kp_scores阈值取0.3到0.5之间具体看你的数据分布。训练集里被遮挡的点如果很多阈值设高了会把原本能用的点也滤掉。从pytorch转onnx到部署整个过程最玄学的就是坐标还原和阈值选择我一般会把推理脚本里加一段可视化参考输出把box和keypoint画在图上人工确认一遍再谈指标。这算我个人的一个习惯每次训练完都先跑一遍可视化和距离误差检查再做评估指标汇报指望loss曲线和单张AP值就下结论很容易被细节坑到。希望帮到你。本文还有配套的精品资源点击获取
返回列表