Task5:基于 Task4 数据复现 Diffusion Policy

一、任务介绍

本次任务使用 Task4 中采集的 LeRobot Dataset v3 数据训练 Diffusion Policy,并部署到 SO-101 真机进行测试。

任务:

Place the bottle next to the national flag

数据集:

feng0724821/so101_test_record

主要输入包括:

observation.state
observation.images.wrist

输出为:

action

其中 observation.stateaction 均为六维,对应 SO-101 的五个关节和夹爪。


二、训练 Diffusion Policy

使用 LeRobot 训练 Diffusion Policy:

lerobot-train \
  --dataset.repo_id=feng0724821/so101_test_record \
  --dataset.eval_split=0.1 \
  --policy.type=diffusion \
  --output_dir=outputs/train/dp_20k \
  --job_name=dp_20k \
  --policy.device=cuda \
  --batch_size=8 \
  --steps=20000 \
  --log_freq=100 \
  --save_freq=5000 \
  --policy.scheduler_warmup_steps=500 \
  --policy.num_train_timesteps=100 \
  --policy.num_inference_steps=10 \
  --num_workers=4 \
  --eval_steps=1000 \
  --wandb.enable=true \
  --wandb.project=so101-dp \
  --policy.push_to_hub=false

主要参数:

训练步数:20000
batch size:8
训练扩散步数:100
推理扩散步数:10
GPU:RTX 4090

训练过程中通过 Weights & Biases 观察:

train_loss
eval_loss
lr
steps

模型大约在 9k step 后逐渐收敛,其中 10k checkpoint 的 eval_loss 最好,因此最终真机推理使用:

outputs/train/dp_20k/checkpoints/010000/pretrained_model/

模型随后上传至 Hugging Face:

hf upload feng0724821/so101_dp_10k \
  outputs/train/dp_20k/checkpoints/010000/pretrained_model/ \
  --repo-type model \
  --private

三、异步真机推理

由于连接 SO-101 的 MacBook Air 没有 NVIDIA GPU,因此没有直接在 Mac 上运行 Diffusion Policy。

最终采用:

MacBook Air:控制机械臂 + 采集相机
RTX 4090 服务器:运行 Diffusion Policy

两端通过 SSH 隧道连接:

Mac robot_client
      ↓
 SSH Tunnel
      ↓
Server policy_server
      ↓
Diffusion Policy

Mac 上建立隧道:

ssh -p 14027 \
  -L 8080:127.0.0.1:8080 \
  fengchangqun@121.48.170.1

服务器启动:

python -m lerobot.async_inference.policy_server \
  --host=127.0.0.1 \
  --port=8080

Mac 再启动 robot_client,负责采集 wrist 相机和机械臂状态,并执行服务器返回的动作。


四、遇到的问题

第一次运行异步推理时,服务器和客户端可以正常连接,但机械臂没有动作。

检查后发现 Diffusion Policy 默认:

n_obs_steps = 2

也就是推理需要连续两帧 observation。

但是原来的 async server 每次只传入一帧,导致输入 shape 不匹配。

因此对 policy_server.py 进行了修改:

  • 初始化 policy queue;
  • 模型加载后调用 policy.reset()
  • 将 wrist 图像映射到 observation.images.wrist
  • 保存上一帧 observation;
  • 第一次输入使用 [obs0, obs0]
  • 后续输入使用 [obs(t-1), obs(t)]

修改后服务器能够正常输出:

torch.Size([1, 32, 6])

即一次生成 32 个六维机械臂动作,SO-101 可以正常执行策略。


五、真机评估

最终使用 10k checkpoint 进行了 20 次独立真机测试。

每次测试大约 30 秒,成功判据为:

成功抓取瓶子并放到国旗旁边

结果:

测试次数 成功 失败 成功率
20 10 10 50%

由于当前使用的是 robot_client + policy_server 异步推理方式,没有自动按照 episode 记录实验结果,因此本次采用人工开始、结束并判断成功或失败。


六、失败原因

测试过程中发现,失败主要出现在:

  • 瓶子距离腕部摄像头较远;
  • 瓶子位于画面边缘;
  • 瓶子摆放角度与训练数据差异较大。

当瓶子在 wrist 相机中比较清晰、位置比较居中时,模型通常能够完成抓取和放置。

因此当前主要问题还是训练数据对不同位置和视角的覆盖不足。

后续可以增加:

边缘位置
较远位置
不同瓶子角度

等困难样本,提高模型的泛化能力。


七、结果

本次完成了:

  • 使用 Task4 数据训练 Diffusion Policy;
  • 完成 20000 steps 训练;
  • 选择 10k checkpoint 进行部署;
  • 将模型上传 Hugging Face;
  • 搭建 Mac + RTX 4090 异步推理架构;
  • 修复 Diffusion Policy 双帧 observation 问题;
  • 成功控制 SO-101 真机;
  • 完成 20 次测试;
  • 成功 10 次,成功率 50%。

数据集:

feng0724821/so101_test_record

模型:

feng0724821/so101_dp_10k

整个流程已经跑通:

SO-101 数据采集
      ↓
LeRobot Dataset
      ↓
Diffusion Policy 训练
      ↓
服务器 GPU 推理
      ↓
SO-101 真机执行

目前主要改进方向是继续增加数据量和困难视角数据,提高真机成功率。

Logo

DAMO开发者矩阵,由阿里巴巴达摩院和中国互联网协会联合发起,致力于探讨最前沿的技术趋势与应用成果,搭建高质量的交流与分享平台,推动技术创新与产业应用链接,围绕“人工智能与新型计算”构建开放共享的开发者生态。

更多推荐