尧图精选

PyTorch混合表示6D姿态估计:遮挡场景实战与自监督优化

🕒 发布时间:2026/9/23 21:28:58 📁 来源:尧图网络
简介这份资源面向计算机视觉方向的研究者与开发者聚焦基于PyTorch的6D物体姿态估计实战解决从2D图像中确定物体三维位置与旋转共6个自由度的核心问题可应用于机器人抓取、虚拟现实与自动驾驶等场景。项目以混合表示方法为主线将像素级特征与局部几何信息结合借助CNN提取图像特征并回归姿态涵盖数据预处理、模型构建、损失函数与优化器选择、MAE/MSE/ARE等评估指标以及测试与可视化全流程。压缩包共1790个文件约4.1MB以cpp与h源码为主辅以py脚本、cu核函数、cmake构建文件及txt说明文档整体呈现一个含Eigen等依赖的完整工程结构。目前已有393人学习下载。通过源码与模型文件读者可掌握深度学习模型设计、训练调优与姿态估计落地的完整思路。1. 混合表示下的 6D 姿态估计为什么单一路径总在遮挡场景翻车做机器人抓取或者 AR 装配引导的同行大概率都经历过这样的场景模型在 LINEMOD 上跑得漂漂亮亮一换到真实工位、零件被机械爪挡住一半位姿直接跳到桌子底下。6D 物体姿态估计要解决的就是从一张 RGB 或 RGB-D 图里同时回归出目标物体的三维旋转和平移也就是那个把物体从模型坐标系搬到相机坐标系的刚体变换。这件事的难点从来不在网络结构本身而在于「用什么表示姿态」——直接回归四元数会遇到旋转的周期性和多解问题逐像素投票又对纹理弱、对称性强的物体极其敏感。混合表示的核心思路是把稠密的对应关系预测和稀疏的关键点几何约束拼在一起让网络既有一像素级的稠密监督又有几何可解释的兜底。这篇笔记就围绕 PyTorch 混合表示下的 6D 姿态估计这条线把环境搭建、数据准备、网络结构、训练参数和推理后处理一步步拆开源码和模型下载的落地路径也会讲清楚适合已经跑过检测、想往位姿方向推进的工程师。2. 混合表示到底混了什么稠密对应与稀疏关键点的分工2.1 从 PVNet 到混合表示两条路线的取舍要理解混合表示得先看清前面两条主流路线各自的死穴。第一条是以 PVNet 为代表的逐像素投票网络对每个属于目标的像素预测一个指向关键点的单位向量再用 RANSAC 投票出关键点二维坐标最后 PnP 解位姿。它的好处是对遮挡鲁棒因为只要还有一部分像素可见投票方向依然能收敛坏处是当物体表面纹理重复或者关键点本身被完全遮住时投票场会出现多个峰值RANSAC 选错峰就全盘皆输。第二条是以 DeepIM 为代表的迭代精化先给一个粗略位姿渲染出对应视角再让网络回归渲染图与真实图之间的相对变换。它精度高但强依赖初始位姿且推理时要反复渲染速度上不去。混合表示做的事情是把这两条路的输出在同一个网络里并行预测一条分支输出稠密的坐标图每个像素回归它在归一化物体坐标系下的 XYZ另一条分支输出稀疏关键点的二维热力图和投票向量。稠密分支提供全局约束稀疏分支提供几何锚点两者在 PnP 阶段融合。这样设计的好处很直接。稠密坐标图在纹理丰富区域精度极高能压住整体平移误差稀疏关键点在遮挡区域依然能靠投票维持方向避免稠密分支在无纹理处乱猜。我一般会把稠密分支的损失权重设得比稀疏分支略高因为稠密监督的梯度更稳定但稀疏分支不能关它是遮挡场景的后悔药。2.2 网络结构共享 backbone 加双头输出落地时网络结构不用搞太复杂共享一个 backbone 再分两个头就够了。backbone 选 ResNet34 或者轻量的 MobileNetV3 都行前者精度稳后者适合边缘部署。下面是一个可以直接抄的最小结构输入是裁剪后的 RGB 图输出稠密坐标图和关键点热力图。import torch import torch.nn as nn import torchvision.models as models class HybridPoseNet(nn.Module): def __init__(self, num_keypoints8, backboneresnet34): super().__init__() # 共享 backbone去掉原始分类头 resnet models.resnet34(weightsmodels.ResNet34_Weights.DEFAULT) self.backbone nn.Sequential(*list(resnet.children())[:-2]) # 输出 1/32 分辨率 feat_dim 512 # 稠密分支每个像素回归物体坐标系下的 XYZ self.dense_head nn.Sequential( nn.Conv2d(feat_dim, 256, 3, padding1), nn.BatchNorm2d(256), nn.ReLU(inplaceTrue), nn.Conv2d(256, 3, 1) # 输出 3 通道 XYZ ) # 稀疏分支关键点热力图 投票向量 self.heatmap_head nn.Conv2d(feat_dim, num_keypoints, 1) self.vote_head nn.Conv2d(feat_dim, num_keypoints * 2, 1) def forward(self, x): feat self.backbone(x) dense_xyz self.dense_head(feat) heatmap self.heatmap_head(feat) vote self.vote_head(feat) return dense_xyz, heatmap, vote这段代码里几个参数值得说清楚。num_keypoints一般取 8对应物体的八个包围盒角点这是 6D 姿态里最通用的关键点定义对称物体可以减到 4 个甚至用对称轴上的点。dense_head最后用 1x1 卷积输出 3 通道是因为每个像素要回归物体坐标系下的 X、Y、Z 三个值注意这里的坐标要归一化到物体包围盒尺度内否则数值范围太大会导致训练发散。vote_head输出num_keypoints * 2是因为每个关键点要预测 x、y 两个方向的单位向量。backbone 输出是 1/32 分辨率如果对小物体精度不满意可以把list(resnet.children())[:-2]改成[:-3]拿 1/16 的特征代价是显存涨一截。2.3 损失函数三个分支怎么配权重混合表示的损失是三项相加稠密坐标用 Smooth L1热力图用 focal loss 处理正负样本极度不平衡投票向量用余弦相似度损失约束方向。权重配置是玄学重灾区我踩过的坑是热力图权重给太高导致网络只顾着找关键点、稠密分支欠拟合。import torch.nn.functional as F def hybrid_loss(dense_pred, dense_gt, hm_pred, hm_gt, vote_pred, vote_gt, w_dense1.0, w_hm0.5, w_vote0.3): # 稠密坐标只在目标 mask 内计算背景不参与 mask dense_gt.sum(dim1, keepdimTrue) ! 0 loss_dense F.smooth_l1_loss(dense_pred * mask, dense_gt * mask, reductionsum) / (mask.sum() 1e-6) # 热力图focal lossalpha 压负样本 b, k, h, w hm_pred.shape hm_pred hm_pred.view(b, k, -1) hm_gt hm_gt.view(b, k, -1) pos hm_gt.eq(1).float() neg 1 - pos loss_hm -(pos * (1 - hm_pred).pow(2) * hm_pred.clamp(1e-6).log() neg * (1 - hm_gt).pow(4) * hm_pred.pow(2) * (1 - hm_pred).clamp(1e-6).log()).mean() # 投票只对关键点邻域内的像素算方向损失 loss_vote (1 - F.cosine_similarity(vote_pred, vote_gt, dim1)).mean() return w_dense * loss_dense w_hm * loss_hm w_vote * loss_votew_dense、w_hm、w_vote这三个权重我一般从 1.0、0.5、0.3 起步如果训练日志里热力图 loss 下降很快但稠密 loss 卡住就把w_hm降到 0.3。mask那一步很关键背景像素的坐标真值是 0如果不 mask 掉网络会花大量精力去拟合背景的零值稠密分支就废了。focal loss 里的alpha和gamma这里用 2 和 4 是 CornerNet 的经典配置对关键点这种稀疏目标很合适。3. 环境搭建与数据准备从零把训练管线跑起来3.1 PyTorch 环境CUDA 版本对齐是第一步环境这块翻车最多的是 CUDA 和 PyTorch 版本对不上。我一般用 conda 建独立环境先确定显卡驱动支持的 CUDA 上限再去 PyTorch 官网找对应的安装命令。假设驱动支持 CUDA 11.8命令如下。conda create -n pose6d python3.10 -y conda activate pose6d # 按官网对应 CUDA 11.8 的命令安装不要凭记忆写版本号 pip install torch torchvision --index-url https://download.pytorch.org/whl/cu118 pip install opencv-python numpy scipy trimesh pyyaml tqdm装完必须验证别急着跑训练。import torch print(torch.__version__) print(torch.cuda.is_available()) print(torch.cuda.get_device_name(0))torch.cuda.is_available()返回 False 的话九成是装成了 CPU 版或者 CUDA 版本和驱动不匹配。这时候别硬跑CPU 上训这个网络一个 epoch 能等到天亮。另外trimesh是用来读物体三维模型的PnP 阶段要用到模型的关键点三维坐标这个包不能少。3.2 数据集格式LINEMOD 与自制数据的对齐公开数据集里 LINEMOD 是 6D 姿态的标配13 个物体每个物体约 1200 张标注图。它的标注格式是每个物体一个gt.yml里面记录每张图里目标的旋转矩阵和平移向量。自制数据的话我建议直接对齐 LINEMOD 的目录结构省得改 dataloader。dataset/ lm/ 000002/ rgb/ # 彩色图 depth/ # 深度图没有可留空 mask/ # 目标分割掩码 gt.yml # 位姿标注 models/ # 物体三维模型 .plydataloader 里要做三件事按 mask 裁剪出目标区域并缩放到固定尺寸、把位姿真值转成稠密坐标图和关键点热力图、做颜色抖动和随机旋转增强。裁剪这一步不能省整图送进去背景占比太大网络学不到东西。import cv2 import numpy as np import yaml def load_linemod_sample(rgb_path, mask_path, gt_path, obj_id, img_size256): img cv2.imread(rgb_path) mask cv2.imread(mask_path, 0) with open(gt_path) as f: gt yaml.safe_load(f) # 取该图里对应物体的位姿 pose None for item in gt: if item[obj_id] obj_id: pose item break R np.array(pose[cam_R_m2c]).reshape(3, 3) t np.array(pose[cam_t_m2c]).reshape(3, 1) # 按 mask 外接框裁剪并 padding 成正方形 ys, xs np.where(mask 0) x0, x1, y0, y1 xs.min(), xs.max(), ys.min(), ys.max() cx, cy (x0 x1) // 2, (y0 y1) // 2 half max(x1 - x0, y1 - y0) // 2 10 crop img[max(0, cy-half):cyhalf, max(0, cx-half):cxhalf] crop cv2.resize(crop, (img_size, img_size)) return crop, R, t裁剪时half多加 10 个像素是留边防止目标贴边被切掉。缩放后位姿不能直接用因为裁剪和缩放改变了相机内参必须同步更新内参矩阵否则 PnP 解出来的平移全是错的。这一步是新手最容易忽略的坑我见过有人训练 loss 降得很好一测位姿偏差半米就是内参没跟着改。3.3 稠密坐标图与热力图真值生成网络要学稠密坐标就得先造出每个像素对应的物体坐标系 XYZ。做法是用位姿真值把物体模型渲染到当前视角渲染出的每个像素在物体坐标系下的坐标就是真值。没有渲染管线的话可以用trimesh加pyrender离线渲染一批或者用深度图反投影近似。def make_dense_gt(depth, K, R, t, obj_model_points): # 深度图反投影到相机坐标系 h, w depth.shape u, v np.meshgrid(np.arange(w), np.arange(h)) z depth / 1000.0 # LINEMOD 深度单位是毫米 x (u - K[0, 2]) * z / K[0, 0] y (v - K[1, 2]) * z / K[1, 1] cam_pts np.stack([x, y, z], axis-1).reshape(-1, 3) # 相机坐标系转到物体坐标系 R_inv R.T t_inv -R_inv t.reshape(3) obj_pts (cam_pts - t_inv) R_inv.T return obj_pts.reshape(h, w, 3)depth / 1000.0是因为 LINEMOD 深度存的是毫米网络里统一用米。R_inv和t_inv是位姿的逆变换把相机系点搬到物体系。这里要注意深度为 0 的像素是无效的生成真值后要把这些位置置零并在 loss 里 mask 掉。热力图真值就是在关键点投影位置画高斯核标准差取 2 个像素太大关键点会糊在一起太小网络学不到。4. 训练参数与推理后处理把位姿从网络输出解出来4.1 训练超参学习率与 batch 的搭配训练这块backbone 用预训练权重的话初始学习率设 1e-4稠密头和稀疏头因为是随机初始化可以给 1e-3用参数组分开设。batch size 看显存256 输入下 8 到 16 都行太小 BN 统计不稳。优化器用 AdamW权重衰减 1e-4。import torch.optim as optim backbone_params list(model.backbone.parameters()) head_params list(model.dense_head.parameters()) \ list(model.heatmap_head.parameters()) \ list(model.vote_head.parameters()) optimizer optim.AdamW([ {params: backbone_params, lr: 1e-4}, {params: head_params, lr: 1e-3} ], weight_decay1e-4) scheduler optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max80, eta_min1e-6)T_max80对应 80 个 epoch余弦退火到 1e-6。如果训练 loss 在前 10 个 epoch 就平了多半是学习率太大或者数据增强太弱。我一般会先跑 5 个 epoch 的小实验看 loss 曲线形状对不对再开全量训练。混合表示的网络收敛比纯检测慢80 到 120 个 epoch 是常态别训 20 个 epoch 就下结论说方法不行。4.2 PnP 解位姿稠密与稀疏结果怎么融合推理阶段稠密分支给出每个像素的物体坐标稀疏分支给出关键点二维位置。融合策略是先用稀疏关键点做一次 PnP 得到粗位姿再用粗位姿把稠密坐标图里的点筛一遍剔除离群点最后用筛选后的稠密对应关系做第二次 PnP 精化。import cv2 import numpy as np def solve_pose_pnp(dense_xyz, dense_mask, kpts_2d, kpts_3d, K, distNone): if dist is None: dist np.zeros(5) # 第一次稀疏关键点 PnP ok, rvec, tvec cv2.solvePnP(kpts_3d, kpts_2d, K, dist, flagscv2.SOLVEPNP_EPNP) if not ok: return None, None # 用粗位姿投影稠密点剔除重投影误差大的 obj_pts dense_xyz[dense_mask].reshape(-1, 3) img_pts np.stack(np.where(dense_mask)[::-1], axis-1).astype(np.float32) proj, _ cv2.projectPoints(obj_pts, rvec, tvec, K, dist) err np.linalg.norm(proj.reshape(-1, 2) - img_pts, axis1) keep err np.percentile(err, 60) # 保留误差最小的 60% # 第二次稠密点精化 ok, rvec, tvec cv2.solvePnP(obj_pts[keep], img_pts[keep], K, dist, rvecrvec, tvectvec, useExtrinsicGuessTrue, flagscv2.SOLVEPNP_ITERATIVE) return rvec, tvecSOLVEPNP_EPNP对初始值不敏感适合第一次粗解SOLVEPNP_ITERATIVE配合useExtrinsicGuessTrue做精化。保留 60% 是经验值遮挡严重时可以降到 40%但别低于 30%点太少 PnP 会不稳。np.percentile(err, 60)这个阈值比固定阈值好因为不同物体的重投影误差量级不一样。4.3 对称物体的特殊处理对称物体是 6D 姿态的老大难一个圆柱体绕轴转任意角度看起来都一样网络学出来的位姿在对称轴上会乱跳。处理办法是在损失里对对称关键点做置换不变或者在推理后把位姿投影到对称等价类里。简单做法是定义对称轴把旋转矩阵绕轴的角度归一化到 [0, 2π/n) 区间。def canonicalize_symmetry(R, axis, n_fold): # 把旋转矩阵绕对称轴的角度折叠到基本域 from scipy.spatial.transform import Rotation as Rot r Rot.from_matrix(R) euler r.as_euler(xyz) # 假设对称轴是 z 轴折叠第三个角 euler[2] euler[2] % (2 * np.pi / n_fold) return Rot.from_euler(xyz, euler).as_matrix()n_fold是旋转对称的阶数圆柱是无穷阶实际取 12 或 24 离散化就够用。这一步放在评估前做能让对称物体的 ADD 指标好看很多但要注意它不改变实际抓取可行性只是让指标可比。5. 避坑与排查那些让位姿突然跳变的细节5.1 现象训练 loss 正常但测试位姿偏差巨大原因裁剪缩放后相机内参没同步更新PnP 用的还是原图内参。解决在 dataloader 里把裁剪偏移和缩放比例记录下来推理时用更新后的内参矩阵公式是fx fx * scalecx (cx - crop_x0) * scale。5.2 现象遮挡场景下位姿偶尔跳到物体背面原因稀疏投票场出现多峰RANSAC 选错峰。解决投票时加一个先验用稠密分支的粗位姿投影关键点位置作为投票中心限制 RANSAC 的搜索范围或者直接把稠密粗解作为唯一初值。5.3 现象稠密坐标图在无纹理区域全是噪声原因无纹理区域本身就没有可区分的特征网络只能瞎猜。解决在 loss 里对高梯度区域加权或者引入深度图作为额外输入深度在无纹理处依然有几何信息。我一般会加一个边缘权重图让网络在物体轮廓附近多学一点。5.4 现象batch size 调大后精度反而下降原因BN 层统计量在混合表示的双头结构下不稳定尤其是稀疏头输出通道少。解决把 BN 换成 GroupNorm或者冻结 backbone 的 BN 只训练 head。GroupNorm 在小 batch 下比 BN 稳代价是训练稍慢。5.5 现象模型下载后加载报 key 不匹配原因训练时用了DataParallel或DistributedDataParallel保存的 state_dict 带module.前缀。解决加载时剥掉前缀state_dict {k.replace(module., ): v for k, v in sd.items()}再load_state_dict。这个坑几乎每个人都会踩一次。6. 进阶技巧用重投影自监督把标注成本压下来真实工位的数据标注是最大的成本项位姿标注尤其贵一张图人工标关键点要几分钟。混合表示有个天然优势稠密分支的坐标预测和稀疏分支的关键点预测可以互相校验这就给了自监督的空间。具体做法是对无标注的 RGB 图先用当前模型推理出粗位姿渲染出物体模型把渲染图和真实图做光度一致性损失反向优化位姿和网络。这就是自训练加渲染精化的路子能把标注需求降一个数量级。实现上分两步。第一步是伪标签生成用训练好的模型在无标注数据上跑推理保留重投影误差小于阈值的样本作为伪标签。第二步是联合训练把伪标签样本和真实标注样本混在一个 batch 里伪标签样本只算稠密分支的损失真实样本算全部三项损失。def self_training_step(model, unlabeled_batch, optimizer, K, threshold5.0): model.train() dense_xyz, hm, vote model(unlabeled_batch[img]) # 用当前预测解粗位姿 rvec, tvec solve_pose_pnp(dense_xyz.detach().cpu().numpy(), unlabeled_batch[mask].cpu().numpy(), unlabeled_batch[kpt2d].cpu().numpy(), unlabeled_batch[kpt3d].cpu().numpy(), K) if rvec is None: return 0.0 # 渲染模型并算光度损失 rendered render_model(unlabeled_batch[model], rvec, tvec, K) photo_loss F.l1_loss(rendered, unlabeled_batch[img]) # 只回传光度损失稠密分支跟着一起优化 optimizer.zero_grad() photo_loss.backward() optimizer.step() return photo_loss.item()threshold5.0是重投影误差的像素阈值超过这个值的伪标签直接丢宁缺毋滥。光度损失用 L1 比 L2 对光照变化更鲁棒。这一步的收益在自制数据上特别明显我做过一个实验用 200 张标注加 2000 张无标注自训练精度能追到 1000 张纯标注的九成标注成本省了八成。验证自训练有没有效果别只看训练 loss要在一个独立的小验证集上看 ADD 或者 5cm5° 指标。如果自训练后指标反而降了多半是伪标签阈值太松把错误位姿也学进去了把threshold从 5.0 收到 3.0 再试。另外渲染器的精度很关键模型和真实物体的尺寸对不上光度损失会一直很高自训练就退化成噪声。我一般会先用卡尺量一下真实物体的关键尺寸把三维模型缩放到一致这个习惯帮我省过好几次返工。这套混合表示加自监督的组合落地门槛其实不高一台带 8G 显存的卡就能跑起来源码结构清晰的话两三天能复现出基线。值不值得投入取决于你的场景是不是遮挡多、纹理弱——如果是混合表示比纯关键点或纯迭代精化的方案稳得多如果场景干净、纹理丰富那用更轻的方案就够了别为了方法而方法。希望帮到你。本文还有配套的精品资源点击获取
上一篇/下一篇内容由系统自动关联 返回资讯列表 →