如何训练高成功率抓取网络?GG-CNN训练参数全解析与实战
【免费下载链接】ggcnnGenerative Grasping CNN from "Closing the Loop for Robotic Grasping: A Real-time, Generative Grasp Synthesis Approach" (RSS 2018)项目地址: https://gitcode.com/gh_mirrors/gg/ggcnn
GG-CNN(Generative Grasping CNN)是一款轻量级机器人抓取网络,能在深度图的每个像素上实时预测抓取质量与位姿,源自 RSS 2018 经典论文的 PyTorch 实现。本文将逐个解析train_ggcnn.py的训练参数,带你从零训练出高成功率的抓取网络模型。
🧠 30 秒认识 GG-CNN:像素级抓取预测网络
GG-CNN 是一个全卷积网络:输入一张 300×300 的深度图,一次前向传播即可输出 4 个通道——
- pos:每个像素的抓取质量(成功率概率)
- cos / sin:抓取角度(用三角函数编码,避免角度不连续问题)
- width:抓取的宽度
这种"生成式"设计让网络可以实时闭环运行:夹爪每次抓取前都重新看一眼场景再决策,即使物体在抓取过程中被碰动也能准确抓取。项目中提供两个网络版本,定义在 models/ggcnn.py 与 models/ggcnn2.py:
| 版本 | 结构特点 | 适用场景 |
|---|---|---|
ggcnn | 卷积编码器 + 转置卷积解码器,32/16/8 通道 | 轻量、快速,经典复现 |
ggcnn2 | 空洞卷积(dilation 2、4)扩大感受野 | 精度更高,适合复杂场景 |
🛠 环境准备:一条命令装好依赖
GG-CNN 训练脚本硬编码使用cuda:0设备,所以需要NVIDIA GPU + CUDA 版 PyTorch。依赖清单非常简洁(requirements.txt):
git clone https://gitcode.com/gh_mirrors/gg/ggcnn cd ggcnn pip install -r requirements.txt核心依赖包括:torch、torchvision、opencv-python、scikit-image、tensorboardX、torchsummary。
📦 数据集准备:Cornell 与 Jacquard 二选一
项目内置两套经典抓取数据集的加载器:
- Cornell 数据集(
utils/data/cornell_data.py):需要先把手工标注的 PCD 文件转换成深度图:
python -m utils.dataset_processing.generate_cornell_depth <Cornell数据集路径>- Jacquard 数据集(
utils/data/jacquard_data.py):规模更大、更接近真实场景,下载解压后可直接使用。
数据读取统一封装在 utils/data/grasp_data.py,每次取样本时会自动做随机旋转(0°/90°/180°/270°)和随机缩放(0.5~1.0 倍),相当于内置的数据增强,无需额外配置。
🚀 快速上手:一行命令启动训练
# 在 Cornell 数据集上训练 GG-CNN python train_ggcnn.py --description my_first_run --network ggcnn \ --dataset cornell --dataset-path <数据集路径> # 在 Jacquard 数据集上训练精度更高的 GG-CNN2 python train_ggcnn.py --description my_ggcnn2 --network ggcnn2 \ --dataset jacquard --dataset-path <数据集路径>训练脚本 train_ggcnn.py 会自动创建带时间戳和实验名的输出目录,模型保存到你无法忽视的位置:output/models/<时间戳_描述>/,文件名直接带有验证集得分,例如epoch_23_iou_0.91,一眼就能挑出最佳模型。
🎛 核心训练参数全解析
运行python train_ggcnn.py --help可查看全部参数,下面是最关键的配置项:
| 参数 | 默认值 | 作用说明 |
|---|---|---|
--network | ggcnn | 选择ggcnn或ggcnn2网络结构 |
--dataset | 必填 | cornell或jacquard |
--dataset-path | 必填 | 数据集根目录路径 |
--use-depth | 1 | 是否使用深度图输入(1 开 / 0 关) |
--use-rgb | 0 | 是否叠加 RGB 图像输入(1 开 / 0 关) |
--split | 0.9 | 训练集占比,剩余作验证集 |
--ds-rotate | 0.0 | 旋转数据集划分起点,用于交叉验证 |
--batch-size | 8 | 批大小 |
--epochs | 50 | 训练轮数 |
--batches-per-epoch | 1000 | 每轮训练取多少个 batch |
--val-batches | 250 | 每轮验证取多少个 batch |
--num-workers | 8 | 数据加载线程数 |
--outdir | output/models/ | 模型输出目录 |
--logdir | tensorboard/ | TensorBoard 日志目录 |
--vis | 关 | 打开 OpenCV 窗口实时观看训练过程 |
关键参数怎么选?
- 输入通道数由
--use-depth和--use-rgb共同决定:1×use_depth + 3×use_rgb。只用深度图是 1 通道(最经典配置);加 RGB 是 4 通道,适合深度传感器有噪声的场景。 --batches-per-epoch是精髓设计:不同数据集规模差异很大(Cornell 约 900 张,Jacquard 上万张),固定"每轮 1000 个 batch"让两套数据集的训练量等价,方便横向比较。--ds-rotate用于交叉验证:它把数据集列表整体平移后再划分训练/验证集,可以检验模型是否只对特定划分"背答案"。
⚙️ 训练过程揭秘:损失函数与最佳模型自动保存
打开train_ggcnn.py的训练循环,你会发现整个流程被设计得很"省心":
- 四通道 MSE 损失:pos、cos、sin、width 各自算均方误差后直接相加(见
compute_loss),无额外权重,简单有效; - Adam 优化器:默认学习率,无调度策略,对轻量网络足够;
- 每轮自动验证:取验证集 250 个 batch,用 IoU 匹配计算 Top-1 抓取成功率(
utils/dataset_processing/evaluation.py); - 最佳模型自动保存:每当验证 IOU 超过历史最佳、第 0 轮、或每 10 轮,都会同时保存完整模型和 state_dict 两份权重;
- TensorBoard 全记录:
train_loss、val_loss、IOU及各分项损失全部写入日志,tensorboard --logdir tensorboard/即可查看训练曲线; - 输出平滑:推理时对 4 个输出做高斯滤波(models/common.py 的
post_process_output),抑制像素级抖动,让抓取框更稳定。
📈 提升抓取成功率的 4 个实战技巧
- 换 GG-CNN2:
--network ggcnn2的空洞卷积带来更大感受野,在 Jacquard 这类复杂场景上通常更稳; - 加入 RGB 输入:深度图有缺失或噪点时,加
--use-rgb 1提供纹理信息兜底; - 多划分交叉验证:用
--ds-rotate 0.33、--ds-rotate 0.66换两批验证集各跑一次,确认成功率不是偶然; - 盯住验证 IOU 而非训练损失:训练损失会持续下降,但验证 IOU 才是真实成功率,以
output/models/中 IOU 最高的那份权重为准。
✅ 训练完成后:评估与可视化
用 eval_ggcnn.py 对训练好的模型做最终评估:
python eval_ggcnn.py --network output/models/xxx/epoch_23_iou_0.91 \ --dataset cornell --dataset-path <数据集路径> --iou-eval --vis--iou-eval:用抓取矩形 IoU 指标评估成功率--vis:绘制网络输出热图与预测抓取框,肉眼直观检查--jacquard-output:生成 Jacquard 官方仿真评测所需的输出格式
📁 关键文件导航
| 文件 | 说明 |
|---|---|
| train_ggcnn.py | 训练入口,所有训练参数在此定义 |
| eval_ggcnn.py | 模型评估与可视化入口 |
| models/ggcnn.py / models/ggcnn2.py | 两个网络结构实现 |
| models/common.py | 输出后处理(角度解码、高斯滤波) |
| utils/data/ | Cornell / Jacquard 数据集加载器 |
| utils/dataset_processing/ | 深度图生成、抓取框绘制与 IoU 评估 |
| requirements.txt | Python 依赖清单 |
🎯 小结
训练 GG-CNN 抓取网络其实只需三步:装依赖 → 备数据 → 跑train_ggcnn.py。真正拉开成功率差距的是参数选择——用ggcnn2换更强感受野、用--use-rgb 1补全信息、用--ds-rotate做交叉验证,最后以验证 IOU 最高的模型作为部署权重。整个项目代码量小、结构清晰,非常适合作为理解"生成式抓取合成"思路的入门项目。
【免费下载链接】ggcnnGenerative Grasping CNN from "Closing the Loop for Robotic Grasping: A Real-time, Generative Grasp Synthesis Approach" (RSS 2018)项目地址: https://gitcode.com/gh_mirrors/gg/ggcnn
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考