基于
Diffusion Policy 训练的 Franka 机器人操作策略。模型预测
绝对关节角度(8维:7个关节 + 夹爪),可直接部署到真实机器人。
.
├── checkpoints/
│ ├── pick_up_milk/
│ │ ├── checkpoints/
│ │ │ ├── epoch=0049-val_loss=0.000.ckpt # 50 epoch checkpoint (~3.3GB)
│ │ │ └── epoch=0099-val_loss=0.000.ckpt # 100 epoch checkpoint (~3.3GB)
│ │ └── .hydra/
│ │ └── config.yaml # Hydra 配置(加载模型必须)
│ ├── stack_cup/
│ │ ├── checkpoints/
│ │ └── .hydra/
│ └── tennis_bucket_upright/
│ ├── checkpoints/
│ └── .hydra/
├── diffusion_policy/ # DP 推理代码(模型架构、normalizer 等)
├── normalizer_stats.json # 全局归一化参数
├── normalizer_stats_pick_up_milk.json
├── normalizer_stats_stack_cup.json
├── normalizer_stats_tennis_bucket_upright.json
├── gripper_conversion.py # 夹爪关节 → CGS 转换工具
├── example_data/
│ └── pick_up_milk/
│ ├── initial_frame.png # 初始帧图像
│ ├── initial_joints.npy # 初始关节角度(8维)
│ └── instruction.txt # 任务指令
├── scripts/
│ ├── dp_policy_server_franka.py # DP 推理服务器(Socket 通信)
│ └── rollout_with_dp_client_franka.py # IRASim 客户端(闭环 rollout)
└── README.md
1git clone https://huggingface.co/ewykric/dp-franka-joint
2cd dp-franka-joint
3pip install torch torchvision hydra-core omegaconf dill diffusers
1# 终端1:启动 DP 服务器
2export CUDA_VISIBLE_DEVICES=0
3python scripts/dp_policy_server_franka.py \
4 --dp_checkpoint checkpoints/pick_up_milk/checkpoints/epoch=0099-val_loss=0.000.ckpt \
5 --task_name pick_up_milk \
6 --port 9966
1# 真机闭环控制伪代码
2dp_server.reset_policy("pick up the milk", initial_joint_positions)
3
4while not done:
5 image = camera.get_image() # (H, W, 3) RGB
6 actions, terminated = dp_server.get_action(image, "pick up the milk")
7 # actions: (15, 8) 绝对关节角度
8
9 for joint_cmd in actions:
10 robot.move_to_joint_positions(joint_cmd[:7])
11 robot.set_gripper(joint_cmd[7])
12 time.sleep(1.0 / control_freq)
13
14 # 更新 DP 的观测缓冲和关节状态
15 dp_server.update_obs(camera.get_image())
16 dp_server.update_joints(robot.get_joint_positions())
1import torch, dill, hydra
2from omegaconf import OmegaConf
3
4OmegaConf.register_new_resolver("eval", eval, replace=True)
5
6# 1. 加载 hydra config
7cfg = OmegaConf.load("checkpoints/pick_up_milk/.hydra/config.yaml")
8
9# 2. 创建模型架构
10cls = hydra.utils.get_class(cfg._target_)
11workspace = cls(cfg)
12
13# 3. 加载权重
14payload = torch.load("checkpoints/pick_up_milk/checkpoints/epoch=0099-val_loss=0.000.ckpt",
15 pickle_module=dill)
16workspace.ema_model.load_state_dict(payload['ema'])
17policy = workspace.ema_model
18
19# 4. 设置 normalizer(从 JSON 加载,参考 dp_policy_server_franka.py)