Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
234 changes: 234 additions & 0 deletions docs/2026-08-07_G2适配训练推理说明.md
Original file line number Diff line number Diff line change
@@ -0,0 +1,234 @@
# G2 适配 DreamZero 训练与推理说明

本文只说明本 PR 保留的最小 G2 适配主路径:G2 joint-space 数据转换、relative action 训练、checkpoint 预检、G2 推理服务和真机客户端解码。

## 适配目标

G2 适配让 DreamZero 使用 G2 的三路视觉、双臂关节状态和双臂动作进行训练,并在推理时输出 G2 可执行的 16 维关节目标。

核心数据契约:

- 视觉输入:`top_head`、`hand_left`、`hand_right` 三路相机。
- 状态输入:16 维,布局为左臂 7 维关节、左夹爪、右臂 7 维关节、右夹爪。
- 动作输出:16 维,布局与状态一致。
- action horizon:24。
- 训练帧数:`num_frames=33`。
- 图像尺寸:`320x176`。
- embodiment tag:`g2`。

## 数据链路

G2 原始 LeRobot 数据先转换为 DreamZero/GEAR 格式:

```bash
python scripts/data/convert_lerobot_g2_to_gear.py \
--source /path/to/g2_lerobot_dataset \
--output /path/to/g2_gear_dataset \
--test-episodes 10 \
--video-width 320 \
--video-height 176
```

转换脚本会:

- 校验 state/action 是否符合固定 16 维 G2 joint layout。
- 把三路相机映射为 `video.top_head`、`video.hand_left`、`video.hand_right`。
- 生成 `train/` 和 `test/` 两个物理隔离 split。
- 写入 `meta/info.json`、`meta/modality.json`、`meta/embodiment.json`、`meta/stats.json`。
- 写入 `meta/relative_stats_dreamzero.json`,用于 relative action 归一化和 motion mask。

训练时 `G2_DATA_ROOT` 必须指向转换后的 `train/` 目录。

## 配置适配

G2 modality 注册在:

```text
groot/vla/configs/data/dreamzero/base_48_wan_fine_aug_relative.yaml
```

G2 的 state/action key:

```yaml
state:
- state.left_joint_position
- state.left_gripper_position
- state.right_joint_position
- state.right_gripper_position
action:
- action.left_joint_position
- action.left_gripper_position
- action.right_joint_position
- action.right_gripper_position
```

所有 state/action key 使用 `q99` 归一化,最后通过 `ConcatTransform` 拼成模型输入。

G2 relative action 数据配置在:

```text
groot/vla/configs/data/dreamzero/g2_relative.yaml
```

关键设置:

```yaml
relative_action: true
relative_action_keys:
- left_joint_position
- right_joint_position
active_hold_index_path: ${g2_data_root}/meta/g2_active_hold_windows.json
active_window_ratio: 0.8
```

含义是:双臂关节学习相对当前状态的增量,夹爪保持 policy space 数值;训练采样时用 active/hold window 提高有动作片段的比例。

## 模型适配

G2 embodiment tag 注册在:

```text
groot/vla/data/schema/embodiment_tags.py
```

```python
G2 = "g2"
```

transform 侧增加 G2 embedding id:

```text
groot/vla/configs/model/dreamzero/transform/base.yaml
```

训练脚本显式覆盖:

```bash
++model_specific_transform.embodiment_tag_mapping.g2=33
```

action head 的核心改动在:

```text
groot/vla/model/dreamzero/action_head/wan_flow_matching_action_tf.py
```

本 PR 保留的核心能力:

- 支持 G2 的 16 维 action 输出契约。
- 支持 action loss / dynamics loss 权重配置。
- 支持 motion mask,根据 `relative_stats_dreamzero.json` 把归一化 action delta 还原到弧度尺度后判断是否为有效运动。
- 支持 LoRA checkpoint 按训练时 base model 加载,避免 LoRA delta 套到错误底座。

checkpoint 加载相关逻辑在:

```text
groot/vla/model/dreamzero/base_vla.py
groot/vla/model/n1_5/sim_policy.py
```

LoRA-only checkpoint 推理时会读取训练配置里的 `pretrained_model_path`,先恢复 DreamZero base,再加载 LoRA 权重。

## 训练主路径

主训练入口:

```text
scripts/train/train_dreamzero_g2_joint_lora.sh
```

典型运行:

```bash
G2_DATA_ROOT=/data/.../g2_gear/train \
OUTPUT_DIR=/data/.../dreamzero_g2_joint_lora \
PRETRAINED_MODEL_PATH=/data/.../DreamZero-AgiBot \
WAN_CKPT_DIR=/data/.../Wan2.1-I2V-14B-480P \
TOKENIZER_DIR=/data/.../umt5-xxl \
GPU_IDS=4,5,6,7 \
bash scripts/train/train_dreamzero_g2_joint_lora.sh
```

关键 Hydra 参数:

```bash
data=dreamzero/g2_relative
train_architecture=lora
num_frames=33
action_horizon=24
num_views=3
image_resolution_width=320
image_resolution_height=176
save_lora_only=true
max_chunk_size=4
```

训练脚本会在启动前检查:

- G2 数据目录和 Wan/T5/CLIP/VAE checkpoint 是否存在。
- `meta/info.json`、`meta/modality.json`、`meta/embodiment.json`、`meta/stats.json`、`meta/relative_stats_dreamzero.json` 是否存在。
- episode 数、parquet 数、视频数、G2 16 维 state/action、三路相机是否符合预期。

## 推理主路径

G2 推理服务入口:

```text
scripts/run_g2_server_9443_final.sh
```

典型运行:

```bash
MODEL_PATH=/data/.../checkpoint-1500 \
WAN_CKPT_DIR=/data/.../Wan2.1-I2V-14B-480P \
TOKENIZER_PATH=/data/.../umt5-xxl \
PORT=9443 \
bash scripts/run_g2_server_9443_final.sh
```

启动前脚本会执行:

```bash
python scripts/audit_g2_checkpoint.py "$MODEL_PATH"
```

预检会确认 checkpoint 的 action horizon、`num_frames`、action 维度、base model 路径和 LoRA/action 权重结构是否满足 G2 契约。

服务端使用:

```bash
python -m torch.distributed.run \
--standalone \
--nproc_per_node=2 \
socket_optimized_AR_g2.py \
--port "$PORT" \
--model-path "$MODEL_PATH" \
--wan-ckpt-dir "$WAN_CKPT_DIR" \
--tokenizer-path "$TOKENIZER_PATH" \
--embodiment-tag g2
```

真机客户端在:

```text
robot_live_client_g2.py
```

relative action checkpoint 的输出不是直接下发的绝对目标。客户端会用当前执行边界的机器人状态还原:

```text
absolute_target = current_state + decoded_relative_delta
```

然后再做 GDK 关节限位裁剪并下发到 G2。

## 最小闭环

1. 用 `scripts/data/convert_lerobot_g2_to_gear.py` 生成 GEAR train/test 数据。
2. 用 `scripts/data/build_g2_active_hold_windows.py` 生成 `meta/g2_active_hold_windows.json`。
3. 确认 `meta/embodiment.json` 为 `{"embodiment_tag": "g2"}`。
4. 用 `scripts/train/train_dreamzero_g2_joint_lora.sh` 训练 relative-action LoRA。
5. 用 `scripts/audit_g2_checkpoint.py` 校验 checkpoint。
6. 用 `scripts/run_g2_server_9443_final.sh` 启动 G2 server。
7. 用 `robot_live_client_g2.py` 按 relative action 模式连接服务端并执行。
3 changes: 3 additions & 0 deletions groot/vla/configs/conf.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -33,6 +33,7 @@ trainer:
profile_record_shapes: false # record tensor shapes (adds overhead)
profile_with_stack: false # record Python stack traces
profile_memory: false # record memory allocation
milestone_save_steps: ${milestone_save_steps}

# === Training Arguments ===

Expand Down Expand Up @@ -67,6 +68,7 @@ num_train_epochs: 1000
max_steps: -1
save_strategy: steps
save_steps: 500
milestone_save_steps: null
eval_strategy: "no" # there has to be a double quote; otherwise a bare `no` will be interpreted as False
save_total_limit: 8
report_to: wandb
Expand All @@ -80,6 +82,7 @@ eval_bf16: true
torch_compile_mode: null

pretrained_model_path: null
pretrained_lora_path: null
only_tune_projectors: false

save_llm: false
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -345,10 +345,86 @@ transform_yam:
# Modality Configs
################################################################################

################################################################################
# g2 (Unitree G2: dual-arm joints + grippers, state/action 16, 3 views)
################################################################################

modality_config_g2:
video:
_target_: groot.vla.data.dataset.ModalityConfig
delta_indices: [0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16, 17, 18, 19, 20, 21, 22, 23, 24]
eval_delta_indices: [-3,-2,-1,0]
modality_keys:
- video.top_head
- video.hand_left
- video.hand_right
state:
_target_: groot.vla.data.dataset.ModalityConfig
delta_indices: [0]
modality_keys:
- state.left_joint_position
- state.left_gripper_position
- state.right_joint_position
- state.right_gripper_position
action:
_target_: groot.vla.data.dataset.ModalityConfig
delta_indices: [0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16, 17, 18, 19, 20, 21, 22, 23]
modality_keys:
- action.left_joint_position
- action.left_gripper_position
- action.right_joint_position
- action.right_gripper_position
language:
_target_: groot.vla.data.dataset.ModalityConfig
delta_indices: [0]
modality_keys:
- annotation.language.action_text

transform_g2:
_target_: groot.vla.data.transform.ComposedModalityTransform
transforms:
- <<: *totensor_cfg
apply_to: ${modality_config_g2.video.modality_keys}
- <<: *crop_cfg
apply_to: ${modality_config_g2.video.modality_keys}
- <<: *resize_cfg
apply_to: ${modality_config_g2.video.modality_keys}
- <<: *color_jitter_cfg
apply_to: ${modality_config_g2.video.modality_keys}
- <<: *to_numpy_cfg
apply_to: ${modality_config_g2.video.modality_keys}

- _target_: groot.vla.data.transform.StateActionToTensor
apply_to: ${modality_config_g2.state.modality_keys}
- _target_: groot.vla.data.transform.StateActionTransform
apply_to: ${modality_config_g2.state.modality_keys}
normalization_modes:
state.left_joint_position: q99
state.left_gripper_position: q99
state.right_joint_position: q99
state.right_gripper_position: q99

- _target_: groot.vla.data.transform.StateActionToTensor
apply_to: ${modality_config_g2.action.modality_keys}
- _target_: groot.vla.data.transform.StateActionTransform
apply_to: ${modality_config_g2.action.modality_keys}
normalization_modes:
action.left_joint_position: q99
action.left_gripper_position: q99
action.right_joint_position: q99
action.right_gripper_position: q99

- _target_: groot.vla.data.transform.ConcatTransform
video_concat_order: ${modality_config_g2.video.modality_keys}
state_concat_order: ${modality_config_g2.state.modality_keys}
action_concat_order: ${modality_config_g2.action.modality_keys}
- ${model_specific_transform}

modality_configs:
oxe_droid: ${modality_config_oxe_droid}
agibot: ${modality_config_agibot}
yam: ${modality_config_yam}
g2: ${modality_config_g2}

################################################################################
# Transforms
Expand All @@ -358,6 +434,7 @@ transforms:
oxe_droid: ${transform_oxe_droid}
agibot: ${transform_agibot}
yam: ${transform_yam}
g2: ${transform_g2}

################################################################################
# Metadata Versions
Expand All @@ -367,10 +444,12 @@ metadata_versions:
oxe_droid: '0221'
agibot: '0221'
yam: '0221'
g2: '0221'

################################################################################
# FPS (per embodiment, null means use dataset default)
################################################################################

fps:
yam: 30
g2: 30
Loading