diff --git a/action_compare_outputs/hist_train_abs_vs_infer_exec.png b/action_compare_outputs/hist_train_abs_vs_infer_exec.png new file mode 100644 index 00000000..043c8430 Binary files /dev/null and b/action_compare_outputs/hist_train_abs_vs_infer_exec.png differ diff --git a/action_compare_outputs/hist_train_rel_vs_infer_raw.png b/action_compare_outputs/hist_train_rel_vs_infer_raw.png new file mode 100644 index 00000000..32b5c51d Binary files /dev/null and b/action_compare_outputs/hist_train_rel_vs_infer_raw.png differ diff --git a/action_compare_outputs/percentile_train_abs_vs_infer_exec.png b/action_compare_outputs/percentile_train_abs_vs_infer_exec.png new file mode 100644 index 00000000..af21281b Binary files /dev/null and b/action_compare_outputs/percentile_train_abs_vs_infer_exec.png differ diff --git a/action_compare_outputs/percentile_train_rel_vs_infer_raw.png b/action_compare_outputs/percentile_train_rel_vs_infer_raw.png new file mode 100644 index 00000000..dd26f4c7 Binary files /dev/null and b/action_compare_outputs/percentile_train_rel_vs_infer_raw.png differ diff --git a/action_compare_outputs/summary_train_abs_vs_infer_exec.csv b/action_compare_outputs/summary_train_abs_vs_infer_exec.csv new file mode 100644 index 00000000..b8cc268a --- /dev/null +++ b/action_compare_outputs/summary_train_abs_vs_infer_exec.csv @@ -0,0 +1,17 @@ +dim,name,train_abs_mean,train_abs_std,train_abs_p01,train_abs_p50,train_abs_p99,infer_exec_mean,infer_exec_std,infer_exec_p01,infer_exec_p50,infer_exec_p99,mean_abs_gap,p99_abs_gap +0,l.joint1,0.641659140586853,0.6947126388549805,-0.5343006920814514,0.5285818874835968,2.0458452701568604,1.158917784690857,0.12374740093946457,0.921103458404541,1.1449467539787292,1.4069665122032167,0.5172586441040039,0.6388787579536437 +1,l.joint2,-0.9733566641807556,0.1690513789653778,-1.3291998982429505,-0.9793720841407776,-0.47714611887931824,-0.9559754729270935,0.04431776702404022,-1.0714688420295715,-0.9545189440250397,-0.8601453870534896,0.01738119125366211,0.3829992681741714 +2,l.joint3,-1.121475338935852,0.4874430000782013,-2.106227159500122,-1.1478887796401978,-0.17000000178813934,-1.5394654273986816,0.09301857650279999,-1.744075183868408,-1.53302800655365,-1.360790419578552,0.4179900884628296,1.1907904177904127 +3,l.joint4,-1.7108030319213867,0.3729032576084137,-2.3499999046325684,-1.72506982088089,-0.888725964426994,-2.1100974082946777,0.06685517728328705,-2.2411738228797913,-2.101929783821106,-1.9408546137809752,0.399294376373291,1.0521286493539812 +4,l.joint5,0.7199459075927734,0.48198315501213074,-0.3437192440032959,0.7494157552719116,1.926978840827947,1.0302972793579102,0.0974796712398529,0.8160326021909714,1.0262407660484314,1.21043949007988,0.3103513717651367,0.7165393507480669 +5,l.joint6,-0.4588630497455597,0.5065776705741882,-1.0123000144958496,-0.6111662983894348,0.7933871150016785,0.26826563477516174,0.11601881682872772,0.05434122681617737,0.2934122681617737,0.5095154047012329,0.7271286845207214,0.28387171030044556 +6,l.joint7,-0.6607242822647095,0.4922950863838196,-1.527162790298462,-0.6384796500205994,0.41203388571739197,-1.4656376838684082,0.08175593614578247,-1.5358970165252686,-1.4959229826927185,-1.2668145632743835,0.8049134016036987,1.6788484489917754 +7,l.gripper,-0.5254634618759155,0.36058521270751953,-0.7850000262260437,-0.7850000262260437,0.0,-0.08208815008401871,0.1205878034234047,-0.29894852489233015,-0.10449077934026718,0.21011711537837985,0.4433753192424774,0.21011711537837985 +8,r.joint1,-1.0014448165893555,0.3315313458442688,-1.7857042229175568,-1.0578718185424805,-0.2088660940527915,-1.0839542150497437,0.05274298042058945,-1.1928240847587586,-1.0835356712341309,-0.9733165013790128,0.08250939846038818,0.7644504073262213 +9,r.joint2,-1.2798432111740112,0.2365627884864807,-1.6668896675109863,-1.3232492208480835,-0.5575107121467586,-0.9333884119987488,0.026572203263640404,-0.9889832156896591,-0.9345415830612183,-0.8817121618986129,0.34645479917526245,0.3242014497518543 +10,r.joint3,1.2204468250274658,0.2114914208650589,0.5723267728090287,1.2206523418426514,1.7621741926670074,1.3478705883026123,0.042895879596471786,1.2373100411891937,1.352783203125,1.4232772827148439,0.12742376327514648,0.3388969099521635 +11,r.joint4,-1.7886172533035278,0.37906575202941895,-2.3499999046325684,-1.889396071434021,-0.6776027393341063,-2.0853042602539062,0.08005916327238083,-2.281415002346039,-2.072542428970337,-1.9328801131248474,0.2966870069503784,1.255277373790741 +12,r.joint5,-0.5560360550880432,0.5796496272087097,-1.6380070447921753,-0.4139392673969269,0.27875876545906064,-1.0406877994537354,0.09973990172147751,-1.292977237701416,-1.0474953651428223,-0.8478468978404997,0.48465174436569214,1.1266056632995602 +13,r.joint6,0.17376446723937988,0.3008327782154083,-1.0065568828582763,0.22868450731039047,0.629759162664414,0.28443607687950134,0.060583099722862244,0.14375218406319618,0.2794297933578491,0.43111039310693755,0.11067160964012146,0.19864876955747646 +14,r.joint7,0.4020680785179138,0.8754451274871826,-0.7782645225524902,0.43664151430130005,1.5358999967575073,1.2508035898208618,0.12017592042684555,0.9389399832487106,1.2523177862167358,1.4777059924602511,0.848735511302948,0.058194004297256186 +15,r.gripper,-0.1922135353088379,0.3254104256629944,-0.7850000262260437,0.0,0.0,-0.23681718111038208,0.1254427582025528,-0.49738632917404174,-0.25897935032844543,0.020387001633644153,0.04460364580154419,0.020387001633644153 diff --git a/action_compare_outputs/summary_train_rel_vs_infer_raw.csv b/action_compare_outputs/summary_train_rel_vs_infer_raw.csv new file mode 100644 index 00000000..0e0a9dcc --- /dev/null +++ b/action_compare_outputs/summary_train_rel_vs_infer_raw.csv @@ -0,0 +1,17 @@ +dim,name,train_rel_mean,train_rel_std,train_rel_p01,train_rel_p50,train_rel_p99,infer_raw_mean,infer_raw_std,infer_raw_p01,infer_raw_p50,infer_raw_p99,mean_abs_gap,p99_abs_gap +0,l.joint1,-0.00039793882751837373,0.016543889418244362,-0.06913042306900025,0.0,0.06047791957855231,1.158917784690857,0.12374740093946457,0.921103458404541,1.1449467539787292,1.4069665122032167,1.1593157052993774,1.3464885926246644 +1,l.joint2,4.492358857532963e-05,0.004986106418073177,-0.011790799498558043,0.0,0.023119817972183238,-0.9559754729270935,0.04431776702404022,-1.0714688420295715,-0.9545189440250397,-0.8601453870534896,0.9560204148292542,0.8832652050256729 +2,l.joint3,0.00027034772210754454,0.013472087681293488,-0.04609204351902008,0.0,0.055739291906356986,-1.5394654273986816,0.09301857650279999,-1.744075183868408,-1.53302800655365,-1.360790419578552,1.5397357940673828,1.4165297114849091 +3,l.joint4,0.00021798511443194002,0.01103353314101696,-0.04429484367370606,0.0,0.03298486232757569,-2.1100974082946777,0.06685517728328705,-2.2411738228797913,-2.101929783821106,-1.9408546137809752,2.1103153228759766,1.973839476108551 +4,l.joint5,-0.00025620264932513237,0.012623071670532227,-0.051544408500194545,0.0,0.03821436434984209,1.0302972793579102,0.0974796712398529,0.8160326021909714,1.0262407660484314,1.21043949007988,1.0305534601211548,1.172225125730038 +5,l.joint6,-0.000564700982067734,0.016357799991965294,-0.0728665655851364,0.0,0.04714033782482148,0.26826563477516174,0.11601881682872772,0.05434122681617737,0.2934122681617737,0.5095154047012329,0.2688303291797638,0.46237506687641144 +6,l.joint7,0.0006646300898864865,0.014079635962843895,-0.04217458128929138,0.0,0.05698370039463045,-1.4940553903579712,0.11484292894601822,-1.7141154992580414,-1.4959229826927185,-1.2668145632743835,1.4947199821472168,1.323798263669014 +7,l.gripper,0.0001746299967635423,0.04241808503866196,-0.070911705493927,0.0,0.05233335494995117,-0.08208815008401871,0.1205878034234047,-0.29894852489233015,-0.10449077934026718,0.21011711537837985,0.08226277679204941,0.15778376042842868 +8,r.joint1,-0.000191343468031846,0.007814115844666958,-0.028782657384872436,0.0,0.027080488801002597,-1.0839542150497437,0.05274298042058945,-1.1928240847587586,-1.0835356712341309,-0.9733165013790128,1.0837628841400146,1.0003969901800154 +9,r.joint2,-0.00012973939010407776,0.003468823852017522,-0.010385245084762573,0.0,0.011857336759567266,-0.9333884119987488,0.026572203263640404,-0.9889832156896591,-0.9345415830612183,-0.8817121618986129,0.9332586526870728,0.8935694986581801 +10,r.joint3,0.00014158380508888513,0.007840425707399845,-0.0273062926530838,0.0,0.030971422195434584,1.3478705883026123,0.042895879596471786,1.2373100411891937,1.352783203125,1.4232772827148439,1.347728967666626,1.3923058605194094 +11,r.joint4,-0.00010751305671874434,0.010465368628501892,-0.04199434399604797,0.0,0.03338115334510813,-2.0853042602539062,0.08005916327238083,-2.281415002346039,-2.072542428970337,-1.9328801131248474,2.0851967334747314,1.9662612664699555 +12,r.joint5,-4.5862423576181754e-05,0.012884396128356457,-0.03616616725921631,0.0,0.04705580264329914,-1.0406877994537354,0.09973990172147751,-1.292977237701416,-1.0474953651428223,-0.8478468978404997,1.0406419038772583,0.8949027004837988 +13,r.joint6,4.1415307350689545e-05,0.010167960077524185,-0.029894600063562392,0.0,0.027249340713024185,0.28443607687950134,0.060583099722862244,0.14375218406319618,0.2794297933578491,0.43111039310693755,0.28439465165138245,0.40386105239391334 +14,r.joint7,-0.00019715055532287806,0.019032930955290794,-0.056884298324584956,0.0,0.052225705683231476,1.2508035898208618,0.12017592042684555,0.9389399832487106,1.2523177862167358,1.4777059924602511,1.2510007619857788,1.4254802867770198 +15,r.gripper,0.10448496043682098,0.08611126989126205,-0.08642586052417753,0.1355040818452835,0.21151389181613922,-0.23681718111038208,0.1254427582025528,-0.49738632917404174,-0.25897935032844543,0.020387001633644153,0.34130215644836426,0.19112689018249507 diff --git a/action_compare_outputs/timeseries_infer_exec_actions.png b/action_compare_outputs/timeseries_infer_exec_actions.png new file mode 100644 index 00000000..cc39d24c Binary files /dev/null and b/action_compare_outputs/timeseries_infer_exec_actions.png differ diff --git a/action_compare_outputs/timeseries_infer_raw_actions.png b/action_compare_outputs/timeseries_infer_raw_actions.png new file mode 100644 index 00000000..44b8419f Binary files /dev/null and b/action_compare_outputs/timeseries_infer_raw_actions.png differ diff --git "a/docs/2026-08-07_G2\351\200\202\351\205\215\350\256\255\347\273\203\346\216\250\347\220\206\350\257\264\346\230\216.md" "b/docs/2026-08-07_G2\351\200\202\351\205\215\350\256\255\347\273\203\346\216\250\347\220\206\350\257\264\346\230\216.md" new file mode 100644 index 00000000..64121c86 --- /dev/null +++ "b/docs/2026-08-07_G2\351\200\202\351\205\215\350\256\255\347\273\203\346\216\250\347\220\206\350\257\264\346\230\216.md" @@ -0,0 +1,394 @@ +# G2 适配 DreamZero 训练与推理说明 + +本文说明当前分支如何把 DreamZero 适配到 G2:数据如何进入训练、模型侧改了什么、checkpoint 如何保存/加载,以及推理服务如何启动。 + +## 目标 + +G2 适配的目标是让 DreamZero 使用 G2 的三路视觉、双臂关节状态和双臂动作进行训练,并在推理时输出可执行的 G2 关节目标。 + +核心约束: + +- 视觉输入:3 路相机,`top_head`、`hand_left`、`hand_right`。 +- 状态输入:16 维,布局为左臂 7 维关节、左夹爪、右臂 7 维关节、右夹爪。 +- 动作输出:16 维,布局与状态一致。 +- action horizon:24。 +- 训练帧数:`num_frames=33`。 +- 图像尺寸:`320x176`。 +- embodiment tag:`g2`。 + +## 数据适配 + +### 1. 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 +``` + +转换脚本做几件事: + +- 校验 G2 state/action 是固定 16 维布局。 +- 将三路相机映射为 DreamZero 使用的 video modality。 +- 将数据拆成物理隔离的 `train/` 和 `test/`。 +- 生成 `meta/info.json`、`meta/modality.json`、`meta/embodiment.json`、`meta/stats.json`。 +- 生成 `meta/relative_stats_dreamzero.json`,用于 relative action 训练和 motion mask。 + +训练时 `G2_DATA_ROOT` 必须指向 GEAR 数据的 `train/` 目录,而不是父目录或 `test/` 目录。 + +### 2. modality 配置 + +G2 的 modality 注册在: + +```text +groot/vla/configs/data/dreamzero/base_48_wan_fine_aug_relative.yaml +``` + +关键配置: + +```yaml +modality_config_g2: + video: + modality_keys: + - video.top_head + - video.hand_left + - video.hand_right + state: + modality_keys: + - state.left_joint_position + - state.left_gripper_position + - state.right_joint_position + - state.right_gripper_position + action: + modality_keys: + - action.left_joint_position + - action.left_gripper_position + - action.right_joint_position + - action.right_gripper_position + language: + modality_keys: + - annotation.language.action_text +``` + +状态和动作都使用 `q99` 归一化;最后通过 `ConcatTransform` 拼成模型看到的三路视频、16 维 state、16 维 action。 + +### 3. dataset 配置 + +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 +``` + +含义是:手臂关节动作按当前状态的相对增量学习,夹爪保持 policy space 数值。 + +G2 absolute action 配置: + +```text +groot/vla/configs/data/dreamzero/g2_absolute.yaml +``` + +它关闭: + +```yaml +relative_action: false +relative_action_keys: [] +``` + +含义是:模型直接预测绝对关节目标。推理客户端也必须按 absolute 模式执行,不能再做 relative decode。 + +## 模型适配 + +### 1. embodiment tag + +G2 已注册到: + +```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 +``` + +部分早期脚本中出现过 `g2=26`,后续训练建议统一使用当前配置中的 G2 id,避免 checkpoint metadata 和推理端不一致。 + +### 2. action head + +主要改动在: + +```text +groot/vla/model/dreamzero/action_head/wan_flow_matching_action_tf.py +``` + +新增能力: + +- `action_only_training`:冻结共享 DiT,只训练 G2 state/action adapter。 +- `dynamics_loss_weight`、`action_loss_weight`:分别控制视频 dynamics loss 和 action flow-matching loss。 +- `action_x0_loss_weight`、`action_endpoint_loss_weight`、`action_start_loss_weight`:给 action 预测增加直接重建/终点/起点一致性约束。 +- `action_video_consistency_loss_weight`:约束未来视频 latent 中恢复的 action 与 action chunk 一致。 +- `motion_mask_enabled`:只在发生明显运动的 timestep 上计算 action loss,减少静止段主导训练。 + +motion mask 依赖 `relative_stats_dreamzero.json` 的 q01/q99,把归一化 action delta 还原到弧度尺度后判断是否超过阈值。 + +### 3. checkpoint 保存与加载 + +主要改动在: + +```text +groot/vla/experiment/base.py +groot/vla/model/dreamzero/base_vla.py +``` + +保存侧: + +- 普通 LoRA 训练仍可保存 LoRA-only checkpoint。 +- `action_only_training=true` 时,只保存 G2 action adapter 权重。 +- action adapter 会额外保存为 `action_expert.safetensors`,方便推理端严格加载。 + +加载侧: + +- LoRA-only checkpoint 必须显式指定训练时使用的 `pretrained_base_model_path`。 +- G2 action adapter 加载流程是:先加载 DreamZero-AgiBot base 的共享权重,再严格加载 G2 的 state/action encoder/decoder。 +- 这样可以避免把 LoRA delta 或 action adapter 套到错误底座上,导致模型能跑但语义不对。 + +## 训练流程 + +### 1. joint 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 +``` + +该路线用于在 DreamZero-AgiBot base 上联合适配 G2 视频/action 表征。关键 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 +``` + +### 2. action adapter only 训练 + +入口脚本: + +```text +scripts/train/train_dreamzero_g2_action_only_stage2.sh +``` + +该路线冻结共享 DiT,只训练 G2 的 state/action adapter: + +```bash +train_architecture=action_only +action_head_cfg.config.action_only_training=true +action_head_cfg.config.tune_diffusion_model=false +save_lora_only=false +``` + +适合在共享视频/动力学能力已经可用时,只修正 G2 action 输出头。 + +### 3. absolute action 训练 + +入口脚本: + +```text +scripts/train/train_dreamzero_g2_absolute_lora.sh +``` + +该路线使用: + +```bash +data=dreamzero/g2_absolute +relative_action=false +``` + +模型直接学习绝对关节目标。训练脚本会检查 action 统计分布,防止把 relative delta 数据误当 absolute action 数据训练。 + +## 推理流程 + +### 1. 启动 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" +``` + +audit 会检查: + +- checkpoint 是否包含 `config.json`、`model.safetensors`、`experiment_cfg/conf.yaml`、`experiment_cfg/metadata.json`。 +- action horizon 是否为 24。 +- `num_frames` 是否为 33。 +- action head 输入/输出维度是否符合 G2 契约。 +- checkpoint 是 LoRA 还是 action adapter。 +- checkpoint 记录的 `pretrained_model_path` 是否存在。 + +### 2. 服务端参数 + +G2 服务端使用: + +```bash +python -m torch.distributed.run \ + --standalone \ + --nproc_per_node=2 \ + socket_optimized_AR_g2.py \ + --port 9443 \ + --model-path "$MODEL_PATH" \ + --wan-ckpt-dir "$WAN_CKPT_DIR" \ + --tokenizer-path "$TOKENIZER_PATH" \ + --embodiment-tag g2 \ + --video-save-mode full \ + --num-inference-timesteps 0 +``` + +`--num-inference-timesteps 0` 表示沿用 checkpoint 里的 diffusion inference step 配置。 + +### 3. G2 relative action 解码 + +G2 真机客户端在: + +```text +robot_live_client_g2.py +``` + +relative action checkpoint 的输出不是直接下发的绝对目标,需要用当前执行边界的机器人状态还原: + +```text +absolute_target = current_state + decoded_relative_delta +``` + +随后客户端会: + +- 将左右臂关节目标限制在 GDK 安全范围内。 +- 保持 16 维 G2 动作布局。 +- 按 G2 控制接口下发目标。 + +absolute action checkpoint 则应直接按绝对关节目标执行,不能再加当前状态。 + +## 评估与排查 + +常用工具: + +```text +scripts/audit_g2_checkpoint.py +scripts/eval/eval_g2_checkpoint_on_testset.py +scripts/eval/rollout_g2_server_windowed.py +scripts/eval/rollout_g2_action_full.py +scripts/eval/action_input_sensitivity.py +scripts/eval/summarize_g2_action_timeline.py +scripts/infer/rollout_g2_server_episode.py +scripts/infer/compare_g2_live_to_dataset.py +scripts/plot_train_vs_infer_actions.py +scripts/stats_action_distribution.py +``` + +建议排查顺序: + +1. 先跑 `scripts/audit_g2_checkpoint.py`,确认 checkpoint 结构和 G2 维度契约正确。 +2. 在 test split 上跑离线评估,确认视频预测和 action 误差是否合理。 +3. 用 rollout 脚本连接本地 server,检查服务端输出 action chunk 是否稳定。 +4. 用 `compare_g2_live_to_dataset.py` 或 action 分布图对比训练数据和 live 输出,重点看 scale、夹爪范围和左右臂是否错位。 +5. 真机执行前确认 checkpoint 是 relative 还是 absolute,并选择对应客户端执行模式。 + +## 常见问题 + +### 1. 训练能跑,但真机动作幅度明显不对 + +优先检查 relative/absolute 是否混用: + +- relative checkpoint:客户端必须用当前状态把 delta 解码成绝对目标。 +- absolute checkpoint:客户端不能再加当前状态。 + +同时检查 `relative_stats_dreamzero.json` 是否来自同一份 G2 训练集。 + +### 2. checkpoint 能加载,但效果像没有适配 G2 + +检查 LoRA 或 action adapter 是否加载到了正确 base: + +- LoRA-only checkpoint 需要训练时的 `pretrained_model_path`。 +- action adapter checkpoint 需要先加载 DreamZero-AgiBot base,再加载 `action_expert.safetensors`。 + +### 3. 静止段太多导致 action 学不动 + +可以使用 active/hold window 采样和 motion mask: + +```bash +active_hold_index_path=$G2_DATA_ROOT/meta/g2_active_hold_windows.json +active_window_ratio=0.8 +++action_head_cfg.config.motion_mask_enabled=true +++action_head_cfg.config.motion_mask_stats_path=$G2_DATA_ROOT/meta/relative_stats_dreamzero.json +``` + +motion mask 只适合 relative action 训练;absolute action 训练默认关闭。 + +## 最小闭环 + +一条最小 G2 适配闭环是: + +1. 用 `convert_lerobot_g2_to_gear.py` 生成 GEAR train/test 数据。 +2. 确认 `meta/embodiment.json` 为 `{"embodiment_tag": "g2"}`。 +3. 用 `data=dreamzero/g2_relative` 或 `data=dreamzero/g2_absolute` 启动训练。 +4. 用 `scripts/audit_g2_checkpoint.py` 校验 checkpoint。 +5. 用 `scripts/run_g2_server_9443_final.sh` 启动 G2 server。 +6. 用 G2 客户端按 checkpoint 类型执行 relative 或 absolute action。 diff --git a/eval_agibot_checkpoint_on_fruit_testset.py b/eval_agibot_checkpoint_on_fruit_testset.py new file mode 100644 index 00000000..86de55f2 --- /dev/null +++ b/eval_agibot_checkpoint_on_fruit_testset.py @@ -0,0 +1,2349 @@ +import dataclasses +import json +import logging +import socket +import asyncio +import os +import http +import logging +import time +import traceback +from pathlib import Path +import glob +import uuid + +import pyarrow.parquet as pq +from omegaconf import OmegaConf +import torch +import tyro +from einops import rearrange +import datetime +import cv2 + +from groot.vla.model.n1_5.sim_policy import GrootSimPolicy +from groot.vla.data.schema import EmbodimentTag +import imageio +import numpy as np + +from openpi_client import base_policy as _base_policy +from openpi_client import msgpack_numpy +import websockets.asyncio.server as _server +import websockets.frames +from tianshou.data import Batch +import torch.distributed as dist +from torch.distributed.device_mesh import DeviceMesh, init_device_mesh + +# Use roboarena policy server interface +from eval_utils.policy_server import WebsocketPolicyServer as RoboarenaServer +from eval_utils.policy_server import PolicyServerConfig + +logger = logging.getLogger(__name__) + +SIGNAL_INFER = 0 +SIGNAL_SHUTDOWN = 1 +SIGNAL_IDLE = 2 +SIGNAL_RESET_CACHE = 3 + + +def _reset_policy_inference_cache(policy: object, reason: str) -> None: + trained_model = getattr(policy, "trained_model", None) + action_head = getattr(trained_model, "action_head", None) + reset_fn = getattr(action_head, "reset_inference_cache", None) + if callable(reset_fn): + reset_fn() + logger.info("Reset action-head inference cache on rank %s (%s)", dist.get_rank() if dist.is_initialized() else "?", reason) + else: + logger.warning("Policy action head does not expose reset_inference_cache(); cache reset skipped (%s)", reason) + +@dataclasses.dataclass +class Args: + # Model paths. + model_path: str = "/data/wangk/checkpoints/dreamzero_agibot_fruit_lora_20k/checkpoint-12000" + wan_ckpt_dir: str = "/data/wangk/checkpoints/Wan2.1-I2V-14B-480P" + tokenizer_path: str = "/data/wangk/checkpoints/umt5-xxl" + + # Packaged AgiBot/G1 fruit held-out split. Override this path explicitly. + test_data_root: str = "/data/training_data/teleop/agibot/fruit/test" + + # Output and sample selection. + output_dir: str = "/data/wangk/dreamzero/agibot_fruit_testset_video_eval/checkpoint" + episode_indices: str = "0" # Comma/range syntax: 0,3,7-9 + frame_index: int = 30 # Row/frame within each episode; -1 chooses the midpoint. + future_frames: int = 33 + prompt_override: str | None = None + + # Inference. + embodiment_tag: str = "agibot" + num_inference_timesteps: int = 4 + enable_dit_cache: bool = False + num_dit_steps: int | None = None + timeout_seconds: int = 50000 + seed: int = 42 + preflight_only: bool = False + + # Kept for compatibility with shared model/path helpers. + port: int = 0 + index: int = 0 + max_chunk_size: int | None = None + video_save_mode: str = "none" + + + +class DistributedRoboarenaPolicyBase: + """Shared distributed inference plumbing for websocket policy wrappers.""" + + def __init__( + self, + groot_policy: GrootSimPolicy, + signal_group: dist.ProcessGroup, + output_dir: str | None = None, + video_save_mode: str = "first", + ) -> None: + self._policy = groot_policy + self._signal_group = signal_group + self._output_dir = output_dir + self._video_save_mode = video_save_mode + self._frame_buffers = self._init_frame_buffers() + self._current_session_id: str | None = None + self.video_across_time = [] + self._msg_index = 0 + + if self._output_dir: + os.makedirs(self._output_dir, exist_ok=True) + + def _init_frame_buffers(self) -> dict[str, list[np.ndarray]]: + return { + "video.top_head": [], + "video.hand_left": [], + "video.hand_right": [], + } + + def _reset_custom_state(self) -> None: + pass + + def _after_infer(self) -> None: + pass + + def _prepare_video_chunk(self, video_pred: torch.Tensor) -> torch.Tensor | None: + if self._video_save_mode == "none": + return None + return video_pred + + def _video_save_fps(self) -> int: + return 5 + + def _convert_observation(self, obs: dict) -> dict: + raise NotImplementedError + + def _convert_action(self, action_dict: dict) -> np.ndarray: + raise NotImplementedError + + def _broadcast_batch_to_workers(self, obs: dict) -> None: + import pickle + + serialized = pickle.dumps(obs) + data_size = len(serialized) + + size_tensor = torch.tensor([data_size], dtype=torch.int64, device='cuda') + dist.broadcast(size_tensor, src=0) + + data_tensor = torch.frombuffer(serialized, dtype=torch.uint8).clone().cuda() + dist.broadcast(data_tensor, src=0) + + def _extract_action_dict(self, action_chunk_dict: object) -> dict[str, object]: + action_dict: dict[str, object] = {} + for key in dir(action_chunk_dict): + if key.startswith('action.'): + action_dict[key] = getattr(action_chunk_dict, key) + return action_dict + + def _broadcast_signal_to_workers(self, signal: int) -> None: + signal_tensor = torch.tensor([signal], dtype=torch.int32, device='cpu') + dist.broadcast(signal_tensor, src=0, group=self._signal_group) + + def _diagnostic_dir(self) -> str | None: + if not self._output_dir: + return None + path = os.path.join(self._output_dir, "diagnostics") + os.makedirs(path, exist_ok=True) + return path + + def _save_predicted_video_chunk(self, video_pred: torch.Tensor) -> None: + """Decode and save exactly one predicted latent chunk. + + This deliberately avoids concatenating independent 33-frame chunks + before VAE decoding, so chunk-boundary artifacts cannot hide the + quality of the current world-model prediction. + """ + diagnostic_dir = self._diagnostic_dir() + if diagnostic_dir is None: + return + + try: + if video_pred.ndim != 5: + raise ValueError( + "video_pred must be (B,C,T,H,W), got " + f"{tuple(video_pred.shape)}" + ) + finite = bool(torch.isfinite(video_pred).all().item()) + latent_float = video_pred.detach().float() + stats = { + "request_index": int(self._msg_index), + "shape": list(video_pred.shape), + "dtype": str(video_pred.dtype), + "finite": finite, + "min": float(latent_float.min().item()), + "max": float(latent_float.max().item()), + "mean": float(latent_float.mean().item()), + "std": float(latent_float.std().item()), + } + with open( + os.path.join(diagnostic_dir, "video_pred_stats.jsonl"), + "a", + encoding="utf-8", + ) as stream: + stream.write(json.dumps(stats, ensure_ascii=False) + "\n") + + if not finite: + raise ValueError(f"video_pred contains NaN/Inf: {stats}") + + action_head = self._policy.trained_model.action_head + device = getattr(action_head, "_device", None) + if device is None: + device = next(self._policy.trained_model.parameters()).device + + latent = video_pred.detach().to( + device=device, + dtype=torch.bfloat16, + ) + frames = action_head.vae.decode( + latent, + tiled=action_head.tiled, + tile_size=( + action_head.tile_size_height, + action_head.tile_size_width, + ), + tile_stride=( + action_head.tile_stride_height, + action_head.tile_stride_width, + ), + ) + frames = rearrange(frames, "B C T H W -> B T H W C")[0] + frames = ( + (frames.float() + 1.0) * 127.5 + ).clamp(0, 255).cpu().numpy().astype(np.uint8) + + output_path = os.path.join( + diagnostic_dir, + f"pred_chunk_request_{self._msg_index:06d}_f{len(frames)}.mp4", + ) + imageio.mimsave( + output_path, + list(frames), + fps=self._video_save_fps(), + codec="libx264", + macro_block_size=None, + ) + imageio.imwrite( + os.path.join( + diagnostic_dir, + f"pred_chunk_request_{self._msg_index:06d}_first.png", + ), + frames[0], + ) + imageio.imwrite( + os.path.join( + diagnostic_dir, + f"pred_chunk_request_{self._msg_index:06d}_last.png", + ), + frames[-1], + ) + + # DreamZero packs three 320x176 views into a 2x2 RGB canvas in + # ConcatTransform order: top_head, hand_left, hand_right, padding. + # Decode the full canvas once, then split decoded RGB only for + # diagnostics. Never split the latent before VAE decoding. + frame_h, frame_w = int(frames.shape[1]), int(frames.shape[2]) + if frame_h % 2 == 0 and frame_w % 2 == 0: + half_h, half_w = frame_h // 2, frame_w // 2 + view_frames = { + "top_head": frames[:, :half_h, :half_w], + "hand_left": frames[:, :half_h, half_w:], + "hand_right": frames[:, half_h:, :half_w], + "padding": frames[:, half_h:, half_w:], + } + view_stats = {} + for view_name, view_video in view_frames.items(): + view_path = os.path.join( + diagnostic_dir, + f"pred_chunk_request_{self._msg_index:06d}_{view_name}_f{len(view_video)}.mp4", + ) + imageio.mimsave( + view_path, + list(view_video), + fps=self._video_save_fps(), + codec="libx264", + macro_block_size=None, + ) + imageio.imwrite( + os.path.join( + diagnostic_dir, + f"pred_chunk_request_{self._msg_index:06d}_{view_name}_first.png", + ), + view_video[0], + ) + imageio.imwrite( + os.path.join( + diagnostic_dir, + f"pred_chunk_request_{self._msg_index:06d}_{view_name}_last.png", + ), + view_video[-1], + ) + view_float = view_video.astype(np.float32) + view_stats[view_name] = { + "shape": list(view_video.shape), + "mean_rgb": [ + float(x) + for x in view_float.mean(axis=(0, 1, 2)).tolist() + ], + "std": float(view_float.std()), + } + stats["decoded_shape"] = list(frames.shape) + stats["decoded_views"] = view_stats + with open( + os.path.join(diagnostic_dir, "decoded_view_stats.jsonl"), + "a", + encoding="utf-8", + ) as stream: + stream.write(json.dumps(stats, ensure_ascii=False) + "\n") + else: + logger.warning( + "Decoded video is not divisible into a 2x2 view grid: %s", + tuple(frames.shape), + ) + + logger.info( + "Saved single predicted video chunk to %s | stats=%s", + output_path, + stats, + ) + except Exception as exc: + logger.exception( + "Failed to save single predicted video chunk: %s", exc + ) + + def infer(self, obs: dict) -> np.ndarray: + session_id = obs.get('session_id') + if session_id is not None and session_id != self._current_session_id: + if self._current_session_id is not None: + logger.info("Session changed from '%s' to '%s', resetting state", self._current_session_id, session_id) + self._broadcast_signal_to_workers(SIGNAL_RESET_CACHE) + self._reset_state() + else: + logger.info("New session started: '%s'", session_id) + self._current_session_id = session_id + + self._msg_index += 1 + converted_obs = self._convert_observation(obs) + + self._broadcast_signal_to_workers(SIGNAL_INFER) + self._broadcast_batch_to_workers(converted_obs) + + batch = Batch(obs=converted_obs) + dist.barrier() + with torch.no_grad(): + result_batch, video_pred = self._policy.lazy_joint_forward_causal(batch) + dist.barrier() + + self._save_predicted_video_chunk(video_pred) + + video_chunk = self._prepare_video_chunk(video_pred) + if video_chunk is not None: + self.video_across_time.append(video_chunk.detach().cpu()) + action = self._convert_action(self._extract_action_dict(result_batch.act)) + self._after_infer() + return action + + def _reset_state(self, save_video: bool = True) -> None: + if save_video and len(self.video_across_time) > 0 and self._output_dir: + try: + frame_list = [] + action_head = self._policy.trained_model.action_head + device = getattr(action_head, "_device", None) + if device is None: + device = next(self._policy.trained_model.parameters()).device + video_across_time_cat = torch.cat(self.video_across_time, dim=2).to(device=device, dtype=torch.bfloat16) + frames = action_head.vae.decode( + video_across_time_cat, + tiled=action_head.tiled, + tile_size=(action_head.tile_size_height, action_head.tile_size_width), + tile_stride=(action_head.tile_stride_height, action_head.tile_stride_width), + ) + frames = rearrange(frames, 'B C T H W -> B T H W C') + frames = frames[0] + frames = ((frames.float() + 1) * 127.5).clip(0, 255).cpu().numpy().astype(np.uint8) + for frame in frames: + frame_list.append(frame) + + if frame_list: + sample_frame = frame_list[0] + if len(sample_frame.shape) == 3 and sample_frame.shape[2] in [1, 3, 4]: + save_dir = self._output_dir + os.makedirs(save_dir, exist_ok=True) + all_mp4_files = [f for f in os.listdir(save_dir) if f.endswith('.mp4')] + timestamp = datetime.datetime.now().strftime('%Y%m%d_%H%M%S') + num_frames = len(frame_list) + output_path = os.path.join(save_dir, f'{timestamp}_{len(all_mp4_files):06}_f{num_frames}.mp4') + imageio.mimsave(output_path, frame_list, fps=self._video_save_fps(), codec='libx264') + logger.info('Saved video on reset to: %s', output_path) + except Exception as exc: + logger.warning('Failed to save video on reset: %s', exc) + + for key in self._frame_buffers: + self._frame_buffers[key] = [] + + self.video_across_time = [] + _reset_policy_inference_cache(self._policy, "wrapper reset_state") + self._reset_custom_state() + + def reset(self, reset_info: dict) -> None: + self._broadcast_signal_to_workers(SIGNAL_RESET_CACHE) + self._reset_state(save_video=True) + + +class ARDroidRoboarenaPolicy(DistributedRoboarenaPolicyBase): + """Wrapper policy that implements roboarena.policy.BasePolicy interface for AR_droid.""" + + FRAMES_PER_CHUNK = 4 + + def __init__( + self, + groot_policy: GrootSimPolicy, + signal_group: dist.ProcessGroup, + output_dir: str | None = None, + video_save_mode: str = "first", + ) -> None: + super().__init__( + groot_policy=groot_policy, + signal_group=signal_group, + output_dir=output_dir, + video_save_mode=video_save_mode, + ) + self._reset_custom_state() + + def _init_frame_buffers(self) -> dict[str, list[np.ndarray]]: + return { + 'video.exterior_image_1_left': [], + 'video.exterior_image_2_left': [], + 'video.wrist_image_left': [], + } + + def _reset_custom_state(self) -> None: + self._is_first_call = True + + def _after_infer(self) -> None: + self._is_first_call = False + + def _convert_observation(self, obs: dict) -> dict: + converted = {} + image_key_mapping = { + 'observation/exterior_image_0_left': 'video.exterior_image_1_left', + 'observation/exterior_image_1_left': 'video.exterior_image_2_left', + 'observation/wrist_image_left': 'video.wrist_image_left', + } + + for roboarena_key, droid_key in image_key_mapping.items(): + if roboarena_key in obs: + data = obs[roboarena_key] + if isinstance(data, np.ndarray): + if data.ndim == 4: + self._frame_buffers[droid_key].extend(list(data)) + else: + self._frame_buffers[droid_key].append(data) + + num_frames = 1 if self._is_first_call else self.FRAMES_PER_CHUNK + + for droid_key, buffer in self._frame_buffers.items(): + if len(buffer) > 0: + if len(buffer) >= num_frames: + frames_to_use = buffer[-num_frames:] + else: + frames_to_use = buffer.copy() + while len(frames_to_use) < num_frames: + frames_to_use.insert(0, buffer[0]) + converted[droid_key] = np.stack(frames_to_use, axis=0) + + joint_pos = obs.get('observation/joint_position', np.zeros(7, dtype=np.float32)) + if joint_pos.ndim == 1: + joint_pos = joint_pos.reshape(1, -1) + converted['state.joint_position'] = joint_pos.astype(np.float64) + + gripper_pos = obs.get('observation/gripper_position', np.zeros(1, dtype=np.float32)) + if gripper_pos.ndim == 1: + gripper_pos = gripper_pos.reshape(1, -1) + converted['state.gripper_position'] = gripper_pos.astype(np.float64) + converted['annotation.language.action_text'] = obs.get('prompt', '') + return converted + + def _convert_action(self, action_dict: dict) -> np.ndarray: + joint_action = None + gripper_action = None + for key, value in action_dict.items(): + if 'joint_position' in key: + joint_action = value + elif 'gripper_position' in key or 'gripper' in key: + gripper_action = value + + if joint_action is None: + return np.zeros((1, 8), dtype=np.float32) + + if isinstance(joint_action, torch.Tensor): + joint_action = joint_action.cpu().numpy() + if joint_action.ndim == 1: + joint_action = joint_action.reshape(1, -1) + + num_steps = joint_action.shape[0] + if gripper_action is not None: + if isinstance(gripper_action, torch.Tensor): + gripper_action = gripper_action.cpu().numpy() + if gripper_action.ndim == 1: + gripper_action = gripper_action.reshape(-1, 1) + elif gripper_action.ndim == 0: + gripper_action = gripper_action.reshape(1, 1) + else: + gripper_action = np.zeros((num_steps, 1), dtype=np.float32) + + return np.concatenate([joint_action, gripper_action], axis=-1).astype(np.float32) + + + +def _load_agibot_fixed_state_from_checkpoint( + model_path: str, + required_state_keys: list[str], +) -> dict[str, np.ndarray]: + """Load nominal AgiBot state values from checkpoint metadata.""" + metadata_path = os.path.join( + os.path.abspath(model_path), + "experiment_cfg", + "metadata.json", + ) + with open(metadata_path, "r", encoding="utf-8") as stream: + metadata = json.load(stream) + + try: + stats = metadata["agibot"]["statistics"]["state"] + except KeyError as exc: + raise KeyError( + f"{metadata_path} does not contain agibot.statistics.state" + ) from exc + + fixed: dict[str, np.ndarray] = {} + for target_key in required_state_keys: + if not target_key.startswith("state."): + raise ValueError( + f"Unexpected AgiBot state modality key: {target_key!r}" + ) + source_key = target_key.split(".", 1)[1] + if source_key not in stats or "mean" not in stats[source_key]: + raise KeyError( + f"Missing checkpoint-mean state for {source_key!r} " + f"in {metadata_path}" + ) + fixed[target_key] = np.asarray( + stats[source_key]["mean"], + dtype=np.float64, + ).reshape(1, -1) + return fixed + + +class AgiBotRoboarenaPolicy(DistributedRoboarenaPolicyBase): + """Run an AgiBot checkpoint on an in-domain AgiBot fruit test split. + + This is intentionally video-only. The checkpoint's own mean AgiBot + state fills required state fields, and native action output is discarded. + """ + + VIDEO_KEY_MAPPING = { + "observation/top_head": "video.top_head", + "observation/hand_left": "video.hand_left", + "observation/hand_right": "video.hand_right", + } + + def __init__( + self, + groot_policy: GrootSimPolicy, + signal_group: dist.ProcessGroup, + model_path: str, + output_dir: str | None = None, + video_save_mode: str = "first", + ) -> None: + super().__init__( + groot_policy=groot_policy, + signal_group=signal_group, + output_dir=output_dir, + video_save_mode=video_save_mode, + ) + self._action_keys = list( + self._policy.modality_configs.action.modality_keys + ) + self._state_keys = list( + self._policy.modality_configs.state.modality_keys + ) + self._language_keys = list( + self._policy.modality_configs.language.modality_keys + ) + eval_indices = list( + getattr( + self._policy.modality_configs.video, + "eval_delta_indices", + [0], + ) + or [0] + ) + self._expected_video_frames = max(1, len(eval_indices)) + self._fixed_state = _load_agibot_fixed_state_from_checkpoint( + model_path, + self._state_keys, + ) + self._safe_g2_hold_action = np.zeros(16, dtype=np.float32) + logger.info( + "[AGIBOT-FRUIT VIDEO TEST] video_frames=%s " + "state_keys=%s language_keys=%s", + self._expected_video_frames, + self._state_keys, + self._language_keys, + ) + + @staticmethod + def _normalize_video( + value: object, + target_key: str, + ) -> np.ndarray: + array_bgr = np.asarray(value) + if array_bgr.ndim == 3: + array_bgr = np.expand_dims(array_bgr, axis=0) + elif array_bgr.ndim != 4: + raise ValueError( + f"AgiBot video input for {target_key} must have " + f"shape (H,W,C) or (T,H,W,C), got {array_bgr.shape}" + ) + if array_bgr.shape[-1] != 3: + raise ValueError( + f"AgiBot video input for {target_key} must have 3 channels, " + f"got {array_bgr.shape}" + ) + # OpenCV decodes the fruit test MP4 files as BGR. DreamZero expects RGB. + return np.ascontiguousarray(array_bgr[..., ::-1]) + + def _prepare_video_chunk( + self, + video_pred: torch.Tensor, + ) -> torch.Tensor | None: + if self._video_save_mode == "none": + return None + if video_pred.ndim != 5: + raise ValueError( + "AgiBot video prediction must be 5D " + f"(B,C,T,H,W), got {tuple(video_pred.shape)}" + ) + if self._video_save_mode == "first": + return video_pred[:, :, :1].contiguous() + if self._video_save_mode == "full": + return video_pred.contiguous() + raise ValueError( + f"Unsupported video_save_mode: {self._video_save_mode!r}" + ) + + def _video_save_fps(self) -> int: + return 30 + + def _save_input_grid( + self, + converted: dict[str, object], + ) -> None: + diagnostic_dir = self._diagnostic_dir() + if diagnostic_dir is None: + return + + camera_keys = ( + "video.top_head", + "video.hand_left", + "video.hand_right", + ) + arrays = [np.asarray(converted[key]) for key in camera_keys] + frame_count = min(array.shape[0] for array in arrays) + grids: list[np.ndarray] = [] + + # The checkpoint may require different raw resolutions per camera. + # Resize only for this human-readable diagnostic grid; the arrays sent + # to the model remain at their native checkpoint-required sizes. + display_size = (320, 176) + for index in range(frame_count): + display_views = [ + cv2.resize( + np.asarray(array[index]), + display_size, + interpolation=cv2.INTER_AREA, + ) + for array in arrays + ] + top_head, hand_left, hand_right = display_views + black = np.zeros_like(top_head) + grids.append( + np.concatenate( + [ + np.concatenate([top_head, hand_left], axis=1), + np.concatenate([hand_right, black], axis=1), + ], + axis=0, + ) + ) + + grid_path = os.path.join( + diagnostic_dir, + f"agibot_fruit_input_request_{self._msg_index:06d}" + f"_f{len(grids)}.mp4", + ) + imageio.mimsave( + grid_path, + grids, + fps=30, + codec="libx264", + macro_block_size=None, + ) + imageio.imwrite( + os.path.join( + diagnostic_dir, + f"agibot_fruit_input_request_{self._msg_index:06d}" + "_last.png", + ), + grids[-1], + ) + logger.info( + "Saved AgiBot model-input display grid to %s | native_shapes=%s", + grid_path, + { + key: list(array.shape) + for key, array in zip(camera_keys, arrays) + }, + ) + + def _convert_observation( + self, + obs: dict, + ) -> dict: + converted: dict[str, object] = {} + missing_video: list[str] = [] + + for source_key, target_key in self.VIDEO_KEY_MAPPING.items(): + value = obs.get(source_key, obs.get(target_key)) + if value is None: + missing_video.append(source_key) + continue + frames = self._normalize_video(value, target_key) + self._frame_buffers[target_key].extend(list(frames)) + history = list( + self._frame_buffers[target_key][ + -self._expected_video_frames: + ] + ) + while len(history) < self._expected_video_frames: + history.insert(0, history[0]) + converted[target_key] = np.stack(history, axis=0) + + if missing_video: + raise ValueError( + "Missing G2 camera inputs for AgiBot video test: " + + ", ".join(sorted(missing_video)) + ) + + packed = np.asarray( + obs.get("observation/state", np.zeros(16, dtype=np.float32)), + dtype=np.float32, + ) + if packed.ndim == 2: + packed = packed[-1] + if packed.ndim == 1 and packed.shape[0] == 16: + self._safe_g2_hold_action = packed.copy() + + for key, value in self._fixed_state.items(): + converted[key] = value.copy() + + prompt = str(obs.get("prompt", "")) + for language_key in self._language_keys: + converted[language_key] = prompt + + self._save_input_grid(converted) + logger.info( + "[AGIBOT-FRUIT VIDEO TEST] request=%s " + "frames_per_view=%s language_keys=%s prompt=%r", + self._msg_index, + { + key: int(np.asarray(converted[key]).shape[0]) + for key in self.VIDEO_KEY_MAPPING.values() + }, + self._language_keys, + prompt, + ) + return converted + + def _convert_action( + self, + action_dict: dict, + ) -> np.ndarray: + # Native AgiBot actions are irrelevant for this diagnostic. Returning + # a G2 hold-shaped tensor keeps the rest of the offline evaluator + # simple and prevents accidental dimension comparisons. + return np.repeat( + self._safe_g2_hold_action.reshape(1, 16), + 48, + axis=0, + ).astype(np.float32) + + +class G2RoboarenaPolicy(DistributedRoboarenaPolicyBase): + """Adapter for the G2 dual-arm joint-space policy.""" + FRAMES_PER_CHUNK = 4 + VIDEO_KEY_MAPPING = { + 'observation/top_head': 'video.top_head', + 'observation/hand_left': 'video.hand_left', + 'observation/hand_right': 'video.hand_right', + } + + STATE_KEY_MAPPING = { + 'observation/left_joint_position': 'state.left_joint_position', + 'observation/left_gripper_position': 'state.left_gripper_position', + 'observation/right_joint_position': 'state.right_joint_position', + 'observation/right_gripper_position': 'state.right_gripper_position', + } + + PACKED_STATE_KEYS = ( + 'observation/state', + 'observation.state', + 'state', + ) + + def __init__( + self, + groot_policy: GrootSimPolicy, + signal_group: dist.ProcessGroup, + output_dir: str | None = None, + video_save_mode: str = "first", + ) -> None: + super().__init__( + groot_policy=groot_policy, + signal_group=signal_group, + output_dir=output_dir, + video_save_mode=video_save_mode, + ) + self._action_keys = list( + self._policy.modality_configs.action.modality_keys + ) + eval_indices = list( + getattr( + self._policy.modality_configs.video, + "eval_delta_indices", + [-3, -2, -1, 0], + ) + or [0] + ) + self._expected_video_frames = max(1, len(eval_indices)) + logger.info( + "G2 wrapper expects %s evaluation frame(s): %s", + self._expected_video_frames, + eval_indices, + ) + + @staticmethod + def _lookup_obs_value( + obs: dict, + source_key: str, + target_key: str, + ) -> object: + if source_key in obs: + return obs[source_key] + return obs.get(target_key) + + @staticmethod + def _normalize_video( + value: object, + target_key: str, + ) -> np.ndarray: + if ( + isinstance(value, dict) + and value.get("__dreamzero_image_encoding__") + == "jpeg_sequence" + ): + frames = [] + expected_shape = tuple(value.get("shape", ())) + expected_dtype = np.dtype( + value.get("dtype", "uint8") + ) + for index, frame_bytes in enumerate( + value.get("frames", []) + ): + encoded = np.frombuffer( + frame_bytes, + dtype=np.uint8, + ) + frame_bgr = cv2.imdecode( + encoded, + cv2.IMREAD_COLOR, + ) + if frame_bgr is None: + raise ValueError( + f"Failed to decode JPEG frame {index} " + f"for {target_key}" + ) + # The G2 client normalizes every camera frame to OpenCV BGR. + # DreamZero's training video path is RGB, so convert exactly + # once at the server/model boundary. + frame_rgb = cv2.cvtColor( + frame_bgr, + cv2.COLOR_BGR2RGB, + ) + frames.append( + frame_rgb.astype( + expected_dtype, + copy=False, + ) + ) + if not frames: + raise ValueError( + f"No JPEG frames were provided for {target_key}" + ) + array = np.stack(frames, axis=0) + if expected_shape and tuple(array.shape) != expected_shape: + raise ValueError( + f"Decoded JPEG video for {target_key} has " + f"shape {array.shape}, expected {expected_shape}" + ) + return array + + array_bgr = np.asarray(value) + if array_bgr.ndim == 3: + array_bgr = np.expand_dims(array_bgr, axis=0) + elif array_bgr.ndim != 4: + raise ValueError( + f"G2 video input for {target_key} must have " + f"shape (H,W,C) or (T,H,W,C), got {array_bgr.shape}" + ) + if array_bgr.shape[-1] != 3: + raise ValueError( + f"G2 video input for {target_key} must have 3 channels, " + f"got {array_bgr.shape}" + ) + # The non-JPEG G2 client path is also BGR because G2Camera._decode() + # converts RGB camera buffers to BGR and leaves BGR buffers unchanged. + # Convert it to the same RGB convention as the JPEG branch. + return np.ascontiguousarray(array_bgr[..., ::-1]) + + @staticmethod + def _normalize_state( + value: object, + target_key: str, + ) -> np.ndarray: + array = np.asarray(value, dtype=np.float64) + if array.ndim == 0: + array = array.reshape(1, 1) + elif array.ndim == 1: + array = array.reshape(1, -1) + elif array.ndim != 2: + raise ValueError( + f"G2 state input for {target_key} must be " + f"1D or 2D, got {array.shape}" + ) + return array + + @staticmethod + def _split_packed_state( + value: object, + ) -> dict[str, np.ndarray]: + packed = np.asarray(value, dtype=np.float64) + if packed.ndim == 1: + packed = packed.reshape(1, -1) + elif packed.ndim != 2: + raise ValueError( + "Packed G2 state must have shape (16,) or " + f"(T,16), got {packed.shape}" + ) + if packed.shape[-1] != 16: + raise ValueError( + f"Packed G2 state must contain 16 values, " + f"got {packed.shape}" + ) + return { + 'state.left_joint_position': packed[:, 0:7], + 'state.left_gripper_position': packed[:, 7:8], + 'state.right_joint_position': packed[:, 8:15], + 'state.right_gripper_position': packed[:, 15:16], + } + + def _prepare_video_chunk( + self, + video_pred: torch.Tensor, + ) -> torch.Tensor | None: + if self._video_save_mode == "none": + return None + if video_pred.ndim != 5: + raise ValueError( + "G2 video prediction must be 5D " + f"(B,C,T,H,W), got {tuple(video_pred.shape)}" + ) + if self._video_save_mode == "first": + return video_pred[:, :, :1].contiguous() + if self._video_save_mode == "full": + return video_pred.contiguous() + raise ValueError( + f"Unsupported video_save_mode: " + f"{self._video_save_mode!r}" + ) + + def _video_save_fps(self) -> int: + return 30 + + def _save_input_grid(self, converted: dict[str, object]) -> None: + """Save the exact four-frame, three-camera history sent to the model.""" + diagnostic_dir = self._diagnostic_dir() + if diagnostic_dir is None: + return + + camera_keys = ( + "video.top_head", + "video.hand_left", + "video.hand_right", + ) + if any(key not in converted for key in camera_keys): + return + + arrays = [np.asarray(converted[key]) for key in camera_keys] + if any(array.ndim != 4 for array in arrays): + raise ValueError( + "Diagnostic input videos must be (T,H,W,C): " + + ", ".join( + f"{key}={array.shape}" + for key, array in zip(camera_keys, arrays) + ) + ) + frame_count = min(array.shape[0] for array in arrays) + grids: list[np.ndarray] = [] + for index in range(frame_count): + top_head, hand_left, hand_right = [ + np.asarray(array[index]) for array in arrays + ] + height, width = top_head.shape[:2] + black = np.zeros_like(top_head) + grid_rgb = np.concatenate( + [ + np.concatenate([top_head, hand_left], axis=1), + np.concatenate([hand_right, black], axis=1), + ], + axis=0, + ) + # _normalize_video() already returns RGB for both JPEG and array + # transport. Keep the diagnostic image byte-for-byte in model + # color order instead of swapping it a second time. + grids.append(np.ascontiguousarray(grid_rgb)) + + output_path = os.path.join( + diagnostic_dir, + f"input_grid_request_{self._msg_index:06d}_f{len(grids)}.mp4", + ) + imageio.mimsave( + output_path, + grids, + fps=5, + codec="libx264", + macro_block_size=None, + ) + imageio.imwrite( + os.path.join( + diagnostic_dir, + f"input_grid_request_{self._msg_index:06d}_last.png", + ), + grids[-1], + ) + channel_means = { + key: [ + float(x) + for x in np.asarray(array, dtype=np.float32).mean( + axis=(0, 1, 2) + ).tolist() + ] + for key, array in zip(camera_keys, arrays) + } + logger.info( + "Saved exact RGB model input grid to %s | shapes=%s | mean_rgb=%s", + output_path, + {key: list(array.shape) for key, array in zip(camera_keys, arrays)}, + channel_means, + ) + + def _convert_observation(self, obs: dict) -> dict: + converted: dict[str, object] = {} + missing_video: list[str] = [] + + for source_key, target_key in self.VIDEO_KEY_MAPPING.items(): + value = self._lookup_obs_value( + obs, + source_key, + target_key, + ) + if value is None: + missing_video.append(source_key) + continue + frames = self._normalize_video( + value, + target_key, + ) + + self._frame_buffers[target_key].extend( + list(frames) + ) + + history = self._frame_buffers[target_key][ + -self._expected_video_frames: + ] + + while len(history) < self._expected_video_frames: + history.insert(0, history[0]) + + converted[target_key] = np.stack(history, axis=0) + + if missing_video: + raise ValueError( + "G2 inference requires video keys: " + + ", ".join(sorted(missing_video)) + ) + + packed_state = None + for key in self.PACKED_STATE_KEYS: + if key in obs: + packed_state = obs[key] + break + + if packed_state is not None: + converted.update( + self._split_packed_state(packed_state) + ) + else: + missing_state: list[str] = [] + for source_key, target_key in self.STATE_KEY_MAPPING.items(): + value = self._lookup_obs_value( + obs, + source_key, + target_key, + ) + if value is None: + missing_state.append(source_key) + continue + converted[target_key] = self._normalize_state( + value, + target_key, + ) + if missing_state: + raise ValueError( + "G2 inference requires a packed 16-D state " + "under observation/state, observation.state, " + "or state; otherwise all split state keys are " + "required: " + + ", ".join(sorted(missing_state)) + ) + + expected_dims = { + 'state.left_joint_position': 7, + 'state.left_gripper_position': 1, + 'state.right_joint_position': 7, + 'state.right_gripper_position': 1, + } + for key, expected_dim in expected_dims.items(): + array = np.asarray(converted[key]) + if array.shape[-1] != expected_dim: + raise ValueError( + f"{key} must have last dimension " + f"{expected_dim}, got {array.shape}" + ) + + converted['annotation.language.action_text'] = obs.get( + 'prompt', + obs.get( + 'annotation.language.action_text', + '', + ), + ) + self._save_input_grid(converted) + return converted + + def _convert_action( + self, + action_dict: dict, + ) -> np.ndarray: + missing = [ + key + for key in self._action_keys + if key not in action_dict + ] + if missing: + raise RuntimeError( + "Missing G2 action outputs: " + + ", ".join(missing) + ) + + arrays: list[np.ndarray] = [] + horizon: int | None = None + for key in self._action_keys: + value = action_dict[key] + if isinstance(value, torch.Tensor): + value = value.detach().cpu().numpy() + array = np.asarray(value) + if array.ndim == 0: + array = array.reshape(1, 1) + elif array.ndim == 1: + array = array.reshape(-1, 1) + else: + array = array.reshape(array.shape[0], -1) + + if horizon is None: + horizon = array.shape[0] + elif array.shape[0] != horizon: + raise RuntimeError( + f"Inconsistent G2 action horizon for {key}: " + f"expected {horizon}, got {array.shape[0]}" + ) + arrays.append(array.astype(np.float32)) + + action = np.concatenate(arrays, axis=-1) + if action.shape[-1] != 16: + raise RuntimeError( + "G2 action must have 16 dimensions ordered as " + "[left_joint(7), left_gripper(1), " + "right_joint(7), right_gripper(1)], " + f"got {action.shape}" + ) + return action.astype(np.float32) + + +class WebsocketPolicyServer: + """Serves a policy using the websocket protocol. See websocket_client_policy.py for a client implementation. + Currently only implements the `load` and `infer` methods. + """ + + def __init__( + self, + policy: _base_policy.BasePolicy, + host: str = "0.0.0.0", + port: int | None = None, + metadata: dict | None = None, + output_dir: str | None = None, + signal_group: dist.ProcessGroup | None = None, + ) -> None: + self._policy = policy + self._host = host + self._port = port + self._metadata = metadata or {} + self._output_dir = output_dir + logging.getLogger("websockets.server").setLevel(logging.INFO) + self.video_across_time = [] + self._msg_index = 0 + self._signal_group = signal_group + if self._output_dir: + os.makedirs(self._output_dir, exist_ok=True) + os.makedirs(os.path.join(self._output_dir, "inputs"), exist_ok=True) + + def serve_forever(self, rank: int = 0) -> None: + asyncio.run(self.run(rank)) + + async def run(self, rank: int = 0): + if rank == 0: + async with _server.serve( + self._handler, + self._host, + self._port, + compression=None, + max_size=None, + process_request=_health_check, + ping_interval=None, + ) as server: + await server.serve_forever() + else: + await self._worker_loop() + + async def _worker_loop(self): + logger.info(f"Worker loop started for rank {dist.get_rank()}") + signal_tensor = torch.zeros(1, dtype=torch.int32, device='cpu') + while True: + try: + dist.broadcast(signal_tensor, src=0, group=self._signal_group) + + signal = signal_tensor.item() + if signal == SIGNAL_SHUTDOWN: + logger.info(f"Rank {dist.get_rank()} received shutdown signal") + break + elif signal == SIGNAL_IDLE: + logger.info(f"Rank {dist.get_rank()} received idle signal. Waiting for next client.") + continue + elif signal == SIGNAL_RESET_CACHE: + logger.info(f"Rank {dist.get_rank()} received inference cache reset signal") + _reset_policy_inference_cache(self._policy, "worker signal") + continue + + batch = self._receive_batch_from_rank0() + dist.barrier() + with torch.no_grad(): + result_batch, video_pred = self._policy.lazy_joint_forward_causal(batch) + dist.barrier() + + except Exception as e: + logger.error(f"Worker loop error on rank {dist.get_rank()}: {e}") + traceback.print_exc() + break + + def _receive_batch_from_rank0(self): + import pickle + + size_tensor = torch.zeros(1, dtype=torch.int64, device='cuda') + dist.broadcast(size_tensor, src=0) + data_size = size_tensor.item() + + data_tensor = torch.zeros(data_size, dtype=torch.uint8, device='cuda') + dist.broadcast(data_tensor, src=0) + + obs = pickle.loads(data_tensor.cpu().numpy().tobytes()) + return Batch(obs=obs) + + def _broadcast_batch_to_workers(self, obs): + import pickle + + serialized = pickle.dumps(obs) + data_size = len(serialized) + + size_tensor = torch.tensor([data_size], dtype=torch.int64, device='cuda') + dist.broadcast(size_tensor, src=0) + + data_tensor = torch.frombuffer(serialized, dtype=torch.uint8).clone().cuda() + dist.broadcast(data_tensor, src=0) + + async def _handler(self, websocket: _server.ServerConnection): + logger.info(f"Connection from {websocket.remote_address} opened") + packer = msgpack_numpy.Packer() + + await websocket.send(packer.pack(self._metadata)) + + signal_tensor = torch.zeros(1, dtype=torch.int32, device='cpu') + + try: + while True: + try: + data = await websocket.recv() + obs = msgpack_numpy.unpackb(data) + self._msg_index += 1 + + signal_tensor.zero_() + dist.broadcast(signal_tensor, src=0, group=self._signal_group) + + self._broadcast_batch_to_workers(obs) + batch = Batch(obs=obs) + + dist.barrier() + with torch.no_grad(): + result_batch, video_pred = self._policy.lazy_joint_forward_causal(batch) + dist.barrier() + + action_chunk_dict = result_batch.act + + def batch_to_dict(batch): + out = {} + for k in dir(batch): + if not k.startswith("action."): + continue + out[k] = getattr(batch, k) + return out + + action_chunk_dict = batch_to_dict(action_chunk_dict) + await websocket.send(packer.pack(action_chunk_dict)) + + except websockets.ConnectionClosed: + logger.info(f"Connection from {websocket.remote_address} closed") + self.video_across_time = [] + break + except Exception: + await websocket.send(traceback.format_exc()) + await websocket.close( + code=websockets.frames.CloseCode.INTERNAL_ERROR, + reason="Internal server error. Traceback included in previous frame.", + ) + raise + finally: + logger.info("Rank 0: Client session ended. Sending idle signal (2) to workers.") + signal_tensor.fill_(2) + dist.broadcast(signal_tensor, src=0, group=self._signal_group) + + +def init_mesh() -> DeviceMesh: + local_rank = int(os.environ.get("LOCAL_RANK", "0")) + + torch.cuda.set_device(local_rank) + _ = torch.cuda.is_available() + _ = torch.cuda.device_count() + + dist.init_process_group("nccl") + rank = dist.get_rank() + world_size = dist.get_world_size() + if world_size not in (1, 2): + raise ValueError( + f"This DreamZero inference path only supports 1 or 2 GPUs, got world_size={world_size}. " + "The action head parallelization code explicitly supports ip_size 1 or 2 only. " + "Please launch with --nproc_per_node=2 (or 1)." + ) + print(f"Rank {rank}/{world_size} (PID: {os.getpid()}) setting device to local_rank={local_rank}") + + torch.cuda.set_device(local_rank) + device = torch.device(f"cuda:{local_rank}") + + mesh = init_device_mesh( + device_type="cuda", + mesh_shape=(world_size,), + mesh_dim_names=("ip",), + ) + print(f"Rank {rank}/{world_size} (PID: {os.getpid()}) using device {device}") + + return mesh + +def _health_check(connection: _server.ServerConnection, request: _server.Request) -> _server.Response | None: + if request.path == "/healthz": + return connection.respond(http.HTTPStatus.OK, "OK\n") + return None + + +def _create_wrapper_policy( + embodiment_tag: str, + groot_policy: GrootSimPolicy, + signal_group: dist.ProcessGroup, + output_dir: str | None, + video_save_mode: str, + model_path: str | None = None, +) -> DistributedRoboarenaPolicyBase: + if embodiment_tag == 'oxe_droid': + return ARDroidRoboarenaPolicy( + groot_policy=groot_policy, + signal_group=signal_group, + output_dir=output_dir, + video_save_mode=video_save_mode, + ) + if embodiment_tag == 'agibot': + if model_path is None: + raise ValueError('model_path is required for AgiBot wrapper') + return AgiBotRoboarenaPolicy( + groot_policy=groot_policy, + signal_group=signal_group, + model_path=model_path, + output_dir=output_dir, + video_save_mode=video_save_mode, + ) + if embodiment_tag == 'g2': + return G2RoboarenaPolicy( + groot_policy=groot_policy, + signal_group=signal_group, + output_dir=output_dir, + video_save_mode=video_save_mode, + ) + raise ValueError(f'Unsupported embodiment_tag: {embodiment_tag}') + + +def _create_server_config(embodiment_tag: str) -> PolicyServerConfig: + if embodiment_tag == 'oxe_droid': + return PolicyServerConfig( + image_resolution=(180, 320), + needs_wrist_camera=True, + n_external_cameras=2, + needs_stereo_camera=False, + needs_session_id=True, + action_space='joint_position', + ) + if embodiment_tag == 'agibot': + return PolicyServerConfig( + image_resolution=(640, 480), + needs_wrist_camera=False, + n_external_cameras=3, + needs_stereo_camera=False, + needs_session_id=True, + action_space='agibot_flattened', + ) + if embodiment_tag == 'g2': + return PolicyServerConfig( + image_resolution=(176, 320), + needs_wrist_camera=False, + n_external_cameras=3, + needs_stereo_camera=False, + needs_session_id=True, + action_space='joint_position', + ) + raise ValueError(f'Unsupported embodiment_tag: {embodiment_tag}') + + +def _build_path_overrides(args: Args) -> tuple[list[str], list[str]]: + model_config_overrides: list[str] = [] + train_config_overrides: list[str] = [] + + if args.wan_ckpt_dir: + wan_ckpt_dir = os.path.abspath(args.wan_ckpt_dir) + required_files = [ + os.path.join(wan_ckpt_dir, "models_t5_umt5-xxl-enc-bf16.pth"), + os.path.join(wan_ckpt_dir, "models_clip_open-clip-xlm-roberta-large-vit-huge-14.pth"), + os.path.join(wan_ckpt_dir, "Wan2.1_VAE.pth"), + ] + missing = [path for path in required_files if not os.path.exists(path)] + if missing: + raise FileNotFoundError( + "Missing Wan checkpoint component(s): " + ", ".join(missing) + ) + model_config_overrides.extend( + [ + f"action_head_cfg.config.diffusion_model_cfg.diffusion_model_pretrained_path={wan_ckpt_dir}", + f"action_head_cfg.config.text_encoder_cfg.text_encoder_pretrained_path={wan_ckpt_dir}/models_t5_umt5-xxl-enc-bf16.pth", + f"action_head_cfg.config.image_encoder_cfg.image_encoder_pretrained_path={wan_ckpt_dir}/models_clip_open-clip-xlm-roberta-large-vit-huge-14.pth", + f"action_head_cfg.config.vae_cfg.vae_pretrained_path={wan_ckpt_dir}/Wan2.1_VAE.pth", + ] + ) + + if args.tokenizer_path: + tokenizer_path = os.path.abspath(args.tokenizer_path) + if not os.path.exists(tokenizer_path): + raise FileNotFoundError(f"Tokenizer path does not exist: {tokenizer_path}") + if args.embodiment_tag.lower() == "agibot": + train_config_overrides.append( + f"transforms.agibot.transforms.10.tokenizer_path={tokenizer_path}" + ) + elif args.embodiment_tag.lower() == "oxe_droid": + train_config_overrides.append( + f"transforms.oxe_droid.transforms.10.tokenizer_path={tokenizer_path}" + ) + elif args.embodiment_tag.lower() == "g2": + train_config_overrides.append( + f"transforms.g2.transforms.10.tokenizer_path={tokenizer_path}" + ) + + return model_config_overrides, train_config_overrides + + + +def _require_file(path: Path, description: str) -> None: + if not path.is_file(): + raise FileNotFoundError(f"Missing {description}: {path}") + + +def _require_dir(path: Path, description: str) -> None: + if not path.is_dir(): + raise FileNotFoundError(f"Missing {description}: {path}") + + +def _parse_episode_indices(spec: str, total_episodes: int) -> list[int]: + values: list[int] = [] + for token in spec.split(','): + token = token.strip() + if not token: + continue + if '-' in token: + left, right = token.split('-', 1) + start = int(left) + end = int(right) + if end < start: + raise ValueError(f"Invalid episode range: {token}") + values.extend(range(start, end + 1)) + else: + values.append(int(token)) + if not values: + raise ValueError("--episode-indices cannot be empty") + unique: list[int] = [] + seen: set[int] = set() + for value in values: + if value < 0 or value >= total_episodes: + raise IndexError( + f"Episode {value} is outside [0, {total_episodes - 1}]" + ) + if value not in seen: + unique.append(value) + seen.add(value) + return unique + + +def _episode_file_from_template( + root: Path, + template: str, + episode_index: int, + chunks_size: int, + video_key: str | None = None, +) -> Path: + values = { + "episode_chunk": episode_index // chunks_size, + "episode_index": episode_index, + } + if video_key is not None: + values["video_key"] = video_key + path = root / template.format(**values) + if path.is_file(): + return path + + # Fallback for datasets whose chunk size metadata was rewritten without + # moving the actual files. + filename = f"episode_{episode_index:06d}" + path.suffix + if video_key is None: + candidates = sorted(root.glob(f"data/chunk-*/{filename}")) + else: + candidates = sorted(root.glob(f"videos/chunk-*/{video_key}/{filename}")) + if len(candidates) != 1: + raise FileNotFoundError( + f"Could not resolve episode file for episode={episode_index}, " + f"video_key={video_key!r}; expected {path}, candidates={candidates}" + ) + return candidates[0] + + +def _column_to_numpy(table: object, name: str, dtype: np.dtype) -> np.ndarray: + if name not in table.column_names: + raise KeyError(f"Parquet column missing: {name}") + values = table[name].combine_chunks().to_pylist() + return np.asarray(values, dtype=dtype) + + +def _unwrap_text(value: object) -> str: + while isinstance(value, (list, tuple)) and len(value) == 1: + value = value[0] + if value is None: + return "" + return str(value) + + +def _decode_video_range_bgr( + video_path: Path, + first_index: int, + last_index: int, +) -> dict[int, np.ndarray]: + if first_index < 0 or last_index < first_index: + raise ValueError( + f"Invalid decode range [{first_index}, {last_index}] for {video_path}" + ) + capture = cv2.VideoCapture(str(video_path)) + if not capture.isOpened(): + raise RuntimeError(f"Failed to open video: {video_path}") + try: + capture.set(cv2.CAP_PROP_POS_FRAMES, float(first_index)) + decoded: dict[int, np.ndarray] = {} + for index in range(first_index, last_index + 1): + ok, frame_bgr = capture.read() + if not ok or frame_bgr is None: + raise RuntimeError( + f"Failed to decode frame {index} from {video_path}" + ) + decoded[index] = np.ascontiguousarray(frame_bgr) + return decoded + finally: + capture.release() + + +def _grid_rgb( + top_bgr: np.ndarray, + left_bgr: np.ndarray, + right_bgr: np.ndarray, +) -> np.ndarray: + top = cv2.cvtColor(top_bgr, cv2.COLOR_BGR2RGB) + left = cv2.cvtColor(left_bgr, cv2.COLOR_BGR2RGB) + right = cv2.cvtColor(right_bgr, cv2.COLOR_BGR2RGB) + if top.shape != left.shape or top.shape != right.shape: + raise ValueError( + f"G2 test views must share one shape, got " + f"top={top.shape}, left={left.shape}, right={right.shape}" + ) + black = np.zeros_like(top) + return np.concatenate( + [ + np.concatenate([top, left], axis=1), + np.concatenate([right, black], axis=1), + ], + axis=0, + ) + + +def _save_rgb_video(path: Path, frames: list[np.ndarray], fps: int = 30) -> None: + if not frames: + raise ValueError(f"No frames to save: {path}") + path.parent.mkdir(parents=True, exist_ok=True) + imageio.mimsave( + path, + frames, + fps=fps, + codec="libx264", + macro_block_size=None, + ) + + +def _read_mp4_rgb(path: Path) -> list[np.ndarray]: + capture = cv2.VideoCapture(str(path)) + if not capture.isOpened(): + raise RuntimeError(f"Failed to open generated video: {path}") + frames: list[np.ndarray] = [] + try: + while True: + ok, frame_bgr = capture.read() + if not ok: + break + frames.append(cv2.cvtColor(frame_bgr, cv2.COLOR_BGR2RGB)) + finally: + capture.release() + if not frames: + raise RuntimeError(f"Generated video has no frames: {path}") + return frames + + +def _add_label_rgb(frame_rgb: np.ndarray, label: str) -> np.ndarray: + frame_bgr = cv2.cvtColor(frame_rgb, cv2.COLOR_RGB2BGR) + cv2.rectangle(frame_bgr, (0, 0), (360, 34), (0, 0, 0), -1) + cv2.putText( + frame_bgr, + label, + (8, 24), + cv2.FONT_HERSHEY_SIMPLEX, + 0.65, + (255, 255, 255), + 2, + cv2.LINE_AA, + ) + return cv2.cvtColor(frame_bgr, cv2.COLOR_BGR2RGB) + + + +def _checkpoint_eval_offsets( + model_path: Path, + embodiment_tag: str, +) -> list[int]: + conf_path = model_path / "experiment_cfg" / "conf.yaml" + _require_file(conf_path, "checkpoint experiment config") + cfg = OmegaConf.to_container(OmegaConf.load(conf_path), resolve=True) + config_key = f"modality_config_{embodiment_tag}" + try: + offsets = list( + cfg[config_key]["video"]["eval_delta_indices"] + ) + except (KeyError, TypeError) as exc: + raise KeyError( + f"Could not read {config_key}.video.eval_delta_indices " + f"from {conf_path}" + ) from exc + if not offsets: + offsets = [0] + return [int(value) for value in offsets] + + +def _checkpoint_agibot_video_resolutions( + model_path: Path, +) -> dict[str, tuple[int, int]]: + """Return required raw input resolutions as camera -> (height, width).""" + metadata_path = model_path / "experiment_cfg" / "metadata.json" + _require_file(metadata_path, "AgiBot checkpoint metadata") + with metadata_path.open("r", encoding="utf-8") as stream: + metadata = json.load(stream) + try: + video_meta = metadata["agibot"]["modalities"]["video"] + except KeyError as exc: + raise KeyError( + f"{metadata_path} lacks agibot.modalities.video" + ) from exc + + result: dict[str, tuple[int, int]] = {} + for name in ("top_head", "hand_left", "hand_right"): + try: + width, height = [ + int(value) + for value in video_meta[name]["resolution"] + ] + except (KeyError, TypeError, ValueError) as exc: + raise ValueError( + f"Invalid AgiBot raw resolution metadata for {name!r}" + ) from exc + if width <= 0 or height <= 0: + raise ValueError( + f"Non-positive AgiBot raw resolution for {name}: " + f"{width}x{height}" + ) + result[name] = (height, width) + return result + + +def _validate_test_dataset(root: Path) -> dict: + _require_dir(root, "AgiBot fruit test split") + info_path = root / "meta" / "info.json" + embodiment_path = root / "meta" / "embodiment.json" + _require_file(info_path, "fruit test info.json") + _require_file(embodiment_path, "fruit test embodiment.json") + + with info_path.open("r", encoding="utf-8") as stream: + info = json.load(stream) + with embodiment_path.open("r", encoding="utf-8") as stream: + embodiment = json.load(stream) + + if embodiment.get("embodiment_tag") != "agibot": + raise ValueError( + "Expected fruit test embodiment_tag='agibot', " + f"got {embodiment}" + ) + + expected_videos = { + "observation.images.top_head", + "observation.images.hand_left", + "observation.images.hand_right", + } + actual_videos = { + key + for key, value in info.get("features", {}).items() + if value.get("dtype") == "video" + } + missing = expected_videos - actual_videos + if missing: + raise ValueError( + "AgiBot fruit test split is missing camera features: " + f"{sorted(missing)}; actual={sorted(actual_videos)}" + ) + + total_episodes = int(info.get("total_episodes", 0)) + if total_episodes <= 0: + raise ValueError( + f"Invalid total_episodes in {info_path}: {total_episodes}" + ) + return info + + + +def _fruit_display_grid_rgb( + top_bgr: np.ndarray, + left_bgr: np.ndarray, + right_bgr: np.ndarray, +) -> np.ndarray: + """Build a 2x2 diagnostic grid even when source cameras differ in size.""" + display_size = (320, 176) + views = [] + for frame_bgr in (top_bgr, left_bgr, right_bgr): + frame_rgb = cv2.cvtColor(frame_bgr, cv2.COLOR_BGR2RGB) + frame_rgb = cv2.resize( + frame_rgb, + display_size, + interpolation=cv2.INTER_AREA, + ) + views.append(frame_rgb) + top, left, right = views + black = np.zeros_like(top) + return np.concatenate( + [ + np.concatenate([top, left], axis=1), + np.concatenate([right, black], axis=1), + ], + axis=0, + ) + + +def _run_one_test_sample( + wrapper: AgiBotRoboarenaPolicy, + args: Args, + info: dict, + episode_index: int, + eval_offsets: list[int], + raw_resolutions: dict[str, tuple[int, int]], +) -> dict: + root = Path(args.test_data_root).resolve() + chunks_size = int(info.get("chunks_size", 1000)) + parquet_path = _episode_file_from_template( + root, + info.get( + "data_path", + "data/chunk-{episode_chunk:03d}/episode_{episode_index:06d}.parquet", + ), + episode_index, + chunks_size, + ) + table = pq.read_table(parquet_path) + num_rows = int(table.num_rows) + if num_rows <= 0: + raise ValueError(f"Empty episode parquet: {parquet_path}") + + row_index = args.frame_index + if row_index < 0: + row_index = num_rows // 2 + minimum = -min(eval_offsets) + maximum_future = max(args.future_frames, 9) + if row_index < minimum: + raise IndexError( + f"frame_index={row_index} needs at least {minimum} previous " + f"frames for offsets {eval_offsets}" + ) + if row_index + maximum_future > num_rows: + row_index = num_rows - maximum_future + if row_index < minimum: + raise IndexError( + f"Episode {episode_index} is too short ({num_rows} rows) for " + f"offsets={eval_offsets}, future_frames={args.future_frames}" + ) + + if args.prompt_override is not None: + prompt = args.prompt_override + else: + prompt = "" + language_candidates = [ + "annotation.detailed_global_instruction_concise", + "annotation.language.action_text", + "annotation.human.action_text", + "task", + ] + for key in language_candidates: + if key in table.column_names: + prompt = _unwrap_text(table[key][row_index].as_py()) + if prompt: + break + + camera_features = { + "top_head": "observation.images.top_head", + "hand_left": "observation.images.hand_left", + "hand_right": "observation.images.hand_right", + } + context_indices = [row_index + offset for offset in eval_offsets] + first_decode = min(context_indices) + last_decode = min( + num_rows - 1, + row_index + args.future_frames - 1, + ) + decoded: dict[str, dict[int, np.ndarray]] = {} + video_paths: dict[str, str] = {} + video_template = info.get( + "video_path", + "videos/chunk-{episode_chunk:03d}/{video_key}/" + "episode_{episode_index:06d}.mp4", + ) + for short_name, feature_key in camera_features.items(): + path = _episode_file_from_template( + root, + video_template, + episode_index, + chunks_size, + video_key=feature_key, + ) + video_paths[short_name] = str(path) + decoded[short_name] = _decode_video_range_bgr( + path, + first_decode, + last_decode, + ) + + sample_name = f"episode_{episode_index:06d}_frame_{row_index:06d}" + sample_dir = Path(args.output_dir).resolve() / sample_name + sample_dir.mkdir(parents=True, exist_ok=True) + + context_grids = [ + _fruit_display_grid_rgb( + decoded["top_head"][index], + decoded["hand_left"][index], + decoded["hand_right"][index], + ) + for index in context_indices + ] + _save_rgb_video( + sample_dir / f"fruit_source_context_f{len(context_grids)}.mp4", + context_grids, + fps=30, + ) + imageio.imwrite( + sample_dir / "fruit_source_context_last.png", + context_grids[-1], + ) + + gt_indices = list(range(row_index, last_decode + 1)) + gt_grids = [ + _fruit_display_grid_rgb( + decoded["top_head"][index], + decoded["hand_left"][index], + decoded["hand_right"][index], + ) + for index in gt_indices + ] + _save_rgb_video( + sample_dir / f"fruit_ground_truth_future_f{len(gt_grids)}.mp4", + gt_grids, + fps=30, + ) + _save_rgb_video( + sample_dir / f"fruit_ground_truth_future_f{len(gt_grids)}_slow.mp4", + gt_grids, + fps=5, + ) + imageio.imwrite( + sample_dir / "fruit_ground_truth_first.png", + gt_grids[0], + ) + imageio.imwrite( + sample_dir / "fruit_ground_truth_last.png", + gt_grids[-1], + ) + + # Resize each source camera independently to the exact raw resolution + # expected by this AgiBot checkpoint. + context_bgr: dict[str, np.ndarray] = {} + for short_name in camera_features: + target_h, target_w = raw_resolutions[short_name] + frames = [ + cv2.resize( + decoded[short_name][index], + (target_w, target_h), + interpolation=cv2.INTER_LINEAR, + ) + for index in context_indices + ] + context_bgr[short_name] = np.stack(frames, axis=0) + + obs = { + "observation/top_head": context_bgr["top_head"], + "observation/hand_left": context_bgr["hand_left"], + "observation/hand_right": context_bgr["hand_right"], + "prompt": prompt, + "session_id": f"offline-fruit-{sample_name}-{uuid.uuid4()}", + } + + wrapper._output_dir = str(sample_dir) + wrapper._msg_index = 0 + logger.info( + "[AGIBOT-FRUIT TESTSET] episode=%s row=%s context=%s " + "prompt=%r raw_resolutions=%s parquet=%s", + episode_index, + row_index, + context_indices, + prompt, + { + key: [value[1], value[0]] + for key, value in raw_resolutions.items() + }, + parquet_path, + ) + ignored_actions = wrapper.infer(obs) + np.save( + sample_dir / "ignored_agibot_actions.npy", + ignored_actions, + ) + + predicted_candidates = sorted( + (sample_dir / "diagnostics").glob( + "pred_chunk_request_000001_f*.mp4" + ) + ) + comparison_native_path: str | None = None + comparison_slow_path: str | None = None + predicted_path: str | None = None + predicted_slow_path: str | None = None + + if predicted_candidates: + predicted_file = predicted_candidates[-1] + predicted_path = str(predicted_file) + predicted_frames = _read_mp4_rgb(predicted_file) + + predicted_slow_file = ( + sample_dir + / f"fruit_predicted_only_f{len(predicted_frames)}_slow.mp4" + ) + _save_rgb_video( + predicted_slow_file, + predicted_frames, + fps=5, + ) + predicted_slow_path = str(predicted_slow_file) + + compare_count = min(len(predicted_frames), len(gt_grids)) + comparison_frames: list[np.ndarray] = [] + for index in range(compare_count): + pred = predicted_frames[index] + gt = gt_grids[index] + if pred.shape[:2] != gt.shape[:2]: + pred = cv2.resize( + pred, + (gt.shape[1], gt.shape[0]), + interpolation=cv2.INTER_LINEAR, + ) + comparison_frames.append( + np.concatenate( + [ + _add_label_rgb(pred, "PREDICTED"), + _add_label_rgb( + gt, + "FRUIT TEST GROUND TRUTH", + ), + ], + axis=1, + ) + ) + + comparison_native = ( + sample_dir + / f"agibot_predicted_vs_fruit_gt_f{compare_count}_30fps.mp4" + ) + comparison_slow = ( + sample_dir + / f"agibot_predicted_vs_fruit_gt_f{compare_count}_slow.mp4" + ) + _save_rgb_video( + comparison_native, + comparison_frames, + fps=30, + ) + _save_rgb_video( + comparison_slow, + comparison_frames, + fps=5, + ) + imageio.imwrite( + sample_dir / "agibot_predicted_vs_fruit_gt_first.png", + comparison_frames[0], + ) + imageio.imwrite( + sample_dir / "agibot_predicted_vs_fruit_gt_last.png", + comparison_frames[-1], + ) + comparison_native_path = str(comparison_native) + comparison_slow_path = str(comparison_slow) + + summary = { + "checkpoint": str(Path(args.model_path).resolve()), + "checkpoint_embodiment": "agibot", + "visual_source_embodiment": "agibot", + "visual_source_task": "fruit", + "test_data_root": str(root), + "episode_index": episode_index, + "row_index": row_index, + "parquet_path": str(parquet_path), + "video_paths": video_paths, + "context_offsets": eval_offsets, + "context_indices": context_indices, + "checkpoint_raw_resolutions": { + key: [value[1], value[0]] + for key, value in raw_resolutions.items() + }, + "prompt": prompt, + "dataset_state_used_for_model_conditioning": False, + "predicted_video": predicted_path, + "predicted_video_slow": predicted_slow_path, + "comparison_video_30fps": comparison_native_path, + "comparison_video_slow": comparison_slow_path, + "generated_frame_count": 9 if predicted_path else 0, + "note": ( + "The model emits one 9-frame decoded video chunk in this evaluator. " + "The slow files change playback speed only; they do not extend the " + "model prediction horizon." + ), + } + with (sample_dir / "summary.json").open( + "w", + encoding="utf-8", + ) as stream: + json.dump(summary, stream, ensure_ascii=False, indent=2) + return summary + + + +def main(args: Args) -> None: + if args.embodiment_tag.lower() != "agibot": + raise ValueError( + "This evaluator runs AgiBot checkpoints on an AgiBot fruit test split; " + "use --embodiment-tag agibot" + ) + if args.future_frames <= 0: + raise ValueError("--future-frames must be positive") + + os.environ["ENABLE_DIT_CACHE"] = ( + "true" if args.enable_dit_cache else "false" + ) + if args.num_dit_steps is not None: + os.environ["NUM_DIT_STEPS"] = str(args.num_dit_steps) + elif args.enable_dit_cache: + os.environ.setdefault("NUM_DIT_STEPS", "8") + os.environ.setdefault("ATTENTION_BACKEND", "FA2") + os.environ.setdefault("TOKENIZERS_PARALLELISM", "false") + torch._dynamo.config.recompile_limit = 800 + + model_path = Path(args.model_path).resolve() + _require_dir(model_path, "AgiBot checkpoint") + metadata_path = model_path / "experiment_cfg" / "metadata.json" + conf_path = model_path / "experiment_cfg" / "conf.yaml" + _require_file(metadata_path, "AgiBot checkpoint metadata") + _require_file(conf_path, "AgiBot checkpoint config") + with metadata_path.open("r", encoding="utf-8") as stream: + checkpoint_metadata = json.load(stream) + if "agibot" not in checkpoint_metadata: + raise KeyError( + f"Checkpoint does not contain AgiBot metadata: " + f"{metadata_path}; keys={list(checkpoint_metadata)}" + ) + + test_info = _validate_test_dataset( + Path(args.test_data_root).resolve() + ) + episode_indices = _parse_episode_indices( + args.episode_indices, + int(test_info["total_episodes"]), + ) + eval_offsets = _checkpoint_eval_offsets( + model_path, + "agibot", + ) + raw_resolutions = _checkpoint_agibot_video_resolutions( + model_path + ) + model_config_overrides, train_config_overrides = ( + _build_path_overrides(args) + ) + + Path(args.output_dir).mkdir(parents=True, exist_ok=True) + logger.info( + "[PREFLIGHT OK] checkpoint=%s checkpoint_embodiment=agibot " + "visual_source=agibot_fruit test=%s episodes=%s eval_offsets=%s " + "raw_resolutions=%s output=%s", + model_path, + Path(args.test_data_root).resolve(), + episode_indices, + eval_offsets, + { + key: [value[1], value[0]] + for key, value in raw_resolutions.items() + }, + Path(args.output_dir).resolve(), + ) + if args.preflight_only: + logger.info( + "Preflight-only validation completed; model was not loaded." + ) + return + + device_mesh = init_mesh() + rank = dist.get_rank() + torch.manual_seed(args.seed) + np.random.seed(args.seed) + + timeout_delta = datetime.timedelta( + seconds=args.timeout_seconds + ) + signal_group = dist.new_group( + backend="gloo", + timeout=timeout_delta, + ) + logger.info( + "Rank %s initialized signal_group (gloo)", + rank, + ) + + policy = GrootSimPolicy( + embodiment_tag=EmbodimentTag("agibot"), + model_path=str(model_path), + device="cuda" if torch.cuda.is_available() else "cpu", + device_mesh=device_mesh, + model_config_overrides=model_config_overrides, + train_config_overrides=train_config_overrides, + ) + action_head = policy.trained_model.action_head + if args.num_inference_timesteps < 0: + raise ValueError( + "--num-inference-timesteps must be non-negative" + ) + if args.num_inference_timesteps > 0: + action_head.num_inference_steps = int( + args.num_inference_timesteps + ) + action_head.num_inference_timesteps = int( + args.num_inference_timesteps + ) + if hasattr(action_head, "config"): + action_head.config.num_inference_timesteps = int( + args.num_inference_timesteps + ) + logger.info( + "[CONFIG CHECK] rank=%s diffusion_steps=%s " + "frame_per_block=%s ENABLE_DIT_CACHE=%s", + rank, + getattr(action_head, "num_inference_steps", None), + getattr(action_head, "num_frame_per_block", None), + os.getenv("ENABLE_DIT_CACHE"), + ) + + if rank == 0: + wrapper = AgiBotRoboarenaPolicy( + groot_policy=policy, + signal_group=signal_group, + model_path=str(model_path), + output_dir=str(Path(args.output_dir).resolve()), + video_save_mode="none", + ) + summaries: list[dict] = [] + try: + for episode_index in episode_indices: + summaries.append( + _run_one_test_sample( + wrapper, + args, + test_info, + episode_index, + eval_offsets, + raw_resolutions, + ) + ) + finally: + wrapper._broadcast_signal_to_workers( + SIGNAL_SHUTDOWN + ) + + report_path = ( + Path(args.output_dir).resolve() + / "testset_report.json" + ) + with report_path.open( + "w", + encoding="utf-8", + ) as stream: + json.dump( + { + "checkpoint": str(model_path), + "checkpoint_embodiment": "agibot", + "visual_source_embodiment": "agibot", + "visual_source_task": "fruit", + "test_data_root": str( + Path(args.test_data_root).resolve() + ), + "episode_indices": episode_indices, + "eval_offsets": eval_offsets, + "checkpoint_raw_resolutions": { + key: [value[1], value[0]] + for key, value in raw_resolutions.items() + }, + "samples": summaries, + }, + stream, + ensure_ascii=False, + indent=2, + ) + logger.info( + "Saved test-set report: %s", + report_path, + ) + dist.barrier() + else: + worker = WebsocketPolicyServer( + policy=policy, + host="127.0.0.1", + port=0, + metadata={}, + output_dir=None, + signal_group=signal_group, + ) + asyncio.run(worker._worker_loop()) + dist.barrier() + + dist.destroy_process_group() + + +def cli() -> None: + logging.basicConfig(level=logging.INFO, force=True) + main(tyro.cli(Args)) + + +if __name__ == "__main__": + cli() diff --git a/eval_agibot_checkpoint_on_g2_testset.py b/eval_agibot_checkpoint_on_g2_testset.py new file mode 100644 index 00000000..dcb22dc5 --- /dev/null +++ b/eval_agibot_checkpoint_on_g2_testset.py @@ -0,0 +1,2290 @@ +import dataclasses +import json +import logging +import socket +import asyncio +import os +import http +import logging +import time +import traceback +from pathlib import Path +import glob +import uuid + +import pyarrow.parquet as pq +from omegaconf import OmegaConf +import torch +import tyro +from einops import rearrange +import datetime +import cv2 + +from groot.vla.model.n1_5.sim_policy import GrootSimPolicy +from groot.vla.data.schema import EmbodimentTag +import imageio +import numpy as np + +from openpi_client import base_policy as _base_policy +from openpi_client import msgpack_numpy +import websockets.asyncio.server as _server +import websockets.frames +from tianshou.data import Batch +import torch.distributed as dist +from torch.distributed.device_mesh import DeviceMesh, init_device_mesh + +# Use roboarena policy server interface +from eval_utils.policy_server import WebsocketPolicyServer as RoboarenaServer +from eval_utils.policy_server import PolicyServerConfig + +logger = logging.getLogger(__name__) + +SIGNAL_INFER = 0 +SIGNAL_SHUTDOWN = 1 +SIGNAL_IDLE = 2 +SIGNAL_RESET_CACHE = 3 + + +def _reset_policy_inference_cache(policy: object, reason: str) -> None: + trained_model = getattr(policy, "trained_model", None) + action_head = getattr(trained_model, "action_head", None) + reset_fn = getattr(action_head, "reset_inference_cache", None) + if callable(reset_fn): + reset_fn() + logger.info("Reset action-head inference cache on rank %s (%s)", dist.get_rank() if dist.is_initialized() else "?", reason) + else: + logger.warning("Policy action head does not expose reset_inference_cache(); cache reset skipped (%s)", reason) + +@dataclasses.dataclass +class Args: + # Model paths. + model_path: str = "/data/wangk/checkpoints/dreamzero_agibot_fruit_lora_20k/checkpoint-12000" + wan_ckpt_dir: str = "/data/wangk/checkpoints/Wan2.1-I2V-14B-480P" + tokenizer_path: str = "/data/wangk/checkpoints/umt5-xxl" + + # G2 held-out split used only as a fixed three-camera visual source. + test_data_root: str = "/data/training_data/teleop/g2/g2_tasks_g1_g7_joint_gear_subtask_v2/test" + + # Output and sample selection. + output_dir: str = "/data/wangk/dreamzero/video_compare_agibot_vs_g1ft_vs_g2ft/g1_finetuned" + episode_indices: str = "0" # Comma/range syntax: 0,3,7-9 + frame_index: int = 30 # Row/frame within each episode; -1 chooses the midpoint. + future_frames: int = 33 + prompt_override: str | None = None + session_id_override: str | None = None + + # Inference. + embodiment_tag: str = "agibot" + num_inference_timesteps: int = 4 + enable_dit_cache: bool = False + num_dit_steps: int | None = None + timeout_seconds: int = 50000 + seed: int = 42 + preflight_only: bool = False + + # Kept for compatibility with shared model/path helpers. + port: int = 0 + index: int = 0 + max_chunk_size: int | None = None + video_save_mode: str = "none" + + + +class DistributedRoboarenaPolicyBase: + """Shared distributed inference plumbing for websocket policy wrappers.""" + + def __init__( + self, + groot_policy: GrootSimPolicy, + signal_group: dist.ProcessGroup, + output_dir: str | None = None, + video_save_mode: str = "first", + ) -> None: + self._policy = groot_policy + self._signal_group = signal_group + self._output_dir = output_dir + self._video_save_mode = video_save_mode + self._frame_buffers = self._init_frame_buffers() + self._current_session_id: str | None = None + self.video_across_time = [] + self._msg_index = 0 + + if self._output_dir: + os.makedirs(self._output_dir, exist_ok=True) + + def _init_frame_buffers(self) -> dict[str, list[np.ndarray]]: + return { + "video.top_head": [], + "video.hand_left": [], + "video.hand_right": [], + } + + def _reset_custom_state(self) -> None: + pass + + def _after_infer(self) -> None: + pass + + def _prepare_video_chunk(self, video_pred: torch.Tensor) -> torch.Tensor | None: + if self._video_save_mode == "none": + return None + return video_pred + + def _video_save_fps(self) -> int: + return 5 + + def _convert_observation(self, obs: dict) -> dict: + raise NotImplementedError + + def _convert_action(self, action_dict: dict) -> np.ndarray: + raise NotImplementedError + + def _broadcast_batch_to_workers(self, obs: dict) -> None: + import pickle + + serialized = pickle.dumps(obs) + data_size = len(serialized) + + size_tensor = torch.tensor([data_size], dtype=torch.int64, device='cuda') + dist.broadcast(size_tensor, src=0) + + data_tensor = torch.frombuffer(serialized, dtype=torch.uint8).clone().cuda() + dist.broadcast(data_tensor, src=0) + + def _extract_action_dict(self, action_chunk_dict: object) -> dict[str, object]: + action_dict: dict[str, object] = {} + for key in dir(action_chunk_dict): + if key.startswith('action.'): + action_dict[key] = getattr(action_chunk_dict, key) + return action_dict + + def _broadcast_signal_to_workers(self, signal: int) -> None: + signal_tensor = torch.tensor([signal], dtype=torch.int32, device='cpu') + dist.broadcast(signal_tensor, src=0, group=self._signal_group) + + def _diagnostic_dir(self) -> str | None: + if not self._output_dir: + return None + path = os.path.join(self._output_dir, "diagnostics") + os.makedirs(path, exist_ok=True) + return path + + def _save_predicted_video_chunk(self, video_pred: torch.Tensor) -> None: + """Decode and save exactly one predicted latent chunk. + + This deliberately avoids concatenating independent 33-frame chunks + before VAE decoding, so chunk-boundary artifacts cannot hide the + quality of the current world-model prediction. + """ + diagnostic_dir = self._diagnostic_dir() + if diagnostic_dir is None: + return + + try: + if video_pred.ndim != 5: + raise ValueError( + "video_pred must be (B,C,T,H,W), got " + f"{tuple(video_pred.shape)}" + ) + finite = bool(torch.isfinite(video_pred).all().item()) + latent_float = video_pred.detach().float() + stats = { + "request_index": int(self._msg_index), + "shape": list(video_pred.shape), + "dtype": str(video_pred.dtype), + "finite": finite, + "min": float(latent_float.min().item()), + "max": float(latent_float.max().item()), + "mean": float(latent_float.mean().item()), + "std": float(latent_float.std().item()), + } + with open( + os.path.join(diagnostic_dir, "video_pred_stats.jsonl"), + "a", + encoding="utf-8", + ) as stream: + stream.write(json.dumps(stats, ensure_ascii=False) + "\n") + + if not finite: + raise ValueError(f"video_pred contains NaN/Inf: {stats}") + + action_head = self._policy.trained_model.action_head + device = getattr(action_head, "_device", None) + if device is None: + device = next(self._policy.trained_model.parameters()).device + + latent = video_pred.detach().to( + device=device, + dtype=torch.bfloat16, + ) + frames = action_head.vae.decode( + latent, + tiled=action_head.tiled, + tile_size=( + action_head.tile_size_height, + action_head.tile_size_width, + ), + tile_stride=( + action_head.tile_stride_height, + action_head.tile_stride_width, + ), + ) + frames = rearrange(frames, "B C T H W -> B T H W C")[0] + frames = ( + (frames.float() + 1.0) * 127.5 + ).clamp(0, 255).cpu().numpy().astype(np.uint8) + + output_path = os.path.join( + diagnostic_dir, + f"pred_chunk_request_{self._msg_index:06d}_f{len(frames)}.mp4", + ) + imageio.mimsave( + output_path, + list(frames), + fps=self._video_save_fps(), + codec="libx264", + macro_block_size=None, + ) + imageio.imwrite( + os.path.join( + diagnostic_dir, + f"pred_chunk_request_{self._msg_index:06d}_first.png", + ), + frames[0], + ) + imageio.imwrite( + os.path.join( + diagnostic_dir, + f"pred_chunk_request_{self._msg_index:06d}_last.png", + ), + frames[-1], + ) + + # DreamZero packs three 320x176 views into a 2x2 RGB canvas in + # ConcatTransform order: top_head, hand_left, hand_right, padding. + # Decode the full canvas once, then split decoded RGB only for + # diagnostics. Never split the latent before VAE decoding. + frame_h, frame_w = int(frames.shape[1]), int(frames.shape[2]) + if frame_h % 2 == 0 and frame_w % 2 == 0: + half_h, half_w = frame_h // 2, frame_w // 2 + view_frames = { + "top_head": frames[:, :half_h, :half_w], + "hand_left": frames[:, :half_h, half_w:], + "hand_right": frames[:, half_h:, :half_w], + "padding": frames[:, half_h:, half_w:], + } + view_stats = {} + for view_name, view_video in view_frames.items(): + view_path = os.path.join( + diagnostic_dir, + f"pred_chunk_request_{self._msg_index:06d}_{view_name}_f{len(view_video)}.mp4", + ) + imageio.mimsave( + view_path, + list(view_video), + fps=self._video_save_fps(), + codec="libx264", + macro_block_size=None, + ) + imageio.imwrite( + os.path.join( + diagnostic_dir, + f"pred_chunk_request_{self._msg_index:06d}_{view_name}_first.png", + ), + view_video[0], + ) + imageio.imwrite( + os.path.join( + diagnostic_dir, + f"pred_chunk_request_{self._msg_index:06d}_{view_name}_last.png", + ), + view_video[-1], + ) + view_float = view_video.astype(np.float32) + view_stats[view_name] = { + "shape": list(view_video.shape), + "mean_rgb": [ + float(x) + for x in view_float.mean(axis=(0, 1, 2)).tolist() + ], + "std": float(view_float.std()), + } + stats["decoded_shape"] = list(frames.shape) + stats["decoded_views"] = view_stats + with open( + os.path.join(diagnostic_dir, "decoded_view_stats.jsonl"), + "a", + encoding="utf-8", + ) as stream: + stream.write(json.dumps(stats, ensure_ascii=False) + "\n") + else: + logger.warning( + "Decoded video is not divisible into a 2x2 view grid: %s", + tuple(frames.shape), + ) + + logger.info( + "Saved single predicted video chunk to %s | stats=%s", + output_path, + stats, + ) + except Exception as exc: + logger.exception( + "Failed to save single predicted video chunk: %s", exc + ) + + def infer(self, obs: dict) -> np.ndarray: + session_id = obs.get('session_id') + if session_id is not None and session_id != self._current_session_id: + if self._current_session_id is not None: + logger.info("Session changed from '%s' to '%s', resetting state", self._current_session_id, session_id) + self._broadcast_signal_to_workers(SIGNAL_RESET_CACHE) + self._reset_state() + else: + logger.info("New session started: '%s'", session_id) + self._current_session_id = session_id + + self._msg_index += 1 + converted_obs = self._convert_observation(obs) + + self._broadcast_signal_to_workers(SIGNAL_INFER) + self._broadcast_batch_to_workers(converted_obs) + + batch = Batch(obs=converted_obs) + dist.barrier() + with torch.no_grad(): + result_batch, video_pred = self._policy.lazy_joint_forward_causal(batch) + dist.barrier() + + self._save_predicted_video_chunk(video_pred) + + video_chunk = self._prepare_video_chunk(video_pred) + if video_chunk is not None: + self.video_across_time.append(video_chunk.detach().cpu()) + action = self._convert_action(self._extract_action_dict(result_batch.act)) + self._after_infer() + return action + + def _reset_state(self, save_video: bool = True) -> None: + if save_video and len(self.video_across_time) > 0 and self._output_dir: + try: + frame_list = [] + action_head = self._policy.trained_model.action_head + device = getattr(action_head, "_device", None) + if device is None: + device = next(self._policy.trained_model.parameters()).device + video_across_time_cat = torch.cat(self.video_across_time, dim=2).to(device=device, dtype=torch.bfloat16) + frames = action_head.vae.decode( + video_across_time_cat, + tiled=action_head.tiled, + tile_size=(action_head.tile_size_height, action_head.tile_size_width), + tile_stride=(action_head.tile_stride_height, action_head.tile_stride_width), + ) + frames = rearrange(frames, 'B C T H W -> B T H W C') + frames = frames[0] + frames = ((frames.float() + 1) * 127.5).clip(0, 255).cpu().numpy().astype(np.uint8) + for frame in frames: + frame_list.append(frame) + + if frame_list: + sample_frame = frame_list[0] + if len(sample_frame.shape) == 3 and sample_frame.shape[2] in [1, 3, 4]: + save_dir = self._output_dir + os.makedirs(save_dir, exist_ok=True) + all_mp4_files = [f for f in os.listdir(save_dir) if f.endswith('.mp4')] + timestamp = datetime.datetime.now().strftime('%Y%m%d_%H%M%S') + num_frames = len(frame_list) + output_path = os.path.join(save_dir, f'{timestamp}_{len(all_mp4_files):06}_f{num_frames}.mp4') + imageio.mimsave(output_path, frame_list, fps=self._video_save_fps(), codec='libx264') + logger.info('Saved video on reset to: %s', output_path) + except Exception as exc: + logger.warning('Failed to save video on reset: %s', exc) + + for key in self._frame_buffers: + self._frame_buffers[key] = [] + + self.video_across_time = [] + _reset_policy_inference_cache(self._policy, "wrapper reset_state") + self._reset_custom_state() + + def reset(self, reset_info: dict) -> None: + self._broadcast_signal_to_workers(SIGNAL_RESET_CACHE) + self._reset_state(save_video=True) + + +class ARDroidRoboarenaPolicy(DistributedRoboarenaPolicyBase): + """Wrapper policy that implements roboarena.policy.BasePolicy interface for AR_droid.""" + + FRAMES_PER_CHUNK = 4 + + def __init__( + self, + groot_policy: GrootSimPolicy, + signal_group: dist.ProcessGroup, + output_dir: str | None = None, + video_save_mode: str = "first", + ) -> None: + super().__init__( + groot_policy=groot_policy, + signal_group=signal_group, + output_dir=output_dir, + video_save_mode=video_save_mode, + ) + self._reset_custom_state() + + def _init_frame_buffers(self) -> dict[str, list[np.ndarray]]: + return { + 'video.exterior_image_1_left': [], + 'video.exterior_image_2_left': [], + 'video.wrist_image_left': [], + } + + def _reset_custom_state(self) -> None: + self._is_first_call = True + + def _after_infer(self) -> None: + self._is_first_call = False + + def _convert_observation(self, obs: dict) -> dict: + converted = {} + image_key_mapping = { + 'observation/exterior_image_0_left': 'video.exterior_image_1_left', + 'observation/exterior_image_1_left': 'video.exterior_image_2_left', + 'observation/wrist_image_left': 'video.wrist_image_left', + } + + for roboarena_key, droid_key in image_key_mapping.items(): + if roboarena_key in obs: + data = obs[roboarena_key] + if isinstance(data, np.ndarray): + if data.ndim == 4: + self._frame_buffers[droid_key].extend(list(data)) + else: + self._frame_buffers[droid_key].append(data) + + num_frames = 1 if self._is_first_call else self.FRAMES_PER_CHUNK + + for droid_key, buffer in self._frame_buffers.items(): + if len(buffer) > 0: + if len(buffer) >= num_frames: + frames_to_use = buffer[-num_frames:] + else: + frames_to_use = buffer.copy() + while len(frames_to_use) < num_frames: + frames_to_use.insert(0, buffer[0]) + converted[droid_key] = np.stack(frames_to_use, axis=0) + + joint_pos = obs.get('observation/joint_position', np.zeros(7, dtype=np.float32)) + if joint_pos.ndim == 1: + joint_pos = joint_pos.reshape(1, -1) + converted['state.joint_position'] = joint_pos.astype(np.float64) + + gripper_pos = obs.get('observation/gripper_position', np.zeros(1, dtype=np.float32)) + if gripper_pos.ndim == 1: + gripper_pos = gripper_pos.reshape(1, -1) + converted['state.gripper_position'] = gripper_pos.astype(np.float64) + converted['annotation.language.action_text'] = obs.get('prompt', '') + return converted + + def _convert_action(self, action_dict: dict) -> np.ndarray: + joint_action = None + gripper_action = None + for key, value in action_dict.items(): + if 'joint_position' in key: + joint_action = value + elif 'gripper_position' in key or 'gripper' in key: + gripper_action = value + + if joint_action is None: + return np.zeros((1, 8), dtype=np.float32) + + if isinstance(joint_action, torch.Tensor): + joint_action = joint_action.cpu().numpy() + if joint_action.ndim == 1: + joint_action = joint_action.reshape(1, -1) + + num_steps = joint_action.shape[0] + if gripper_action is not None: + if isinstance(gripper_action, torch.Tensor): + gripper_action = gripper_action.cpu().numpy() + if gripper_action.ndim == 1: + gripper_action = gripper_action.reshape(-1, 1) + elif gripper_action.ndim == 0: + gripper_action = gripper_action.reshape(1, 1) + else: + gripper_action = np.zeros((num_steps, 1), dtype=np.float32) + + return np.concatenate([joint_action, gripper_action], axis=-1).astype(np.float32) + + + +def _load_agibot_fixed_state_from_checkpoint( + model_path: str, + required_state_keys: list[str], +) -> dict[str, np.ndarray]: + """Load nominal AgiBot state values from checkpoint metadata.""" + metadata_path = os.path.join( + os.path.abspath(model_path), + "experiment_cfg", + "metadata.json", + ) + with open(metadata_path, "r", encoding="utf-8") as stream: + metadata = json.load(stream) + + try: + stats = metadata["agibot"]["statistics"]["state"] + except KeyError as exc: + raise KeyError( + f"{metadata_path} does not contain agibot.statistics.state" + ) from exc + + fixed: dict[str, np.ndarray] = {} + for target_key in required_state_keys: + if not target_key.startswith("state."): + raise ValueError( + f"Unexpected AgiBot state modality key: {target_key!r}" + ) + source_key = target_key.split(".", 1)[1] + if source_key not in stats or "mean" not in stats[source_key]: + raise KeyError( + f"Missing checkpoint-mean state for {source_key!r} " + f"in {metadata_path}" + ) + fixed[target_key] = np.asarray( + stats[source_key]["mean"], + dtype=np.float64, + ).reshape(1, -1) + return fixed + + +class AgiBotRoboarenaPolicy(DistributedRoboarenaPolicyBase): + """Run an AgiBot checkpoint on fixed G2 test-set camera frames. + + This is intentionally video-only. G2 proprioception is not used for + conditioning; the checkpoint's own mean AgiBot state fills its required + state fields. The native AgiBot action output is discarded. + """ + + VIDEO_KEY_MAPPING = { + "observation/top_head": "video.top_head", + "observation/hand_left": "video.hand_left", + "observation/hand_right": "video.hand_right", + } + + def __init__( + self, + groot_policy: GrootSimPolicy, + signal_group: dist.ProcessGroup, + model_path: str, + output_dir: str | None = None, + video_save_mode: str = "first", + ) -> None: + super().__init__( + groot_policy=groot_policy, + signal_group=signal_group, + output_dir=output_dir, + video_save_mode=video_save_mode, + ) + self._action_keys = list( + self._policy.modality_configs.action.modality_keys + ) + self._state_keys = list( + self._policy.modality_configs.state.modality_keys + ) + self._language_keys = list( + self._policy.modality_configs.language.modality_keys + ) + eval_indices = list( + getattr( + self._policy.modality_configs.video, + "eval_delta_indices", + [0], + ) + or [0] + ) + self._expected_video_frames = max(1, len(eval_indices)) + self._fixed_state = _load_agibot_fixed_state_from_checkpoint( + model_path, + self._state_keys, + ) + self._safe_g2_hold_action = np.zeros(16, dtype=np.float32) + logger.info( + "[AGIBOT-ON-G2 VIDEO TEST] video_frames=%s " + "state_keys=%s language_keys=%s", + self._expected_video_frames, + self._state_keys, + self._language_keys, + ) + + @staticmethod + def _normalize_video( + value: object, + target_key: str, + ) -> np.ndarray: + array_bgr = np.asarray(value) + if array_bgr.ndim == 3: + array_bgr = np.expand_dims(array_bgr, axis=0) + elif array_bgr.ndim != 4: + raise ValueError( + f"AgiBot video input for {target_key} must have " + f"shape (H,W,C) or (T,H,W,C), got {array_bgr.shape}" + ) + if array_bgr.shape[-1] != 3: + raise ValueError( + f"AgiBot video input for {target_key} must have 3 channels, " + f"got {array_bgr.shape}" + ) + # OpenCV decodes the G2 test MP4 files as BGR. DreamZero expects RGB. + return np.ascontiguousarray(array_bgr[..., ::-1]) + + def _prepare_video_chunk( + self, + video_pred: torch.Tensor, + ) -> torch.Tensor | None: + if self._video_save_mode == "none": + return None + if video_pred.ndim != 5: + raise ValueError( + "AgiBot video prediction must be 5D " + f"(B,C,T,H,W), got {tuple(video_pred.shape)}" + ) + if self._video_save_mode == "first": + return video_pred[:, :, :1].contiguous() + if self._video_save_mode == "full": + return video_pred.contiguous() + raise ValueError( + f"Unsupported video_save_mode: {self._video_save_mode!r}" + ) + + def _video_save_fps(self) -> int: + return 30 + + def _save_input_grid( + self, + converted: dict[str, object], + ) -> None: + diagnostic_dir = self._diagnostic_dir() + if diagnostic_dir is None: + return + + camera_keys = ( + "video.top_head", + "video.hand_left", + "video.hand_right", + ) + arrays = [np.asarray(converted[key]) for key in camera_keys] + frame_count = min(array.shape[0] for array in arrays) + grids: list[np.ndarray] = [] + + # The checkpoint may require different raw resolutions per camera. + # Resize only for this human-readable diagnostic grid; the arrays sent + # to the model remain at their native checkpoint-required sizes. + display_size = (320, 176) + for index in range(frame_count): + display_views = [ + cv2.resize( + np.asarray(array[index]), + display_size, + interpolation=cv2.INTER_AREA, + ) + for array in arrays + ] + top_head, hand_left, hand_right = display_views + black = np.zeros_like(top_head) + grids.append( + np.concatenate( + [ + np.concatenate([top_head, hand_left], axis=1), + np.concatenate([hand_right, black], axis=1), + ], + axis=0, + ) + ) + + grid_path = os.path.join( + diagnostic_dir, + f"agibot_from_g2_input_request_{self._msg_index:06d}" + f"_f{len(grids)}.mp4", + ) + imageio.mimsave( + grid_path, + grids, + fps=30, + codec="libx264", + macro_block_size=None, + ) + imageio.imwrite( + os.path.join( + diagnostic_dir, + f"agibot_from_g2_input_request_{self._msg_index:06d}" + "_last.png", + ), + grids[-1], + ) + logger.info( + "Saved AgiBot model-input display grid to %s | native_shapes=%s", + grid_path, + { + key: list(array.shape) + for key, array in zip(camera_keys, arrays) + }, + ) + + def _convert_observation( + self, + obs: dict, + ) -> dict: + converted: dict[str, object] = {} + missing_video: list[str] = [] + + for source_key, target_key in self.VIDEO_KEY_MAPPING.items(): + value = obs.get(source_key, obs.get(target_key)) + if value is None: + missing_video.append(source_key) + continue + frames = self._normalize_video(value, target_key) + self._frame_buffers[target_key].extend(list(frames)) + history = list( + self._frame_buffers[target_key][ + -self._expected_video_frames: + ] + ) + while len(history) < self._expected_video_frames: + history.insert(0, history[0]) + converted[target_key] = np.stack(history, axis=0) + + if missing_video: + raise ValueError( + "Missing G2 camera inputs for AgiBot video test: " + + ", ".join(sorted(missing_video)) + ) + + packed = np.asarray( + obs.get("observation/state", np.zeros(16, dtype=np.float32)), + dtype=np.float32, + ) + if packed.ndim == 2: + packed = packed[-1] + if packed.ndim == 1 and packed.shape[0] == 16: + self._safe_g2_hold_action = packed.copy() + + for key, value in self._fixed_state.items(): + converted[key] = value.copy() + + prompt = str(obs.get("prompt", "")) + for language_key in self._language_keys: + converted[language_key] = prompt + + self._save_input_grid(converted) + logger.info( + "[AGIBOT-ON-G2 VIDEO TEST] request=%s " + "frames_per_view=%s language_keys=%s prompt=%r", + self._msg_index, + { + key: int(np.asarray(converted[key]).shape[0]) + for key in self.VIDEO_KEY_MAPPING.values() + }, + self._language_keys, + prompt, + ) + return converted + + def _convert_action( + self, + action_dict: dict, + ) -> np.ndarray: + # Native AgiBot actions are irrelevant for this diagnostic. Returning + # a G2 hold-shaped tensor keeps the rest of the offline evaluator + # simple and prevents accidental dimension comparisons. + return np.repeat( + self._safe_g2_hold_action.reshape(1, 16), + 48, + axis=0, + ).astype(np.float32) + + +class G2RoboarenaPolicy(DistributedRoboarenaPolicyBase): + """Adapter for the G2 dual-arm joint-space policy.""" + FRAMES_PER_CHUNK = 4 + VIDEO_KEY_MAPPING = { + 'observation/top_head': 'video.top_head', + 'observation/hand_left': 'video.hand_left', + 'observation/hand_right': 'video.hand_right', + } + + STATE_KEY_MAPPING = { + 'observation/left_joint_position': 'state.left_joint_position', + 'observation/left_gripper_position': 'state.left_gripper_position', + 'observation/right_joint_position': 'state.right_joint_position', + 'observation/right_gripper_position': 'state.right_gripper_position', + } + + PACKED_STATE_KEYS = ( + 'observation/state', + 'observation.state', + 'state', + ) + + def __init__( + self, + groot_policy: GrootSimPolicy, + signal_group: dist.ProcessGroup, + output_dir: str | None = None, + video_save_mode: str = "first", + ) -> None: + super().__init__( + groot_policy=groot_policy, + signal_group=signal_group, + output_dir=output_dir, + video_save_mode=video_save_mode, + ) + self._action_keys = list( + self._policy.modality_configs.action.modality_keys + ) + eval_indices = list( + getattr( + self._policy.modality_configs.video, + "eval_delta_indices", + [-3, -2, -1, 0], + ) + or [0] + ) + self._expected_video_frames = max(1, len(eval_indices)) + logger.info( + "G2 wrapper expects %s evaluation frame(s): %s", + self._expected_video_frames, + eval_indices, + ) + + @staticmethod + def _lookup_obs_value( + obs: dict, + source_key: str, + target_key: str, + ) -> object: + if source_key in obs: + return obs[source_key] + return obs.get(target_key) + + @staticmethod + def _normalize_video( + value: object, + target_key: str, + ) -> np.ndarray: + if ( + isinstance(value, dict) + and value.get("__dreamzero_image_encoding__") + == "jpeg_sequence" + ): + frames = [] + expected_shape = tuple(value.get("shape", ())) + expected_dtype = np.dtype( + value.get("dtype", "uint8") + ) + for index, frame_bytes in enumerate( + value.get("frames", []) + ): + encoded = np.frombuffer( + frame_bytes, + dtype=np.uint8, + ) + frame_bgr = cv2.imdecode( + encoded, + cv2.IMREAD_COLOR, + ) + if frame_bgr is None: + raise ValueError( + f"Failed to decode JPEG frame {index} " + f"for {target_key}" + ) + # The G2 client normalizes every camera frame to OpenCV BGR. + # DreamZero's training video path is RGB, so convert exactly + # once at the server/model boundary. + frame_rgb = cv2.cvtColor( + frame_bgr, + cv2.COLOR_BGR2RGB, + ) + frames.append( + frame_rgb.astype( + expected_dtype, + copy=False, + ) + ) + if not frames: + raise ValueError( + f"No JPEG frames were provided for {target_key}" + ) + array = np.stack(frames, axis=0) + if expected_shape and tuple(array.shape) != expected_shape: + raise ValueError( + f"Decoded JPEG video for {target_key} has " + f"shape {array.shape}, expected {expected_shape}" + ) + return array + + array_bgr = np.asarray(value) + if array_bgr.ndim == 3: + array_bgr = np.expand_dims(array_bgr, axis=0) + elif array_bgr.ndim != 4: + raise ValueError( + f"G2 video input for {target_key} must have " + f"shape (H,W,C) or (T,H,W,C), got {array_bgr.shape}" + ) + if array_bgr.shape[-1] != 3: + raise ValueError( + f"G2 video input for {target_key} must have 3 channels, " + f"got {array_bgr.shape}" + ) + # The non-JPEG G2 client path is also BGR because G2Camera._decode() + # converts RGB camera buffers to BGR and leaves BGR buffers unchanged. + # Convert it to the same RGB convention as the JPEG branch. + return np.ascontiguousarray(array_bgr[..., ::-1]) + + @staticmethod + def _normalize_state( + value: object, + target_key: str, + ) -> np.ndarray: + array = np.asarray(value, dtype=np.float64) + if array.ndim == 0: + array = array.reshape(1, 1) + elif array.ndim == 1: + array = array.reshape(1, -1) + elif array.ndim != 2: + raise ValueError( + f"G2 state input for {target_key} must be " + f"1D or 2D, got {array.shape}" + ) + return array + + @staticmethod + def _split_packed_state( + value: object, + ) -> dict[str, np.ndarray]: + packed = np.asarray(value, dtype=np.float64) + if packed.ndim == 1: + packed = packed.reshape(1, -1) + elif packed.ndim != 2: + raise ValueError( + "Packed G2 state must have shape (16,) or " + f"(T,16), got {packed.shape}" + ) + if packed.shape[-1] != 16: + raise ValueError( + f"Packed G2 state must contain 16 values, " + f"got {packed.shape}" + ) + return { + 'state.left_joint_position': packed[:, 0:7], + 'state.left_gripper_position': packed[:, 7:8], + 'state.right_joint_position': packed[:, 8:15], + 'state.right_gripper_position': packed[:, 15:16], + } + + def _prepare_video_chunk( + self, + video_pred: torch.Tensor, + ) -> torch.Tensor | None: + if self._video_save_mode == "none": + return None + if video_pred.ndim != 5: + raise ValueError( + "G2 video prediction must be 5D " + f"(B,C,T,H,W), got {tuple(video_pred.shape)}" + ) + if self._video_save_mode == "first": + return video_pred[:, :, :1].contiguous() + if self._video_save_mode == "full": + return video_pred.contiguous() + raise ValueError( + f"Unsupported video_save_mode: " + f"{self._video_save_mode!r}" + ) + + def _video_save_fps(self) -> int: + return 30 + + def _save_input_grid(self, converted: dict[str, object]) -> None: + """Save the exact four-frame, three-camera history sent to the model.""" + diagnostic_dir = self._diagnostic_dir() + if diagnostic_dir is None: + return + + camera_keys = ( + "video.top_head", + "video.hand_left", + "video.hand_right", + ) + if any(key not in converted for key in camera_keys): + return + + arrays = [np.asarray(converted[key]) for key in camera_keys] + if any(array.ndim != 4 for array in arrays): + raise ValueError( + "Diagnostic input videos must be (T,H,W,C): " + + ", ".join( + f"{key}={array.shape}" + for key, array in zip(camera_keys, arrays) + ) + ) + frame_count = min(array.shape[0] for array in arrays) + grids: list[np.ndarray] = [] + for index in range(frame_count): + top_head, hand_left, hand_right = [ + np.asarray(array[index]) for array in arrays + ] + height, width = top_head.shape[:2] + black = np.zeros_like(top_head) + grid_rgb = np.concatenate( + [ + np.concatenate([top_head, hand_left], axis=1), + np.concatenate([hand_right, black], axis=1), + ], + axis=0, + ) + # _normalize_video() already returns RGB for both JPEG and array + # transport. Keep the diagnostic image byte-for-byte in model + # color order instead of swapping it a second time. + grids.append(np.ascontiguousarray(grid_rgb)) + + output_path = os.path.join( + diagnostic_dir, + f"input_grid_request_{self._msg_index:06d}_f{len(grids)}.mp4", + ) + imageio.mimsave( + output_path, + grids, + fps=5, + codec="libx264", + macro_block_size=None, + ) + imageio.imwrite( + os.path.join( + diagnostic_dir, + f"input_grid_request_{self._msg_index:06d}_last.png", + ), + grids[-1], + ) + channel_means = { + key: [ + float(x) + for x in np.asarray(array, dtype=np.float32).mean( + axis=(0, 1, 2) + ).tolist() + ] + for key, array in zip(camera_keys, arrays) + } + logger.info( + "Saved exact RGB model input grid to %s | shapes=%s | mean_rgb=%s", + output_path, + {key: list(array.shape) for key, array in zip(camera_keys, arrays)}, + channel_means, + ) + + def _convert_observation(self, obs: dict) -> dict: + converted: dict[str, object] = {} + missing_video: list[str] = [] + + for source_key, target_key in self.VIDEO_KEY_MAPPING.items(): + value = self._lookup_obs_value( + obs, + source_key, + target_key, + ) + if value is None: + missing_video.append(source_key) + continue + frames = self._normalize_video( + value, + target_key, + ) + + self._frame_buffers[target_key].extend( + list(frames) + ) + + history = self._frame_buffers[target_key][ + -self._expected_video_frames: + ] + + while len(history) < self._expected_video_frames: + history.insert(0, history[0]) + + converted[target_key] = np.stack(history, axis=0) + + if missing_video: + raise ValueError( + "G2 inference requires video keys: " + + ", ".join(sorted(missing_video)) + ) + + packed_state = None + for key in self.PACKED_STATE_KEYS: + if key in obs: + packed_state = obs[key] + break + + if packed_state is not None: + converted.update( + self._split_packed_state(packed_state) + ) + else: + missing_state: list[str] = [] + for source_key, target_key in self.STATE_KEY_MAPPING.items(): + value = self._lookup_obs_value( + obs, + source_key, + target_key, + ) + if value is None: + missing_state.append(source_key) + continue + converted[target_key] = self._normalize_state( + value, + target_key, + ) + if missing_state: + raise ValueError( + "G2 inference requires a packed 16-D state " + "under observation/state, observation.state, " + "or state; otherwise all split state keys are " + "required: " + + ", ".join(sorted(missing_state)) + ) + + expected_dims = { + 'state.left_joint_position': 7, + 'state.left_gripper_position': 1, + 'state.right_joint_position': 7, + 'state.right_gripper_position': 1, + } + for key, expected_dim in expected_dims.items(): + array = np.asarray(converted[key]) + if array.shape[-1] != expected_dim: + raise ValueError( + f"{key} must have last dimension " + f"{expected_dim}, got {array.shape}" + ) + + converted['annotation.language.action_text'] = obs.get( + 'prompt', + obs.get( + 'annotation.language.action_text', + '', + ), + ) + self._save_input_grid(converted) + return converted + + def _convert_action( + self, + action_dict: dict, + ) -> np.ndarray: + missing = [ + key + for key in self._action_keys + if key not in action_dict + ] + if missing: + raise RuntimeError( + "Missing G2 action outputs: " + + ", ".join(missing) + ) + + arrays: list[np.ndarray] = [] + horizon: int | None = None + for key in self._action_keys: + value = action_dict[key] + if isinstance(value, torch.Tensor): + value = value.detach().cpu().numpy() + array = np.asarray(value) + if array.ndim == 0: + array = array.reshape(1, 1) + elif array.ndim == 1: + array = array.reshape(-1, 1) + else: + array = array.reshape(array.shape[0], -1) + + if horizon is None: + horizon = array.shape[0] + elif array.shape[0] != horizon: + raise RuntimeError( + f"Inconsistent G2 action horizon for {key}: " + f"expected {horizon}, got {array.shape[0]}" + ) + arrays.append(array.astype(np.float32)) + + action = np.concatenate(arrays, axis=-1) + if action.shape[-1] != 16: + raise RuntimeError( + "G2 action must have 16 dimensions ordered as " + "[left_joint(7), left_gripper(1), " + "right_joint(7), right_gripper(1)], " + f"got {action.shape}" + ) + return action.astype(np.float32) + + +class WebsocketPolicyServer: + """Serves a policy using the websocket protocol. See websocket_client_policy.py for a client implementation. + Currently only implements the `load` and `infer` methods. + """ + + def __init__( + self, + policy: _base_policy.BasePolicy, + host: str = "0.0.0.0", + port: int | None = None, + metadata: dict | None = None, + output_dir: str | None = None, + signal_group: dist.ProcessGroup | None = None, + ) -> None: + self._policy = policy + self._host = host + self._port = port + self._metadata = metadata or {} + self._output_dir = output_dir + logging.getLogger("websockets.server").setLevel(logging.INFO) + self.video_across_time = [] + self._msg_index = 0 + self._signal_group = signal_group + if self._output_dir: + os.makedirs(self._output_dir, exist_ok=True) + os.makedirs(os.path.join(self._output_dir, "inputs"), exist_ok=True) + + def serve_forever(self, rank: int = 0) -> None: + asyncio.run(self.run(rank)) + + async def run(self, rank: int = 0): + if rank == 0: + async with _server.serve( + self._handler, + self._host, + self._port, + compression=None, + max_size=None, + process_request=_health_check, + ping_interval=None, + ) as server: + await server.serve_forever() + else: + await self._worker_loop() + + async def _worker_loop(self): + logger.info(f"Worker loop started for rank {dist.get_rank()}") + signal_tensor = torch.zeros(1, dtype=torch.int32, device='cpu') + while True: + try: + dist.broadcast(signal_tensor, src=0, group=self._signal_group) + + signal = signal_tensor.item() + if signal == SIGNAL_SHUTDOWN: + logger.info(f"Rank {dist.get_rank()} received shutdown signal") + break + elif signal == SIGNAL_IDLE: + logger.info(f"Rank {dist.get_rank()} received idle signal. Waiting for next client.") + continue + elif signal == SIGNAL_RESET_CACHE: + logger.info(f"Rank {dist.get_rank()} received inference cache reset signal") + _reset_policy_inference_cache(self._policy, "worker signal") + continue + + batch = self._receive_batch_from_rank0() + dist.barrier() + with torch.no_grad(): + result_batch, video_pred = self._policy.lazy_joint_forward_causal(batch) + dist.barrier() + + except Exception as e: + logger.error(f"Worker loop error on rank {dist.get_rank()}: {e}") + traceback.print_exc() + break + + def _receive_batch_from_rank0(self): + import pickle + + size_tensor = torch.zeros(1, dtype=torch.int64, device='cuda') + dist.broadcast(size_tensor, src=0) + data_size = size_tensor.item() + + data_tensor = torch.zeros(data_size, dtype=torch.uint8, device='cuda') + dist.broadcast(data_tensor, src=0) + + obs = pickle.loads(data_tensor.cpu().numpy().tobytes()) + return Batch(obs=obs) + + def _broadcast_batch_to_workers(self, obs): + import pickle + + serialized = pickle.dumps(obs) + data_size = len(serialized) + + size_tensor = torch.tensor([data_size], dtype=torch.int64, device='cuda') + dist.broadcast(size_tensor, src=0) + + data_tensor = torch.frombuffer(serialized, dtype=torch.uint8).clone().cuda() + dist.broadcast(data_tensor, src=0) + + async def _handler(self, websocket: _server.ServerConnection): + logger.info(f"Connection from {websocket.remote_address} opened") + packer = msgpack_numpy.Packer() + + await websocket.send(packer.pack(self._metadata)) + + signal_tensor = torch.zeros(1, dtype=torch.int32, device='cpu') + + try: + while True: + try: + data = await websocket.recv() + obs = msgpack_numpy.unpackb(data) + self._msg_index += 1 + + signal_tensor.zero_() + dist.broadcast(signal_tensor, src=0, group=self._signal_group) + + self._broadcast_batch_to_workers(obs) + batch = Batch(obs=obs) + + dist.barrier() + with torch.no_grad(): + result_batch, video_pred = self._policy.lazy_joint_forward_causal(batch) + dist.barrier() + + action_chunk_dict = result_batch.act + + def batch_to_dict(batch): + out = {} + for k in dir(batch): + if not k.startswith("action."): + continue + out[k] = getattr(batch, k) + return out + + action_chunk_dict = batch_to_dict(action_chunk_dict) + await websocket.send(packer.pack(action_chunk_dict)) + + except websockets.ConnectionClosed: + logger.info(f"Connection from {websocket.remote_address} closed") + self.video_across_time = [] + break + except Exception: + await websocket.send(traceback.format_exc()) + await websocket.close( + code=websockets.frames.CloseCode.INTERNAL_ERROR, + reason="Internal server error. Traceback included in previous frame.", + ) + raise + finally: + logger.info("Rank 0: Client session ended. Sending idle signal (2) to workers.") + signal_tensor.fill_(2) + dist.broadcast(signal_tensor, src=0, group=self._signal_group) + + +def init_mesh() -> DeviceMesh: + local_rank = int(os.environ.get("LOCAL_RANK", "0")) + + torch.cuda.set_device(local_rank) + _ = torch.cuda.is_available() + _ = torch.cuda.device_count() + + dist.init_process_group("nccl") + rank = dist.get_rank() + world_size = dist.get_world_size() + if world_size not in (1, 2): + raise ValueError( + f"This DreamZero inference path only supports 1 or 2 GPUs, got world_size={world_size}. " + "The action head parallelization code explicitly supports ip_size 1 or 2 only. " + "Please launch with --nproc_per_node=2 (or 1)." + ) + print(f"Rank {rank}/{world_size} (PID: {os.getpid()}) setting device to local_rank={local_rank}") + + torch.cuda.set_device(local_rank) + device = torch.device(f"cuda:{local_rank}") + + mesh = init_device_mesh( + device_type="cuda", + mesh_shape=(world_size,), + mesh_dim_names=("ip",), + ) + print(f"Rank {rank}/{world_size} (PID: {os.getpid()}) using device {device}") + + return mesh + +def _health_check(connection: _server.ServerConnection, request: _server.Request) -> _server.Response | None: + if request.path == "/healthz": + return connection.respond(http.HTTPStatus.OK, "OK\n") + return None + + +def _create_wrapper_policy( + embodiment_tag: str, + groot_policy: GrootSimPolicy, + signal_group: dist.ProcessGroup, + output_dir: str | None, + video_save_mode: str, + model_path: str | None = None, +) -> DistributedRoboarenaPolicyBase: + if embodiment_tag == 'oxe_droid': + return ARDroidRoboarenaPolicy( + groot_policy=groot_policy, + signal_group=signal_group, + output_dir=output_dir, + video_save_mode=video_save_mode, + ) + if embodiment_tag == 'agibot': + if model_path is None: + raise ValueError('model_path is required for AgiBot wrapper') + return AgiBotRoboarenaPolicy( + groot_policy=groot_policy, + signal_group=signal_group, + model_path=model_path, + output_dir=output_dir, + video_save_mode=video_save_mode, + ) + if embodiment_tag == 'g2': + return G2RoboarenaPolicy( + groot_policy=groot_policy, + signal_group=signal_group, + output_dir=output_dir, + video_save_mode=video_save_mode, + ) + raise ValueError(f'Unsupported embodiment_tag: {embodiment_tag}') + + +def _create_server_config(embodiment_tag: str) -> PolicyServerConfig: + if embodiment_tag == 'oxe_droid': + return PolicyServerConfig( + image_resolution=(180, 320), + needs_wrist_camera=True, + n_external_cameras=2, + needs_stereo_camera=False, + needs_session_id=True, + action_space='joint_position', + ) + if embodiment_tag == 'agibot': + return PolicyServerConfig( + image_resolution=(640, 480), + needs_wrist_camera=False, + n_external_cameras=3, + needs_stereo_camera=False, + needs_session_id=True, + action_space='agibot_flattened', + ) + if embodiment_tag == 'g2': + return PolicyServerConfig( + image_resolution=(176, 320), + needs_wrist_camera=False, + n_external_cameras=3, + needs_stereo_camera=False, + needs_session_id=True, + action_space='joint_position', + ) + raise ValueError(f'Unsupported embodiment_tag: {embodiment_tag}') + + +def _build_path_overrides(args: Args) -> tuple[list[str], list[str]]: + model_config_overrides: list[str] = [] + train_config_overrides: list[str] = [] + + if args.wan_ckpt_dir: + wan_ckpt_dir = os.path.abspath(args.wan_ckpt_dir) + required_files = [ + os.path.join(wan_ckpt_dir, "models_t5_umt5-xxl-enc-bf16.pth"), + os.path.join(wan_ckpt_dir, "models_clip_open-clip-xlm-roberta-large-vit-huge-14.pth"), + os.path.join(wan_ckpt_dir, "Wan2.1_VAE.pth"), + ] + missing = [path for path in required_files if not os.path.exists(path)] + if missing: + raise FileNotFoundError( + "Missing Wan checkpoint component(s): " + ", ".join(missing) + ) + model_config_overrides.extend( + [ + f"action_head_cfg.config.diffusion_model_cfg.diffusion_model_pretrained_path={wan_ckpt_dir}", + f"action_head_cfg.config.text_encoder_cfg.text_encoder_pretrained_path={wan_ckpt_dir}/models_t5_umt5-xxl-enc-bf16.pth", + f"action_head_cfg.config.image_encoder_cfg.image_encoder_pretrained_path={wan_ckpt_dir}/models_clip_open-clip-xlm-roberta-large-vit-huge-14.pth", + f"action_head_cfg.config.vae_cfg.vae_pretrained_path={wan_ckpt_dir}/Wan2.1_VAE.pth", + ] + ) + + if args.tokenizer_path: + tokenizer_path = os.path.abspath(args.tokenizer_path) + if not os.path.exists(tokenizer_path): + raise FileNotFoundError(f"Tokenizer path does not exist: {tokenizer_path}") + if args.embodiment_tag.lower() == "agibot": + train_config_overrides.append( + f"transforms.agibot.transforms.10.tokenizer_path={tokenizer_path}" + ) + elif args.embodiment_tag.lower() == "oxe_droid": + train_config_overrides.append( + f"transforms.oxe_droid.transforms.10.tokenizer_path={tokenizer_path}" + ) + elif args.embodiment_tag.lower() == "g2": + train_config_overrides.append( + f"transforms.g2.transforms.10.tokenizer_path={tokenizer_path}" + ) + + return model_config_overrides, train_config_overrides + + + +def _require_file(path: Path, description: str) -> None: + if not path.is_file(): + raise FileNotFoundError(f"Missing {description}: {path}") + + +def _require_dir(path: Path, description: str) -> None: + if not path.is_dir(): + raise FileNotFoundError(f"Missing {description}: {path}") + + +def _parse_episode_indices(spec: str, total_episodes: int) -> list[int]: + values: list[int] = [] + for token in spec.split(','): + token = token.strip() + if not token: + continue + if '-' in token: + left, right = token.split('-', 1) + start = int(left) + end = int(right) + if end < start: + raise ValueError(f"Invalid episode range: {token}") + values.extend(range(start, end + 1)) + else: + values.append(int(token)) + if not values: + raise ValueError("--episode-indices cannot be empty") + unique: list[int] = [] + seen: set[int] = set() + for value in values: + if value < 0 or value >= total_episodes: + raise IndexError( + f"Episode {value} is outside [0, {total_episodes - 1}]" + ) + if value not in seen: + unique.append(value) + seen.add(value) + return unique + + +def _episode_file_from_template( + root: Path, + template: str, + episode_index: int, + chunks_size: int, + video_key: str | None = None, +) -> Path: + values = { + "episode_chunk": episode_index // chunks_size, + "episode_index": episode_index, + } + if video_key is not None: + values["video_key"] = video_key + path = root / template.format(**values) + if path.is_file(): + return path + + # Fallback for datasets whose chunk size metadata was rewritten without + # moving the actual files. + filename = f"episode_{episode_index:06d}" + path.suffix + if video_key is None: + candidates = sorted(root.glob(f"data/chunk-*/{filename}")) + else: + candidates = sorted(root.glob(f"videos/chunk-*/{video_key}/{filename}")) + if len(candidates) != 1: + raise FileNotFoundError( + f"Could not resolve episode file for episode={episode_index}, " + f"video_key={video_key!r}; expected {path}, candidates={candidates}" + ) + return candidates[0] + + +def _column_to_numpy(table: object, name: str, dtype: np.dtype) -> np.ndarray: + if name not in table.column_names: + raise KeyError(f"Parquet column missing: {name}") + values = table[name].combine_chunks().to_pylist() + return np.asarray(values, dtype=dtype) + + +def _unwrap_text(value: object) -> str: + while isinstance(value, (list, tuple)) and len(value) == 1: + value = value[0] + if value is None: + return "" + return str(value) + + +def _decode_video_range_bgr( + video_path: Path, + first_index: int, + last_index: int, +) -> dict[int, np.ndarray]: + if first_index < 0 or last_index < first_index: + raise ValueError( + f"Invalid decode range [{first_index}, {last_index}] for {video_path}" + ) + capture = cv2.VideoCapture(str(video_path)) + if not capture.isOpened(): + raise RuntimeError(f"Failed to open video: {video_path}") + try: + capture.set(cv2.CAP_PROP_POS_FRAMES, float(first_index)) + decoded: dict[int, np.ndarray] = {} + for index in range(first_index, last_index + 1): + ok, frame_bgr = capture.read() + if not ok or frame_bgr is None: + raise RuntimeError( + f"Failed to decode frame {index} from {video_path}" + ) + decoded[index] = np.ascontiguousarray(frame_bgr) + return decoded + finally: + capture.release() + + +def _grid_rgb( + top_bgr: np.ndarray, + left_bgr: np.ndarray, + right_bgr: np.ndarray, +) -> np.ndarray: + top = cv2.cvtColor(top_bgr, cv2.COLOR_BGR2RGB) + left = cv2.cvtColor(left_bgr, cv2.COLOR_BGR2RGB) + right = cv2.cvtColor(right_bgr, cv2.COLOR_BGR2RGB) + if top.shape != left.shape or top.shape != right.shape: + raise ValueError( + f"G2 test views must share one shape, got " + f"top={top.shape}, left={left.shape}, right={right.shape}" + ) + black = np.zeros_like(top) + return np.concatenate( + [ + np.concatenate([top, left], axis=1), + np.concatenate([right, black], axis=1), + ], + axis=0, + ) + + +def _save_rgb_video(path: Path, frames: list[np.ndarray], fps: int = 30) -> None: + if not frames: + raise ValueError(f"No frames to save: {path}") + path.parent.mkdir(parents=True, exist_ok=True) + imageio.mimsave( + path, + frames, + fps=fps, + codec="libx264", + macro_block_size=None, + ) + + +def _read_mp4_rgb(path: Path) -> list[np.ndarray]: + capture = cv2.VideoCapture(str(path)) + if not capture.isOpened(): + raise RuntimeError(f"Failed to open generated video: {path}") + frames: list[np.ndarray] = [] + try: + while True: + ok, frame_bgr = capture.read() + if not ok: + break + frames.append(cv2.cvtColor(frame_bgr, cv2.COLOR_BGR2RGB)) + finally: + capture.release() + if not frames: + raise RuntimeError(f"Generated video has no frames: {path}") + return frames + + +def _add_label_rgb(frame_rgb: np.ndarray, label: str) -> np.ndarray: + frame_bgr = cv2.cvtColor(frame_rgb, cv2.COLOR_RGB2BGR) + cv2.rectangle(frame_bgr, (0, 0), (360, 34), (0, 0, 0), -1) + cv2.putText( + frame_bgr, + label, + (8, 24), + cv2.FONT_HERSHEY_SIMPLEX, + 0.65, + (255, 255, 255), + 2, + cv2.LINE_AA, + ) + return cv2.cvtColor(frame_bgr, cv2.COLOR_BGR2RGB) + + + +def _checkpoint_eval_offsets( + model_path: Path, + embodiment_tag: str, +) -> list[int]: + conf_path = model_path / "experiment_cfg" / "conf.yaml" + _require_file(conf_path, "checkpoint experiment config") + cfg = OmegaConf.to_container(OmegaConf.load(conf_path), resolve=True) + config_key = f"modality_config_{embodiment_tag}" + try: + offsets = list( + cfg[config_key]["video"]["eval_delta_indices"] + ) + except (KeyError, TypeError) as exc: + raise KeyError( + f"Could not read {config_key}.video.eval_delta_indices " + f"from {conf_path}" + ) from exc + if not offsets: + offsets = [0] + return [int(value) for value in offsets] + + +def _checkpoint_agibot_video_resolutions( + model_path: Path, +) -> dict[str, tuple[int, int]]: + """Return required raw input resolutions as camera -> (height, width).""" + metadata_path = model_path / "experiment_cfg" / "metadata.json" + _require_file(metadata_path, "AgiBot checkpoint metadata") + with metadata_path.open("r", encoding="utf-8") as stream: + metadata = json.load(stream) + try: + video_meta = metadata["agibot"]["modalities"]["video"] + except KeyError as exc: + raise KeyError( + f"{metadata_path} lacks agibot.modalities.video" + ) from exc + + result: dict[str, tuple[int, int]] = {} + for name in ("top_head", "hand_left", "hand_right"): + try: + width, height = [ + int(value) + for value in video_meta[name]["resolution"] + ] + except (KeyError, TypeError, ValueError) as exc: + raise ValueError( + f"Invalid AgiBot raw resolution metadata for {name!r}" + ) from exc + if width <= 0 or height <= 0: + raise ValueError( + f"Non-positive AgiBot raw resolution for {name}: " + f"{width}x{height}" + ) + result[name] = (height, width) + return result + + +def _validate_test_dataset(root: Path) -> dict: + _require_dir(root, "G2 test split") + info_path = root / "meta" / "info.json" + modality_path = root / "meta" / "modality.json" + embodiment_path = root / "meta" / "embodiment.json" + for path, description in ( + (info_path, "test info.json"), + (modality_path, "test modality.json"), + (embodiment_path, "test embodiment.json"), + ): + _require_file(path, description) + + with info_path.open("r", encoding="utf-8") as stream: + info = json.load(stream) + with embodiment_path.open("r", encoding="utf-8") as stream: + embodiment = json.load(stream) + + if embodiment.get("embodiment_tag") != "g2": + raise ValueError( + f"Expected test embodiment_tag='g2', got {embodiment}" + ) + for key in ("observation.state", "action"): + shape = info.get("features", {}).get(key, {}).get("shape") + if shape != [16]: + raise ValueError(f"{key} must be 16-D, got {shape}") + expected_videos = { + "observation.images.top_head", + "observation.images.hand_left", + "observation.images.hand_right", + } + actual_videos = { + key + for key, value in info.get("features", {}).items() + if value.get("dtype") == "video" + } + if actual_videos != expected_videos: + raise ValueError( + f"Unexpected G2 test camera keys: {sorted(actual_videos)}" + ) + return info + + + +def _run_one_test_sample( + wrapper: AgiBotRoboarenaPolicy, + args: Args, + info: dict, + episode_index: int, + eval_offsets: list[int], + raw_resolutions: dict[str, tuple[int, int]], +) -> dict: + root = Path(args.test_data_root).resolve() + chunks_size = int(info.get("chunks_size", 1000)) + parquet_path = _episode_file_from_template( + root, + info.get( + "data_path", + "data/chunk-{episode_chunk:03d}/episode_{episode_index:06d}.parquet", + ), + episode_index, + chunks_size, + ) + table = pq.read_table(parquet_path) + state = _column_to_numpy(table, "observation.state", np.float32) + num_rows = int(len(state)) + if state.shape != (num_rows, 16): + raise ValueError( + f"Unexpected G2 test state shape for episode {episode_index}: " + f"{state.shape}" + ) + + row_index = args.frame_index + if row_index < 0: + row_index = num_rows // 2 + minimum = -min(eval_offsets) + maximum_future = max(args.future_frames, 9) + if row_index < minimum: + raise IndexError( + f"frame_index={row_index} needs at least {minimum} previous " + f"frames for offsets {eval_offsets}" + ) + if row_index + maximum_future > num_rows: + row_index = num_rows - maximum_future + if row_index < minimum: + raise IndexError( + f"Episode {episode_index} is too short ({num_rows} rows) for " + f"offsets={eval_offsets}, future_frames={args.future_frames}" + ) + + language_key = "annotation.language.action_text" + if args.prompt_override is not None: + prompt = args.prompt_override + elif language_key in table.column_names: + prompt = _unwrap_text(table[language_key][row_index].as_py()) + else: + prompt = "" + + camera_features = { + "top_head": "observation.images.top_head", + "hand_left": "observation.images.hand_left", + "hand_right": "observation.images.hand_right", + } + context_indices = [row_index + offset for offset in eval_offsets] + first_decode = min(context_indices) + last_decode = min( + num_rows - 1, + row_index + args.future_frames - 1, + ) + decoded: dict[str, dict[int, np.ndarray]] = {} + video_paths: dict[str, str] = {} + video_template = info.get( + "video_path", + "videos/chunk-{episode_chunk:03d}/{video_key}/" + "episode_{episode_index:06d}.mp4", + ) + for short_name, feature_key in camera_features.items(): + path = _episode_file_from_template( + root, + video_template, + episode_index, + chunks_size, + video_key=feature_key, + ) + video_paths[short_name] = str(path) + decoded[short_name] = _decode_video_range_bgr( + path, + first_decode, + last_decode, + ) + + sample_name = f"episode_{episode_index:06d}_frame_{row_index:06d}" + sample_dir = Path(args.output_dir).resolve() / sample_name + sample_dir.mkdir(parents=True, exist_ok=True) + + # Human-readable source context and source ground truth remain at the + # original G2 test-set resolution. + context_grids = [ + _grid_rgb( + decoded["top_head"][index], + decoded["hand_left"][index], + decoded["hand_right"][index], + ) + for index in context_indices + ] + _save_rgb_video( + sample_dir / f"g2_source_context_f{len(context_grids)}.mp4", + context_grids, + fps=30, + ) + imageio.imwrite( + sample_dir / "g2_source_context_last.png", + context_grids[-1], + ) + + gt_indices = list(range(row_index, last_decode + 1)) + gt_grids = [ + _grid_rgb( + decoded["top_head"][index], + decoded["hand_left"][index], + decoded["hand_right"][index], + ) + for index in gt_indices + ] + _save_rgb_video( + sample_dir / f"g2_ground_truth_future_f{len(gt_grids)}.mp4", + gt_grids, + fps=30, + ) + imageio.imwrite( + sample_dir / "g2_ground_truth_first.png", + gt_grids[0], + ) + imageio.imwrite( + sample_dir / "g2_ground_truth_last.png", + gt_grids[-1], + ) + + # Resize the same G2 source images to the exact raw resolution contract + # stored in the AgiBot checkpoint metadata. + context_bgr: dict[str, np.ndarray] = {} + for short_name in camera_features: + target_h, target_w = raw_resolutions[short_name] + frames = [ + cv2.resize( + decoded[short_name][index], + (target_w, target_h), + interpolation=cv2.INTER_LINEAR, + ) + for index in context_indices + ] + context_bgr[short_name] = np.stack(frames, axis=0) + + current_state = state[row_index] + obs = { + "observation/top_head": context_bgr["top_head"], + "observation/hand_left": context_bgr["hand_left"], + "observation/hand_right": context_bgr["hand_right"], + "observation/state": current_state, + "prompt": prompt, + "session_id": ( + args.session_id_override + if args.session_id_override is not None + else f"offline-{sample_name}-{uuid.uuid4()}" + ), + } + + wrapper._output_dir = str(sample_dir) + wrapper._msg_index = 0 + logger.info( + "[AGIBOT-ON-G2 TESTSET] episode=%s row=%s context=%s " + "prompt=%r raw_resolutions=%s parquet=%s", + episode_index, + row_index, + context_indices, + prompt, + { + key: [value[1], value[0]] + for key, value in raw_resolutions.items() + }, + parquet_path, + ) + ignored_actions = wrapper.infer(obs) + np.save( + sample_dir / "ignored_hold_actions.npy", + ignored_actions, + ) + + predicted_candidates = sorted( + (sample_dir / "diagnostics").glob( + "pred_chunk_request_000001_f*.mp4" + ) + ) + comparison_path: str | None = None + predicted_path: str | None = None + if predicted_candidates: + predicted_file = predicted_candidates[-1] + predicted_path = str(predicted_file) + predicted_frames = _read_mp4_rgb(predicted_file) + compare_count = min(len(predicted_frames), len(gt_grids)) + comparison_frames: list[np.ndarray] = [] + for index in range(compare_count): + pred = predicted_frames[index] + gt = gt_grids[index] + if pred.shape[:2] != gt.shape[:2]: + pred = cv2.resize( + pred, + (gt.shape[1], gt.shape[0]), + interpolation=cv2.INTER_LINEAR, + ) + comparison_frames.append( + np.concatenate( + [ + _add_label_rgb(pred, "PREDICTED"), + _add_label_rgb( + gt, + "G2 TEST GROUND TRUTH", + ), + ], + axis=1, + ) + ) + comparison_file = ( + sample_dir + / f"agibot_predicted_vs_g2_gt_f{compare_count}.mp4" + ) + _save_rgb_video( + comparison_file, + comparison_frames, + fps=30, + ) + imageio.imwrite( + sample_dir / "agibot_predicted_vs_g2_gt_first.png", + comparison_frames[0], + ) + imageio.imwrite( + sample_dir / "agibot_predicted_vs_g2_gt_last.png", + comparison_frames[-1], + ) + comparison_path = str(comparison_file) + + summary = { + "checkpoint": str(Path(args.model_path).resolve()), + "checkpoint_embodiment": "agibot", + "visual_source_embodiment": "g2", + "test_data_root": str(root), + "episode_index": episode_index, + "row_index": row_index, + "parquet_path": str(parquet_path), + "video_paths": video_paths, + "context_offsets": eval_offsets, + "context_indices": context_indices, + "checkpoint_raw_resolutions": { + key: [value[1], value[0]] + for key, value in raw_resolutions.items() + }, + "prompt": prompt, + "g2_state_used_for_model_conditioning": False, + "predicted_video": predicted_path, + "comparison_video": comparison_path, + "action_note": ( + "Native AgiBot actions were intentionally discarded; " + "this is a video-only diagnostic." + ), + } + with (sample_dir / "summary.json").open( + "w", + encoding="utf-8", + ) as stream: + json.dump(summary, stream, ensure_ascii=False, indent=2) + return summary + + + +def main(args: Args) -> None: + if args.embodiment_tag.lower() != "agibot": + raise ValueError( + "This evaluator runs AgiBot checkpoints on G2 test-set images; " + "use --embodiment-tag agibot" + ) + if args.future_frames <= 0: + raise ValueError("--future-frames must be positive") + + os.environ["ENABLE_DIT_CACHE"] = ( + "true" if args.enable_dit_cache else "false" + ) + if args.num_dit_steps is not None: + os.environ["NUM_DIT_STEPS"] = str(args.num_dit_steps) + elif args.enable_dit_cache: + os.environ.setdefault("NUM_DIT_STEPS", "8") + os.environ.setdefault("ATTENTION_BACKEND", "FA2") + os.environ.setdefault("TOKENIZERS_PARALLELISM", "false") + torch._dynamo.config.recompile_limit = 800 + + model_path = Path(args.model_path).resolve() + _require_dir(model_path, "AgiBot checkpoint") + metadata_path = model_path / "experiment_cfg" / "metadata.json" + conf_path = model_path / "experiment_cfg" / "conf.yaml" + _require_file(metadata_path, "AgiBot checkpoint metadata") + _require_file(conf_path, "AgiBot checkpoint config") + with metadata_path.open("r", encoding="utf-8") as stream: + checkpoint_metadata = json.load(stream) + if "agibot" not in checkpoint_metadata: + raise KeyError( + f"Checkpoint does not contain AgiBot metadata: " + f"{metadata_path}; keys={list(checkpoint_metadata)}" + ) + + test_info = _validate_test_dataset( + Path(args.test_data_root).resolve() + ) + episode_indices = _parse_episode_indices( + args.episode_indices, + int(test_info["total_episodes"]), + ) + eval_offsets = _checkpoint_eval_offsets( + model_path, + "agibot", + ) + raw_resolutions = _checkpoint_agibot_video_resolutions( + model_path + ) + model_config_overrides, train_config_overrides = ( + _build_path_overrides(args) + ) + + Path(args.output_dir).mkdir(parents=True, exist_ok=True) + logger.info( + "[PREFLIGHT OK] checkpoint=%s checkpoint_embodiment=agibot " + "visual_source=g2 test=%s episodes=%s eval_offsets=%s " + "raw_resolutions=%s output=%s", + model_path, + Path(args.test_data_root).resolve(), + episode_indices, + eval_offsets, + { + key: [value[1], value[0]] + for key, value in raw_resolutions.items() + }, + Path(args.output_dir).resolve(), + ) + if args.preflight_only: + logger.info( + "Preflight-only validation completed; model was not loaded." + ) + return + + device_mesh = init_mesh() + rank = dist.get_rank() + torch.manual_seed(args.seed) + np.random.seed(args.seed) + + timeout_delta = datetime.timedelta( + seconds=args.timeout_seconds + ) + signal_group = dist.new_group( + backend="gloo", + timeout=timeout_delta, + ) + logger.info( + "Rank %s initialized signal_group (gloo)", + rank, + ) + + policy = GrootSimPolicy( + embodiment_tag=EmbodimentTag("agibot"), + model_path=str(model_path), + device="cuda" if torch.cuda.is_available() else "cpu", + device_mesh=device_mesh, + model_config_overrides=model_config_overrides, + train_config_overrides=train_config_overrides, + ) + action_head = policy.trained_model.action_head + if args.num_inference_timesteps < 0: + raise ValueError( + "--num-inference-timesteps must be non-negative" + ) + if args.num_inference_timesteps > 0: + action_head.num_inference_steps = int( + args.num_inference_timesteps + ) + action_head.num_inference_timesteps = int( + args.num_inference_timesteps + ) + if hasattr(action_head, "config"): + action_head.config.num_inference_timesteps = int( + args.num_inference_timesteps + ) + logger.info( + "[CONFIG CHECK] rank=%s diffusion_steps=%s " + "frame_per_block=%s ENABLE_DIT_CACHE=%s", + rank, + getattr(action_head, "num_inference_steps", None), + getattr(action_head, "num_frame_per_block", None), + os.getenv("ENABLE_DIT_CACHE"), + ) + + if rank == 0: + wrapper = AgiBotRoboarenaPolicy( + groot_policy=policy, + signal_group=signal_group, + model_path=str(model_path), + output_dir=str(Path(args.output_dir).resolve()), + video_save_mode="none", + ) + summaries: list[dict] = [] + try: + for episode_index in episode_indices: + summaries.append( + _run_one_test_sample( + wrapper, + args, + test_info, + episode_index, + eval_offsets, + raw_resolutions, + ) + ) + finally: + wrapper._broadcast_signal_to_workers( + SIGNAL_SHUTDOWN + ) + + report_path = ( + Path(args.output_dir).resolve() + / "testset_report.json" + ) + with report_path.open( + "w", + encoding="utf-8", + ) as stream: + json.dump( + { + "checkpoint": str(model_path), + "checkpoint_embodiment": "agibot", + "visual_source_embodiment": "g2", + "test_data_root": str( + Path(args.test_data_root).resolve() + ), + "episode_indices": episode_indices, + "eval_offsets": eval_offsets, + "checkpoint_raw_resolutions": { + key: [value[1], value[0]] + for key, value in raw_resolutions.items() + }, + "samples": summaries, + }, + stream, + ensure_ascii=False, + indent=2, + ) + logger.info( + "Saved test-set report: %s", + report_path, + ) + dist.barrier() + else: + worker = WebsocketPolicyServer( + policy=policy, + host="127.0.0.1", + port=0, + metadata={}, + output_dir=None, + signal_group=signal_group, + ) + asyncio.run(worker._worker_loop()) + dist.barrier() + + dist.destroy_process_group() + + +def cli() -> None: + logging.basicConfig(level=logging.INFO, force=True) + main(tyro.cli(Args)) + + +if __name__ == "__main__": + cli() diff --git a/eval_g2_checkpoint_on_testset.py b/eval_g2_checkpoint_on_testset.py new file mode 100644 index 00000000..e41a595d --- /dev/null +++ b/eval_g2_checkpoint_on_testset.py @@ -0,0 +1,1279 @@ +"""Offline G2 checkpoint video diagnostic on the held-out G2 test split. + +This module reuses the dataset/video plumbing from +``eval_agibot_checkpoint_on_g2_testset.py`` but constructs a native G2 policy. +Predicted actions are deliberately ignored: the purpose of this evaluator is +to compare world-model video quality across the base, joint-LoRA, and +video-only-LoRA checkpoints without connecting to a robot. +""" + +from __future__ import annotations + +import asyncio +import csv +import dataclasses +import datetime +import json +import logging +import os +import shutil +from pathlib import Path + +import numpy as np +import pyarrow.parquet as pq +import torch +import torch.distributed as dist +import tyro + +import eval_agibot_checkpoint_on_g2_testset as shared +from groot.vla.data.schema import EmbodimentTag +from groot.vla.model.n1_5.sim_policy import GrootSimPolicy + + +@dataclasses.dataclass +class Args(shared.Args): + model_path: str = ( + "/data/wangk/checkpoints/" + "dreamzero_g2_video_only_lora_1k/checkpoint-1000" + ) + output_dir: str = ( + "/data/wangk/dreamzero/g2_testset_video_eval/" + "video_only_checkpoint-1000" + ) + embodiment_tag: str = "g2" + diagnostic_fps: int = 3 + full_episode: bool = False + rollout_stride: int = 24 + action_eval: bool = False + # Windowed teacher-forced diagnostic. Each window is made of four + # consecutive 4-frame observations; four causal forwards are concatenated + # to the 33 decoded frames used for the video comparison. + windowed: bool = False + window_history: int = 4 + window_stride: int = 4 + window_starts: str = "0,4,8,12,16,20,24,28" + rollout_future_frames: int = 33 + rollout_blocks: int = 4 + # Keep the default report compact: one stitched comparison MP4 and two + # action CSV tables. Per-window input/prediction artifacts are opt-in. + save_window_artifacts: bool = False + + +class G2VideoDiagnosticPolicy(shared.G2RoboarenaPolicy): + """Run the G2 video branch while making actions non-observable outputs.""" + + def _convert_action(self, action_dict: dict) -> np.ndarray: + del action_dict + return np.zeros((24, 16), dtype=np.float32) + + +def _checkpoint_g2_video_resolutions( + model_path: Path, +) -> dict[str, tuple[int, int]]: + metadata_path = model_path / "experiment_cfg" / "metadata.json" + shared._require_file(metadata_path, "G2 checkpoint metadata") + with metadata_path.open("r", encoding="utf-8") as stream: + metadata = json.load(stream) + try: + video_meta = metadata["g2"]["modalities"]["video"] + except KeyError as exc: + raise KeyError( + f"{metadata_path} lacks g2.modalities.video" + ) from exc + + result: dict[str, tuple[int, int]] = {} + for name in ("top_head", "hand_left", "hand_right"): + try: + width, height = [ + int(value) + for value in video_meta[name]["resolution"] + ] + except (KeyError, TypeError, ValueError) as exc: + raise ValueError( + f"Invalid G2 raw resolution metadata for {name!r}" + ) from exc + if width <= 0 or height <= 0: + raise ValueError( + f"Non-positive G2 raw resolution for {name}: " + f"{width}x{height}" + ) + result[name] = (height, width) + return result + + +def _save_readable_comparison_views( + summary: dict, + diagnostic_fps: int, +) -> None: + """Save a slow MP4 and a 3x3 sheet for a short predicted chunk.""" + comparison_value = summary.get("comparison_video") + if not comparison_value: + return + comparison_path = Path(comparison_value) + frames = shared._read_mp4_rgb(comparison_path) + if not frames: + return + if diagnostic_fps <= 0: + raise ValueError( + f"diagnostic_fps must be positive, got {diagnostic_fps}" + ) + + g2_comparison_path = comparison_path.with_name( + "g2_predicted_vs_g2_gt_f9.mp4" + ) + shutil.copy2(comparison_path, g2_comparison_path) + for position in ("first", "last"): + legacy_path = comparison_path.with_name( + f"agibot_predicted_vs_g2_gt_{position}.png" + ) + if legacy_path.is_file(): + shutil.copy2( + legacy_path, + comparison_path.with_name( + f"g2_predicted_vs_g2_gt_{position}.png" + ), + ) + + slow_path = comparison_path.with_name( + "g2_predicted_vs_gt_f9_slow_3fps.mp4" + ) + shared._save_rgb_video( + slow_path, + frames, + fps=diagnostic_fps, + ) + + target_width = 640 + resized = [ + shared.cv2.resize( + frame, + ( + target_width, + max( + 1, + round(frame.shape[0] * target_width / frame.shape[1]), + ), + ), + interpolation=shared.cv2.INTER_AREA, + ) + for frame in frames + ] + frame_height, frame_width = resized[0].shape[:2] + black = np.zeros( + (frame_height, frame_width, 3), + dtype=np.uint8, + ) + cells = resized[:9] + [black] * max(0, 9 - len(resized)) + rows = [ + np.concatenate(cells[index:index + 3], axis=1) + for index in range(0, 9, 3) + ] + sheet_path = comparison_path.with_name( + "g2_predicted_vs_gt_f9_contact_sheet.png" + ) + shared.imageio.imwrite( + sheet_path, + np.concatenate(rows, axis=0), + ) + summary["comparison_video"] = str(g2_comparison_path) + summary["slow_comparison_video"] = str(slow_path) + summary["comparison_contact_sheet"] = str(sheet_path) + + +def _parse_window_starts(value: str) -> list[int]: + starts: list[int] = [] + for token in str(value).split(","): + token = token.strip() + if not token: + continue + try: + start = int(token) + except ValueError as exc: + raise ValueError( + f"window_starts must be comma-separated integers, got {value!r}" + ) from exc + if start < 0: + raise ValueError(f"window start must be non-negative, got {start}") + starts.append(start) + if not starts: + raise ValueError("window_starts must contain at least one start index") + return starts + + +def _decode_window_latents( + wrapper: shared.G2RoboarenaPolicy, + latent_chunks: list[torch.Tensor], + expected_frames: int, +) -> tuple[list[np.ndarray], list[int]]: + """Decode the concatenated causal chunks and return RGB frames plus shape.""" + if not latent_chunks: + raise RuntimeError("The window produced no video latent chunks") + + latent = torch.cat(latent_chunks, dim=2) + action_head = wrapper._policy.trained_model.action_head + device = getattr(action_head, "_device", None) + if device is None: + device = next(action_head.parameters()).device + latent = latent.to(device=device, dtype=torch.bfloat16) + with torch.no_grad(): + decoded = action_head.vae.decode( + latent, + tiled=action_head.tiled, + tile_size=(action_head.tile_size_height, action_head.tile_size_width), + tile_stride=(action_head.tile_stride_height, action_head.tile_stride_width), + ) + decoded_shape = [int(value) for value in decoded.shape] + frames = decoded[0].permute(1, 2, 3, 0) + frames = ( + # Wan VAE returns RGB in [-1, 1]. Convert to uint8 exactly as the + # online G2 server does; omitting the 127.5 scale makes every + # predicted frame appear black (values collapse to 0 or 1). + (frames.float() + 1.0) + * 127.5 + ).clamp(0.0, 255.0) + frames = ( + frames + .cpu() + .numpy() + .astype(np.uint8) + ) + if frames.shape[0] < expected_frames: + raise RuntimeError( + "Causal rollout decoded fewer frames than requested: " + f"decoded={frames.shape[0]} requested={expected_frames} " + f"latent_shape={tuple(latent.shape)}" + ) + return [np.ascontiguousarray(frame) for frame in frames[:expected_frames]], decoded_shape + + +def _save_window_action_plot( + path: Path, + predicted: np.ndarray, + ground_truth: np.ndarray, +) -> None: + """Save a compact action plot without making plotting a hard dependency.""" + try: + import matplotlib + + matplotlib.use("Agg") + import matplotlib.pyplot as plt + except Exception as exc: # pragma: no cover - diagnostic-only fallback + logging.warning("Could not create action plot %s: %s", path, exc) + return + + # Plot the first causal block; all four blocks are retained in the npy + # files and the JSON metrics below. + pred0 = predicted[0] + gt0 = ground_truth[0] + horizon = np.arange(1, pred0.shape[0] + 1) + fig, axes = plt.subplots(2, 1, figsize=(12, 7), sharex=True) + axes[0].plot( + horizon, + np.abs(pred0[:, :7] - gt0[:, :7]).mean(axis=1), + label="arm MAE", + color="tab:blue", + linewidth=2, + ) + axes[0].plot( + horizon, + np.abs(pred0[:, [7, 15]] - gt0[:, [7, 15]]).mean(axis=1), + label="gripper MAE", + color="tab:orange", + linewidth=2, + ) + axes[0].set_ylabel("absolute error") + axes[0].set_title("Window 0 action error (same forward as video prediction)") + axes[0].grid(alpha=0.25) + axes[0].legend() + + axes[1].plot( + horizon, + gt0[:, :7].mean(axis=1), + label="GT arm mean", + color="tab:green", + ) + axes[1].plot( + horizon, + pred0[:, :7].mean(axis=1), + label="PRED arm mean", + color="tab:red", + ) + axes[1].set_xlabel("action horizon step") + axes[1].set_ylabel("joint mean") + axes[1].grid(alpha=0.25) + axes[1].legend() + fig.tight_layout() + fig.savefig(path, dpi=140) + plt.close(fig) + + +def _run_windowed_episode_comparison( + wrapper: shared.G2RoboarenaPolicy, + args: Args, + info: dict, + episode_index: int, + raw_resolutions: dict[str, tuple[int, int]], +) -> dict: + """Run contiguous 4-frame packets and compare a 33-frame causal rollout. + + One window is intentionally teacher-forced at the observation-packet level, + matching deployment: packets ``0:4``, ``4:8``, ``8:12`` and ``12:16`` are + sent through the *same* policy cache. The first packet's latent output has + three frames and the next three have two each; VAE decoding therefore gives + exactly 33 frames. The next requested window resets the cache and starts + at the next four-frame packet (by default 4:8). + """ + if args.window_history != 4: + raise ValueError( + "The causal G2 evaluator currently requires window_history=4; " + f"got {args.window_history}" + ) + if args.rollout_blocks != 4: + raise ValueError( + "The 33-frame causal rollout requires rollout_blocks=4; " + f"got {args.rollout_blocks}" + ) + if args.rollout_future_frames != 33: + raise ValueError( + "This diagnostic is defined for rollout_future_frames=33; " + f"got {args.rollout_future_frames}" + ) + + root = Path(args.test_data_root).resolve() + chunks_size = int(info.get("chunks_size", 1000)) + parquet_path = shared._episode_file_from_template( + root, + info.get( + "data_path", + "data/chunk-{episode_chunk:03d}/episode_{episode_index:06d}.parquet", + ), + episode_index, + chunks_size, + ) + table = pq.read_table(parquet_path) + state = shared._column_to_numpy(table, "observation.state", np.float32) + actions = shared._column_to_numpy(table, "action", np.float32) + num_rows = int(len(state)) + if state.shape != (num_rows, 16) or actions.shape != (num_rows, 16): + raise ValueError( + "Windowed G2 diagnostic expects state/action shapes " + f"({num_rows},16), got state={state.shape} action={actions.shape}" + ) + + starts = _parse_window_starts(args.window_starts) + video_template = info.get( + "video_path", + "videos/chunk-{episode_chunk:03d}/{video_key}/" + "episode_{episode_index:06d}.mp4", + ) + camera_features = { + "top_head": "observation.images.top_head", + "hand_left": "observation.images.hand_left", + "hand_right": "observation.images.hand_right", + } + video_paths = { + short_name: str( + shared._episode_file_from_template( + root, + video_template, + episode_index, + chunks_size, + video_key=feature_key, + ) + ) + for short_name, feature_key in camera_features.items() + } + + output_root = Path(args.output_dir).resolve() + output_root.mkdir(parents=True, exist_ok=True) + full_comparison_frames: list[np.ndarray] = [] + all_predicted_actions: list[np.ndarray] = [] + all_ground_truth_actions: list[np.ndarray] = [] + language_key = "annotation.language.action_text" + window_summaries: list[dict] = [] + for window_start in starts: + first_packet = window_start + last_packet = window_start + args.window_history * args.rollout_blocks - 1 + anchor = window_start + args.window_history - 1 + gt_video_last = anchor + args.rollout_future_frames - 1 + gt_action_last = max( + window_start + (args.rollout_blocks - 1) * args.window_history + + args.window_history - 1 + + 24 - 1, + anchor + 24 - 1, + ) + last_needed = max(last_packet, gt_video_last, gt_action_last) + if first_packet < 0 or last_needed >= num_rows: + raise IndexError( + f"Window start={window_start} needs rows through {last_needed}, " + f"but episode {episode_index} has {num_rows} rows" + ) + + decoded: dict[str, dict[int, np.ndarray]] = {} + for short_name, path_value in video_paths.items(): + decoded[short_name] = shared._decode_video_range_bgr( + Path(path_value), + first_packet, + last_needed, + ) + + window_dir = ( + output_root / f"episode_{episode_index:06d}_window_{window_start:06d}" + ) + if args.save_window_artifacts: + window_dir.mkdir(parents=True, exist_ok=True) + + # Keep each four-frame packet visible so there is no ambiguity about + # what was sent to the model for each causal block. + packet_indices: list[list[int]] = [] + for block_index in range(args.rollout_blocks): + packet_start = window_start + block_index * args.window_history + indices = list( + range(packet_start, packet_start + args.window_history) + ) + packet_indices.append(indices) + if args.save_window_artifacts: + input_grids = [ + shared._grid_rgb( + decoded["top_head"][index], + decoded["hand_left"][index], + decoded["hand_right"][index], + ) + for index in indices + ] + shared._save_rgb_video( + window_dir + / f"input_context_block_{block_index:02d}_f{len(input_grids)}.mp4", + input_grids, + fps=30, + ) + + gt_indices = list( + range(anchor, anchor + args.rollout_future_frames) + ) + gt_grids = [ + shared._grid_rgb( + decoded["top_head"][index], + decoded["hand_left"][index], + decoded["hand_right"][index], + ) + for index in gt_indices + ] + gt_path = window_dir / f"ground_truth_anchor_{anchor:06d}_f33.mp4" + if args.save_window_artifacts: + shared._save_rgb_video(gt_path, gt_grids, fps=30) + + # Force the actual four-frame history for this diagnostic even when a + # checkpoint metadata file advertises eval_delta_indices=[0]. The + # online G2 client sends four frames; this makes the comparison match + # that packet contract while retaining the checkpoint for the model. + wrapper._expected_video_frames = args.window_history + wrapper._output_dir = str(window_dir) if args.save_window_artifacts else None + wrapper._msg_index = 0 + wrapper._reset_state(save_video=False) + + predicted_actions: list[np.ndarray] = [] + latent_chunks: list[torch.Tensor] = [] + block_prompts: list[str] = [] + window_prompt: str | None = None + session_id = f"g2-windowed-episode-{episode_index:06d}" + for block_index, indices in enumerate(packet_indices): + block_anchor = indices[-1] + if args.prompt_override is not None: + raw_prompt = args.prompt_override + elif language_key in table.column_names: + raw_prompt = shared._unwrap_text( + table[language_key][block_anchor].as_py() + ) + else: + raw_prompt = "" + if window_prompt is None: + window_prompt = raw_prompt + # Keep language fixed across the four causal blocks. A language + # change resets WAN's current_start_frame and would invalidate the + # intended 3+2+2+2 latent concatenation. + prompt = window_prompt + block_prompts.append(prompt) + obs = { + "observation/top_head": np.stack( + [decoded["top_head"][index] for index in indices] + ), + "observation/hand_left": np.stack( + [decoded["hand_left"][index] for index in indices] + ), + "observation/hand_right": np.stack( + [decoded["hand_right"][index] for index in indices] + ), + "observation/state": state[block_anchor], + "prompt": prompt, + "session_id": session_id, + } + predicted_actions.append( + np.asarray(wrapper.infer(obs), dtype=np.float32) + ) + + latent_chunks = [ + chunk.detach().cpu().clone() + for chunk in wrapper.video_across_time + ] + predicted_frames, decoded_shape = _decode_window_latents( + wrapper, + latent_chunks, + args.rollout_future_frames, + ) + predicted_path = window_dir / "predicted_rollout_f33.mp4" + if args.save_window_artifacts: + shared._save_rgb_video(predicted_path, predicted_frames, fps=30) + + comparison_frames = [ + np.concatenate( + [ + shared._add_label_rgb(predicted, "PREDICTED"), + shared._add_label_rgb(ground_truth, "G2 GROUND TRUTH"), + ], + axis=1, + ) + for predicted, ground_truth in zip(predicted_frames, gt_grids) + ] + comparison_path = window_dir / "predicted_vs_ground_truth_f33.mp4" + if args.save_window_artifacts: + shared._save_rgb_video(comparison_path, comparison_frames, fps=30) + shared.imageio.imwrite( + window_dir / "predicted_vs_ground_truth_first.png", + comparison_frames[0], + ) + shared.imageio.imwrite( + window_dir / "predicted_vs_ground_truth_last.png", + comparison_frames[-1], + ) + + # Keep one chronological report video instead of exposing four tiny + # request videos per window. The windows are deliberately + # teacher-forced, so each appended segment has a matching GT segment. + full_comparison_frames.extend(comparison_frames) + + predicted_action_array = np.stack(predicted_actions) + gt_action_array = np.stack( + [ + actions[indices[-1] : indices[-1] + 24] + for indices in packet_indices + ] + ) + if predicted_action_array.shape != gt_action_array.shape: + raise RuntimeError( + "Predicted/GT action shape mismatch: " + f"pred={predicted_action_array.shape} gt={gt_action_array.shape}" + ) + if args.save_window_artifacts: + np.save(window_dir / "predicted_actions_4x24x16.npy", predicted_action_array) + np.save(window_dir / "ground_truth_actions_4x24x16.npy", gt_action_array) + action_abs_error = np.abs(predicted_action_array - gt_action_array) + if args.save_window_artifacts: + np.save(window_dir / "action_abs_error_4x24x16.npy", action_abs_error) + _save_window_action_plot( + window_dir / "action_pred_vs_gt.png", + predicted_action_array, + gt_action_array, + ) + + all_predicted_actions.append(predicted_action_array) + all_ground_truth_actions.append(gt_action_array) + + action_metrics = { + "overall_mae": float(action_abs_error.mean()), + "arm_mae": float( + action_abs_error[:, :, [*range(7), *range(8, 15)]].mean() + ), + "gripper_mae": float( + action_abs_error[:, :, [7, 15]].mean() + ), + "block_mae": action_abs_error.mean(axis=(1, 2)).tolist(), + "horizon_mae": action_abs_error.mean(axis=(0, 2)).tolist(), + } + summary = { + "checkpoint": str(Path(args.model_path).resolve()), + "test_data_root": str(root), + "episode_index": episode_index, + "window_start": window_start, + "packet_indices": packet_indices, + "anchor_indices": [indices[-1] for indices in packet_indices], + "gt_video_indices": gt_indices, + "prompt_per_block": block_prompts, + "checkpoint_eval_offsets": _checkpoint_eval_offsets( + Path(args.model_path).resolve(), "g2" + ), + "forced_model_context_frames": args.window_history, + "rollout_blocks": args.rollout_blocks, + "rollout_future_frames": args.rollout_future_frames, + "latent_chunks": [list(chunk.shape) for chunk in latent_chunks], + "decoded_shape_bcthw": decoded_shape, + "action_metrics": action_metrics, + "same_forward_video_action": True, + "teacher_forced_observation_packets": True, + } + if args.save_window_artifacts: + summary.update( + { + "predicted_video": str(predicted_path), + "ground_truth_video": str(gt_path), + "comparison_video": str(comparison_path), + "predicted_actions": str( + window_dir / "predicted_actions_4x24x16.npy" + ), + "ground_truth_actions": str( + window_dir / "ground_truth_actions_4x24x16.npy" + ), + } + ) + with (window_dir / "summary.json").open( + "w", encoding="utf-8" + ) as stream: + json.dump(summary, stream, ensure_ascii=False, indent=2) + window_summaries.append(summary) + + # Reset the causal cache and frame buffer before the next independent + # window, without writing the wrapper's generic concatenated MP4. The + # reset signal must reach rank 1 as well; otherwise the next window + # mixes rank-0 and rank-1 KV caches and can stall or abort. + wrapper._broadcast_signal_to_workers(shared.SIGNAL_RESET_CACHE) + wrapper._reset_state(save_video=False) + + if not all_predicted_actions: + raise RuntimeError("Windowed rollout produced no action predictions") + + predicted_sequence = np.stack(all_predicted_actions, axis=0) + ground_truth_sequence = np.stack(all_ground_truth_actions, axis=0) + sequence_error = np.abs(predicted_sequence - ground_truth_sequence) + action_names = ( + [f"left_joint_{index}" for index in range(7)] + + ["left_gripper"] + + [f"right_joint_{index}" for index in range(7)] + + ["right_gripper"] + ) + + # One detailed, machine-readable table for the complete selected rollout. + detail_csv = output_root / ( + f"episode_{episode_index:06d}_action_pred_vs_gt.csv" + ) + detail_fields = [ + "window_start", + "block_index", + "anchor_frame", + "horizon_step", + ] + [f"gt_{name}" for name in action_names] + [ + f"pred_{name}" for name in action_names + ] + with detail_csv.open("w", newline="", encoding="utf-8") as stream: + writer = csv.DictWriter(stream, fieldnames=detail_fields) + writer.writeheader() + for window_index, window_start in enumerate(starts): + for block_index in range(args.rollout_blocks): + anchor_index = ( + window_start + block_index * args.window_history + + args.window_history - 1 + ) + for horizon_index in range(24): + gt_row = ground_truth_sequence[ + window_index, block_index, horizon_index + ] + pred_row = predicted_sequence[ + window_index, block_index, horizon_index + ] + row = { + "window_start": window_start, + "block_index": block_index, + "anchor_frame": anchor_index, + "horizon_step": horizon_index + 1, + } + row.update( + { + f"gt_{name}": float(value) + for name, value in zip(action_names, gt_row) + } + ) + row.update( + { + f"pred_{name}": float(value) + for name, value in zip(action_names, pred_row) + } + ) + writer.writerow(row) + + def _metric_row(label: str, error: np.ndarray) -> dict[str, object]: + return { + "scope": label, + "overall_mae": float(error.mean()), + "arm_mae": float( + error[..., [*range(7), *range(8, 15)]].mean() + ), + "gripper_mae": float(error[..., [7, 15]].mean()), + "left_arm_mae": float(error[..., :7].mean()), + "right_arm_mae": float(error[..., 8:15].mean()), + "num_action_rows": int(np.prod(error.shape[:-1])), + } + + metrics_csv = output_root / ( + f"episode_{episode_index:06d}_action_metrics.csv" + ) + metric_fields = [ + "scope", + "window_start", + "block_index", + "anchor_frame", + "overall_mae", + "arm_mae", + "gripper_mae", + "left_arm_mae", + "right_arm_mae", + "num_action_rows", + ] + metric_rows: list[dict[str, object]] = [] + for window_index, window_summary in enumerate(window_summaries): + window_error = sequence_error[window_index] + for block_index in range(args.rollout_blocks): + block_error = window_error[block_index] + metric = _metric_row( + f"window_{window_summary['window_start']:06d}_block_{block_index:02d}", + block_error[None, ...], + ) + metric.update( + { + "window_start": window_summary["window_start"], + "block_index": block_index, + "anchor_frame": window_summary["anchor_indices"][block_index], + } + ) + metric_rows.append(metric) + overall_metric = _metric_row("overall", sequence_error) + overall_metric.update( + {"window_start": "", "block_index": "", "anchor_frame": ""} + ) + metric_rows.append(overall_metric) + with metrics_csv.open("w", newline="", encoding="utf-8") as stream: + writer = csv.DictWriter(stream, fieldnames=metric_fields) + writer.writeheader() + writer.writerows(metric_rows) + + # Compact visual summary: means for each arm plus the absolute error. It + # is intentionally one figure for the whole selected rollout, not one + # figure per 24-step request. + action_plot = output_root / ( + f"episode_{episode_index:06d}_action_pred_vs_gt.png" + ) + try: + import matplotlib + + matplotlib.use("Agg") + import matplotlib.pyplot as plt + + flat_pred = predicted_sequence.reshape(-1, 16) + flat_gt = ground_truth_sequence.reshape(-1, 16) + flat_error = sequence_error.reshape(-1, 16) + x = np.arange(1, len(flat_pred) + 1) + fig, axes = plt.subplots(3, 1, figsize=(16, 10), sharex=True) + axes[0].plot(x, flat_gt[:, :7].mean(axis=1), label="GT left arm", color="tab:blue") + axes[0].plot(x, flat_pred[:, :7].mean(axis=1), label="PRED left arm", color="tab:orange") + axes[0].set_ylabel("left arm mean") + axes[0].legend(loc="upper right") + axes[1].plot(x, flat_gt[:, 8:15].mean(axis=1), label="GT right arm", color="tab:green") + axes[1].plot(x, flat_pred[:, 8:15].mean(axis=1), label="PRED right arm", color="tab:red") + axes[1].set_ylabel("right arm mean") + axes[1].legend(loc="upper right") + axes[2].plot(x, flat_error[:, [*range(7), *range(8, 15)]].mean(axis=1), label="arm MAE") + axes[2].plot(x, flat_error[:, [7, 15]].mean(axis=1), label="gripper MAE") + axes[2].set_ylabel("absolute error") + axes[2].set_xlabel("concatenated action row (window/block/horizon)") + axes[2].legend(loc="upper right") + for axis in axes: + axis.grid(alpha=0.25) + fig.suptitle( + f"G2 checkpoint action: predicted vs ground truth, episode {episode_index}" + ) + fig.tight_layout() + fig.savefig(action_plot, dpi=140) + plt.close(fig) + except Exception as exc: # pragma: no cover - diagnostic-only fallback + logging.warning("Could not create sequence action plot %s: %s", action_plot, exc) + action_plot = None + + full_comparison_path = output_root / ( + f"episode_{episode_index:06d}_full_predicted_vs_ground_truth_" + f"f{len(full_comparison_frames)}.mp4" + ) + shared._save_rgb_video(full_comparison_path, full_comparison_frames, fps=30) + + report = { + "checkpoint": str(Path(args.model_path).resolve()), + "test_data_root": str(root), + "episode_index": episode_index, + "window_starts": starts, + "window_history": args.window_history, + "window_stride": args.window_stride, + "rollout_blocks": args.rollout_blocks, + "rollout_future_frames": args.rollout_future_frames, + "video_paths": video_paths, + "full_comparison_video": str(full_comparison_path), + "action_detail_csv": str(detail_csv), + "action_metrics_csv": str(metrics_csv), + "action_plot": str(action_plot) if action_plot else None, + "num_comparison_frames": len(full_comparison_frames), + "save_window_artifacts": bool(args.save_window_artifacts), + "overall_action_metrics": overall_metric, + "windows": window_summaries, + } + report_path = ( + Path(args.output_dir).resolve() + / f"episode_{episode_index:06d}_windowed_report.json" + ) + with report_path.open("w", encoding="utf-8") as stream: + json.dump(report, stream, ensure_ascii=False, indent=2) + return report + + +def _run_full_episode_comparison( + wrapper: shared.G2RoboarenaPolicy, + args: Args, + test_info: dict, + episode_index: int, + eval_offsets: list[int], + raw_resolutions: dict[str, tuple[int, int]], +) -> dict: + """Roll through one complete held-out episode and keep one final MP4.""" + root = Path(args.test_data_root).resolve() + chunks_size = int(test_info.get("chunks_size", 1000)) + parquet_path = shared._episode_file_from_template( + root, + test_info.get( + "data_path", + "data/chunk-{episode_chunk:03d}/episode_{episode_index:06d}.parquet", + ), + episode_index, + chunks_size, + ) + table = pq.read_table(parquet_path) + num_rows = table.num_rows + ground_truth_actions = shared._column_to_numpy( + table, "action", np.float32 + ) + if ground_truth_actions.shape != (num_rows, 16): + raise ValueError( + f"Expected G2 ground-truth action shape {(num_rows, 16)}, " + f"got {ground_truth_actions.shape}" + ) + first_row = max(0, -min(eval_offsets)) + last_row = max(first_row, num_rows - max(args.future_frames, 9)) + stride = max(1, int(args.rollout_stride)) + + output_root = Path(args.output_dir).resolve() + temp_root = output_root / f".full_rollout_episode_{episode_index:06d}" + if temp_root.exists(): + shutil.rmtree(temp_root) + temp_root.mkdir(parents=True, exist_ok=True) + + frames: list[np.ndarray] = [] + predicted_horizons: list[np.ndarray] = [] + ground_truth_horizons: list[np.ndarray] = [] + rollout_rows: list[int] = [] + chunks = 0 + session_id = f"g2-full-rollout-episode-{episode_index:06d}" + for row_index in range(first_row, last_row + 1, stride): + chunk_args = dataclasses.replace( + args, + frame_index=row_index, + output_dir=str(temp_root), + full_episode=False, + session_id_override=session_id, + ) + summary = shared._run_one_test_sample( + wrapper, + chunk_args, + test_info, + episode_index, + eval_offsets, + raw_resolutions, + ) + action_path = ( + temp_root + / f"episode_{episode_index:06d}_frame_{row_index:06d}" + / "ignored_hold_actions.npy" + ) + predicted = np.load(action_path) + if predicted.shape != (24, 16): + raise ValueError( + f"Expected predicted action shape (24, 16), got {predicted.shape}" + ) + ground_truth_horizon = ground_truth_actions[ + row_index:row_index + 24 + ] + if ground_truth_horizon.shape != (24, 16): + raise ValueError( + "Expected ground-truth action horizon shape (24, 16), " + f"got {ground_truth_horizon.shape} at row {row_index}" + ) + predicted_horizons.append(predicted.astype(np.float32)) + ground_truth_horizons.append( + ground_truth_horizon.astype(np.float32) + ) + rollout_rows.append(row_index) + comparison = summary.get("comparison_video") + if comparison and not args.action_eval: + frames.extend(shared._read_mp4_rgb(Path(comparison))) + chunks += 1 + + if args.action_eval: + import matplotlib + matplotlib.use("Agg") + import matplotlib.pyplot as plt + + pred = np.stack(predicted_horizons) + gt = np.stack(ground_truth_horizons) + rows = np.asarray(rollout_rows) + checkpoint_name = Path(args.model_path).name.replace("checkpoint-", "") + plot_specs = [ + ("left_arm", list(range(7))), + ("right_arm", list(range(8, 15))), + ("grippers", [7, 15]), + ] + plot_paths: list[str] = [] + for label, indices in plot_specs: + motion_score = np.ptp( + gt[:, :, indices], axis=1 + ).mean(axis=1) + window_index = int(np.argmax(motion_score)) + selected_row = int(rows[window_index]) + horizon_x = np.arange(1, 25) + fig, axes = plt.subplots( + len(indices), 1, figsize=(14, max(4, 2.2 * len(indices))), + sharex=True, + ) + axes = np.atleast_1d(axes) + for axis, dim in zip(axes, indices): + axis.plot( + horizon_x, + gt[window_index, :, dim], + color="tab:blue", + linewidth=2.0, + label="GT", + ) + axis.plot( + horizon_x, + pred[window_index, :, dim], + color="tab:orange", + linewidth=1.8, + label="PRED", + ) + axis.set_ylabel(f"dim {dim}") + axis.grid(alpha=0.25) + axes[0].legend() + axes[-1].set_xticks([1, 4, 8, 12, 16, 20, 24]) + axes[-1].set_xlabel("future action step") + fig.suptitle( + f"G2 checkpoint-{checkpoint_name}: {label}, " + f"episode row {selected_row}" + ) + fig.tight_layout() + plot_path = output_root / ( + f"g2_checkpoint{checkpoint_name}_action_horizon_" + f"{label}_pred_vs_gt.png" + ) + fig.savefig(plot_path, dpi=150) + plt.close(fig) + plot_paths.append(str(plot_path)) + + absolute_error = np.abs(pred - gt) + arm_indices = list(range(7)) + list(range(8, 15)) + fig, axis = plt.subplots(figsize=(10, 5)) + horizon_steps = np.arange(1, 25) + axis.plot( + horizon_steps, + absolute_error[:, :, arm_indices].mean(axis=(0, 2)), + linewidth=2.0, + label="arm MAE", + ) + axis.plot( + horizon_steps, + absolute_error[:, :, [7, 15]].mean(axis=(0, 2)), + linewidth=2.0, + label="gripper MAE", + ) + axis.set_xticks([1, 4, 8, 12, 16, 20, 24]) + axis.set_xlabel("predicted horizon step") + axis.set_ylabel("mean absolute error") + axis.set_title( + f"G2 checkpoint-{checkpoint_name}: 24-step action horizon error" + ) + axis.grid(alpha=0.3) + axis.legend() + fig.tight_layout() + error_path = output_root / ( + f"g2_checkpoint{checkpoint_name}_action_horizon_mae.png" + ) + fig.savefig(error_path, dpi=150) + plt.close(fig) + plot_paths.append(str(error_path)) + shutil.rmtree(temp_root) + return { + "episode_index": episode_index, + "rollout_chunks": chunks, + "action_plots": plot_paths, + } + + if not frames: + raise RuntimeError( + f"Full rollout produced no comparison frames for episode {episode_index}" + ) + checkpoint_name = Path(args.model_path).name.replace("checkpoint-", "") + final_path = output_root / ( + f"g2_checkpoint{checkpoint_name}_full_pred_vs_gt.mp4" + ) + shared._save_rgb_video(final_path, frames, fps=args.diagnostic_fps) + shutil.rmtree(temp_root) + return { + "episode_index": episode_index, + "rollout_chunks": chunks, + "comparison_video": str(final_path), + } + + +def main(args: Args) -> None: + if args.embodiment_tag.lower() != "g2": + raise ValueError("This evaluator requires --embodiment-tag g2") + if args.future_frames <= 0: + raise ValueError("--future-frames must be positive") + + os.environ["ENABLE_DIT_CACHE"] = ( + "true" if args.enable_dit_cache else "false" + ) + if args.num_dit_steps is not None: + os.environ["NUM_DIT_STEPS"] = str(args.num_dit_steps) + elif args.enable_dit_cache: + os.environ.setdefault("NUM_DIT_STEPS", "8") + os.environ.setdefault("ATTENTION_BACKEND", "FA2") + os.environ.setdefault("TOKENIZERS_PARALLELISM", "false") + torch._dynamo.config.recompile_limit = 800 + + model_path = Path(args.model_path).resolve() + shared._require_dir(model_path, "G2 checkpoint") + metadata_path = model_path / "experiment_cfg" / "metadata.json" + conf_path = model_path / "experiment_cfg" / "conf.yaml" + shared._require_file(metadata_path, "G2 checkpoint metadata") + shared._require_file(conf_path, "G2 checkpoint config") + with metadata_path.open("r", encoding="utf-8") as stream: + checkpoint_metadata = json.load(stream) + if "g2" not in checkpoint_metadata: + raise KeyError( + f"Checkpoint does not contain G2 metadata: " + f"{metadata_path}; keys={list(checkpoint_metadata)}" + ) + + test_info = shared._validate_test_dataset( + Path(args.test_data_root).resolve() + ) + episode_indices = shared._parse_episode_indices( + args.episode_indices, + int(test_info["total_episodes"]), + ) + eval_offsets = shared._checkpoint_eval_offsets(model_path, "g2") + raw_resolutions = _checkpoint_g2_video_resolutions(model_path) + model_config_overrides, train_config_overrides = ( + shared._build_path_overrides(args) + ) + + Path(args.output_dir).mkdir(parents=True, exist_ok=True) + logging.info( + "[PREFLIGHT OK] checkpoint=%s checkpoint_embodiment=g2 " + "visual_source=g2 test=%s episodes=%s eval_offsets=%s " + "raw_resolutions=%s output=%s", + model_path, + Path(args.test_data_root).resolve(), + episode_indices, + eval_offsets, + { + key: [value[1], value[0]] + for key, value in raw_resolutions.items() + }, + Path(args.output_dir).resolve(), + ) + if args.preflight_only: + logging.info( + "Preflight-only validation completed; model was not loaded." + ) + return + + device_mesh = shared.init_mesh() + rank = dist.get_rank() + torch.manual_seed(args.seed) + np.random.seed(args.seed) + + signal_group = dist.new_group( + backend="gloo", + timeout=datetime.timedelta(seconds=args.timeout_seconds), + ) + policy = GrootSimPolicy( + embodiment_tag=EmbodimentTag("g2"), + model_path=str(model_path), + device="cuda" if torch.cuda.is_available() else "cpu", + device_mesh=device_mesh, + model_config_overrides=model_config_overrides, + train_config_overrides=train_config_overrides, + ) + action_head = policy.trained_model.action_head + if args.num_inference_timesteps < 0: + raise ValueError( + "--num-inference-timesteps must be non-negative" + ) + if args.num_inference_timesteps > 0: + action_head.num_inference_steps = int( + args.num_inference_timesteps + ) + action_head.num_inference_timesteps = int( + args.num_inference_timesteps + ) + if hasattr(action_head, "config"): + action_head.config.num_inference_timesteps = int( + args.num_inference_timesteps + ) + + if rank == 0: + if args.windowed: + # The windowed diagnostic intentionally keeps video and action in + # one policy object and one sequence of causal forwards. + wrapper = shared.G2RoboarenaPolicy( + groot_policy=policy, + signal_group=signal_group, + output_dir=str(Path(args.output_dir).resolve()), + video_save_mode="full", + ) + try: + reports = [ + _run_windowed_episode_comparison( + wrapper, + args, + test_info, + episode_index, + raw_resolutions, + ) + for episode_index in episode_indices + ] + finally: + wrapper._broadcast_signal_to_workers(shared.SIGNAL_SHUTDOWN) + + report_path = ( + Path(args.output_dir).resolve() / "windowed_report.json" + ) + with report_path.open("w", encoding="utf-8") as stream: + json.dump( + { + "checkpoint": str(model_path), + "checkpoint_embodiment": "g2", + "visual_source_embodiment": "g2", + "test_data_root": str( + Path(args.test_data_root).resolve() + ), + "episode_indices": episode_indices, + "reports": reports, + }, + stream, + ensure_ascii=False, + indent=2, + ) + logging.info("Saved windowed report: %s", report_path) + dist.barrier() + dist.destroy_process_group() + return + + wrapper_class = ( + shared.G2RoboarenaPolicy + if args.action_eval + else G2VideoDiagnosticPolicy + ) + wrapper = wrapper_class( + groot_policy=policy, + signal_group=signal_group, + output_dir=str(Path(args.output_dir).resolve()), + video_save_mode="none", + ) + summaries: list[dict] = [] + try: + for episode_index in episode_indices: + if args.full_episode: + summary = _run_full_episode_comparison( + wrapper, + args, + test_info, + episode_index, + eval_offsets, + raw_resolutions, + ) + else: + summary = shared._run_one_test_sample( + wrapper, + args, + test_info, + episode_index, + eval_offsets, + raw_resolutions, + ) + summary["checkpoint_embodiment"] = "g2" + summary["g2_state_used_for_model_conditioning"] = True + summary["action_note"] = ( + "Actions were deliberately ignored; this run evaluates " + "only the predicted video." + ) + if not args.full_episode: + _save_readable_comparison_views( + summary, + args.diagnostic_fps, + ) + summaries.append(summary) + finally: + wrapper._broadcast_signal_to_workers(shared.SIGNAL_SHUTDOWN) + + report_path = ( + Path(args.output_dir).resolve() / "testset_report.json" + ) + with report_path.open("w", encoding="utf-8") as stream: + json.dump( + { + "checkpoint": str(model_path), + "checkpoint_embodiment": "g2", + "visual_source_embodiment": "g2", + "test_data_root": str( + Path(args.test_data_root).resolve() + ), + "episode_indices": episode_indices, + "eval_offsets": eval_offsets, + "checkpoint_raw_resolutions": { + key: [value[1], value[0]] + for key, value in raw_resolutions.items() + }, + "samples": summaries, + }, + stream, + ensure_ascii=False, + indent=2, + ) + logging.info("Saved test-set report: %s", report_path) + dist.barrier() + else: + worker = shared.WebsocketPolicyServer( + policy=policy, + host="127.0.0.1", + port=0, + metadata={}, + output_dir=None, + signal_group=signal_group, + ) + asyncio.run(worker._worker_loop()) + dist.barrier() + + dist.destroy_process_group() + + +def cli() -> None: + logging.basicConfig(level=logging.INFO, force=True) + main(tyro.cli(Args)) + + +if __name__ == "__main__": + cli() diff --git a/eval_g2_server_initial_actions.py b/eval_g2_server_initial_actions.py new file mode 100644 index 00000000..89ac8c78 --- /dev/null +++ b/eval_g2_server_initial_actions.py @@ -0,0 +1,190 @@ +"""Evaluate production causal actions on every G2 episode start.""" + +from __future__ import annotations + +import argparse +import json +import logging +from pathlib import Path + +import cv2 +import numpy as np +import pyarrow.parquet as pq + +from eval_utils.policy_client import WebsocketClientPolicy + + +CAMERAS = ("top_head", "hand_left", "hand_right") + + +def _read_history(path: Path, end_row: int = 3) -> np.ndarray: + capture = cv2.VideoCapture(str(path)) + frames: list[np.ndarray] = [] + try: + while len(frames) <= end_row: + ok, bgr = capture.read() + if not ok: + raise RuntimeError(f"Could not read frame {len(frames)}: {path}") + frames.append(cv2.cvtColor(bgr, cv2.COLOR_BGR2RGB)) + finally: + capture.release() + return np.stack(frames[-4:], axis=0).astype(np.uint8) + + +def _arm_metrics( + action: np.ndarray, state: np.ndarray, ground_truth: np.ndarray +) -> dict[str, object]: + pred_motion = [ + float(np.max(np.abs(action[:, 0:7] - state[None, 0:7]))), + float(np.max(np.abs(action[:, 8:15] - state[None, 8:15]))), + ] + gt_motion = [ + float(np.max(np.abs(ground_truth[:, 0:7] - state[None, 0:7]))), + float(np.max(np.abs(ground_truth[:, 8:15] - state[None, 8:15]))), + ] + pred_dominant = int(np.argmax(pred_motion)) + gt_dominant = int(np.argmax(gt_motion)) + gt_peak = max(gt_motion) + pred_on_gt_arm = pred_motion[gt_dominant] + return { + "pred_motion_left_rad": pred_motion[0], + "pred_motion_right_rad": pred_motion[1], + "gt_motion_left_rad": gt_motion[0], + "gt_motion_right_rad": gt_motion[1], + "pred_dominant_arm": ("left", "right")[pred_dominant], + "gt_dominant_arm": ("left", "right")[gt_dominant], + "dominant_arm_correct": pred_dominant == gt_dominant, + "motion_ratio_on_gt_arm": ( + pred_on_gt_arm / gt_peak if gt_peak > 1e-8 else None + ), + "mae_left_rad": float( + np.mean(np.abs(action[:, 0:7] - ground_truth[:, 0:7])) + ), + "mae_right_rad": float( + np.mean(np.abs(action[:, 8:15] - ground_truth[:, 8:15])) + ), + } + + +def main() -> None: + parser = argparse.ArgumentParser() + parser.add_argument("dataset_root", type=Path) + parser.add_argument("--host", default="127.0.0.1") + parser.add_argument("--port", type=int, required=True) + parser.add_argument("--output-dir", type=Path, required=True) + parser.add_argument("--max-episodes", type=int, default=0) + args = parser.parse_args() + + with (args.dataset_root / "meta/info.json").open( + "r", encoding="utf-8" + ) as stream: + info = json.load(stream) + total = int(info["total_episodes"]) + if args.max_episodes > 0: + total = min(total, args.max_episodes) + + client = WebsocketClientPolicy(host=args.host, port=args.port) + records: list[dict[str, object]] = [] + row = 3 + for episode in range(total): + parquet_path = ( + args.dataset_root + / "data" + / f"chunk-{episode // 1000:03d}" + / f"episode_{episode:06d}.parquet" + ) + table = pq.read_table(parquet_path) + states = np.asarray( + table["observation.state"].to_pylist(), dtype=np.float32 + ) + actions = np.asarray(table["action"].to_pylist(), dtype=np.float32) + if len(actions) < row + 24: + logging.warning("Skipping short episode %d", episode) + continue + prompt = str(table["annotation.language.action_text"][row].as_py()) + observation: dict[str, object] = { + "observation/left_joint_position": states[row, 0:7], + "observation/left_gripper_position": states[row, 7:8], + "observation/right_joint_position": states[row, 8:15], + "observation/right_gripper_position": states[row, 15:16], + "prompt": prompt, + "session_id": f"initial-action-audit-episode-{episode:06d}", + } + for camera in CAMERAS: + video_path = ( + args.dataset_root + / "videos" + / f"chunk-{episode // 1000:03d}" + / f"observation.images.{camera}" + / f"episode_{episode:06d}.mp4" + ) + observation[f"observation/{camera}"] = _read_history(video_path) + + predicted = np.asarray( + client.infer(observation), dtype=np.float32 + ) + if predicted.shape != (24, 16): + raise ValueError( + f"Episode {episode}: expected (24,16), got {predicted.shape}" + ) + record = { + "episode": episode, + **_arm_metrics( + predicted, + states[row], + actions[row : row + 24], + ), + } + records.append(record) + logging.info( + "episode=%d dominant=%s/%s correct=%s pred_LR=(%.4f,%.4f) " + "gt_LR=(%.4f,%.4f)", + episode, + record["pred_dominant_arm"], + record["gt_dominant_arm"], + record["dominant_arm_correct"], + record["pred_motion_left_rad"], + record["pred_motion_right_rad"], + record["gt_motion_left_rad"], + record["gt_motion_right_rad"], + ) + + if not records: + raise RuntimeError("No episodes were evaluated") + args.output_dir.mkdir(parents=True, exist_ok=True) + with (args.output_dir / "per_episode.json").open( + "w", encoding="utf-8" + ) as stream: + json.dump(records, stream, ensure_ascii=False, indent=2) + + ratios = [ + float(record["motion_ratio_on_gt_arm"]) + for record in records + if record["motion_ratio_on_gt_arm"] is not None + ] + summary = { + "checkpoint_server": f"{args.host}:{args.port}", + "dataset_root": str(args.dataset_root), + "episodes": len(records), + "dominant_arm_accuracy": float( + np.mean([record["dominant_arm_correct"] for record in records]) + ), + "median_motion_ratio_on_gt_arm": float(np.median(ratios)), + "mean_motion_ratio_on_gt_arm": float(np.mean(ratios)), + "mean_mae_left_rad": float( + np.mean([record["mae_left_rad"] for record in records]) + ), + "mean_mae_right_rad": float( + np.mean([record["mae_right_rad"] for record in records]) + ), + } + with (args.output_dir / "summary.json").open( + "w", encoding="utf-8" + ) as stream: + json.dump(summary, stream, ensure_ascii=False, indent=2) + print(json.dumps(summary, ensure_ascii=False, indent=2)) + + +if __name__ == "__main__": + logging.basicConfig(level=logging.INFO, force=True) + main() diff --git a/groot/vla/configs/conf.yaml b/groot/vla/configs/conf.yaml index 48da05bc..2aef043a 100644 --- a/groot/vla/configs/conf.yaml +++ b/groot/vla/configs/conf.yaml @@ -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 === @@ -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 @@ -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 diff --git a/groot/vla/configs/data/dreamzero/base_48_wan_fine_aug_relative.yaml b/groot/vla/configs/data/dreamzero/base_48_wan_fine_aug_relative.yaml index cfcdd449..5a13e22f 100644 --- a/groot/vla/configs/data/dreamzero/base_48_wan_fine_aug_relative.yaml +++ b/groot/vla/configs/data/dreamzero/base_48_wan_fine_aug_relative.yaml @@ -345,10 +345,174 @@ transform_yam: # Modality Configs ################################################################################ +# BEGIN G2_EE_QUATERNION_CONFIG +# G2 Cartesian end-effector data: xyz + quaternion(wxyz) + gripper per arm. +modality_config_g2_ee: + 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_effector_position + - state.left_effector_rotation + - state.left_gripper_position + - state.right_effector_position + - state.right_effector_rotation + - 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_effector_position + - action.left_effector_rotation + - action.left_gripper_position + - action.right_effector_position + - action.right_effector_rotation + - action.right_gripper_position + language: + _target_: groot.vla.data.dataset.ModalityConfig + delta_indices: [0] + modality_keys: + - annotation.language.action_text + +transform_g2_ee: + _target_: groot.vla.data.transform.ComposedModalityTransform + transforms: + - <<: *totensor_cfg + apply_to: ${modality_config_g2_ee.video.modality_keys} + - <<: *crop_cfg + apply_to: ${modality_config_g2_ee.video.modality_keys} + - <<: *resize_cfg + apply_to: ${modality_config_g2_ee.video.modality_keys} + - <<: *color_jitter_cfg + apply_to: ${modality_config_g2_ee.video.modality_keys} + - <<: *to_numpy_cfg + apply_to: ${modality_config_g2_ee.video.modality_keys} + + - _target_: groot.vla.data.transform.StateActionToTensor + apply_to: ${modality_config_g2_ee.state.modality_keys} + - _target_: groot.vla.data.transform.StateActionTransform + apply_to: ${modality_config_g2_ee.state.modality_keys} + target_rotations: + state.left_effector_rotation: rotation_6d + state.right_effector_rotation: rotation_6d + normalization_modes: + state.left_effector_position: q99 + state.left_effector_rotation: min_max + state.left_gripper_position: q99 + state.right_effector_position: q99 + state.right_effector_rotation: min_max + state.right_gripper_position: q99 + + - _target_: groot.vla.data.transform.StateActionToTensor + apply_to: ${modality_config_g2_ee.action.modality_keys} + - _target_: groot.vla.data.transform.StateActionTransform + apply_to: ${modality_config_g2_ee.action.modality_keys} + target_rotations: + action.left_effector_rotation: rotation_6d + action.right_effector_rotation: rotation_6d + normalization_modes: + action.left_effector_position: q99 + action.left_effector_rotation: min_max + action.left_gripper_position: q99 + action.right_effector_position: q99 + action.right_effector_rotation: min_max + action.right_gripper_position: q99 + + - _target_: groot.vla.data.transform.ConcatTransform + video_concat_order: ${modality_config_g2_ee.video.modality_keys} + state_concat_order: ${modality_config_g2_ee.state.modality_keys} + action_concat_order: ${modality_config_g2_ee.action.modality_keys} + - ${model_specific_transform} +# END G2_EE_QUATERNION_CONFIG + +################################################################################ +# 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 @@ -358,6 +522,7 @@ transforms: oxe_droid: ${transform_oxe_droid} agibot: ${transform_agibot} yam: ${transform_yam} + g2: ${transform_g2} ################################################################################ # Metadata Versions @@ -367,6 +532,7 @@ metadata_versions: oxe_droid: '0221' agibot: '0221' yam: '0221' + g2: '0221' ################################################################################ # FPS (per embodiment, null means use dataset default) @@ -374,3 +540,4 @@ metadata_versions: fps: yam: 30 + g2: 30 diff --git a/groot/vla/configs/data/dreamzero/base_48_wan_fine_aug_relative_1.yaml b/groot/vla/configs/data/dreamzero/base_48_wan_fine_aug_relative_1.yaml new file mode 100644 index 00000000..3ec62aad --- /dev/null +++ b/groot/vla/configs/data/dreamzero/base_48_wan_fine_aug_relative_1.yaml @@ -0,0 +1,464 @@ +# @package _global_ + +# Assume `model_specific_transform` is defined in the model config file + +################################################################################ +# Normalization Statistics +# By default, we compute the normalization statistics for the datasets actually +# used in the mixture. If you want to use the global metadata, set this to true. +################################################################################ + +use_global_metadata: false + +################################################################################ +# Dimension Information +################################################################################ + +num_frames: 49 +action_horizon: 48 +state_horizon: 1 + +# image_resolution_width: 832 +# image_resolution_height: 480 + +image_resolution_width: 480 +image_resolution_height: 256 + +image_resolution_width_single_frame: 256 +image_resolution_height_single_frame: 256 + +################################################################################ +# Anchored Video Transforms +################################################################################ +totensor_cfg: &totensor_cfg + _target_: groot.vla.data.transform.VideoToTensor + apply_to: ??? + +crop_cfg: &crop_cfg + _target_: groot.vla.data.transform.VideoCrop + apply_to: ??? + scale: 0.95 + mode: random + + +resize_cfg: &resize_cfg + _target_: groot.vla.data.transform.VideoResize + apply_to: ??? + height: ${image_resolution_height} + width: ${image_resolution_width} + interpolation: linear + +resize_cfg_single_frame: &resize_cfg_single_frame + _target_: groot.vla.data.transform.VideoResize + apply_to: ??? + height: ${image_resolution_height_single_frame} + width: ${image_resolution_width_single_frame} + interpolation: linear + +color_jitter_cfg: &color_jitter_cfg + _target_: groot.vla.data.transform.VideoColorJitter + apply_to: ??? + brightness: 0.3 + contrast: 0.4 + saturation: 0.5 + hue: 0.08 + +random_grayscale_cfg: &random_grayscale_cfg + _target_: groot.vla.data.transform.VideoRandomGrayscale + apply_to: ??? + p: 0.1 + +random_posterize_cfg: &random_posterize_cfg + _target_: groot.vla.data.transform.VideoRandomPosterize + apply_to: ??? + bits: 4 + p: 0.1 + +normalize_cfg: &normalize_cfg + _target_: groot.vla.data.transform.VideoNormalize + apply_to: ??? + mean: [0.5, 0.5, 0.5] + std: [0.5, 0.5, 0.5] + + +to_numpy_cfg: &to_numpy_cfg + _target_: groot.vla.data.transform.VideoToNumpy + apply_to: ??? + + +################################################################################ +# oxe_droid (OXE Droid) +################################################################################ + +# Modality Configs +modality_config_oxe_droid: + 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: [0] + modality_keys: + - video.exterior_image_1_left + - video.exterior_image_2_left + - video.wrist_image_left + state: + _target_: groot.vla.data.dataset.ModalityConfig + delta_indices: [0] + modality_keys: + - state.joint_position + - state.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.joint_position + - action.gripper_position + language: + _target_: groot.vla.data.dataset.ModalityConfig + delta_indices: [0] + modality_keys: + - annotation.language.language_instruction + - annotation.language.language_instruction_2 + - annotation.language.language_instruction_3 + lapa_action: + _target_: groot.vla.data.dataset.ModalityConfig + delta_indices: [0] + modality_keys: + - lapa_action + +# Transforms +transform_oxe_droid: + _target_: groot.vla.data.transform.ComposedModalityTransform + transforms: + # Video transforms + - <<: *totensor_cfg + apply_to: ${modality_config_oxe_droid.video.modality_keys} + - <<: *crop_cfg + apply_to: ${modality_config_oxe_droid.video.modality_keys} + - <<: *resize_cfg + apply_to: ${modality_config_oxe_droid.video.modality_keys} + - <<: *color_jitter_cfg + apply_to: ${modality_config_oxe_droid.video.modality_keys} + - <<: *to_numpy_cfg + apply_to: ${modality_config_oxe_droid.video.modality_keys} + + # State transforms + - _target_: groot.vla.data.transform.StateActionToTensor + apply_to: ${modality_config_oxe_droid.state.modality_keys} + - _target_: groot.vla.data.transform.StateActionTransform + apply_to: ${modality_config_oxe_droid.state.modality_keys} + normalization_modes: + state.joint_position: q99 + state.gripper_position: q99 + + # Action transforms + - _target_: groot.vla.data.transform.StateActionToTensor + apply_to: ${modality_config_oxe_droid.action.modality_keys} + - _target_: groot.vla.data.transform.StateActionTransform + apply_to: ${modality_config_oxe_droid.action.modality_keys} + normalization_modes: + action.joint_position: q99 + action.gripper_position: q99 + + # ConcatTransform + - _target_: groot.vla.data.transform.ConcatTransform + video_concat_order: ${modality_config_oxe_droid.video.modality_keys} + state_concat_order: ${modality_config_oxe_droid.state.modality_keys} + action_concat_order: ${modality_config_oxe_droid.action.modality_keys} + + # Model-specific transform + - ${model_specific_transform} + +################################################################################ +# agibot (AGIbot: state 32, action 22, 3 views) +################################################################################ + +modality_config_agibot: + 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_arm_joint_position + - state.right_arm_joint_position + - state.left_effector_position + - state.right_effector_position + - state.head_position + - state.waist_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_arm_joint_position + - action.right_arm_joint_position + - action.left_effector_position + - action.right_effector_position + - action.head_position + - action.waist_position + - action.robot_velocity + language: + _target_: groot.vla.data.dataset.ModalityConfig + delta_indices: [0] + modality_keys: + - annotation.language.action_text + +transform_agibot: + _target_: groot.vla.data.transform.ComposedModalityTransform + transforms: + # Video transforms + - <<: *totensor_cfg + apply_to: ${modality_config_agibot.video.modality_keys} + - <<: *crop_cfg + apply_to: ${modality_config_agibot.video.modality_keys} + - <<: *resize_cfg + apply_to: ${modality_config_agibot.video.modality_keys} + - <<: *color_jitter_cfg + apply_to: ${modality_config_agibot.video.modality_keys} + - <<: *to_numpy_cfg + apply_to: ${modality_config_agibot.video.modality_keys} + + # State transforms + - _target_: groot.vla.data.transform.StateActionToTensor + apply_to: ${modality_config_agibot.state.modality_keys} + - _target_: groot.vla.data.transform.StateActionTransform + apply_to: ${modality_config_agibot.state.modality_keys} + normalization_modes: + state.left_arm_joint_position: q99 + state.right_arm_joint_position: q99 + state.left_effector_position: q99 + state.right_effector_position: q99 + state.head_position: q99 + state.waist_position: q99 + + # Action transforms + - _target_: groot.vla.data.transform.StateActionToTensor + apply_to: ${modality_config_agibot.action.modality_keys} + - _target_: groot.vla.data.transform.StateActionTransform + apply_to: ${modality_config_agibot.action.modality_keys} + normalization_modes: + action.left_arm_joint_position: q99 + action.right_arm_joint_position: q99 + action.left_effector_position: q99 + action.right_effector_position: q99 + action.head_position: q99 + action.waist_position: q99 + action.robot_velocity: q99 + + # ConcatTransform + - _target_: groot.vla.data.transform.ConcatTransform + video_concat_order: ${modality_config_agibot.video.modality_keys} + state_concat_order: ${modality_config_agibot.state.modality_keys} + action_concat_order: ${modality_config_agibot.action.modality_keys} + + # Model-specific transform + - ${model_specific_transform} + +################################################################################ +# yam (YAM: joint+gripper only from Dataset/YAM_play_data/meta/modality.json; +# state 14 dims, action 14 dims, 3 views) +################################################################################ + +modality_config_yam: + 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: [0] + modality_keys: + - video.top_camera-images-rgb + - video.left_camera-images-rgb + - video.right_camera-images-rgb + state: + _target_: groot.vla.data.dataset.ModalityConfig + delta_indices: [0] + modality_keys: + - state.left_joint_pos + - state.left_gripper_pos + - state.right_joint_pos + - state.right_gripper_pos + 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_pos + - action.left_gripper_pos + - action.right_joint_pos + - action.right_gripper_pos + language: + _target_: groot.vla.data.dataset.ModalityConfig + delta_indices: [0] + modality_keys: + - annotation.task + +transform_yam: + _target_: groot.vla.data.transform.ComposedModalityTransform + transforms: + # Video transforms + - <<: *totensor_cfg + apply_to: ${modality_config_yam.video.modality_keys} + - <<: *crop_cfg + apply_to: ${modality_config_yam.video.modality_keys} + - <<: *resize_cfg + apply_to: ${modality_config_yam.video.modality_keys} + - <<: *color_jitter_cfg + apply_to: ${modality_config_yam.video.modality_keys} + - <<: *to_numpy_cfg + apply_to: ${modality_config_yam.video.modality_keys} + + # State transforms + - _target_: groot.vla.data.transform.StateActionToTensor + apply_to: ${modality_config_yam.state.modality_keys} + - _target_: groot.vla.data.transform.StateActionTransform + apply_to: ${modality_config_yam.state.modality_keys} + normalization_modes: + state.left_joint_pos: q99 + state.left_gripper_pos: q99 + state.right_joint_pos: q99 + state.right_gripper_pos: q99 + + # Action transforms + - _target_: groot.vla.data.transform.StateActionToTensor + apply_to: ${modality_config_yam.action.modality_keys} + - _target_: groot.vla.data.transform.StateActionTransform + apply_to: ${modality_config_yam.action.modality_keys} + normalization_modes: + action.left_joint_pos: q99 + action.left_gripper_pos: q99 + action.right_joint_pos: q99 + action.right_gripper_pos: q99 + + # ConcatTransform + - _target_: groot.vla.data.transform.ConcatTransform + video_concat_order: ${modality_config_yam.video.modality_keys} + state_concat_order: ${modality_config_yam.state.modality_keys} + action_concat_order: ${modality_config_yam.action.modality_keys} + + # Model-specific transform + - ${model_specific_transform} + +################################################################################ +# Modality Configs +################################################################################ + +# BEGIN G2_EE_QUATERNION_CONFIG +# G2 Cartesian end-effector data: xyz + quaternion(wxyz) + gripper per arm. +modality_config_g2_ee: + 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_effector_position + - state.left_effector_rotation + - state.left_gripper_position + - state.right_effector_position + - state.right_effector_rotation + - 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_effector_position + - action.left_effector_rotation + - action.left_gripper_position + - action.right_effector_position + - action.right_effector_rotation + - action.right_gripper_position + language: + _target_: groot.vla.data.dataset.ModalityConfig + delta_indices: [0] + modality_keys: + - annotation.language.action_text + +transform_g2_ee: + _target_: groot.vla.data.transform.ComposedModalityTransform + transforms: + - <<: *totensor_cfg + apply_to: ${modality_config_g2_ee.video.modality_keys} + - <<: *crop_cfg + apply_to: ${modality_config_g2_ee.video.modality_keys} + - <<: *resize_cfg + apply_to: ${modality_config_g2_ee.video.modality_keys} + - <<: *color_jitter_cfg + apply_to: ${modality_config_g2_ee.video.modality_keys} + - <<: *to_numpy_cfg + apply_to: ${modality_config_g2_ee.video.modality_keys} + + - _target_: groot.vla.data.transform.StateActionToTensor + apply_to: ${modality_config_g2_ee.state.modality_keys} + - _target_: groot.vla.data.transform.StateActionTransform + apply_to: ${modality_config_g2_ee.state.modality_keys} + target_rotations: + state.left_effector_rotation: rotation_6d + state.right_effector_rotation: rotation_6d + normalization_modes: + state.left_effector_position: q99 + state.left_effector_rotation: min_max + state.left_gripper_position: q99 + state.right_effector_position: q99 + state.right_effector_rotation: min_max + state.right_gripper_position: q99 + + - _target_: groot.vla.data.transform.StateActionToTensor + apply_to: ${modality_config_g2_ee.action.modality_keys} + - _target_: groot.vla.data.transform.StateActionTransform + apply_to: ${modality_config_g2_ee.action.modality_keys} + target_rotations: + action.left_effector_rotation: rotation_6d + action.right_effector_rotation: rotation_6d + normalization_modes: + action.left_effector_position: q99 + action.left_effector_rotation: min_max + action.left_gripper_position: q99 + action.right_effector_position: q99 + action.right_effector_rotation: min_max + action.right_gripper_position: q99 + + - _target_: groot.vla.data.transform.ConcatTransform + video_concat_order: ${modality_config_g2_ee.video.modality_keys} + state_concat_order: ${modality_config_g2_ee.state.modality_keys} + action_concat_order: ${modality_config_g2_ee.action.modality_keys} + - ${model_specific_transform} +# END G2_EE_QUATERNION_CONFIG + +modality_configs: + oxe_droid: ${modality_config_oxe_droid} + agibot: ${modality_config_agibot} + yam: ${modality_config_yam} + +################################################################################ +# Transforms +################################################################################ + +transforms: + oxe_droid: ${transform_oxe_droid} + agibot: ${transform_agibot} + yam: ${transform_yam} + +################################################################################ +# Metadata Versions +################################################################################ + +metadata_versions: + oxe_droid: '0221' + agibot: '0221' + yam: '0221' + +################################################################################ +# FPS (per embodiment, null means use dataset default) +################################################################################ + +fps: + yam: 30 diff --git a/groot/vla/configs/data/dreamzero/g2_absolute.yaml b/groot/vla/configs/data/dreamzero/g2_absolute.yaml new file mode 100644 index 00000000..54bceeef --- /dev/null +++ b/groot/vla/configs/data/dreamzero/g2_absolute.yaml @@ -0,0 +1,57 @@ +# @package _global_ + +# G2 absolute-action training (rich action space / canonical subtask data). +# Difference vs. g2_relative.yaml: relative_action=false, so the model learns +# to predict ABSOLUTE joint targets directly (no current-state delta +# conversion). Deployment must use the client with --g2-arm-mode absolute. +# +# Set the data root via env: +# G2_DATA_ROOT=/data/training_data/teleop/g2/g2_tasks_g1_g7_joint_gear_subtask_v4_stratified_canonical/train + +defaults: + - dreamzero/base_48_wan_fine_aug_relative + - _self_ + +max_state_dim: 64 +use_global_metadata: false +relative_action: false +relative_action_per_horizon: false +relative_action_keys: [] +max_chunk_size: 5 +dataset_shard_sampling_rate: 0.1 +active_hold_index_path: ${g2_data_root}/meta/g2_active_hold_windows.json +active_window_ratio: 0.8 +mixture_dataset_cls: groot.vla.data.dataset.lerobot_sharded.ShardedLeRobotMixtureDataset.from_mixture_spec +single_dataset_cls: groot.vla.data.dataset.lerobot_sharded.ShardedLeRobotSubLangSingleActionChunkDatasetDROID + +g2_data_root: ??? + +train_dataset: + _target_: ${mixture_dataset_cls} + _convert_: object + mixture_spec: + - dataset_path: + g2: + - ${g2_data_root} + dataset_weight: 1.0 + distribute_weights: true + + dataset_class: ${single_dataset_cls} + all_modality_configs: ${modality_configs} + all_transforms: ${transforms} + metadata_versions: ${metadata_versions} + fps: ${fps} + dataset_kwargs: + video_backend: decord + use_global_metadata: ${use_global_metadata} + max_chunk_size: ${max_chunk_size} + relative_action: ${relative_action} + relative_action_keys: ${relative_action_keys} + relative_action_per_horizon: ${relative_action_per_horizon} + mixture_kwargs: + training: true + balance_dataset_weights: false + seed: 42 + shard_sampling_rate: ${dataset_shard_sampling_rate} + active_hold_index_path: ${active_hold_index_path} + active_window_ratio: ${active_window_ratio} diff --git a/groot/vla/configs/data/dreamzero/g2_relative.yaml b/groot/vla/configs/data/dreamzero/g2_relative.yaml new file mode 100644 index 00000000..d935912a --- /dev/null +++ b/groot/vla/configs/data/dreamzero/g2_relative.yaml @@ -0,0 +1,52 @@ +# @package _global_ + +defaults: + - dreamzero/base_48_wan_fine_aug_relative + - _self_ + +max_state_dim: 64 +use_global_metadata: false +relative_action: true +relative_action_per_horizon: false +relative_action_keys: + - left_joint_position + - right_joint_position +max_chunk_size: 5 +dataset_shard_sampling_rate: 0.1 +active_hold_index_path: ${g2_data_root}/meta/g2_active_hold_windows.json +active_window_ratio: 0.8 +mixture_dataset_cls: groot.vla.data.dataset.lerobot_sharded.ShardedLeRobotMixtureDataset.from_mixture_spec +single_dataset_cls: groot.vla.data.dataset.lerobot_sharded.ShardedLeRobotSubLangSingleActionChunkDatasetDROID + +# Override with g2_data_root=/path/to/the/GEAR/train/directory. +g2_data_root: ??? + +train_dataset: + _target_: ${mixture_dataset_cls} + _convert_: object + mixture_spec: + - dataset_path: + g2: + - ${g2_data_root} + dataset_weight: 1.0 + distribute_weights: true + + dataset_class: ${single_dataset_cls} + all_modality_configs: ${modality_configs} + all_transforms: ${transforms} + metadata_versions: ${metadata_versions} + fps: ${fps} + dataset_kwargs: + video_backend: decord + use_global_metadata: ${use_global_metadata} + max_chunk_size: ${max_chunk_size} + relative_action: ${relative_action} + relative_action_keys: ${relative_action_keys} + relative_action_per_horizon: ${relative_action_per_horizon} + mixture_kwargs: + training: true + balance_dataset_weights: false + seed: 42 + shard_sampling_rate: ${dataset_shard_sampling_rate} + active_hold_index_path: ${active_hold_index_path} + active_window_ratio: ${active_window_ratio} diff --git a/groot/vla/configs/model/dreamzero/action_head/wan_flow_matching_action_tf.yaml b/groot/vla/configs/model/dreamzero/action_head/wan_flow_matching_action_tf.yaml index 8cccdc63..fb133969 100644 --- a/groot/vla/configs/model/dreamzero/action_head/wan_flow_matching_action_tf.yaml +++ b/groot/vla/configs/model/dreamzero/action_head/wan_flow_matching_action_tf.yaml @@ -38,7 +38,7 @@ action_head_cfg: model_dtype: float32 max_state_dim: ${max_state_dim} max_action_dim: ${max_action_dim} - action_loss_embodiment_ids: [26, 17, 32] + action_loss_embodiment_ids: [26, 17, 32, 33] hidden_size: ${hidden_size} input_embedding_dim: 1536 backbone_embedding_dim: ${backbone_hidden_size} @@ -106,3 +106,12 @@ action_head_cfg: tune_projector: true tune_diffusion_model: true + action_only_training: false + dynamics_loss_weight: 1.0 + action_loss_weight: 1.0 + # Auxiliary action constraints (keep small to avoid disrupting flow-matching): + # L_start: first action of the chunk anchors to the current joint state. + # L_video_to_action: actions recovered from predicted video latent must + # match the action chunk (action depends on implicit dynamics). + action_start_loss_weight: 0.05 + action_video_consistency_loss_weight: 0.10 diff --git a/groot/vla/configs/model/dreamzero/transform/base.yaml b/groot/vla/configs/model/dreamzero/transform/base.yaml index 3873b9ab..cd6439dd 100644 --- a/groot/vla/configs/model/dreamzero/transform/base.yaml +++ b/groot/vla/configs/model/dreamzero/transform/base.yaml @@ -33,6 +33,9 @@ embodiment_tag_to_projector_index: oxe_plex: 30 dream: 31 yam: 32 + # Logical G2 identity. The current DreamZero-AgiBot action adapter is + # single-embodiment internally and maps this to local adapter slot 0. + g2: 33 xdof: 22 gr1_unified_segmentation: 14 language_table_sim: 7 diff --git a/groot/vla/data/dataset/lerobot_sharded.py b/groot/vla/data/dataset/lerobot_sharded.py old mode 100755 new mode 100644 index f8aa338c..3d8df304 --- a/groot/vla/data/dataset/lerobot_sharded.py +++ b/groot/vla/data/dataset/lerobot_sharded.py @@ -1291,6 +1291,8 @@ def __init__( shard_sampling_rate: float = 0.5, num_shards_to_sample: int = 2**20, allow_padding_at_end: bool = False, + active_hold_index_path: str | None = None, + active_window_ratio: float = 0.8, ): """ Initialize the mixture dataset. @@ -1317,6 +1319,25 @@ def __init__( # Set properties self.shard_sampling_rate = shard_sampling_rate self.num_shards_to_sample = num_shards_to_sample + self.active_window_ratio = float(active_window_ratio) + if not 0.0 <= self.active_window_ratio <= 1.0: + raise ValueError("active_window_ratio must be in [0, 1]") + self._active_steps: set[tuple[int, int]] | None = None + if active_hold_index_path: + index_path = Path(active_hold_index_path) + with index_path.open() as stream: + active_hold = json.load(stream) + self._active_steps = { + (int(episode), int(step)) + for episode, steps in active_hold["active"].items() + for step in steps + } + print( + "[G2 ACTIVE/HOLD SAMPLING] " + f"index={index_path} active={len(self._active_steps)} " + f"target_ratio={self.active_window_ratio:.2f} " + f"threshold={active_hold['arm_motion_threshold']:.8f}" + ) # Calculate shard sampling weights all_shard_sampling_weights = [] @@ -1501,9 +1522,37 @@ def __iter__(self): allowed_indices = allowed_indices[allowed_indices <= allowed_length] for i in allowed_indices: all_steps.append((trajectory_id, i)) + sample_count = int(dataset.num_steps_per_shard * self.shard_sampling_rate) if self.training: rng.shuffle(all_steps) - sampled_steps = all_steps[: int(dataset.num_steps_per_shard * self.shard_sampling_rate)] + if self._active_steps is None: + sampled_steps = all_steps[:sample_count] + else: + active_steps = [ + step for step in all_steps if step in self._active_steps + ] + hold_steps = [ + step for step in all_steps if step not in self._active_steps + ] + if self.training: + rng.shuffle(active_steps) + rng.shuffle(hold_steps) + active_count = min( + len(active_steps), + int(round(sample_count * self.active_window_ratio)), + ) + hold_count = min(len(hold_steps), sample_count - active_count) + sampled_steps = ( + active_steps[:active_count] + hold_steps[:hold_count] + ) + if len(sampled_steps) < sample_count: + used = set(sampled_steps) + sampled_steps.extend( + step for step in all_steps if step not in used + ) + sampled_steps = sampled_steps[:sample_count] + if self.training: + rng.shuffle(sampled_steps) for trajectory_id, step_index in sampled_steps: # print( # f"Loading step data from rank {self.rank}, worker {self.worker_id}: {dataset_index} {trajectory_id}, {step_index}" diff --git a/groot/vla/data/schema/embodiment_tags.py b/groot/vla/data/schema/embodiment_tags.py old mode 100755 new mode 100644 index e7eaec04..dd09e353 --- a/groot/vla/data/schema/embodiment_tags.py +++ b/groot/vla/data/schema/embodiment_tags.py @@ -153,6 +153,8 @@ class EmbodimentTag(Enum): """ AGIBOT = "agibot" + G2 = "g2" + YAM = "yam" DREAM = "dream" diff --git a/groot/vla/experiment/base.py b/groot/vla/experiment/base.py index 5e9a7d28..faaf8dc4 100644 --- a/groot/vla/experiment/base.py +++ b/groot/vla/experiment/base.py @@ -144,6 +144,11 @@ def on_save(self, args, state, control, **kwargs): print(f"Copying wandb_config.json from {wandb_config_src} to {wandb_config_dst}") shutil.copy2(wandb_config_src, wandb_config_dst) + trainer_state = checkpoint_dir / "trainer_state.json" + training_state = checkpoint_dir / "training_state.json" + if trainer_state.exists(): + shutil.copy2(trainer_state, training_state) + class ProfCallback(transformers.TrainerCallback): """Callback to manage PyTorch profiler during training. @@ -481,14 +486,77 @@ def save_model(self, output_dir: Optional[str], _internal_call: bool): else: state_dict = self.model.state_dict() - if self.base_cfg.save_lora_only: - # Save only the trainable parameters - train_key = [k for k, v in self.model.named_parameters() if v.requires_grad] - lora_state_dict = {k: v for k, v in self.model.state_dict().items() if k in train_key} + action_head_config = getattr( + getattr(self.model, "action_head", None), + "config", + None, + ) + action_only_training = bool( + getattr(action_head_config, "action_only_training", False) + ) + + if action_only_training: + adapter_prefixes = ( + "action_head.model.state_encoder.", + "action_head.model.action_encoder.", + "action_head.model.action_decoder.", + ) + adapter_state_dict = { + key: value + for key, value in state_dict.items() + if key.startswith(adapter_prefixes) + } + expected_adapter_keys = { + key + for key in self.model.state_dict() + if key.startswith(adapter_prefixes) + } + if set(adapter_state_dict) != expected_adapter_keys: + missing = sorted( + expected_adapter_keys - set(adapter_state_dict) + ) + unexpected = sorted( + set(adapter_state_dict) - expected_adapter_keys + ) + raise RuntimeError( + "Action adapter save contract mismatch: " + f"missing={missing[:20]} unexpected={unexpected[:20]}" + ) + state_dict = adapter_state_dict + elif self.base_cfg.save_lora_only: + # Save trainable parameters. During the action-only second stage, + # also retain the frozen stage-1 LoRA tensors: inference needs + # those adapters to preserve the already-recovered video model. + train_key = { + k for k, v in self.model.named_parameters() if v.requires_grad + } + lora_state_dict = { + k: v + for k, v in self.model.state_dict().items() + if k in train_key + } state_dict = lora_state_dict if self.args.should_save: ret = self.model.save_pretrained(output_dir, state_dict=state_dict) + if action_only_training: + from safetensors.torch import save_file + + action_adapter_path = os.path.join( + output_dir, + "action_expert.safetensors", + ) + save_file( + { + key: value.detach().cpu().contiguous() + for key, value in state_dict.items() + }, + action_adapter_path, + ) + print( + "[ACTION ADAPTER SAVE] " + f"{action_adapter_path} keys={len(state_dict)}" + ) # can separately save the VLM model for downstream evalualtion if self.base_cfg.save_llm: @@ -701,6 +769,20 @@ def create_model(self, cfg, training_args): safetensors_index_path = os.path.join(ckpt_dir, "model.safetensors.index.json") safetensors_path = os.path.join(ckpt_dir, "model.safetensors") + model_state = model.state_dict() + loaded_base_keys = set() + skipped_base_keys = set() + + def load_compatible_base_weights(weights): + compatible = {} + for key, value in weights.items(): + if key in model_state and model_state[key].shape == value.shape: + compatible[key] = value + loaded_base_keys.add(key) + else: + skipped_base_keys.add(key) + model.load_state_dict(compatible, strict=False) + if os.path.exists(safetensors_index_path): with open(safetensors_index_path, 'r') as f: index = json.load(f) @@ -708,23 +790,74 @@ def create_model(self, cfg, training_args): shard_path = os.path.join(ckpt_dir, shard_file) mprint(f"Loading shard: {shard_path}") shard_state_dict = load_file(shard_path) - model.load_state_dict(shard_state_dict, strict=False) + load_compatible_base_weights(shard_state_dict) del shard_state_dict gc.collect() elif os.path.exists(safetensors_path): state_dict = load_file(safetensors_path) - model.load_state_dict(state_dict, strict=False) + load_compatible_base_weights(state_dict) else: raise FileNotFoundError( f"No weights found at '{ckpt_dir}'. " "Expected 'model.safetensors' or 'model.safetensors.index.json'." ) + if getattr( + model.action_head.config, + "action_only_training", + False, + ): + adapter_prefixes = ( + "action_head.model.state_encoder.", + "action_head.model.action_encoder.", + "action_head.model.action_decoder.", + ) + required_shared_dit_keys = { + key + for key in model_state + if key.startswith("action_head.model.") + and not key.startswith(adapter_prefixes) + } + missing_shared_dit = sorted( + required_shared_dit_keys - loaded_base_keys + ) + if missing_shared_dit: + raise RuntimeError( + "Incomplete DreamZero-AgiBot shared DiT load: " + + ", ".join(missing_shared_dit[:20]) + ) + mprint( + "[BASE LOAD] DreamZero-AgiBot loaded | " + f"shared_dit_keys={len(required_shared_dit_keys)} " + f"skipped_incompatible={len(skipped_base_keys)}" + ) + mprint( + "[EMBODIMENT] g2 logical_id=33 " + "action_adapter_local_slot=0" + ) + if (hasattr(model, 'action_head') and hasattr(model.action_head, 'inject_lora_after_loading') and model.action_head.config.defer_lora_injection): model.action_head.inject_lora_after_loading() + if cfg.pretrained_lora_path is not None: + mprint( + f"Loading pretrained LoRA/action weights from: " + f"{cfg.pretrained_lora_path}" + ) + model.load_lora_weight(cfg.pretrained_lora_path) + + if ( + hasattr(model, "action_head") + and getattr( + model.action_head.config, + "action_only_training", + False, + ) + ): + model.action_head.configure_action_only_training() + mprint("Successfully loaded pretrained weights") model.config.resume_path = model.config._name_or_path = training_args.output_dir diff --git a/groot/vla/experiment/experiment.py b/groot/vla/experiment/experiment.py index e02d36c5..f76d9f69 100644 --- a/groot/vla/experiment/experiment.py +++ b/groot/vla/experiment/experiment.py @@ -7,6 +7,7 @@ import numpy as np from omegaconf import DictConfig import torch +from transformers import TrainerCallback from groot.vla.experiment.base import BaseExperiment, BaseTrainer from groot.vla.utils.action_args_override_utils import apply_action_overrides @@ -17,6 +18,18 @@ INITIAL_ACTIONS_FILENAME = "initial_actions.npz" +class MilestoneSaveCallback(TrainerCallback): + """Request checkpoints only at configured optimizer-step milestones.""" + + def __init__(self, milestones): + self.milestones = frozenset(int(step) for step in milestones) + + def on_step_end(self, args, state, control, **kwargs): + if state.global_step in self.milestones: + control.should_save = True + return control + + class ForceRestart(ValueError): pass @@ -35,8 +48,16 @@ def __init__(self, **kwargs): self.rank = dist.get_rank() self.micro_global_step = 0 + milestones = kwargs.pop("milestone_save_steps", None) super().__init__(**kwargs) + if milestones: + self.add_callback(MilestoneSaveCallback(milestones)) + if self.rank == 0: + print( + "[CHECKPOINT MILESTONES] " + + ",".join(str(int(step)) for step in milestones) + ) def training_step(self, model, inputs, *args, **kwargs): self.micro_global_step += 1 diff --git a/groot/vla/model/dreamzero/action_head/wan_flow_matching_action_tf.py b/groot/vla/model/dreamzero/action_head/wan_flow_matching_action_tf.py index db8fe0f8..27e5e7e0 100644 --- a/groot/vla/model/dreamzero/action_head/wan_flow_matching_action_tf.py +++ b/groot/vla/model/dreamzero/action_head/wan_flow_matching_action_tf.py @@ -128,6 +128,60 @@ class WANPolicyHeadConfig(PretrainedConfig): tune_diffusion_model: bool = field( default=True, metadata={"help": "Whether to tune the diffusion model."} ) + action_only_training: bool = field( + default=False, + metadata={"help": "Freeze the shared DiT and train only state/action adapters."}, + ) + dynamics_loss_weight: float = field( + default=1.0, metadata={"help": "Multiplier for the video dynamics loss."} + ) + action_loss_weight: float = field( + default=1.0, metadata={"help": "Multiplier for the action flow-matching loss."} + ) + action_x0_loss_weight: float = field( + default=0.0, + metadata={ + "help": "Multiplier for direct decoded-action reconstruction loss." + }, + ) + action_endpoint_loss_weight: float = field( + default=0.0, + metadata={ + "help": "Multiplier for direct final-step action reconstruction loss." + }, + ) + action_start_loss_weight: float = field( + default=0.0, + metadata={ + "help": "Multiplier for the current-state start consistency loss " + "(first action of each chunk should be close to the current joint " + "state). Soft anchor: state is a training target, not a model input." + }, + ) + action_video_consistency_loss_weight: float = field( + default=0.0, + metadata={ + "help": "Multiplier for the video->action consistency loss " + "(actions recovered from the predicted future video latent should " + "match the action chunk). Increases the action's dependence on the " + "video/dynamics representation. Keep small initially to avoid " + "disrupting flow-matching." + }, + ) + motion_mask_enabled: bool = field( + default=False, + metadata={"help": "Mask stationary action timesteps out of the action loss."}, + ) + motion_mask_threshold_rad: float = field( + default=0.03, + metadata={"help": "Joint movement threshold in radians for motion masking."}, + ) + motion_mask_stats_path: str | None = field( + default=None, + metadata={ + "help": "Relative-action stats used to convert normalized action deltas to radians." + }, + ) load_pretrained_det_decode_layer_path: str = field( default=None, metadata={"help": "Path to pretrained detection model."} ) @@ -238,6 +292,13 @@ def __init__( self.model = instantiate(config.diffusion_model_cfg) self.action_dim = config.action_dim self.action_horizon = config.action_horizon + + # Optional video->action consistency head (created lazily on first + # forward once the video-latent channel count is known). Recovering + # actions from the predicted future video latent makes the action + # depend on the learned implicit dynamics. + self.video_to_action_head: nn.Module | None = None + self._video_head_latent_channels: int | None = None self.num_inference_timesteps = config.num_inference_timesteps text_enc_path = ensure_file( @@ -313,10 +374,129 @@ def __init__( # self.num_timestep_buckets = config.num_timestep_buckets self.config = config self._noise_logged = False + self._loss_contract_logged = False + self._motion_mask_logged = False self.defer_lora_injection = config.defer_lora_injection + self._motion_mask_arm_ranges = self._load_motion_mask_arm_ranges(config) print("defer_lora_injection@@", self.defer_lora_injection) self.set_trainable_parameters(config.tune_projector, config.tune_diffusion_model) + @staticmethod + def _load_motion_mask_arm_ranges(config: WANPolicyHeadConfig) -> torch.Tensor | None: + """Load q99 normalization ranges for the 14 arm joints. + + The action tensor reaching this head is q99-normalized. For a q99 + transform, normalized_delta * (q99-q01) / 2 is the original action + delta, so the configured radian threshold remains meaningful without + changing the dataset or model inputs. + """ + if not config.motion_mask_enabled: + return None + + stats_path = config.motion_mask_stats_path + if not stats_path: + raise ValueError( + "motion_mask_stats_path is required when motion_mask_enabled=true" + ) + if not os.path.isfile(stats_path): + raise FileNotFoundError( + f"Motion-mask stats file does not exist: {stats_path}" + ) + + with open(stats_path, "r", encoding="utf-8") as f: + stats = json.load(f) + + ranges = [] + for key in ("left_joint_position", "right_joint_position"): + key_stats = stats.get(key) + if key_stats is None: + raise KeyError(f"Missing {key} in motion-mask stats: {stats_path}") + q01 = key_stats.get("q01") + q99 = key_stats.get("q99") + if q01 is None or q99 is None or len(q01) != 7 or len(q99) != 7: + raise ValueError( + f"Expected 7-dimensional q01/q99 stats for {key}: {stats_path}" + ) + ranges.extend(float(high) - float(low) for low, high in zip(q01, q99)) + + ranges_tensor = torch.tensor(ranges, dtype=torch.float32) + if not torch.all(ranges_tensor > 0): + raise ValueError( + f"Motion-mask q99 ranges must be positive, got {ranges_tensor.tolist()}" + ) + print( + "[MOTION MASK] enabled=true " + f"threshold_rad={config.motion_mask_threshold_rad:.4f} " + f"stats={stats_path}" + ) + return ranges_tensor + + def _build_motion_loss_mask(self, actions: torch.Tensor) -> torch.Tensor: + """Return a [B, T] mask for action timesteps involved in motion. + + Actions are normalized before reaching this head, so use the loaded + q99 ranges to evaluate the 14-arm-joint delta in radians. The action + stream can contain multiple 24-step chunks; deltas never cross a + chunk boundary. + """ + if self._motion_mask_arm_ranges is None: + raise RuntimeError("Motion-mask ranges were not initialized") + if actions.ndim != 3: + raise RuntimeError(f"Expected actions with shape [B,T,D], got {actions.shape}") + + batch_size, action_steps, _ = actions.shape + if action_steps % self.action_horizon != 0: + raise RuntimeError( + f"Motion mask requires action length divisible by horizon " + f"{self.action_horizon}, got {action_steps}" + ) + + arm_indices = torch.tensor( + [0, 1, 2, 3, 4, 5, 6, 8, 9, 10, 11, 12, 13, 14], + device=actions.device, + ) + num_chunks = action_steps // self.action_horizon + arm_actions = actions.index_select(dim=2, index=arm_indices).reshape( + batch_size, num_chunks, self.action_horizon, 14 + ) + delta_normalized = arm_actions[:, :, 1:] - arm_actions[:, :, :-1] + ranges = self._motion_mask_arm_ranges.to( + device=actions.device, dtype=actions.dtype + ).reshape(1, 1, 1, 14) + delta_rad = delta_normalized * ranges / 2.0 + max_joint_move = delta_rad.abs().max(dim=-1).values + moving_transition = max_joint_move > float(self.config.motion_mask_threshold_rad) + + # A transition contributes to both endpoints: this preserves the + # target immediately before movement and the target at the new pose. + motion_mask = torch.zeros( + batch_size, + num_chunks, + self.action_horizon, + dtype=actions.dtype, + device=actions.device, + ) + motion_mask[:, :, :-1] = moving_transition.to(dtype=actions.dtype) + motion_mask[:, :, 1:] = torch.maximum( + motion_mask[:, :, 1:], moving_transition.to(dtype=actions.dtype) + ) + motion_mask = motion_mask.reshape(batch_size, action_steps) + + if not self._motion_mask_logged: + active_steps = motion_mask.sum().detach() + total_steps = torch.tensor( + motion_mask.numel(), dtype=motion_mask.dtype, device=motion_mask.device + ) + print( + "[MOTION MASK] action_steps=" + f"{action_steps} chunks={num_chunks} " + f"active_ratio={float(active_steps / total_steps.clamp_min(1.0)):.4f} " + f"active_steps={int(active_steps.item())}/{motion_mask.numel()}" + ) + self._motion_mask_logged = True + + return motion_mask + def reset_inference_cache(self) -> None: self.kv_cache1 = None self.kv_cache_neg = None @@ -344,7 +524,9 @@ def set_trainable_parameters(self, tune_projector: bool, tune_diffusion_model: b if not any(p.requires_grad for p in self.parameters()): print("Warning: No action head trainable parameters found.") - if self.train_architecture == "lora" and not self.defer_lora_injection: + if self.train_architecture == "action_only": + self.configure_action_only_training() + elif self.train_architecture == "lora" and not self.defer_lora_injection: print("Adding LoRA to model") for p in self.parameters(): p.requires_grad = False @@ -416,6 +598,53 @@ def inject_lora_after_loading(self): else: print("LoRA injection not needed (train_architecture != 'lora')") + def configure_action_only_training(self): + """Freeze the shared video DiT and tune only the three action adapters.""" + if self.train_architecture != "action_only": + raise RuntimeError( + "action_only_training requires train_architecture=action_only, " + f"got {self.train_architecture!r}" + ) + lora_parameters = [ + name for name, _ in self.model.named_parameters() + if "lora_" in name + ] + if lora_parameters: + raise RuntimeError( + "Action-adapter-only training forbids shared DiT LoRA; " + f"found {len(lora_parameters)} LoRA parameters" + ) + for parameter in self.parameters(): + parameter.requires_grad = False + self.model.state_encoder.requires_grad_(True) + self.model.action_encoder.requires_grad_(True) + self.model.action_decoder.requires_grad_(True) + self.text_encoder.requires_grad_(False) + self.image_encoder.requires_grad_(False) + self.vae.requires_grad_(False) + allowed_fragments = ( + "action_head.model.state_encoder.", + "action_head.model.action_encoder.", + "action_head.model.action_decoder.", + ) + unexpected = [ + name + for name, parameter in self.named_parameters() + if parameter.requires_grad + and not any(fragment in f"action_head.{name}" for fragment in allowed_fragments) + ] + if unexpected: + raise RuntimeError( + "Action-adapter trainable whitelist violation: " + + ", ".join(unexpected[:20]) + ) + print( + "[TRAINABLE WHITELIST] state_encoder, action_encoder, " + "action_decoder only" + ) + print("[SHARED DIT LORA] disabled") + self.print_trainable_params() + def set_frozen_modules_to_eval_mode(self): """ Huggingface will call model.train() at each training_step. To ensure @@ -761,17 +990,298 @@ def forward(self, backbone_output: BatchFeature, action_input: BatchFeature) -> weighted_dynamics_loss = weight_dynamics.mean() if actions.numel() > 0: - action_loss_per_sample = torch.nn.functional.mse_loss( + action_loss_per_element = torch.nn.functional.mse_loss( action_noise_pred.float(), training_target_action.float(), reduction='none' - ) * action_mask # shape: [B, ...] - action_loss_per_sample = has_real_action[:, None].float() * action_loss_per_sample # apply has_real_action - weight_action = action_loss_per_sample.mean(dim=2) * self.scheduler.training_weight( + ) + valid_action_mask = action_mask.to( + dtype=action_loss_per_element.dtype, + device=action_loss_per_element.device, + ) + # Do not divide G2's 16 valid dimensions by max_action_dim=32. + # The previous masked-then-mean reduction silently halved the + # G2 action objective because its padded dimensions are zero. + valid_dim_count = valid_action_mask.sum(dim=2).clamp_min(1.0) + valid_dims = int(valid_dim_count.max().item()) + if self.config.action_only_training and valid_dims != 16: + raise RuntimeError( + "G2 action adapter requires exactly 16 valid action " + f"dimensions, got {valid_dims}" + ) + + arm_indices = torch.tensor( + [0, 1, 2, 3, 4, 5, 6, 8, 9, 10, 11, 12, 13, 14], + device=action_loss_per_element.device, + ) + gripper_indices = torch.tensor( + [7, 15], + device=action_loss_per_element.device, + ) + arm_element_loss = action_loss_per_element.index_select( + dim=2, index=arm_indices + ) + arm_mask = valid_action_mask.index_select( + dim=2, index=arm_indices + ) + gripper_element_loss = action_loss_per_element.index_select( + dim=2, index=gripper_indices + ) + gripper_mask = valid_action_mask.index_select( + dim=2, index=gripper_indices + ) + arm_loss = ( + arm_element_loss * arm_mask + ).sum(dim=2) / arm_mask.sum(dim=2).clamp_min(1.0) + gripper_loss = ( + gripper_element_loss * gripper_mask + ).sum(dim=2) / gripper_mask.sum(dim=2).clamp_min(1.0) + action_loss_per_timestep = ( + (14.0 / 16.0) * arm_loss + + (2.0 / 16.0) * gripper_loss + ) + action_loss_per_timestep = ( + has_real_action[:, None].float() + * action_loss_per_timestep + ) + if not self._loss_contract_logged: + old_reduction = ( + action_loss_per_element * valid_action_mask + ).mean(dim=2).mean().detach() + corrected_reduction = ( + action_loss_per_timestep.mean().detach() + ) + ratio = corrected_reduction / old_reduction.clamp_min(1e-12) + print( + "[ACTION LOSS CONTRACT] padded_action_dim=" + f"{action_loss_per_element.shape[2]} " + f"valid_action_dim={valid_dims} " + f"corrected_loss/old_loss={float(ratio):.4f}" + ) + self._loss_contract_logged = True + weight_action = action_loss_per_timestep * self.scheduler.training_weight( timestep_action.flatten(0, 1), ).unflatten(0, (noise_action.shape[0], noise_action.shape[1])).to(self._device) - weighted_action_loss = weight_action.mean() - loss = weighted_dynamics_loss + weighted_action_loss + action_time_mask = has_real_action[:, None].float() + if self.config.motion_mask_enabled: + action_time_mask = action_time_mask * self._build_motion_loss_mask(actions) + weighted_action_loss = ( + weight_action * action_time_mask + ).sum() / action_time_mask.sum().clamp_min(1.0) + else: + weighted_action_loss = weight_action.mean() + + # The flow-matching objective supervises the velocity/noise + # field. Add an optional direct x0 objective so the decoded + # action trajectory and its final pose are also constrained. + action_x0_loss = torch.tensor(0.0, device=self._device) + action_endpoint_loss = torch.tensor(0.0, device=self._device) + action_start_loss = torch.tensor(0.0, device=self._device) + if ( + float(self.config.action_x0_loss_weight) > 0.0 + or float(self.config.action_endpoint_loss_weight) > 0.0 + or float(self.config.action_start_loss_weight) > 0.0 + ): + sigma_action = self.scheduler.sigmas.to( + device=noisy_actions.device, dtype=noisy_actions.dtype + )[timestep_action_id.to(device=noisy_actions.device)].unsqueeze(-1) + decoded_actions = ( + noisy_actions.float() + - sigma_action.float() * action_noise_pred.float() + ) + direct_action_element_loss = torch.nn.functional.smooth_l1_loss( + decoded_actions, + actions.float(), + beta=0.05, + reduction="none", + ) + direct_arm_loss = ( + direct_action_element_loss.index_select(dim=2, index=arm_indices) + * valid_action_mask.index_select(dim=2, index=arm_indices) + ).sum(dim=2) / valid_action_mask.index_select( + dim=2, index=arm_indices + ).sum(dim=2).clamp_min(1.0) + direct_gripper_loss = ( + direct_action_element_loss.index_select( + dim=2, index=gripper_indices + ) + * valid_action_mask.index_select(dim=2, index=gripper_indices) + ).sum(dim=2) / valid_action_mask.index_select( + dim=2, index=gripper_indices + ).sum(dim=2).clamp_min(1.0) + direct_action_per_timestep = ( + (14.0 / 16.0) * direct_arm_loss + + (2.0 / 16.0) * direct_gripper_loss + ) + direct_action_per_timestep = ( + has_real_action[:, None].float() + * direct_action_per_timestep + ) + action_x0_loss = ( + direct_action_per_timestep * action_time_mask + ).sum() / action_time_mask.sum().clamp_min(1.0) + + if float(self.config.action_endpoint_loss_weight) > 0.0: + if actions.shape[1] % self.action_horizon != 0: + raise RuntimeError( + "Endpoint action loss requires action length divisible " + f"by horizon {self.action_horizon}, got {actions.shape[1]}" + ) + num_chunks = actions.shape[1] // self.action_horizon + decoded_chunks = decoded_actions.reshape( + actions.shape[0], num_chunks, self.action_horizon, -1 + ) + target_chunks = actions.float().reshape( + actions.shape[0], num_chunks, self.action_horizon, -1 + ) + valid_chunks = valid_action_mask.reshape( + actions.shape[0], num_chunks, self.action_horizon, -1 + ) + endpoint_element_loss = torch.nn.functional.smooth_l1_loss( + decoded_chunks[:, :, -1], + target_chunks[:, :, -1], + beta=0.05, + reduction="none", + ) + endpoint_valid = valid_chunks[:, :, -1] + endpoint_arm_loss = ( + endpoint_element_loss.index_select(dim=2, index=arm_indices) + * endpoint_valid.index_select(dim=2, index=arm_indices) + ).sum(dim=2) / endpoint_valid.index_select( + dim=2, index=arm_indices + ).sum(dim=2).clamp_min(1.0) + endpoint_gripper_loss = ( + endpoint_element_loss.index_select( + dim=2, index=gripper_indices + ) + * endpoint_valid.index_select(dim=2, index=gripper_indices) + ).sum(dim=2) / endpoint_valid.index_select( + dim=2, index=gripper_indices + ).sum(dim=2).clamp_min(1.0) + endpoint_per_chunk = ( + (14.0 / 16.0) * endpoint_arm_loss + + (2.0 / 16.0) * endpoint_gripper_loss + ) + endpoint_chunk_mask = action_time_mask.reshape( + actions.shape[0], num_chunks, self.action_horizon + ).amax(dim=2) + action_endpoint_loss = ( + endpoint_per_chunk * endpoint_chunk_mask + ).sum() / endpoint_chunk_mask.sum().clamp_min(1.0) + + # Current-state start consistency (L_start): the FIRST predicted + # action should be close to the current joint state. state is a + # TRAINING TARGET here, not a model input, so this anchors the + # model's trajectory to the true start pose without re-enabling + # the state-input shortcut. Only the sequence-start state is + # available, so the anchor applies to the first action step. + if float(self.config.action_start_loss_weight) > 0.0: + # action is padded to action_dim (32), state to + # max_state_dim (64); compare only the valid G2 dims (16). + start_pred = decoded_actions[:, 0, :16] # (B, 16) + state_0 = state_features[:, 0, :16] # (B, 16) + start_element = torch.nn.functional.smooth_l1_loss( + start_pred.float(), state_0.float(), beta=0.05, reduction="none" + ) + start_valid = valid_action_mask[:, 0] + start_arm = ( + start_element.index_select(dim=1, index=arm_indices) + * start_valid.index_select(dim=1, index=arm_indices) + ).sum(dim=1) / start_valid.index_select( + dim=1, index=arm_indices + ).sum(dim=1).clamp_min(1.0) + start_grip = ( + start_element.index_select(dim=1, index=gripper_indices) + * start_valid.index_select(dim=1, index=gripper_indices) + ).sum(dim=1) / start_valid.index_select( + dim=1, index=gripper_indices + ).sum(dim=1).clamp_min(1.0) + start_per_sample = ( + (14.0 / 16.0) * start_arm + (2.0 / 16.0) * start_grip + ) + action_start_loss = ( + start_per_sample * has_real_action.float() + ).sum() / has_real_action.sum().clamp_min(1.0) + + # Video->action consistency (L_video_to_action): recover the + # action from the PREDICTED clean video latent so the action + # chunk depends on the learned implicit dynamics encoded in the + # video. Implemented defensively: if the latent/head shapes are + # not as expected the term is skipped (logged) instead of + # crashing the run. + action_video_consistency_loss = torch.tensor(0.0, device=self._device) + if float(self.config.action_video_consistency_loss_weight) > 0.0 and actions.numel() > 0: + try: + sigma_video = self.scheduler.sigmas.to( + device=noisy_latents.device, dtype=noisy_latents.dtype + )[timestep_id.to(device=noisy_latents.device)] + while sigma_video.ndim < noisy_latents.ndim: + sigma_video = sigma_video.unsqueeze(-1) + decoded_latents = ( + noisy_latents.float() - sigma_video.float() * video_noise_pred.float() + ) + # pool spatial + time, keep channel dim (dim 1) + pooled = decoded_latents.mean(dim=tuple(range(2, decoded_latents.ndim))) + if self.video_to_action_head is None: + self._video_head_latent_channels = int(pooled.shape[1]) + self.video_to_action_head = nn.Linear( + self._video_head_latent_channels, self.action_dim + ).to(device=pooled.device, dtype=pooled.dtype) + pred_action = self.video_to_action_head(pooled.float()) + vid_element = torch.nn.functional.mse_loss( + pred_action, actions[:, 0].float(), reduction="none" + ) + vid_valid = valid_action_mask[:, 0] + vid_arm = ( + vid_element.index_select(dim=1, index=arm_indices) + * vid_valid.index_select(dim=1, index=arm_indices) + ).sum(dim=1) / vid_valid.index_select( + dim=1, index=arm_indices + ).sum(dim=1).clamp_min(1.0) + vid_grip = ( + vid_element.index_select(dim=1, index=gripper_indices) + * vid_valid.index_select(dim=1, index=gripper_indices) + ).sum(dim=1) / vid_valid.index_select( + dim=1, index=gripper_indices + ).sum(dim=1).clamp_min(1.0) + vid_per_sample = ( + (14.0 / 16.0) * vid_arm + (2.0 / 16.0) * vid_grip + ) + action_video_consistency_loss = ( + vid_per_sample * has_real_action.float() + ).sum() / has_real_action.sum().clamp_min(1.0) + except Exception as exc: # noqa: BLE001 - defensive skip + logging.warning( + "video->action consistency loss skipped (%s); " + "check latent/head shapes.", + exc, + ) + + action_objective = ( + float(self.config.action_loss_weight) * weighted_action_loss + + float(self.config.action_x0_loss_weight) * action_x0_loss + + float(self.config.action_endpoint_loss_weight) + * action_endpoint_loss + + float(self.config.action_start_loss_weight) * action_start_loss + + float(self.config.action_video_consistency_loss_weight) + * action_video_consistency_loss + ) + if self.config.action_only_training: + loss = action_objective + weighted_dynamics_loss = weighted_dynamics_loss.detach() + else: + loss = ( + float(self.config.dynamics_loss_weight) + * weighted_dynamics_loss + + action_objective + ) else: weighted_action_loss = torch.tensor(0.0, device=self._device) + action_x0_loss = torch.tensor(0.0, device=self._device) + action_endpoint_loss = torch.tensor(0.0, device=self._device) + if self.config.action_only_training: + raise RuntimeError( + "Action-adapter-only training received an empty action " + "tensor; refusing to optimize video loss" + ) loss = weighted_dynamics_loss # loss = dynamics_loss_per_sample.mean() @@ -780,6 +1290,10 @@ def forward(self, backbone_output: BatchFeature, action_input: BatchFeature) -> "loss": loss, "dynamics_loss": weighted_dynamics_loss, "action_loss": weighted_action_loss, + "action_x0_loss": action_x0_loss, + "action_endpoint_loss": action_endpoint_loss, + "action_start_loss": action_start_loss, + "action_video_consistency_loss": action_video_consistency_loss, } return BatchFeature(data=output_dict) diff --git a/groot/vla/model/dreamzero/base_vla.py b/groot/vla/model/dreamzero/base_vla.py index d81401a4..f89dd185 100644 --- a/groot/vla/model/dreamzero/base_vla.py +++ b/groot/vla/model/dreamzero/base_vla.py @@ -102,6 +102,106 @@ def validate_inputs(self, inputs): if detected_error: raise ValueError(error_msg) + @classmethod + def load_action_adapter( + cls, + adapter_model_path: str, + pretrained_base_model_path: str, + config: VLAConfig | None = None, + ): + """Load a clean DreamZero base and then a strict G2 action adapter.""" + import json + import os + from safetensors.torch import load_file + + if config is None: + config = cls.config_class.from_pretrained(adapter_model_path) + model = cls(config) + model_state = model.state_dict() + adapter_prefixes = ( + "action_head.model.state_encoder.", + "action_head.model.action_encoder.", + "action_head.model.action_decoder.", + ) + + index_path = os.path.join( + pretrained_base_model_path, + "model.safetensors.index.json", + ) + single_path = os.path.join( + pretrained_base_model_path, + "model.safetensors", + ) + if os.path.exists(index_path): + with open(index_path) as stream: + index = json.load(stream) + base_files = [ + os.path.join(pretrained_base_model_path, filename) + for filename in sorted(set(index["weight_map"].values())) + ] + elif os.path.exists(single_path): + base_files = [single_path] + else: + raise FileNotFoundError( + f"No DreamZero base weights at {pretrained_base_model_path}" + ) + + loaded_base_keys: set[str] = set() + skipped_base_keys: set[str] = set() + for base_file in base_files: + shard = load_file(base_file) + compatible = {} + for key, value in shard.items(): + if key in model_state and model_state[key].shape == value.shape: + compatible[key] = value + loaded_base_keys.add(key) + else: + skipped_base_keys.add(key) + model.load_state_dict(compatible, strict=False) + + required_shared_keys = { + key + for key in model_state + if not key.startswith(adapter_prefixes) + } + missing_shared = sorted(required_shared_keys - loaded_base_keys) + if missing_shared: + raise RuntimeError( + "Incomplete shared DreamZero base load; missing shared keys: " + + ", ".join(missing_shared[:20]) + ) + print( + "[BASE LOAD] DreamZero-AgiBot loaded " + f"shared_keys={len(required_shared_keys)} " + f"skipped_incompatible={len(skipped_base_keys)}" + ) + print("[EMBODIMENT] g2 (logical id), action-adapter local slot=0") + + adapter_path = os.path.join( + adapter_model_path, + "action_expert.safetensors", + ) + adapter_state = load_file(adapter_path) + expected_adapter_keys = { + key for key in model_state if key.startswith(adapter_prefixes) + } + if set(adapter_state) != expected_adapter_keys: + missing = sorted(expected_adapter_keys - set(adapter_state)) + unexpected = sorted(set(adapter_state) - expected_adapter_keys) + raise RuntimeError( + "Action adapter key contract mismatch: " + f"missing={missing[:20]} unexpected={unexpected[:20]}" + ) + model.load_state_dict(adapter_state, strict=False) + for module_name in ( + "state_encoder", + "action_encoder", + "action_decoder", + ): + print(f"[ACTION ADAPTER] {module_name} loaded") + print("[SHARED DIT LORA] disabled") + return model + def validate_data(self, action_head_outputs, backbone_outputs, is_training): fail_backbone = ( @@ -336,78 +436,170 @@ def from_pretrained_for_tuning( @classmethod def load_lora( - cls, - pretrained_model_name_or_path: str - ): + cls, + pretrained_model_name_or_path: str, + pretrained_base_model_path: str | None = None, + ): + """Load a LoRA checkpoint on the exact full-model training base. + + A LoRA delta is not a standalone model. G2 training first instantiated + the checkpoint architecture, loaded ``DreamZero-AgiBot``, and only then + injected LoRA adapters. Inference must mirror that order. Loading the + raw Wan component and applying the LoRA delta to it produces a + numerically valid but semantically wrong model. + """ from safetensors.torch import load_file import os import json - print("loading lora@@@@@") + import gc + + if pretrained_base_model_path is None: + raise ValueError( + "A LoRA-only checkpoint requires pretrained_base_model_path. " + "Use the same full checkpoint recorded as " + "pretrained_model_path in experiment_cfg/conf.yaml." + ) + if not os.path.isdir(pretrained_base_model_path): + raise FileNotFoundError( + f"LoRA base model directory does not exist: " + f"{pretrained_base_model_path}" + ) + + print( + "Loading LoRA checkpoint " + f"{pretrained_model_name_or_path} on base " + f"{pretrained_base_model_path}" + ) - # Check for different checkpoint formats safetensors_path = os.path.join(pretrained_model_name_or_path, "model.safetensors") safetensors_index_path = os.path.join(pretrained_model_name_or_path, "model.safetensors.index.json") - - state_dict = {} + + lora_state_dict = {} if os.path.exists(safetensors_index_path): - # Handle sharded safetensors - print(f"Loading sharded safetensors using index: {safetensors_index_path}") - with open(safetensors_index_path, 'r') as f: index = json.load(f) - - # Load each shard - for shard_file in set(index["weight_map"].values()): + for shard_file in sorted(set(index["weight_map"].values())): shard_path = os.path.join(pretrained_model_name_or_path, shard_file) - print(f"Loading shard: {shard_path}") - shard_state_dict = load_file(shard_path) - state_dict.update(shard_state_dict) - + lora_state_dict.update(load_file(shard_path)) elif os.path.exists(safetensors_path): - # Handle single safetensors file - print(f"Loading weights from safetensors: {safetensors_path}") - state_dict.update(load_file(safetensors_path)) - - # Load config - print("loading config@@") + lora_state_dict.update(load_file(safetensors_path)) + else: + raise FileNotFoundError( + f"No LoRA weights found at {pretrained_model_name_or_path}" + ) + config_path = os.path.join(pretrained_model_name_or_path, "config.json") with open(config_path, "r") as f: config_dict = json.load(f) config = VLAConfig(**config_dict) - print("loading model") - # Disable defer_lora_injection so LoRA layers are created during init, - # matching the PEFT key hierarchy (base_model.model.*) in the checkpoint. + # Mirror training: build the target G2 architecture without loading raw + # Wan component weights, load the full DreamZero base, then inject LoRA. ah_cfg = config.action_head_cfg inner = ah_cfg.get('config', ah_cfg) if isinstance(ah_cfg.get('config'), dict) else ah_cfg - if 'defer_lora_injection' in inner: - inner['defer_lora_injection'] = False - print("defer_lora_injection disabled for load_lora") - # Enable component loading so DiT base weights are loaded from pretrained - if 'skip_component_loading' in inner: - inner['skip_component_loading'] = False - print("skip_component_loading disabled for load_lora") - - # Instantiate model (LoRA layers now exist from init) + inner['defer_lora_injection'] = True + inner['skip_component_loading'] = True model = cls(config) - # Remove .base_layer from keys if present - has_base_layer = any(".base_layer." in key for key in state_dict.keys()) + base_index_path = os.path.join( + pretrained_base_model_path, + "model.safetensors.index.json", + ) + base_single_path = os.path.join( + pretrained_base_model_path, + "model.safetensors", + ) + loaded_base_keys: set[str] = set() + unexpected_base_keys: set[str] = set() + + def load_base_state_dict(base_state_dict: dict) -> None: + model_keys = set(model.state_dict()) + loaded_base_keys.update(set(base_state_dict) & model_keys) + _, unexpected = model.load_state_dict( + base_state_dict, + strict=False, + ) + unexpected_base_keys.update(unexpected) + + if os.path.exists(base_index_path): + with open(base_index_path, "r") as f: + base_index = json.load(f) + for shard_file in sorted(set(base_index["weight_map"].values())): + shard_path = os.path.join( + pretrained_base_model_path, + shard_file, + ) + print(f"Loading DreamZero base shard: {shard_path}") + base_state_dict = load_file(shard_path) + load_base_state_dict(base_state_dict) + del base_state_dict + gc.collect() + elif os.path.exists(base_single_path): + base_state_dict = load_file(base_single_path) + load_base_state_dict(base_state_dict) + del base_state_dict + gc.collect() + else: + raise FileNotFoundError( + "No full base weights found at " + f"{pretrained_base_model_path}" + ) + + if not loaded_base_keys: + raise RuntimeError( + "The full DreamZero base checkpoint did not match any " + "parameters in the G2 model architecture" + ) + if unexpected_base_keys: + print( + "Ignoring base-only keys that are not present in the G2 " + f"architecture: count={len(unexpected_base_keys)}" + ) + + if not ( + hasattr(model, "action_head") + and hasattr(model.action_head, "inject_lora_after_loading") + ): + raise RuntimeError( + "G2 action head does not support deferred LoRA injection" + ) + model.action_head.inject_lora_after_loading() + + has_base_layer = any( + ".base_layer." in key for key in lora_state_dict + ) if has_base_layer: print("Removing '.base_layer' from state dict keys") - state_dict = {k.replace(".base_layer.", "."): v for k, v in state_dict.items()} + lora_state_dict = { + key.replace(".base_layer.", "."): value + for key, value in lora_state_dict.items() + } + + model_keys = set(model.state_dict()) + unknown_lora_keys = set(lora_state_dict) - model_keys + if unknown_lora_keys: + sample = sorted(unknown_lora_keys)[:20] + raise RuntimeError( + "LoRA checkpoint keys do not match the model constructed on " + f"the DreamZero base: count={len(unknown_lora_keys)}, " + f"sample={sample}" + ) - # Load weights - missing_keys, unexpected_keys = model.load_state_dict(state_dict, strict=False) - - if missing_keys: - print(f"Missing keys when loading pretrained weights: {missing_keys}") + missing_keys, unexpected_keys = model.load_state_dict( + lora_state_dict, + strict=False, + ) if unexpected_keys: - print(f"Unexpected keys when loading pretrained weights: {unexpected_keys}") - - print("Successfully loaded pretrained weights") + raise RuntimeError( + f"Unexpected LoRA keys after validation: {unexpected_keys}" + ) - print(f"{cls}\n") + print( + "Successfully loaded full DreamZero base and LoRA delta | " + f"base_keys={len(loaded_base_keys)} " + f"lora_keys={len(lora_state_dict)} " + f"model_only_keys={len(missing_keys)}" + ) return model def load_lora_weight(self, pretrained_model_name_or_path: str): diff --git a/groot/vla/model/dreamzero/modules/train_dreamzero_g2_absolute_lora.sh b/groot/vla/model/dreamzero/modules/train_dreamzero_g2_absolute_lora.sh new file mode 100644 index 00000000..9e8bbe18 --- /dev/null +++ b/groot/vla/model/dreamzero/modules/train_dreamzero_g2_absolute_lora.sh @@ -0,0 +1,184 @@ +#!/usr/bin/env bash +set -Eeuo pipefail + +PROJECT_ROOT=${PROJECT_ROOT:-/home/ubuntu/projects/wangk/dreamzero} +G2_DATA_ROOT=${G2_DATA_ROOT:-/data/training_data/teleop/g2/g2_mock_light_module_joint_gear_policy_gripper/train} +OUTPUT_DIR=${OUTPUT_DIR:-/data/wangk/checkpoints/dreamzero_g2_absolute_video1_action10_v1} +WAN_CKPT_DIR=${WAN_CKPT_DIR:-/data/wangk/checkpoints/Wan2.1-I2V-14B-480P} +TOKENIZER_DIR=${TOKENIZER_DIR:-/data/wangk/checkpoints/umt5-xxl} +PRETRAINED_MODEL_PATH=${PRETRAINED_MODEL_PATH:-/data/wangk/checkpoints/DreamZero-AgiBot} +PYTHON_BIN=${PYTHON_BIN:-/data/wangk/conda/envs/dreamzero/bin/python} + +GPU_IDS=${GPU_IDS:-0,1,2,3} +EXPECTED_EPISODES=${EXPECTED_EPISODES:-110} +MAX_STEPS=${MAX_STEPS:-4000} +SAVE_STEPS=${SAVE_STEPS:-500} +WANDB_MODE=${WANDB_MODE:-offline} +HYDRA_FULL_ERROR=${HYDRA_FULL_ERROR:-1} + +export G2_DATA_ROOT OUTPUT_DIR WAN_CKPT_DIR TOKENIZER_DIR +export PRETRAINED_MODEL_PATH GPU_IDS EXPECTED_EPISODES MAX_STEPS SAVE_STEPS +export WANDB_MODE HYDRA_FULL_ERROR CUDA_VISIBLE_DEVICES="$GPU_IDS" + +fail() { + echo "[ERROR] $*" >&2 + exit 1 +} + +[[ "$GPU_IDS" == "0,1,2,3" ]] || fail "This run is pinned to GPUs 0,1,2,3" +for path in "$PROJECT_ROOT" "$G2_DATA_ROOT" "$WAN_CKPT_DIR" \ + "$TOKENIZER_DIR" "$PRETRAINED_MODEL_PATH"; do + [[ -d "$path" ]] || fail "Missing directory: $path" +done +[[ -x "$PYTHON_BIN" ]] || fail "Missing Python interpreter: $PYTHON_BIN" + +"$PYTHON_BIN" - <<'PY' +import json +import os +from pathlib import Path + +root = Path(os.environ["G2_DATA_ROOT"]) +info = json.loads((root / "meta/info.json").read_text()) +expected = int(os.environ["EXPECTED_EPISODES"]) +if info["total_episodes"] != expected: + raise RuntimeError(f"Expected {expected} episodes, got {info['total_episodes']}") +if info["features"]["observation.state"]["shape"] != [16]: + raise RuntimeError("G2 state must be 16D") +if info["features"]["action"]["shape"] != [16]: + raise RuntimeError("G2 action must be 16D") +modality = json.loads((root / "meta/modality.json").read_text()) +expected_video = {"top_head", "hand_left", "hand_right"} +if set(modality["video"]) != expected_video: + raise RuntimeError(f"Unexpected video modalities: {sorted(modality['video'])}") +stats = json.loads((root / "meta/stats.json").read_text()) +for feature_name in ("observation.state", "action"): + values = stats[feature_name] + for index, key in ((7, "left_gripper_position"), (15, "right_gripper_position")): + if values["min"][index] < -1e-4 or values["max"][index] > 1.0001: + raise RuntimeError( + f"{feature_name} {key} is not policy-space [0,1]: " + f"min={values['min'][index]} max={values['max'][index]}" + ) +# All-absolute action contract: joint positions must be absolute (not near-zero +# relative deltas). An absolute joint position spans a wide [q01, q99] range, +# whereas a relative delta sits in a narrow band around 0. Check the SPREAD. +for feature_name in ("observation.state", "action"): + values = stats[feature_name] + arm_indices = [*range(0, 7), *range(8, 15)] + for i in arm_indices: + spread = values["q99"][i] - values["q01"][i] + if spread < 0.1: + raise RuntimeError( + f"{feature_name} dim {i} looks like a relative delta " + f"(q01={values['q01'][i]:.4f} q99={values['q99'][i]:.4f} " + f"spread={spread:.4f}); absolute joint-position stats expected." + ) +active_hold = json.loads((root / "meta/g2_active_hold_windows.json").read_text()) +if active_hold["action_horizon"] != 24: + raise RuntimeError("Active/hold index must use a 24-step horizon") +if active_hold["active_count"] <= 0 or active_hold["hold_count"] <= 0: + raise RuntimeError("Active/hold index is empty") +tasks_path = root / "meta/tasks.jsonl" +tasks = [json.loads(line) for line in tasks_path.read_text().splitlines() if line.strip()] +if len(tasks) != 1 or not str(tasks[0].get("task", "")).strip(): + raise RuntimeError("Expected exactly one non-empty G2 task prompt") +print(f"Prompt contract: {tasks[0]['task']}") +parquets = list(root.glob("data/chunk-*/*.parquet")) +videos = list(root.glob("videos/chunk-*/*/*.mp4")) +if len(parquets) != expected or len(videos) != expected * 3: + raise RuntimeError( + f"File count mismatch: parquets={len(parquets)} videos={len(videos)}" + ) +print( + "Validated G2 all-absolute dataset: " + f"episodes={expected} frames={info['total_frames']} tasks={info['total_tasks']} " + f"videos={len(videos)} active={active_hold['active_count']} " + f"hold={active_hold['hold_count']}" +) +PY + +if [[ -d "$OUTPUT_DIR" ]] && [[ -n "$(find "$OUTPUT_DIR" -mindepth 1 -maxdepth 1 -print -quit)" ]]; then + fail "Output directory is not empty: $OUTPUT_DIR" +fi +mkdir -p "$OUTPUT_DIR" +TRAIN_LOG="$OUTPUT_DIR/train.log" +cd "$PROJECT_ROOT" + +echo "G2 all-absolute joint training: Future-Video→Action bottleneck + state late gated residual" +echo "GPUs: $CUDA_VISIBLE_DEVICES" +echo "Base: $PRETRAINED_MODEL_PATH" +echo "Dataset: $G2_DATA_ROOT" +echo "Output: $OUTPUT_DIR" +echo "Objective: LoRA rank8 + flow action + decoded-action + endpoint (motion mask OFF)" + +"$PYTHON_BIN" -m torch.distributed.run \ + --nproc_per_node 4 \ + --standalone \ + groot/vla/experiment/experiment.py \ + report_to=wandb \ + data=dreamzero/g2_absolute \ + wandb_project=dreamzero \ + train_architecture=lora \ + num_frames=33 \ + action_horizon=24 \ + num_views=3 \ + model=dreamzero/vla \ + model/dreamzero/action_head=wan_flow_matching_action_tf \ + model/dreamzero/transform=dreamzero_cotrain \ + num_frame_per_block=2 \ + num_action_per_block=24 \ + num_state_per_block=1 \ + seed=42 \ + training_args.learning_rate=1e-5 \ + training_args.deepspeed="groot/vla/configs/deepspeed/zero2.json" \ + training_args.gradient_accumulation_steps=2 \ + save_steps="$SAVE_STEPS" \ + training_args.warmup_ratio=0.05 \ + output_dir="$OUTPUT_DIR" \ + per_device_train_batch_size=1 \ + max_steps="$MAX_STEPS" \ + weight_decay=1e-5 \ + save_total_limit=9 \ + upload_checkpoints=false \ + bf16=true \ + tf32=true \ + eval_bf16=true \ + dataloader_pin_memory=false \ + dataloader_num_workers=1 \ + image_resolution_width=320 \ + image_resolution_height=176 \ + save_lora_only=true \ + max_chunk_size=4 \ + frame_seqlen=880 \ + save_strategy=steps \ + active_window_ratio=0.9 \ + g2_data_root="$G2_DATA_ROOT" \ + dit_version="$WAN_CKPT_DIR" \ + text_encoder_pretrained_path="$WAN_CKPT_DIR/models_t5_umt5-xxl-enc-bf16.pth" \ + image_encoder_pretrained_path="$WAN_CKPT_DIR/models_clip_open-clip-xlm-roberta-large-vit-huge-14.pth" \ + vae_pretrained_path="$WAN_CKPT_DIR/Wan2.1_VAE.pth" \ + tokenizer_path="$TOKENIZER_DIR" \ + pretrained_model_path="$PRETRAINED_MODEL_PATH" \ + pretrained_lora_path=null \ + ++model_specific_transform.embodiment_tag_mapping.g2=33 \ + ++model_specific_transform.always_use_default_instruction=false \ + language_dropout_prob=0.0 \ + ++action_head_cfg.config.skip_component_loading=true \ + ++action_head_cfg.config.defer_lora_injection=true \ + action_head_cfg.config.action_only_training=false \ + action_head_cfg.config.tune_diffusion_model=true \ + ++action_head_cfg.config.lora_rank=8 \ + ++action_head_cfg.config.lora_alpha=8 \ + action_head_cfg.config.decouple_video_action_noise=false \ + action_head_cfg.config.dynamics_loss_weight=1.0 \ + action_head_cfg.config.action_loss_weight=10.0 \ + ++action_head_cfg.config.action_x0_loss_weight=1.0 \ + ++action_head_cfg.config.action_endpoint_loss_weight=0.5 \ + ++action_head_cfg.config.action_velocity_loss_weight=0.5 \ + ++action_head_cfg.config.action_acc_loss_weight=0.1 \ + ++action_head_cfg.config.state_dropout=0.0 \ + ++action_head_cfg.config.cut_state_attention=true \ + ++action_head_cfg.config.motion_mask_enabled=false \ + 2>&1 | tee "$TRAIN_LOG" + +echo "Completed G2 all-absolute joint training" diff --git a/groot/vla/model/dreamzero/modules/wan_video_dit_action_casual_chunk.py b/groot/vla/model/dreamzero/modules/wan_video_dit_action_casual_chunk.py index 02009be9..78def9b4 100644 --- a/groot/vla/model/dreamzero/modules/wan_video_dit_action_casual_chunk.py +++ b/groot/vla/model/dreamzero/modules/wan_video_dit_action_casual_chunk.py @@ -197,7 +197,8 @@ def __init__(self, qk_norm=True, eps=1e-6, num_action_per_block=32, - num_state_per_block=1): + num_state_per_block=1, + cut_state_attention=True): assert dim % num_heads == 0 super().__init__() self.dim = dim @@ -212,6 +213,7 @@ def __init__(self, self.frame_seqlen = frame_seqlen self.num_action_per_block = num_action_per_block self.num_state_per_block = num_state_per_block + self.cut_state_attention = cut_state_attention # layers self.q = nn.Linear(dim, dim) self.k = nn.Linear(dim, dim) @@ -479,7 +481,11 @@ def _blockwise_causal_flash_attn(self, q, k, v, frame_seqlen, num_frame_per_bloc state_block_starts = [state_start + i * num_state_per_block for i in range(num_state_blocks)] state_block_ends = [state_start + (i + 1) * num_state_per_block for i in range(num_state_blocks)] - # Process each image block + # Process each image block. + # Context: first image + image blocks + current action block (joint + # video<->action dynamics). When cut_state_attention=True (default) state + # is NOT in the context: it only conditions the shared blocks via the e0 + # modulation. When False, state joins the context. for block_idx in range(num_image_blocks): block_start = image_block_starts[block_idx] block_end = image_block_ends[block_idx] @@ -488,53 +494,61 @@ def _blockwise_causal_flash_attn(self, q, k, v, frame_seqlen, num_frame_per_bloc action_block_end = action_block_ends[block_idx] state_block_start = state_block_starts[block_idx] state_block_end = state_block_ends[block_idx] - - # Build context: first image + relevant image blocks + current action + current state - k_context = torch.cat([ + + k_parts = [ k[:, first_image_start:first_image_end], # First image k[:, image_kv_start:block_end], # Image blocks k[:, action_block_start:action_block_end], # Current action block - k[:, state_block_start:state_block_end] # Current state block - ], dim=1) - v_context = torch.cat([ + ] + v_parts = [ v[:, first_image_start:first_image_end], v[:, image_kv_start:block_end], v[:, action_block_start:action_block_end], - v[:, state_block_start:state_block_end] - ], dim=1) - + ] + if not self.cut_state_attention: + k_parts.append(k[:, state_block_start:state_block_end]) + v_parts.append(v[:, state_block_start:state_block_end]) + k_context = torch.cat(k_parts, dim=1) + v_context = torch.cat(v_parts, dim=1) + output[:, block_start:block_end] = self.attn( q[:, block_start:block_end], k_context, v_context ) - - # Process each action block + + # Process each action block. + # Context: first image + image blocks + current action block. When + # cut_state_attention=True (default) state is NOT in the action context — + # the action chunk must be driven by the future video, not by a direct + # current-state copy. When False, state joins the context. for block_idx in range(num_action_blocks): action_block_start = action_block_starts[block_idx] action_block_end = action_block_ends[block_idx] image_block_end = image_block_ends[block_idx] state_block_start = state_block_starts[block_idx] state_block_end = state_block_ends[block_idx] - + # Determine image context range if self.local_attn_size != -1: image_kv_start = max(image_blocks_start, image_block_end - self.local_attn_size * frame_seqlen) else: image_kv_start = image_blocks_start - - # Build context - k_context = torch.cat([ + + k_parts = [ k[:, first_image_start:first_image_end], # First image k[:, image_kv_start:image_block_end], # Image blocks k[:, action_block_start:action_block_end], # Current action block - k[:, state_block_start:state_block_end] # Current state block - ], dim=1) - v_context = torch.cat([ + ] + v_parts = [ v[:, first_image_start:first_image_end], v[:, image_kv_start:image_block_end], v[:, action_block_start:action_block_end], - v[:, state_block_start:state_block_end] - ], dim=1) - + ] + if not self.cut_state_attention: + k_parts.append(k[:, state_block_start:state_block_end]) + v_parts.append(v[:, state_block_start:state_block_end]) + k_context = torch.cat(k_parts, dim=1) + v_context = torch.cat(v_parts, dim=1) + output[:, action_block_start:action_block_end] = self.attn( q[:, action_block_start:action_block_end], k_context, v_context ) @@ -693,8 +707,13 @@ def _process_noisy_image_blocks(self, noisy_image_q, noisy_image_k, noisy_image_ action_block_ends = [start + self.num_action_per_block for start in action_block_starts] state_block_starts = [i * self.num_state_per_block for i in range(num_blocks)] state_block_ends = [start + self.num_state_per_block for start in state_block_starts] - - # Process noisy image blocks + + # Process noisy image blocks. + # Context: first_clean_frame + clean_blocks[0:i] + current_noisy_block + action[i]. + # When cut_state_attention=True (default), state is NOT part of the video + # context: the video branch is its own dynamics model conditioned on the + # shared blocks (state enters via e0), and state must not leak into the + # future-video prediction either. When False, state joins the context. for block_idx in range(num_blocks): noisy_start = noisy_block_starts[block_idx] noisy_end = noisy_block_ends[block_idx] @@ -703,23 +722,25 @@ def _process_noisy_image_blocks(self, noisy_image_q, noisy_image_k, noisy_image_ action_end = action_block_ends[block_idx] state_start = state_block_starts[block_idx] state_end = state_block_ends[block_idx] - + q_block = noisy_image_q[:, noisy_start:noisy_end] - - # Build context: first_clean_frame + clean_blocks[0:i] + current_noisy_block + action[i] + state[i] - k_context = torch.cat([ + + k_parts = [ clean_image_k[:, :clean_end], noisy_image_k[:, noisy_start:noisy_end], noisy_action_k[:, action_start:action_end], - noisy_state_k[:, state_start:state_end] - ], dim=1) - v_context = torch.cat([ + ] + v_parts = [ clean_image_v[:, :clean_end], noisy_image_v[:, noisy_start:noisy_end], noisy_action_v[:, action_start:action_end], - noisy_state_v[:, state_start:state_end] - ], dim=1) - + ] + if not self.cut_state_attention: + k_parts.append(noisy_state_k[:, state_start:state_end]) + v_parts.append(noisy_state_v[:, state_start:state_end]) + k_context = torch.cat(k_parts, dim=1) + v_context = torch.cat(v_parts, dim=1) + output[:, noisy_start:noisy_end] = self.attn(q_block, k_context, v_context) return output @@ -752,8 +773,13 @@ def _process_noisy_action_blocks(self, noisy_action_q, noisy_action_k, noisy_act noisy_image_block_ends = [start + self.frame_seqlen * self.num_frame_per_block for start in noisy_image_block_starts] state_block_starts = [i * self.num_state_per_block for i in range(num_blocks)] state_block_ends = [start + self.num_state_per_block for start in state_block_starts] - - # Process noisy action blocks + + # Process noisy action blocks. + # Context: first_clean_frame + clean_blocks[0:i] + noisy_image[i] + action[i]. + # When cut_state_attention=True (default), state is NOT part of the action + # context: the action chunk must derive "what happens next" from the future + # video, not from a raw current-state copy. State only conditions the shared + # blocks via e0. When False, the original state-in-context behavior returns. for block_idx in range(num_blocks): action_start = action_block_starts[block_idx] action_end = action_block_ends[block_idx] @@ -762,23 +788,25 @@ def _process_noisy_action_blocks(self, noisy_action_q, noisy_action_k, noisy_act noisy_img_end = noisy_image_block_ends[block_idx] state_start = state_block_starts[block_idx] state_end = state_block_ends[block_idx] - + q_block = noisy_action_q[:, action_start:action_end] - - # Build context: first_clean_frame + clean_blocks[0:i] + noisy_image[i] + action[i] + state[i] - k_context = torch.cat([ + + k_parts = [ clean_image_k[:, :clean_end], noisy_image_k[:, noisy_img_start:noisy_img_end], noisy_action_k[:, action_start:action_end], - noisy_state_k[:, state_start:state_end] - ], dim=1) - v_context = torch.cat([ + ] + v_parts = [ clean_image_v[:, :clean_end], noisy_image_v[:, noisy_img_start:noisy_img_end], noisy_action_v[:, action_start:action_end], - noisy_state_v[:, state_start:state_end] - ], dim=1) - + ] + if not self.cut_state_attention: + k_parts.append(noisy_state_k[:, state_start:state_end]) + v_parts.append(noisy_state_v[:, state_start:state_end]) + k_context = torch.cat(k_parts, dim=1) + v_context = torch.cat(v_parts, dim=1) + output[:, action_start:action_end] = self.attn(q_block, k_context, v_context) return output @@ -1099,7 +1127,8 @@ def __init__(self, cross_attn_norm=False, eps=1e-6, num_action_per_block=32, - num_state_per_block=1): + num_state_per_block=1, + cut_state_attention=True): super().__init__() self.dim = dim self.ffn_dim = ffn_dim @@ -1122,6 +1151,7 @@ def __init__(self, eps=eps, num_action_per_block=num_action_per_block, num_state_per_block=num_state_per_block, + cut_state_attention=cut_state_attention, ) self.norm3 = WanLayerNorm( dim, eps, @@ -1253,7 +1283,9 @@ def __init__(self, hidden_size=1024, diffusion_model_pretrained_path=None, num_action_per_block=32, - num_state_per_block=1): + num_state_per_block=1, + state_dropout=0.0, + cut_state_attention=True): r""" Initialize the diffusion model backbone. @@ -1342,6 +1374,24 @@ def __init__(self, output_dim=action_dim, ) + # State conditioning (AdaLN-style). The current proprioceptive state is + # encoded once into a moderate bottleneck and added to the per-block + # condition embedding `e0`. This symmetrically conditions BOTH video and + # action tokens through the shared DiT blocks. State never appears as + # register tokens that action can attend to, and there is no additive + # state->action path in the output heads. + self.state_cond_proj = nn.Sequential( + nn.Linear(max_state_dim, 256), + nn.SiLU(), + nn.Linear(256, dim * 6), + ) + self.state_dropout = state_dropout + # When True, action/video attention excludes the state register segment + # (state only conditions via the e0 AdaLN modulation). The state token + # structure/shape is preserved; the register state slots are zeroed so + # they cannot act as an attention shortcut. + self.cut_state_attention = cut_state_attention + # embeddings self.patch_embedding = nn.Conv3d( in_dim, dim, kernel_size=patch_size, stride=patch_size) @@ -1359,7 +1409,8 @@ def __init__(self, self.blocks = nn.ModuleList([ CausalWanAttentionBlock(cross_attn_type, dim, ffn_dim, num_heads, frame_seqlen, self.local_attn_size, sink_size, num_frame_per_block, qk_norm, cross_attn_norm, eps, - num_action_per_block, num_state_per_block) + num_action_per_block, num_state_per_block, + cut_state_attention=self.cut_state_attention) for _ in range(num_layers) ]) @@ -1383,6 +1434,12 @@ def __init__(self, # initialize weights self.init_weights() + # Zero-init the state conditioning's output projection so the initial + # AdaLN state modulation is ~0 (training starts video-driven; the state + # conditioning grows only if it helps). + nn.init.zeros_(self.state_cond_proj[-1].weight) + nn.init.zeros_(self.state_cond_proj[-1].bias) + self.gradient_checkpointing = True self.independent_first_frame = False if self.num_frame_per_block == 1 else True @@ -1393,7 +1450,8 @@ def _set_gradient_checkpointing(self, module, value=False): @staticmethod def _prepare_blockwise_causal_attn_mask( device: torch.device | str, num_frames: int = 21, - frame_seqlen: int = 1560, num_frame_per_block=1, local_attn_size=-1, action_horizon=1, state_horizon=1, num_action_per_block=30, num_state_per_block=1 + frame_seqlen: int = 1560, num_frame_per_block=1, local_attn_size=-1, action_horizon=1, state_horizon=1, num_action_per_block=30, num_state_per_block=1, + cut_state_attention: bool = True ) -> BlockMask: """ We will divide the token sequence into the following format: @@ -1498,16 +1556,25 @@ def attention_mask(b, h, q_idx, kv_idx): image_to_first = q_is_image_block & kv_is_first_image # Image block to first image: always allowed image_to_image = q_is_image_block & kv_is_image_block & (kv_block <= q_block) # Image block to image block: can attend to current and previous image blocks image_to_action = q_is_image_block & kv_is_action & (kv_block == q_block) # Image block to action: can attend to current action block - image_to_state = q_is_image_block & kv_is_state & (kv_block == q_block) # Image block to state: can attend to current state block - + # When cut_state_attention=True, the state edges are removed: neither + # video nor action may read the state register directly; state only + # conditions the shared blocks via the e0 AdaLN modulation. + image_to_state = ( + (not cut_state_attention) + and q_is_image_block and kv_is_state and (kv_block == q_block) + ) + image_block_mask = image_to_first | image_to_image | image_to_action | image_to_state - + # Action query action_to_image = q_is_action & kv_is_image_block & (kv_block <= q_block) # Action to image block: can attend to current and all previous image blocks action_to_action = q_is_action & kv_is_action & (kv_block == q_block) # Action to action: only same block - action_to_state = q_is_action & kv_is_state & (kv_block == q_block) # Action to state: only same block + action_to_state = ( + (not cut_state_attention) + and q_is_action and kv_is_state and (kv_block == q_block) + ) # Action to state: only same block (disabled when cut_state_attention) action_to_first = q_is_action & kv_is_first_image # Action to first image: always allowed - + action_mask = action_to_image | action_to_action | action_to_state | action_to_first # State query (conditioning) - cannot attend to anything @@ -1714,7 +1781,14 @@ def _forward_blocks( if action is not None: embodiment_id = torch.tensor([0], device=x.device).repeat(x.shape[0]) action_features = self.action_encoder(action, timestep_action, embodiment_id) - state_features = self.state_encoder(state, embodiment_id) + # State no longer enters the DiT sequence as register tokens that + # video/action could attend to (that was the shortcut). The register + # state slots are zeroed; real state influence is injected as an + # AdaLN-style modulation on `e0` below. + state_features = torch.zeros( + (state.shape[0], state.shape[1], self.dim), + device=action_features.device, dtype=action_features.dtype, + ) action_register = torch.cat([action_features, state_features], dim=1) action_length = action_features.shape[1] action_register_length = action_register.shape[1] @@ -1741,6 +1815,11 @@ def _forward_blocks( e0 = self.time_projection(e) e0 = e0.unflatten(dim=2, sizes=(6, self.dim)) + # Proprioceptive state is intentionally NOT used in this model variant: + # the register state slots are zeroed, attention edges are cut + # (cut_state_attention), and no AdaLN state modulation / dropout is + # applied. The action chunk must be predicted purely from the video. + # context context = self.text_embedding(context) @@ -2011,7 +2090,14 @@ def _forward_train( embodiment_id = torch.tensor([0]).repeat(x.shape[0]).to(device=embodiment_id.device) action_features = self.action_encoder(action, timestep_action, embodiment_id) action_length = action_features.shape[1] - state_features = self.state_encoder(state, embodiment_id) + # State no longer enters the DiT sequence as register tokens that + # video/action could attend to (that was the shortcut). The register + # state slots are zeroed; real state influence is injected as an + # AdaLN-style modulation on `e0` below. + state_features = torch.zeros( + (state.shape[0], state.shape[1], self.dim), + device=action_features.device, dtype=action_features.dtype, + ) action_register = torch.cat([action_features, state_features], dim=1) action_register_length = action_register.shape[1] x = torch.cat([x, action_register], dim=1) @@ -2039,6 +2125,11 @@ def _forward_train( e0 = self.time_projection(e) e0 = e0.unflatten(dim=2, sizes=(6, self.dim)) + # Proprioceptive state is intentionally NOT used in this model variant: + # the register state slots are zeroed, attention edges are cut + # (cut_state_attention), and no AdaLN state modulation / dropout is + # applied. The action chunk must be predicted purely from the video. + # context assert context.shape[1] == self.text_len context = self.text_embedding(context) diff --git a/groot/vla/model/dreamzero/transform/dreamzero_cotrain.py b/groot/vla/model/dreamzero/transform/dreamzero_cotrain.py index 0e303b29..b1d6ef16 100644 --- a/groot/vla/model/dreamzero/transform/dreamzero_cotrain.py +++ b/groot/vla/model/dreamzero/transform/dreamzero_cotrain.py @@ -101,7 +101,10 @@ def collate(features: List[dict], tokenizer: AutoTokenizer, num_views=3, embodim # If it's already a scalar (string, float, int, etc.), convert to string processed_item = str(parsed_item) - if num_views > 1 and elem["embodiment_id"] == embodiment_tag_mapping[EmbodimentTag.AGIBOT.value]: + if num_views > 1 and elem["embodiment_id"] in ( + embodiment_tag_mapping.get(EmbodimentTag.AGIBOT.value), + embodiment_tag_mapping.get(EmbodimentTag.G2.value), + ): processed_item = "A multi-view video shows that a robot " + processed_item.lower() + " The video is split into four views: The top-left view shows the camera view from the robot's head, the top-right view shows the camera view from the right hand, the bottom-left view shows the camera view from the left hand, and the bottom-right view is a black screen (inactive view). The robot " + processed_item.lower() elif elem["embodiment_id"] == embodiment_tag_mapping[EmbodimentTag.OXE_DROID.value]: processed_item = ( @@ -123,7 +126,10 @@ def collate(features: List[dict], tokenizer: AutoTokenizer, num_views=3, embodim output_values.append(processed_item) except (ValueError, SyntaxError, TypeError): # If parsing fails or item is already a string, use it directly - if num_views > 1 and elem["embodiment_id"] == embodiment_tag_mapping[EmbodimentTag.AGIBOT.value]: + if num_views > 1 and elem["embodiment_id"] in ( + embodiment_tag_mapping.get(EmbodimentTag.AGIBOT.value), + embodiment_tag_mapping.get(EmbodimentTag.G2.value), + ): item = "A multi-view video shows that a robot " + str(item).lower() + " The video is split into four views: The top-left view shows the camera view from the robot's head, the top-right view shows the camera view from the right hand, the bottom-left view shows the camera view from the left hand, and the bottom-right view is a black screen (inactive view). The robot " + str(item).lower() elif elem["embodiment_id"] == embodiment_tag_mapping[EmbodimentTag.OXE_DROID.value]: item = ( diff --git a/groot/vla/model/n1_5/sim_policy.py b/groot/vla/model/n1_5/sim_policy.py old mode 100755 new mode 100644 index 77ff9296..3adf73a7 --- a/groot/vla/model/n1_5/sim_policy.py +++ b/groot/vla/model/n1_5/sim_policy.py @@ -254,6 +254,7 @@ def __init__( model_target = train_cfg.model._target_ self.model_target = model_target + action_adapter_path = model_dir / "action_expert.safetensors" if model_config_overrides is not None and len(model_config_overrides) != 0: print(f"Applying model config overrides: {model_config_overrides}") @@ -277,9 +278,24 @@ def __init__( model_config = model_config_class.from_dict(model_config) # Instantiate the model - if hasattr(train_cfg, "save_lora_only") and train_cfg.save_lora_only is True: + if action_adapter_path.exists(): + print("Loading G2 action adapter on clean DreamZero base") + base_model_path = train_cfg.get("pretrained_model_path", None) + model = model_class.load_action_adapter( + str(model_dir), + pretrained_base_model_path=base_model_path, + config=model_config, + ) + elif hasattr(train_cfg, "save_lora_only") and train_cfg.save_lora_only is True: print(f"Loading LoRA weights from pretrained") - model = model_class.load_lora(model_path) + base_model_path = train_cfg.get( + "pretrained_model_path", + None, + ) + model = model_class.load_lora( + model_path, + pretrained_base_model_path=base_model_path, + ) else: print(f"Loading model from pretrained directly") model = model_class.from_pretrained(model_path, config=model_config) @@ -289,10 +305,25 @@ def __init__( cls_module, cls_name = model_target.rsplit(".", 1) if 'lora' in cls_name: cls_module, cls_name = cls_module.rsplit(".", 1) - if hasattr(train_cfg, "save_lora_only") and train_cfg.save_lora_only is True: + if action_adapter_path.exists(): + print("Loading G2 action adapter on clean DreamZero base") + cls = getattr(importlib.import_module(cls_module), cls_name) + base_model_path = train_cfg.get("pretrained_model_path", None) + model = cls.load_action_adapter( + str(model_dir), + pretrained_base_model_path=base_model_path, + ) + elif hasattr(train_cfg, "save_lora_only") and train_cfg.save_lora_only is True: print(f"Loading LoRA weights from pretrained") cls = getattr(importlib.import_module(cls_module), cls_name) - model = cls.load_lora(model_path) + base_model_path = train_cfg.get( + "pretrained_model_path", + None, + ) + model = cls.load_lora( + model_path, + pretrained_base_model_path=base_model_path, + ) else: print(f"Loading model from pretrained directly") cls = getattr(importlib.import_module(cls_module), cls_name) diff --git a/make_g2_dataset_snapshot.py b/make_g2_dataset_snapshot.py new file mode 100644 index 00000000..93bdb242 --- /dev/null +++ b/make_g2_dataset_snapshot.py @@ -0,0 +1,92 @@ +"""Build a live-replay-compatible snapshot from a G2 Gear episode.""" + +from __future__ import annotations + +import argparse +from pathlib import Path + +import cv2 +import numpy as np +import pyarrow.parquet as pq + + +CAMERAS = ("top_head", "hand_left", "hand_right") + + +def _read_video_prefix(path: Path, end_row: int) -> np.ndarray: + capture = cv2.VideoCapture(str(path)) + frames: list[np.ndarray] = [] + try: + while len(frames) <= end_row: + ok, bgr = capture.read() + if not ok: + raise RuntimeError( + f"Could not read frame {len(frames)} from {path}" + ) + frames.append(cv2.cvtColor(bgr, cv2.COLOR_BGR2RGB)) + finally: + capture.release() + start = max(0, end_row - 3) + history = frames[start : end_row + 1] + while len(history) < 4: + history.insert(0, history[0]) + return np.stack(history, axis=0).astype(np.uint8) + + +def main() -> None: + parser = argparse.ArgumentParser() + parser.add_argument("root", type=Path) + parser.add_argument("--episode", type=int, default=0) + parser.add_argument("--row", type=int, default=3) + parser.add_argument("--output", type=Path, required=True) + args = parser.parse_args() + + chunk = args.episode // 1000 + parquet = ( + args.root + / "data" + / f"chunk-{chunk:03d}" + / f"episode_{args.episode:06d}.parquet" + ) + table = pq.read_table(parquet) + states = np.asarray( + table["observation.state"].to_pylist(), dtype=np.float32 + ) + actions = np.asarray(table["action"].to_pylist(), dtype=np.float32) + if not 0 <= args.row < len(states): + raise IndexError(f"row {args.row} is outside episode length {len(states)}") + if args.row + 24 > len(actions): + raise IndexError("Not enough future actions for a 24-step horizon") + + videos: dict[str, np.ndarray] = {} + for camera in CAMERAS: + video_path = ( + args.root + / "videos" + / f"chunk-{chunk:03d}" + / f"observation.images.{camera}" + / f"episode_{args.episode:06d}.mp4" + ) + videos[camera] = _read_video_prefix(video_path, args.row) + + prompt = str( + table["annotation.language.action_text"][args.row].as_py() + ) + args.output.parent.mkdir(parents=True, exist_ok=True) + np.savez_compressed( + args.output, + top_head_jpeg_decoded=videos["top_head"], + hand_left_jpeg_decoded=videos["hand_left"], + hand_right_jpeg_decoded=videos["hand_right"], + state_16=states[args.row], + raw_server_actions=np.zeros((24, 16), dtype=np.float32), + ground_truth_actions=actions[args.row : args.row + 24], + prompt=np.asarray(prompt), + episode=np.asarray(args.episode, dtype=np.int32), + row=np.asarray(args.row, dtype=np.int32), + ) + print(args.output) + + +if __name__ == "__main__": + main() diff --git a/plot_dreamzero_training_curves_global_step.py b/plot_dreamzero_training_curves_global_step.py new file mode 100644 index 00000000..8c2f767e --- /dev/null +++ b/plot_dreamzero_training_curves_global_step.py @@ -0,0 +1,630 @@ +#!/usr/bin/env python3 +""" +Plot DreamZero training metrics with the real global-step axis. + +Default run: + /data/wangk/checkpoints/dreamzero_g2_joint_lora_full + +Data-source priority: + 1. loss_log.jsonl + 2. trainer_state.json + 3. TensorBoard events + 4. train.log + +If a metric log contains N records without an explicit step while the completed +run has global_step=5000, the script maps the records over the real training +range instead of incorrectly using 0..N as the x-axis. +""" + +from __future__ import annotations + +import argparse +import ast +import csv +import json +import math +import re +from collections import defaultdict +from pathlib import Path +from typing import Any + +import matplotlib.pyplot as plt + + +DEFAULT_RUN_DIR = Path( + "/data/wangk/checkpoints/dreamzero_g2_joint_lora_full" +) + + +def finite_float(value: Any) -> float | None: + try: + number = float(value) + except (TypeError, ValueError): + return None + return number if math.isfinite(number) else None + + +def normalize_metric_name(name: str) -> str: + aliases = { + "train/loss": "loss", + "train_loss": "loss", + "training_loss": "loss", + "train/action_loss": "action_loss_avg", + "action_loss": "action_loss_avg", + "action/loss": "action_loss_avg", + "lr": "learning_rate", + } + name = str(name).strip() + return aliases.get(name, name) + + +def load_json(path: Path) -> dict[str, Any]: + try: + data = json.loads(path.read_text(encoding="utf-8")) + except Exception: + return {} + return data if isinstance(data, dict) else {} + + +def read_run_metadata(run_dir: Path) -> dict[str, int | float | None]: + state_files = [run_dir / "trainer_state.json"] + state_files.extend(sorted(run_dir.glob("checkpoint-*/trainer_state.json"))) + + best: dict[str, Any] = {} + best_step = -1 + for path in state_files: + if not path.exists(): + continue + data = load_json(path) + step = int(data.get("global_step") or 0) + if step >= best_step: + best = data + best_step = step + + global_step = int(best.get("global_step") or 0) or None + max_steps = int(best.get("max_steps") or 0) or None + + logging_steps = best.get("logging_steps") + try: + logging_steps = int(logging_steps) + except (TypeError, ValueError): + logging_steps = None + + return { + "global_step": global_step, + "max_steps": max_steps, + "logging_steps": logging_steps, + } + + +def row_step(row: dict[str, Any]) -> int | None: + for key in ( + "global_step", + "step", + "training_step", + "train_step", + "iteration", + "iter", + ): + value = row.get(key) + if value is None: + continue + try: + return int(value) + except (TypeError, ValueError): + continue + return None + + +NON_METRIC_KEYS = { + "global_step", + "step", + "training_step", + "train_step", + "iteration", + "iter", + "epoch", + "timestamp", + "time", + "rank", +} + + +def infer_steps( + rows: list[dict[str, Any]], + global_step: int | None, + logging_steps: int | None, +) -> list[int]: + explicit = [row_step(row) for row in rows] + if rows and all(step is not None for step in explicit): + return [int(step) for step in explicit if step is not None] + + count = len(rows) + if count == 0: + return [] + + # Prefer trainer logging_steps when it exactly or nearly spans the run. + if logging_steps and logging_steps > 0: + steps = [(index + 1) * logging_steps for index in range(count)] + if global_step: + steps = [min(step, global_step) for step in steps] + if steps[-1] < global_step and global_step / count > logging_steps * 1.25: + # The saved logging_steps does not explain the number of rows. + steps = [] + if steps: + return steps + + # Main fallback for this run: + # 2500 scalar rows over a completed 5000-step run -> 2, 4, ..., 5000. + if global_step and global_step > 0: + scale = global_step / count + steps = [ + max(1, min(global_step, int(round((index + 1) * scale)))) + for index in range(count) + ] + steps[-1] = global_step + return steps + + return list(range(1, count + 1)) + + +def read_loss_jsonl( + path: Path, + global_step: int | None, + logging_steps: int | None, +) -> dict[str, list[tuple[int, float]]]: + records: dict[str, list[tuple[int, float]]] = defaultdict(list) + if not path.exists(): + return records + + rows: list[dict[str, Any]] = [] + with path.open("r", encoding="utf-8", errors="replace") as handle: + for line_number, line in enumerate(handle, start=1): + line = line.strip() + if not line: + continue + try: + value = json.loads(line) + except json.JSONDecodeError: + continue + if isinstance(value, dict): + rows.append(value) + + steps = infer_steps(rows, global_step, logging_steps) + for step, row in zip(steps, rows): + for key, value in row.items(): + if key in NON_METRIC_KEYS: + continue + number = finite_float(value) + if number is not None: + records[normalize_metric_name(key)].append((step, number)) + + return records + + +def read_trainer_history( + run_dir: Path, +) -> dict[str, list[tuple[int, float]]]: + records: dict[str, list[tuple[int, float]]] = defaultdict(list) + path = run_dir / "trainer_state.json" + if not path.exists(): + return records + + data = load_json(path) + for row in data.get("log_history", []): + if not isinstance(row, dict): + continue + step = row_step(row) + if step is None: + continue + for key, value in row.items(): + if key in NON_METRIC_KEYS: + continue + number = finite_float(value) + if number is not None: + records[normalize_metric_name(key)].append((step, number)) + return records + + +def read_tensorboard( + run_dir: Path, +) -> dict[str, list[tuple[int, float]]]: + records: dict[str, list[tuple[int, float]]] = defaultdict(list) + files = sorted(run_dir.rglob("events.out.tfevents.*")) + if not files: + return records + + try: + from tensorboard.backend.event_processing.event_accumulator import ( + EventAccumulator, + ) + except Exception as exc: + print(f"[WARN] TensorBoard unavailable: {exc}") + return records + + for path in files: + try: + accumulator = EventAccumulator( + str(path), + size_guidance={"scalars": 0}, + ) + accumulator.Reload() + for tag in accumulator.Tags().get("scalars", []): + name = normalize_metric_name(tag) + for event in accumulator.Scalars(tag): + number = finite_float(event.value) + if number is not None: + records[name].append((int(event.step), number)) + except Exception as exc: + print(f"[WARN] Cannot read {path}: {exc}") + + return records + + +DICT_RE = re.compile(r"\{.*\}") +PROGRESS_RE = re.compile(r"\b(\d+)\s*/\s*(\d+)\b") + + +def read_train_log( + path: Path, +) -> dict[str, list[tuple[int, float]]]: + records: dict[str, list[tuple[int, float]]] = defaultdict(list) + if not path.exists(): + return records + + current_step: int | None = None + with path.open("r", encoding="utf-8", errors="replace") as handle: + for line in handle: + progress = PROGRESS_RE.search(line) + if progress: + current_step = int(progress.group(1)) + + match = DICT_RE.search(line) + if not match: + continue + try: + row = ast.literal_eval(match.group(0)) + except Exception: + continue + if not isinstance(row, dict): + continue + + explicit = row_step(row) + if explicit is not None: + current_step = explicit + if current_step is None: + continue + + for key, value in row.items(): + if key in NON_METRIC_KEYS: + continue + number = finite_float(value) + if number is not None: + records[normalize_metric_name(key)].append( + (current_step, number) + ) + return records + + +def merge_sources( + sources: list[dict[str, list[tuple[int, float]]]], +) -> dict[str, list[tuple[int, float]]]: + """ + Sources are ordered from highest to lowest priority. + + A higher-priority source owns a metric when it contains that metric. This + prevents a lower-priority source with an incorrect synthetic axis from + overwriting loss_log.jsonl's corrected 0..5000 axis. + """ + result: dict[str, list[tuple[int, float]]] = {} + all_names = { + name + for source in sources + for name, values in source.items() + if values + } + + for name in sorted(all_names): + selected: list[tuple[int, float]] = [] + for source in sources: + values = source.get(name, []) + if values: + selected = values + break + + deduplicated: dict[int, float] = {} + for step, value in selected: + deduplicated[int(step)] = float(value) + result[name] = sorted(deduplicated.items()) + + return result + + +def moving_average(values: list[float], window: int) -> list[float]: + if window <= 1: + return values[:] + + result: list[float] = [] + queue: list[float] = [] + running = 0.0 + for value in values: + queue.append(value) + running += value + if len(queue) > window: + running -= queue.pop(0) + result.append(running / len(queue)) + return result + + +def find_metric( + records: dict[str, list[tuple[int, float]]], + names: list[str], +) -> str | None: + for name in names: + if records.get(name): + return name + for actual, values in records.items(): + lower = actual.lower() + if values and any(name.lower() in lower for name in names): + return actual + return None + + +def plot_metric( + values: list[tuple[int, float]], + path: Path, + title: str, + ylabel: str, + smooth_window: int, + final_step: int | None, +) -> None: + steps = [step for step, _ in values] + raw = [value for _, value in values] + smoothed = moving_average(raw, smooth_window) + + plt.figure(figsize=(12, 6)) + plt.plot(steps, raw, linewidth=0.8, alpha=0.25, label="raw") + plt.plot( + steps, + smoothed, + linewidth=2.0, + label=f"moving average ({smooth_window})", + ) + plt.xlabel("Global training step") + plt.ylabel(ylabel) + plt.title(title) + if final_step: + plt.xlim(0, final_step) + plt.grid(True, alpha=0.25) + plt.legend() + plt.tight_layout() + plt.savefig(path, dpi=180) + plt.close() + + +def write_csv( + records: dict[str, list[tuple[int, float]]], + path: Path, +) -> None: + metric_names = sorted(records) + all_steps = sorted( + { + step + for values in records.values() + for step, _ in values + } + ) + maps = { + name: dict(records[name]) + for name in metric_names + } + + with path.open("w", newline="", encoding="utf-8") as handle: + writer = csv.writer(handle) + writer.writerow(["global_step", *metric_names]) + for step in all_steps: + writer.writerow( + [step, *[maps[name].get(step, "") for name in metric_names]] + ) + + +def make_dashboard( + records: dict[str, list[tuple[int, float]]], + path: Path, + smooth_window: int, + final_step: int | None, +) -> None: + specs = [ + (["model_forward_time"], "model_forward_time"), + (["dynamics_loss_avg"], "dynamics_loss_avg"), + (["action_loss_avg"], "action_loss_avg"), + (["training_step_time"], "training_step_time"), + (["loss"], "loss"), + (["grad_norm"], "grad_norm"), + (["learning_rate"], "learning_rate"), + ] + + available = [] + for candidates, title in specs: + key = find_metric(records, candidates) + if key: + available.append((key, title)) + + if not available: + return + + fig, axes = plt.subplots( + len(available), + 1, + figsize=(14, max(4, 3.1 * len(available))), + sharex=True, + ) + if len(available) == 1: + axes = [axes] + + for axis, (key, title) in zip(axes, available): + values = records[key] + steps = [step for step, _ in values] + raw = [value for _, value in values] + smooth = moving_average(raw, smooth_window) + + axis.plot(steps, raw, linewidth=0.7, alpha=0.22) + axis.plot(steps, smooth, linewidth=1.7) + axis.set_title(title) + axis.grid(True, alpha=0.25) + if final_step: + axis.set_xlim(0, final_step) + + axes[-1].set_xlabel("Global training step") + fig.suptitle( + f"DreamZero G2 Training Metrics — final global step: " + f"{final_step if final_step else 'unknown'}", + y=0.995, + ) + fig.tight_layout() + fig.savefig(path, dpi=180) + plt.close(fig) + + +def main() -> int: + parser = argparse.ArgumentParser() + parser.add_argument("--run-dir", type=Path, default=DEFAULT_RUN_DIR) + parser.add_argument("--log-file", type=Path, default=None) + parser.add_argument("--output-dir", type=Path, default=None) + parser.add_argument("--smooth-window", type=int, default=25) + args = parser.parse_args() + + run_dir = args.run_dir.resolve() + log_file = ( + args.log_file.resolve() + if args.log_file + else run_dir / "train.log" + ) + output_dir = ( + args.output_dir.resolve() + if args.output_dir + else run_dir / "training_plots" + ) + output_dir.mkdir(parents=True, exist_ok=True) + + metadata = read_run_metadata(run_dir) + final_step = ( + metadata["global_step"] + or metadata["max_steps"] + ) + final_step = int(final_step) if final_step else None + + print("[INFO] run directory:", run_dir) + print("[INFO] trainer global_step:", metadata["global_step"]) + print("[INFO] trainer max_steps:", metadata["max_steps"]) + print("[INFO] trainer logging_steps:", metadata["logging_steps"]) + + loss_jsonl = read_loss_jsonl( + run_dir / "loss_log.jsonl", + final_step, + ( + int(metadata["logging_steps"]) + if metadata["logging_steps"] + else None + ), + ) + trainer = read_trainer_history(run_dir) + events = read_tensorboard(run_dir) + plain_log = read_train_log(log_file) + + records = merge_sources( + [loss_jsonl, trainer, events, plain_log] + ) + records = { + name: values + for name, values in records.items() + if values + } + if not records: + print("[ERROR] No scalar training metrics were found.") + return 1 + + print("[INFO] selected metrics:") + for name, values in records.items(): + print( + f" {name}: {len(values)} points, " + f"step {values[0][0]} -> {values[-1][0]}" + ) + + write_csv(records, output_dir / "training_metrics_global_step.csv") + + individual_specs = [ + ( + ["loss"], + "01_loss_global_step.png", + "DreamZero G2 Total Loss", + "Loss", + ), + ( + ["action_loss_avg"], + "02_action_loss_global_step.png", + "DreamZero G2 Action Loss", + "Action loss", + ), + ( + ["dynamics_loss_avg"], + "03_dynamics_loss_global_step.png", + "DreamZero G2 Dynamics Loss", + "Dynamics loss", + ), + ( + ["learning_rate"], + "04_learning_rate_global_step.png", + "Learning Rate", + "Learning rate", + ), + ( + ["grad_norm"], + "05_grad_norm_global_step.png", + "Gradient Norm", + "Gradient norm", + ), + ( + ["training_step_time"], + "06_training_step_time_global_step.png", + "Training Step Time", + "Seconds", + ), + ( + ["model_forward_time"], + "07_model_forward_time_global_step.png", + "Model Forward Time", + "Seconds", + ), + ] + + for candidates, filename, title, ylabel in individual_specs: + key = find_metric(records, candidates) + if not key: + continue + plot_metric( + records[key], + output_dir / filename, + title, + ylabel, + args.smooth_window, + final_step, + ) + + make_dashboard( + records, + output_dir / "training_metrics_global_step.png", + args.smooth_window, + final_step, + ) + + print("[OK] plots written to:", output_dir) + print( + "[OK] main figure:", + output_dir / "training_metrics_global_step.png", + ) + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/replay_g2_live_snapshot.py b/replay_g2_live_snapshot.py new file mode 100644 index 00000000..c5213e44 --- /dev/null +++ b/replay_g2_live_snapshot.py @@ -0,0 +1,128 @@ +"""Replay one captured G2 observation against a websocket policy server. + +This diagnostic never talks to the robot SDK and never executes actions. It +can repeat the exact same observation under one session ID to expose how the +server's causal cache changes consecutive action chunks. +""" + +from __future__ import annotations + +import argparse +import logging +from pathlib import Path +import uuid + +import numpy as np + +from eval_utils.policy_client import WebsocketClientPolicy + + +def _parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser() + parser.add_argument("snapshot", type=Path) + parser.add_argument("--host", default="127.0.0.1") + parser.add_argument("--port", type=int, required=True) + parser.add_argument("--repeats", type=int, default=3) + parser.add_argument("--output", type=Path, required=True) + parser.add_argument( + "--session-id", + default=None, + help="Fresh UUID by default so the server resets its causal cache.", + ) + parser.add_argument( + "--left-gripper", + type=float, + default=None, + help="Optional counterfactual override for state_16[7].", + ) + return parser.parse_args() + + +def _build_observation( + snapshot: np.lib.npyio.NpzFile, + *, + session_id: str, + left_gripper: float | None, +) -> tuple[dict[str, object], np.ndarray]: + state = np.asarray(snapshot["state_16"], dtype=np.float32) + if state.shape != (16,): + raise ValueError(f"Expected state_16 shape (16,), got {state.shape}") + state = state.copy() + if left_gripper is not None: + state[7] = np.float32(left_gripper) + return { + "observation/top_head": np.asarray( + snapshot["top_head_jpeg_decoded"], dtype=np.uint8 + ), + "observation/hand_left": np.asarray( + snapshot["hand_left_jpeg_decoded"], dtype=np.uint8 + ), + "observation/hand_right": np.asarray( + snapshot["hand_right_jpeg_decoded"], dtype=np.uint8 + ), + "observation/left_joint_position": state[0:7], + "observation/left_gripper_position": state[7:8], + "observation/right_joint_position": state[8:15], + "observation/right_gripper_position": state[15:16], + "prompt": str(snapshot["prompt"].item()), + "session_id": session_id, + }, state + + +def main() -> None: + args = _parse_args() + if args.repeats < 1: + raise ValueError("--repeats must be positive") + + snapshot = np.load(args.snapshot, allow_pickle=False) + session_id = args.session_id or f"snapshot-replay-{uuid.uuid4()}" + observation, state = _build_observation( + snapshot, + session_id=session_id, + left_gripper=args.left_gripper, + ) + client = WebsocketClientPolicy(host=args.host, port=args.port) + + chunks: list[np.ndarray] = [] + for index in range(args.repeats): + action = np.asarray(client.infer(dict(observation)), dtype=np.float32) + if action.shape != (24, 16): + raise ValueError( + f"Replay {index} returned {action.shape}, expected (24,16)" + ) + chunks.append(action) + arm_delta = np.concatenate( + (action[:, 0:7] - state[None, 0:7], + action[:, 8:15] - state[None, 8:15]), + axis=1, + ) + logging.info( + "chunk=%d first_arm_max_delta=%.6f " + "full_arm_max_delta=%.6f left_grip=[%.6f,%.6f] " + "right_grip=[%.6f,%.6f]", + index, + float(np.max(np.abs(arm_delta[0]))), + float(np.max(np.abs(arm_delta))), + float(action[:, 7].min()), + float(action[:, 7].max()), + float(action[:, 15].min()), + float(action[:, 15].max()), + ) + + args.output.parent.mkdir(parents=True, exist_ok=True) + np.savez_compressed( + args.output, + replay_actions=np.stack(chunks), + production_actions=np.asarray( + snapshot["raw_server_actions"], dtype=np.float32 + ), + state_16=state, + prompt=np.asarray(str(snapshot["prompt"].item())), + session_id=np.asarray(session_id), + ) + logging.info("Saved replay comparison to %s", args.output) + + +if __name__ == "__main__": + logging.basicConfig(level=logging.INFO, force=True) + main() diff --git a/robot_live_client_g2.py b/robot_live_client_g2.py new file mode 100644 index 00000000..a8d4d6f7 --- /dev/null +++ b/robot_live_client_g2.py @@ -0,0 +1,677 @@ +#!/usr/bin/env python3 +"""DreamZero live-policy client for AgiBot G2. + +The policy/inference loop is shared with :mod:`robot_live_client`. This file +adapts the current ``agibot_gdk`` G2 API to the small camera/robot interface +used by that loop. Run after sourcing ``~/.cache/agibot/app/env.sh``. + +G2 has a three-axis head while the trained policy has two head dimensions. +They are mapped to yaw (idx11) and pitch (idx13); roll (idx12) is preserved. +G2's five-motor parallel waist needs an inverse-kinematics API which is not +exposed by the supplied examples, so waist commands are deliberately ignored. + +python robot_live_client_g2.py \ + --host 111.0.22.33 \ + --port 30001 \ + --prompt "机器人右臂先从网口扩展坞物料区抓取网口扩展坞,左臂随后从网线物料区依次抓取两根网线并逐一插入" \ + --sdk-arm-order left_right \ + --observation-fps 5 \ + --observation-history 1 \ + --image-transport jpeg \ + --arm-execution-mode direct-48 \ + --direct-control-hz 15 \ + --arm-delta-limit 0.04 \ + --arm-velocity-limit 0.35 \ + --arm-acceleration-limit 0.80 \ + --arm-close-timeout 0.15 \ + --apply-actions + + +""" + +from __future__ import annotations + +import logging +import sys +import threading +import time +from typing import Any + +import cv2 +import numpy as np + +try: + import agibot_gdk +except Exception as exc: # pragma: no cover - only available in G2 runtime + raise RuntimeError( + "Failed to import agibot_gdk. Source ~/.cache/agibot/app/env.sh and " + "use the Python version shipped with the G2 GDK." + ) from exc + +import robot_live_client as live + +_shared_execute_arm_trajectory = live._execute_arm_trajectory +_shared_execute_arm_direct48 = live._execute_arm_direct48 +_shared_maybe_smooth_actions = live._maybe_smooth_actions +_SharedWebsocketClientPolicy = live.WebsocketClientPolicy + + +ARM_JOINT_NAMES = [ + *(f"idx2{i}_arm_l_joint{i}" for i in range(1, 8)), + *(f"idx6{i}_arm_r_joint{i}" for i in range(1, 8)), +] +HEAD_JOINT_NAMES = ["idx11_head_joint1", "idx12_head_joint2", "idx13_head_joint3"] +BODY_JOINT_NAMES = [f"idx0{i}_body_joint{i}" for i in range(1, 6)] +G2_ACTION_DIM = 16 + + +def _decode_g2_relative_action( + actions: np.ndarray, + current_arm: np.ndarray, + current_gripper: np.ndarray, +) -> np.ndarray: + """ + Convert DreamZero relative action into absolute G2 joint targets. + + Training: + action = target_joint - current_joint + + Deployment: + target_joint = predicted_delta + current_joint + + G2 action layout: + [left_arm7, + left_gripper, + right_arm7, + right_gripper] + """ + + actions = np.asarray(actions, dtype=np.float32) + if actions.ndim == 1: + actions = actions.reshape(1, -1) + if actions.ndim != 2 or actions.shape[1] != G2_ACTION_DIM: + raise ValueError( + f"Expected G2 relative action shape (T, {G2_ACTION_DIM}), " + f"got {actions.shape}" + ) + + current_arm = np.asarray(current_arm, dtype=np.float32).reshape(14) + current_gripper = np.asarray(current_gripper, dtype=np.float32).reshape(2) + decoded = actions.copy() + + decoded[:, 0:7] += current_arm[0:7] + decoded[:, 7] += current_gripper[0] + decoded[:, 8:15] += current_arm[7:14] + decoded[:, 15] += current_gripper[1] + return decoded + + +G2_GRIPPER_OPEN_POSITION = -0.785 +G2_GRIPPER_CLOSED_POSITION = 0.0 +G2_ARM_GRIPPER_MIN_INTERVAL_S = 0.050 +# Source: ~/.cache/agibot/app/gdk/config/mc_impl_config.json. +# Keep targets just inside the GDK boundary to avoid float32 round-off at the +# exact limit. Values are ordered like ARM_JOINT_NAMES (left 7, then right 7). +G2_ARM_JOINT_LIMIT_EPSILON = 1e-3 +G2_ARM_JOINT_MIN = np.asarray( + [ + -3.071796, + -2.059505, + -3.071796, + -2.495838, + -3.071796, + -1.012308, + -1.535907, + -3.071796, + -2.059505, + -3.071796, + -2.495838, + -3.071796, + -1.012308, + -1.535907, + ], + dtype=np.float32, +) +G2_ARM_JOINT_MAX = np.asarray( + [ + 3.071796, + 2.059505, + 3.071796, + 1.012308, + 3.071796, + 1.012308, + 1.535907, + 3.071796, + 2.059505, + 3.071796, + 1.012308, + 3.071796, + 1.012308, + 1.535907, + ], + dtype=np.float32, +) + + +class G2WebsocketClientPolicy(_SharedWebsocketClientPolicy): + """WebSocket policy client used by the G2 live client.""" + + +def _position_by_name(robot: Any) -> tuple[dict[str, float], int]: + response = robot.get_joint_states() + states = response.get("states", []) + positions = { + str(state["name"]): float(state.get("motor_position", state.get("position", 0.0))) + for state in states + } + timestamp = int(response.get("timestamp", time.time_ns())) + return positions, timestamp + + +class G2Camera: + """Present the G2 Camera API as the three named streams used by DreamZero.""" + + _TYPES: dict[str, Any] = { + "head": agibot_gdk.CameraType.kHeadColor, + "hand_left": agibot_gdk.CameraType.kHandLeftColor, + "hand_right": agibot_gdk.CameraType.kHandRightColor, + } + + def __init__(self, names: list[str]) -> None: + unknown = set(names) - self._TYPES.keys() + if unknown: + raise ValueError(f"Unsupported G2 camera names: {sorted(unknown)}") + self._camera = agibot_gdk.Camera() + + @staticmethod + def _decode(image: Any) -> np.ndarray | None: + if image is None or not hasattr(image, "data"): + return None + data = image.data + if data is None or np.asarray(data).size == 0: + return None + if image.encoding in (agibot_gdk.Encoding.JPEG, agibot_gdk.Encoding.PNG): + return cv2.imdecode(np.frombuffer(data, np.uint8), cv2.IMREAD_COLOR) + if image.encoding != agibot_gdk.Encoding.UNCOMPRESSED: + raise RuntimeError(f"Unsupported G2 image encoding: {image.encoding}") + + raw = np.frombuffer(data, dtype=np.uint8) + if image.color_format == agibot_gdk.ColorFormat.GRAY8: + gray = raw.reshape((image.height, image.width)) + return cv2.cvtColor(gray, cv2.COLOR_GRAY2BGR) + frame = raw.reshape((image.height, image.width, 3)) + if image.color_format == agibot_gdk.ColorFormat.RGB: + frame = cv2.cvtColor(frame, cv2.COLOR_RGB2BGR) + elif image.color_format != agibot_gdk.ColorFormat.BGR: + raise RuntimeError(f"Unsupported G2 color format: {image.color_format}") + return frame + + def get_latest_image(self, name: str) -> tuple[np.ndarray | None, int]: + image = self._camera.get_latest_image(self._TYPES[name], 1000.0) + if image is None: + return None, 0 + timestamp = int(getattr(image, "timestamp_ns", getattr(image, "timestamp", 0))) + return self._decode(image), timestamp + + def close(self) -> None: + self._camera.close_camera() + + +class G2Robot: + """Compatibility adapter around ``agibot_gdk.Robot`` and ``Pnc``.""" + + def __init__(self) -> None: + self._robot = agibot_gdk.Robot() + self._pnc = None + self._head_roll = 0.0 + self._waist_warning_emitted = False + self._last_arm_gripper_command_at: float | None = None + self._last_arm_gripper_command_kind: str | None = None + # Arm and gripper commands share one G2 control channel. Serialize both + # paths so commands cannot overlap across client worker threads. + self._arm_gripper_lock = threading.Lock() + + def _wait_arm_gripper_switch(self, command_kind: str) -> None: + if ( + self._last_arm_gripper_command_at is not None + and self._last_arm_gripper_command_kind != command_kind + ): + remaining = ( + G2_ARM_GRIPPER_MIN_INTERVAL_S + - (time.monotonic() - self._last_arm_gripper_command_at) + ) + if remaining > 0.0: + time.sleep(remaining) + + def _mark_arm_gripper_command(self, command_kind: str) -> None: + self._last_arm_gripper_command_at = time.monotonic() + self._last_arm_gripper_command_kind = command_kind + + def arm_joint_states(self) -> tuple[list[float], int]: + positions, timestamp = _position_by_name(self._robot) + missing = [name for name in ARM_JOINT_NAMES if name not in positions] + if missing: + raise RuntimeError(f"G2 joint-state response is missing arm joints: {missing}") + return [positions[name] for name in ARM_JOINT_NAMES], timestamp + + def head_joint_states(self) -> tuple[list[float], int]: + positions, timestamp = _position_by_name(self._robot) + missing = [name for name in HEAD_JOINT_NAMES if name not in positions] + if missing: + raise RuntimeError(f"G2 joint-state response is missing head joints: {missing}") + self._head_roll = positions[HEAD_JOINT_NAMES[1]] + return [positions[HEAD_JOINT_NAMES[0]], positions[HEAD_JOINT_NAMES[2]]], timestamp + + def waist_joint_states(self) -> tuple[list[float], int]: + # Raw body motor angles are not equivalent to policy [pitch, lift]. + _, timestamp = _position_by_name(self._robot) + return [0.0, 0.0], timestamp + + def gripper_states(self) -> tuple[list[float], int]: + response = self._robot.get_end_state() + values: list[float] = [] + for side in ("left", "right"): + end = response.get(f"{side}_end_state", {}) + states = end.get("end_states", []) + position = float(states[0].get("position", 0.0)) if states else 0.0 + # G2 training data, policy output, and the omnipicker SDK all use + # the same physical range: -0.785 is open and 0 is closed. + values.append( + float( + np.clip( + position, + G2_GRIPPER_OPEN_POSITION, + G2_GRIPPER_CLOSED_POSITION, + ) + ) + ) + return values, int(time.time_ns()) + + def _joint_request(self, names: list[str], positions: list[float], speed: float = 0.3) -> None: + request = agibot_gdk.JointControlReq() + request.life_time = 1.0 + request.joint_names = names + logging.info( + "FINAL SDK ARM CMD=%s", + positions + ) + request.joint_positions = [float(value) for value in positions] + request.joint_velocities = [float(speed)] * len(names) + result = self._robot.joint_control_request(request) + if result not in (None, 0): + raise RuntimeError(f"G2 joint_control_request failed with result {result}") + + def move_arm(self, positions: list[float]) -> None: + if len(positions) != 14: + raise ValueError(f"G2 arm command must contain 14 positions, got {len(positions)}") + + # The shared executor applies interpolation and delta limits after the + # policy action was clipped. Clamp the final values again immediately + # before the SDK call so a boundary joint cannot drift out of range. + values = np.asarray(positions, dtype=np.float64) + safe_min = ( + G2_ARM_JOINT_MIN.astype(np.float64) + + G2_ARM_JOINT_LIMIT_EPSILON + ) + safe_max = ( + G2_ARM_JOINT_MAX.astype(np.float64) + - G2_ARM_JOINT_LIMIT_EPSILON + ) + clipped = np.clip(values, safe_min, safe_max) + clipped_mask = clipped != values + if np.any(clipped_mask): + affected = [ + ARM_JOINT_NAMES[index] + for index in np.flatnonzero(clipped_mask) + ] + logging.warning( + "Clipped final G2 SDK arm command inside joint limits; affected joints=%s", + affected, + ) + + with self._arm_gripper_lock: + self._wait_arm_gripper_switch("arm") + try: + self._joint_request(ARM_JOINT_NAMES, clipped.tolist()) + finally: + self._mark_arm_gripper_command("arm") + + def move_head(self, positions: list[float]) -> None: + if len(positions) != 2: + raise ValueError(f"G2 policy head command must contain 2 positions, got {len(positions)}") + # G2 order is yaw, roll, pitch. Preserve the unmodelled roll axis. + self._robot.move_head_joint( + [float(positions[0]), self._head_roll, float(positions[1])], + [0.3, 0.3, 0.3], + ) + + def move_waist(self, positions: list[float]) -> None: + if not self._waist_warning_emitted: + logging.warning( + "Ignoring waist action: G2 uses five coupled body motors and the supplied SDK " + "examples do not expose a safe pitch/lift command API." + ) + self._waist_warning_emitted = True + + def move_gripper(self, positions: list[float]) -> None: + if len(positions) != 2: + raise ValueError(f"G2 gripper command must contain 2 positions, got {len(positions)}") + command = agibot_gdk.JointStates() + command.group = "dual_tool" + command.target_type = "omnipicker" + states: list[Any] = [] + for value in positions: + state = agibot_gdk.JointState() + state.position = float( + np.clip( + value, + G2_GRIPPER_OPEN_POSITION, + G2_GRIPPER_CLOSED_POSITION, + ) + ) + states.append(state) + command.states = states + command.nums = len(states) + with self._arm_gripper_lock: + self._wait_arm_gripper_switch("gripper") + try: + result = self._robot.move_ee_pos(command) + finally: + self._mark_arm_gripper_command("gripper") + if result != 0: + raise RuntimeError(f"G2 dual gripper move_ee_pos failed with result {result}") + + def move_wheel(self, linear: float, angular: float) -> None: + if self._pnc is None: + self._pnc = agibot_gdk.Pnc() + self._pnc.request_chassis_control(0) + time.sleep(0.5) + twist = agibot_gdk.Twist() + twist.linear = agibot_gdk.Vector3() + twist.angular = agibot_gdk.Vector3() + twist.linear.x = float(linear) + twist.angular.z = float(angular) + self._pnc.move_chassis(twist) + + def shutdown(self) -> None: + # G2 objects have no shutdown method in the supplied SDK examples. + if self._pnc is not None: + self.move_wheel(0.0, 0.0) + + +def _build_g2_obs( + head_img: np.ndarray, + left_img: np.ndarray, + right_img: np.ndarray, + arm_pos: list[float], + head_pos: list[float], + waist_pos: list[float], + gripper_pos: list[float], + prompt: str, + session_id: str, + sdk_arm_order: live.ArmOrder, + obs_flip_config: live.ObsFlipConfig, + image_transport: live.ImageTransportMode, + image_jpeg_quality: int, +) -> dict[str, object]: + """Build the observation keys declared by modality_config_g2.""" + del head_pos, waist_pos + policy_arm = live._sdk_to_policy_arm(arm_pos, sdk_arm_order) + policy_arm = live._apply_obs_joint_sign_flips(policy_arm, obs_flip_config) + gripper = np.asarray(gripper_pos, dtype=np.float32) + if gripper.shape != (2,): + raise ValueError(f"G2 gripper state must contain 2 values, got {gripper.shape}") + return { + "observation/top_head": live._encode_video_observation( + head_img, image_transport=image_transport, jpeg_quality=image_jpeg_quality + ), + "observation/hand_left": live._encode_video_observation( + left_img, image_transport=image_transport, jpeg_quality=image_jpeg_quality + ), + "observation/hand_right": live._encode_video_observation( + right_img, image_transport=image_transport, jpeg_quality=image_jpeg_quality + ), + "observation/left_joint_position": policy_arm[:7], + "observation/left_gripper_position": gripper[:1], + "observation/right_joint_position": policy_arm[7:], + "observation/right_gripper_position": gripper[1:], + "prompt": prompt, + "session_id": session_id, + } + + +def _parse_g2_action_row(row: np.ndarray, horizon: int) -> dict[str, np.ndarray]: + """Parse [left arm 7, left grip 1, right arm 7, right grip 1].""" + row = np.asarray(row, dtype=np.float32).reshape(-1) + if row.shape[0] == 22: + # Internal representation used only by the shared 22-D executor. + return { + "left_arm": row[0:7], + "right_arm": row[7:14], + "gripper": row[14:16], + "head": row[16:18], + "waist": row[18:20], + "wheel": row[20:22], + "horizon": np.asarray([horizon], dtype=np.int32), + } + if row.shape[0] != G2_ACTION_DIM: + raise ValueError(f"Expected G2 action dimension 16, got {row.shape[0]}") + return { + "left_arm": row[0:7], + "right_arm": row[8:15], + "gripper": np.asarray([row[7], row[15]], dtype=np.float32), + # NaN marks modalities which do not exist in the G2 policy output, so + # the shared executor skips their command paths. + "head": np.full(2, np.nan, dtype=np.float32), + "waist": np.full(2, np.nan, dtype=np.float32), + "wheel": np.full(2, np.nan, dtype=np.float32), + "horizon": np.asarray([horizon], dtype=np.int32), + } + + +def _parse_g2_action_first(actions: np.ndarray) -> dict[str, np.ndarray]: + actions = np.asarray(actions, dtype=np.float32) + if actions.ndim == 1: + actions = actions.reshape(1, -1) + if actions.ndim != 2 or actions.shape[1] != G2_ACTION_DIM: + raise ValueError(f"Expected G2 action shape (T, 16), got {actions.shape}") + return _parse_g2_action_row(actions[0], actions.shape[0]) + + +def _select_g2_gripper_command( + actions: np.ndarray, gripper_config: live.GripperConfig +) -> dict[str, np.ndarray]: + actions = np.asarray(actions, dtype=np.float32) + if actions.ndim == 1: + actions = actions.reshape(1, -1) + if actions.ndim != 2 or actions.shape[1] != G2_ACTION_DIM: + raise ValueError(f"Expected G2 action shape (T, 16), got {actions.shape}") + first = actions[0, [7, 15]].astype(np.float32) + last = actions[-1, [7, 15]].astype(np.float32) + command = live._gripper_policy_to_command(last, gripper_config) + return {"first_policy": first, "last_policy": last, "command_policy": command} + + +def _g2_gripper_policy_to_command( + policy_values: np.ndarray, gripper_config: live.GripperConfig +) -> np.ndarray: + """Pass G2 omnipicker positions through in the native SDK scale. + + G2 inference returns physical positions in [-0.785, 0], so thresholding + them into the shared client's normalized open/closed values would reverse + or discard the command. ``gripper_config`` is intentionally unused on G2. + """ + del gripper_config + pair = np.asarray(policy_values, dtype=np.float32).reshape(-1) + if pair.shape[0] != 2: + raise ValueError(f"Expected G2 gripper pair with 2 dims, got {pair.shape[0]}") + return np.clip( + pair, + G2_GRIPPER_OPEN_POSITION, + G2_GRIPPER_CLOSED_POSITION, + ).astype(np.float32) + + +def _g2_gripper_state_for_log(values: list | np.ndarray) -> np.ndarray: + """Keep G2 diagnostic state in the native range used on the wire.""" + pair = np.asarray(values, dtype=np.float32).reshape(-1) + if pair.shape[0] != 2: + raise ValueError(f"Expected G2 gripper state with 2 dims, got {pair.shape[0]}") + return np.clip( + pair, + G2_GRIPPER_OPEN_POSITION, + G2_GRIPPER_CLOSED_POSITION, + ).astype(np.float32) + + +def _clip_g2_arm_joint_limits(actions: np.ndarray) -> np.ndarray: + """Clamp G2 arm targets to the absolute limits enforced by the GDK.""" + arm_targets = np.concatenate((actions[:, 0:7], actions[:, 8:15]), axis=1) + safe_min = G2_ARM_JOINT_MIN + G2_ARM_JOINT_LIMIT_EPSILON + safe_max = G2_ARM_JOINT_MAX - G2_ARM_JOINT_LIMIT_EPSILON + clipped_targets = np.clip(arm_targets, safe_min, safe_max).astype(np.float32) + clipped_mask = clipped_targets != arm_targets + if np.any(clipped_mask): + affected = [ + ARM_JOINT_NAMES[index] + for index in range(len(ARM_JOINT_NAMES)) + if np.any(clipped_mask[:, index]) + ] + logging.warning( + "Clipped %s G2 arm target value(s) to GDK absolute joint limits; affected joints=%s", + int(np.count_nonzero(clipped_mask)), + affected, + ) + return clipped_targets + + +def _g2_to_shared_actions(actions: np.ndarray) -> np.ndarray: + """Reorder external G2 actions into the shared executor's 22-D layout.""" + actions = np.asarray(actions, dtype=np.float32) + if actions.ndim == 1: + actions = actions.reshape(1, -1) + if actions.ndim != 2 or actions.shape[1] != G2_ACTION_DIM: + raise ValueError(f"Expected G2 action shape (T, 16), got {actions.shape}") + clipped_arm_targets = _clip_g2_arm_joint_limits(actions) + shared = np.full((actions.shape[0], 22), np.nan, dtype=np.float32) + shared[:, 0:14] = clipped_arm_targets + shared[:, 14] = actions[:, 7] + shared[:, 15] = actions[:, 15] + return shared + + +def _shared_to_g2_actions(actions: np.ndarray) -> np.ndarray: + actions = np.asarray(actions, dtype=np.float32) + g2 = np.empty((actions.shape[0], G2_ACTION_DIM), dtype=np.float32) + g2[:, 0:7] = actions[:, 0:7] + g2[:, 7] = actions[:, 14] + g2[:, 8:15] = actions[:, 7:14] + g2[:, 15] = actions[:, 15] + return g2 + + +def _decode_g2_actions_for_execution( + *, + robot: G2Robot, + actions: np.ndarray, + sdk_arm_order: live.ArmOrder, +) -> np.ndarray: + """Decode checkpoint-relative actions against the latest robot state.""" + current_arm_sdk, _ = robot.arm_joint_states() + current_gripper, _ = robot.gripper_states() + current_arm_policy = live._sdk_to_policy_arm( + np.asarray(current_arm_sdk, dtype=np.float32), + sdk_arm_order, + ) + decoded = _decode_g2_relative_action( + actions, + current_arm=current_arm_policy, + current_gripper=np.asarray(current_gripper, dtype=np.float32), + ) + logging.info( + "Decoded G2 relative action against execution-boundary state | " + "delta_first_left[:3]=%s delta_first_right[:3]=%s " + "target_first_left[:3]=%s target_first_right[:3]=%s", + np.round(np.asarray(actions, dtype=np.float32).reshape(-1, G2_ACTION_DIM)[0, 0:3], 4).tolist(), + np.round(np.asarray(actions, dtype=np.float32).reshape(-1, G2_ACTION_DIM)[0, 8:11], 4).tolist(), + np.round(decoded[0, 0:3], 4).tolist(), + np.round(decoded[0, 8:11], 4).tolist(), + ) + return decoded + + +def _execute_g2_trajectory( + *, + robot: G2Robot, + actions: np.ndarray, + sdk_arm_order: live.ArmOrder, + **kwargs: Any, +) -> dict[str, Any]: + decoded = _decode_g2_actions_for_execution( + robot=robot, + actions=actions, + sdk_arm_order=sdk_arm_order, + ) + return _shared_execute_arm_trajectory( + robot=robot, + actions=_g2_to_shared_actions(decoded), + sdk_arm_order=sdk_arm_order, + **kwargs, + ) + + +def _execute_g2_direct48( + *, + robot: G2Robot, + actions: np.ndarray, + sdk_arm_order: live.ArmOrder, + **kwargs: Any, +) -> dict[str, Any]: + decoded = _decode_g2_actions_for_execution( + robot=robot, + actions=actions, + sdk_arm_order=sdk_arm_order, + ) + return _shared_execute_arm_direct48( + robot=robot, + actions=_g2_to_shared_actions(decoded), + sdk_arm_order=sdk_arm_order, + **kwargs, + ) + + +def _smooth_g2_actions( + actions: np.ndarray, config: live.ActionSmoothingConfig +) -> tuple[np.ndarray, float]: + shared, duration = _shared_maybe_smooth_actions(_g2_to_shared_actions(actions), config) + return _shared_to_g2_actions(shared), duration + + +def main() -> None: + result = agibot_gdk.gdk_init() + success = getattr(getattr(agibot_gdk, "GDKRes", object), "kSuccess", None) + if success is not None and result != success: + raise RuntimeError(f"agibot_gdk.gdk_init() failed: {result}") + + live.Camera = G2Camera + live.Robot = G2Robot + live.WebsocketClientPolicy = G2WebsocketClientPolicy + live._build_obs = _build_g2_obs + live._parse_action_row = _parse_g2_action_row + live._parse_action_first = _parse_g2_action_first + live._select_gripper_command = _select_g2_gripper_command + live._gripper_policy_to_command = _g2_gripper_policy_to_command + live._sdk_gripper_to_policy_obs = _g2_gripper_state_for_log + live._execute_arm_trajectory = _execute_g2_trajectory + live._execute_arm_direct48 = _execute_g2_direct48 + live._maybe_smooth_actions = _smooth_g2_actions + + # modality_config_g2.eval_delta_indices is [0], unlike the four-frame G1 + # evaluation history. Respect an explicit CLI override when one is given. + if "--observation-history" not in sys.argv: + sys.argv.extend(["--observation-history", "1"]) + live.main() + + +if __name__ == "__main__": + main() diff --git a/rollout_g2_server_episode.py b/rollout_g2_server_episode.py new file mode 100644 index 00000000..d06c423b --- /dev/null +++ b/rollout_g2_server_episode.py @@ -0,0 +1,447 @@ +#!/usr/bin/env python3 +"""Teacher-forced full-episode rollout against one persistent G2 server. + +The model server is deliberately not started here. This client connects to +the already-loaded checkpoint, sends one 4-frame observation packet every +``stride`` frames, and records every returned 24x16 action chunk against the +dataset action horizon at the same anchor. At the end it sends one reset +message so the server writes its accumulated predicted video. +""" + +from __future__ import annotations + +import argparse +import csv +import json +import logging +import time +import uuid +from pathlib import Path + +import cv2 +import imageio.v2 as imageio +import numpy as np +import pyarrow.parquet as pq + +from eval_utils.policy_client import WebsocketClientPolicy + + +ACTION_NAMES = ( + [f"left_joint_{i}" for i in range(7)] + + ["left_gripper"] + + [f"right_joint_{i}" for i in range(7)] + + ["right_gripper"] +) + + +def _episode_path(root: Path, template: str, episode: int, chunks_size: int, **kwargs: object) -> Path: + return root / template.format( + episode_chunk=episode // chunks_size, + episode_index=episode, + **kwargs, + ) + + +def _read_video(path: Path, expected_rows: int) -> np.ndarray: + capture = cv2.VideoCapture(str(path)) + if not capture.isOpened(): + raise RuntimeError(f"Failed to open G2 video: {path}") + frames: list[np.ndarray] = [] + try: + while True: + ok, frame = capture.read() + if not ok: + break + # The online G2 client sends camera arrays in this BGR convention; + # keep it identical to the existing dataset evaluator. + frames.append(np.ascontiguousarray(frame)) + finally: + capture.release() + if len(frames) < expected_rows: + raise RuntimeError( + f"{path} contains {len(frames)} frames, expected at least {expected_rows}" + ) + return np.stack(frames[:expected_rows], axis=0) + + +def _grid_rgb(top_bgr: np.ndarray, left_bgr: np.ndarray, right_bgr: np.ndarray) -> np.ndarray: + top = cv2.cvtColor(top_bgr, cv2.COLOR_BGR2RGB) + left = cv2.cvtColor(left_bgr, cv2.COLOR_BGR2RGB) + right = cv2.cvtColor(right_bgr, cv2.COLOR_BGR2RGB) + black = np.zeros_like(top) + return np.concatenate( + [np.concatenate([top, left], axis=1), np.concatenate([right, black], axis=1)], + axis=0, + ) + + +def _label(frame_rgb: np.ndarray, text: str) -> np.ndarray: + frame_bgr = cv2.cvtColor(frame_rgb, cv2.COLOR_RGB2BGR) + cv2.rectangle(frame_bgr, (0, 0), (360, 34), (0, 0, 0), -1) + cv2.putText( + frame_bgr, + text, + (8, 24), + cv2.FONT_HERSHEY_SIMPLEX, + 0.65, + (255, 255, 255), + 2, + cv2.LINE_AA, + ) + return cv2.cvtColor(frame_bgr, cv2.COLOR_BGR2RGB) + + +def _save_video(path: Path, frames: list[np.ndarray], fps: int = 30) -> None: + if not frames: + raise RuntimeError(f"No frames to save: {path}") + path.parent.mkdir(parents=True, exist_ok=True) + imageio.mimsave(path, frames, fps=fps, codec="libx264", macro_block_size=None) + + +def _metric(error: np.ndarray) -> dict[str, float | int]: + arm = error[..., [*range(7), *range(8, 15)]] + grip = error[..., [7, 15]] + return { + "overall_mae": float(error.mean()), + "arm_mae": float(arm.mean()), + "gripper_mae": float(grip.mean()), + "left_arm_mae": float(error[..., :7].mean()), + "right_arm_mae": float(error[..., 8:15].mean()), + "num_action_rows": int(np.prod(error.shape[:-1])), + } + + +def _write_action_reports( + output_dir: Path, + episode: int, + anchors: list[int], + predicted: np.ndarray, + ground_truth: np.ndarray, +) -> dict[str, object]: + error = np.abs(predicted - ground_truth) + detail_path = output_dir / f"episode_{episode:06d}_full_action_pred_vs_gt.csv" + fields = ["block_index", "anchor_frame", "horizon_step"] + fields += [f"gt_{name}" for name in ACTION_NAMES] + fields += [f"pred_{name}" for name in ACTION_NAMES] + with detail_path.open("w", newline="", encoding="utf-8") as stream: + writer = csv.DictWriter(stream, fieldnames=fields) + writer.writeheader() + for block_index, anchor in enumerate(anchors): + for horizon in range(predicted.shape[1]): + gt_row = ground_truth[block_index, horizon] + pred_row = predicted[block_index, horizon] + row: dict[str, object] = { + "block_index": block_index, + "anchor_frame": anchor, + "horizon_step": horizon + 1, + } + row.update({f"gt_{name}": float(v) for name, v in zip(ACTION_NAMES, gt_row)}) + row.update({f"pred_{name}": float(v) for name, v in zip(ACTION_NAMES, pred_row)}) + writer.writerow(row) + + metrics_path = output_dir / f"episode_{episode:06d}_full_action_metrics.csv" + metric_fields = [ + "scope", + "block_index", + "anchor_frame", + "overall_mae", + "arm_mae", + "gripper_mae", + "left_arm_mae", + "right_arm_mae", + "num_action_rows", + ] + rows: list[dict[str, object]] = [] + for block_index, anchor in enumerate(anchors): + row: dict[str, object] = _metric(error[block_index][None, ...]) + row.update({ + "scope": f"block_{block_index:04d}", + "block_index": block_index, + "anchor_frame": anchor, + }) + rows.append(row) + overall: dict[str, object] = _metric(error) + overall.update({"scope": "overall", "block_index": "", "anchor_frame": ""}) + rows.append(overall) + with metrics_path.open("w", newline="", encoding="utf-8") as stream: + writer = csv.DictWriter(stream, fieldnames=metric_fields) + writer.writeheader() + writer.writerows(rows) + + plot_path = output_dir / f"episode_{episode:06d}_full_action_pred_vs_gt.png" + try: + import matplotlib + + matplotlib.use("Agg") + import matplotlib.pyplot as plt + + x = np.arange(predicted.shape[0]) + first_pred = predicted[:, 0] + first_gt = ground_truth[:, 0] + first_error = np.abs(first_pred - first_gt) + fig, axes = plt.subplots(3, 1, figsize=(16, 10), sharex=True) + axes[0].plot(x, first_gt[:, :7].mean(axis=1), label="GT left arm", color="tab:blue") + axes[0].plot(x, first_pred[:, :7].mean(axis=1), label="PRED left arm", color="tab:orange") + axes[0].set_ylabel("left arm mean") + axes[0].legend() + axes[1].plot(x, first_gt[:, 8:15].mean(axis=1), label="GT right arm", color="tab:green") + axes[1].plot(x, first_pred[:, 8:15].mean(axis=1), label="PRED right arm", color="tab:red") + axes[1].set_ylabel("right arm mean") + axes[1].legend() + axes[2].plot(x, first_error[:, [*range(7), *range(8, 15)]].mean(axis=1), label="arm MAE") + axes[2].plot(x, first_error[:, [7, 15]].mean(axis=1), label="gripper MAE") + axes[2].set_xlabel("teacher-forced inference block (4 frames per block)") + axes[2].set_ylabel("absolute error") + axes[2].legend() + for axis in axes: + axis.grid(alpha=0.25) + fig.suptitle(f"G2 full-task first-action comparison, episode {episode}") + fig.tight_layout() + fig.savefig(plot_path, dpi=140) + plt.close(fig) + except Exception as exc: # pragma: no cover + logging.warning("Could not save action plot: %s", exc) + plot_path = None + + return { + "action_detail_csv": str(detail_path), + "action_metrics_csv": str(metrics_path), + "action_plot": str(plot_path) if plot_path else None, + "overall_action_metrics": overall, + } + + +def _parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser() + parser.add_argument("--host", default="127.0.0.1") + parser.add_argument("--port", type=int, default=30002) + parser.add_argument("--test-data-root", type=Path, required=True) + parser.add_argument("--episode-index", type=int, default=0) + parser.add_argument("--output-dir", type=Path, required=True) + parser.add_argument( + "--server-video-dir", + type=Path, + required=True, + help="Persistent server's VIDEO_SAVE_MODE=full output directory.", + ) + parser.add_argument("--start-frame", type=int, default=0) + parser.add_argument("--end-frame", type=int, default=-1) + parser.add_argument("--block-stride", type=int, default=4) + parser.add_argument("--max-blocks", type=int, default=None) + parser.add_argument("--prompt", default=None) + return parser.parse_args() + + +def main() -> None: + args = _parse_args() + if args.block_stride < 1: + raise ValueError("--block-stride must be positive") + args.output_dir.mkdir(parents=True, exist_ok=True) + root = args.test_data_root.resolve() + info = json.loads((root / "meta/info.json").read_text(encoding="utf-8")) + chunks_size = int(info.get("chunks_size", 1000)) + episode = int(args.episode_index) + parquet_path = _episode_path( + root, + info.get("data_path", "data/chunk-{episode_chunk:03d}/episode_{episode_index:06d}.parquet"), + episode, + chunks_size, + ) + table = pq.read_table(parquet_path) + # PyArrow returns an object array for the nested fixed-size columns; use + # Python rows before converting so the 16-D vectors are not treated as + # scalar objects. + state = np.asarray(table["observation.state"].to_pylist(), dtype=np.float32) + action = np.asarray(table["action"].to_pylist(), dtype=np.float32) + state = state.reshape(table.num_rows, 16) + action = action.reshape(table.num_rows, 16) + num_rows = int(table.num_rows) + video_template = info.get( + "video_path", + "videos/chunk-{episode_chunk:03d}/{video_key}/episode_{episode_index:06d}.mp4", + ) + videos = {} + for name, key in (("top_head", "observation.images.top_head"), ("hand_left", "observation.images.hand_left"), ("hand_right", "observation.images.hand_right")): + path = _episode_path(root, video_template, episode, chunks_size, video_key=key) + videos[name] = _read_video(path, num_rows) + + first_block = max(0, int(args.start_frame)) + last_anchor = num_rows - 24 + if args.end_frame >= 0: + last_anchor = min(last_anchor, int(args.end_frame)) + # A packet starts at frame s and is anchored at s+3. Keep a complete + # 24-step GT horizon for every request. + starts = list(range(first_block, max(first_block, last_anchor - 3) + 1, args.block_stride)) + if args.max_blocks is not None: + starts = starts[: max(0, int(args.max_blocks))] + if not starts: + raise RuntimeError("No valid full-task inference blocks") + + prompt_key = "annotation.language.action_text" + default_prompt = "" + if prompt_key in table.column_names: + default_prompt = str(table[prompt_key][0].as_py()) + prompt = args.prompt if args.prompt is not None else default_prompt + session_id = f"g2-full-task-episode-{episode:06d}-{uuid.uuid4()}" + + logging.info( + "Connecting to persistent server %s:%s; episode=%d rows=%d blocks=%d start=%d end=%d stride=%d", + args.host, + args.port, + episode, + num_rows, + len(starts), + starts[0], + starts[-1], + args.block_stride, + ) + client = WebsocketClientPolicy(host=args.host, port=args.port) + logging.info("Server metadata: %s", client.get_server_metadata()) + + predicted_chunks: list[np.ndarray] = [] + anchors: list[int] = [] + started = time.time() + try: + for block_index, packet_start in enumerate(starts): + anchor = packet_start + 3 + observation = { + "observation/top_head": videos["top_head"][packet_start : packet_start + 4], + "observation/hand_left": videos["hand_left"][packet_start : packet_start + 4], + "observation/hand_right": videos["hand_right"][packet_start : packet_start + 4], + "observation/state": state[anchor], + "prompt": prompt, + "session_id": session_id, + } + result = np.asarray(client.infer(observation), dtype=np.float32) + if result.shape != (24, 16): + raise RuntimeError( + f"server block {block_index} returned {result.shape}, expected (24,16)" + ) + predicted_chunks.append(result) + anchors.append(anchor) + if block_index == 0 or (block_index + 1) % 10 == 0 or block_index + 1 == len(starts): + logging.info( + "full-task block %d/%d anchor=%d elapsed=%.1fs first_arm=[%.4f,%.4f,%.4f]", + block_index + 1, + len(starts), + anchor, + time.time() - started, + result[0, 0], + result[0, 1], + result[0, 2], + ) + finally: + # reset makes the persistent server flush its accumulated predicted + # video; it does not unload the checkpoint or terminate the server. + try: + client.reset({"session_id": session_id}) + except Exception: + logging.exception("Server reset/video flush failed") + + predicted = np.stack(predicted_chunks, axis=0) + ground_truth = np.stack( + [action[anchor : anchor + 24] for anchor in anchors], + axis=0, + ) + reports = _write_action_reports( + args.output_dir, + episode, + anchors, + predicted, + ground_truth, + ) + np.savez_compressed( + args.output_dir / f"episode_{episode:06d}_full_action_arrays.npz", + predicted=predicted, + ground_truth=ground_truth, + anchors=np.asarray(anchors, dtype=np.int32), + ) + + # The policy is queried every four real frames. Use exactly those four + # frames per block for the comparison timeline; comparing the server's + # raw concatenated latent frames 1:1 against the 30-Hz GT makes the + # prediction look artificially fast (one block decodes 9/12 frames while + # the real observation stream advances only four frames). + gt_rollout_frames: list[np.ndarray] = [] + for anchor in anchors: + gt_rollout_frames.extend( + _grid_rgb( + videos["top_head"][index], + videos["hand_left"][index], + videos["hand_right"][index], + ) + for index in range(anchor, anchor + 4) + ) + gt_path = args.output_dir / f"episode_{episode:06d}_ground_truth_rollout_aligned.mp4" + _save_video(gt_path, gt_rollout_frames) + + # The server's reset flushes its generated video under its configured + # output directory. Copy the latest file into this report directory and + # build the side-by-side video over the common frame range. + server_video_dir = args.server_video_dir.resolve() + server_videos = sorted(server_video_dir.glob("*.mp4"), key=lambda p: p.stat().st_mtime) + predicted_video_path = server_videos[-1] if server_videos else None + comparison_path = None + aligned_predicted_frames = 0 + raw_predicted_frames = 0 + if predicted_video_path is not None: + predicted_reader = imageio.get_reader(predicted_video_path) + raw_predicted = [np.asarray(frame) for frame in predicted_reader] + predicted_reader.close() + raw_predicted_frames = len(raw_predicted) + if not raw_predicted: + raise RuntimeError(f"Server video has no decoded frames: {predicted_video_path}") + target_frames = len(gt_rollout_frames) + if raw_predicted_frames == target_frames: + aligned_predicted = raw_predicted + else: + # Preserve the complete predicted trajectory but put it on the + # real 4-frames-per-policy-block clock. + indices = np.rint( + np.linspace(0, raw_predicted_frames - 1, target_frames) + ).astype(np.int64) + aligned_predicted = [raw_predicted[int(index)] for index in indices] + aligned_predicted_frames = len(aligned_predicted) + comparison_path = args.output_dir / f"episode_{episode:06d}_predicted_vs_ground_truth_full.mp4" + writer = imageio.get_writer(comparison_path, fps=30, codec="libx264", macro_block_size=None) + try: + for pred_rgb, gt_rgb in zip(aligned_predicted, gt_rollout_frames): + writer.append_data( + np.concatenate( + [_label(pred_rgb, "PREDICTED"), _label(gt_rgb, "G2 GROUND TRUTH")], + axis=1, + ) + ) + finally: + writer.close() + logging.info( + "Saved time-aligned comparison video raw_pred=%d aligned=%d gt=%d: %s", + raw_predicted_frames, + aligned_predicted_frames, + len(gt_rollout_frames), + comparison_path, + ) + + report = { + "episode_index": episode, + "num_rows": num_rows, + "num_blocks": len(starts), + "block_stride": args.block_stride, + "anchors": anchors, + "session_id": session_id, + "server_video": str(predicted_video_path) if predicted_video_path else None, + "server_video_raw_frames": raw_predicted_frames, + "comparison_video_aligned_frames": aligned_predicted_frames, + "video_alignment": "4 real frames per teacher-forced policy block; predicted raw latent video uniformly resampled", + "ground_truth_video": str(gt_path), + "comparison_video": str(comparison_path) if comparison_path else None, + **reports, + } + report_path = args.output_dir / f"episode_{episode:06d}_full_task_report.json" + report_path.write_text(json.dumps(report, ensure_ascii=False, indent=2), encoding="utf-8") + logging.info("Saved full-task report: %s", report_path) + + +if __name__ == "__main__": + logging.basicConfig(level=logging.INFO, force=True) + main() diff --git a/run_agibot_checkpoint_on_g2_testset.sh b/run_agibot_checkpoint_on_g2_testset.sh new file mode 100755 index 00000000..a933fa8c --- /dev/null +++ b/run_agibot_checkpoint_on_g2_testset.sh @@ -0,0 +1,89 @@ +#!/usr/bin/env bash +set -Eeuo pipefail + +# Replay one fixed G2 held-out test sample through an AgiBot/G1 checkpoint. +# Only future-video quality is evaluated. Native AgiBot actions are discarded. + +PROJECT_ROOT=${PROJECT_ROOT:-/home/ubuntu/projects/wangk/dreamzero} +EVAL_SCRIPT=${EVAL_SCRIPT:-$PROJECT_ROOT/eval_agibot_checkpoint_on_g2_testset.py} + +MODEL_PATH=${MODEL_PATH:-/data/wangk/checkpoints/dreamzero_agibot_fruit_lora_20k/checkpoint-12000} +TEST_DATA_ROOT=${TEST_DATA_ROOT:-/data/training_data/teleop/g2/g2_tasks_g1_g7_joint_gear_subtask_v2/test} +WAN_CKPT_DIR=${WAN_CKPT_DIR:-/data/wangk/checkpoints/Wan2.1-I2V-14B-480P} +TOKENIZER_PATH=${TOKENIZER_PATH:-/data/wangk/checkpoints/umt5-xxl} + +CHECKPOINT_NAME=$(basename "$MODEL_PATH") +OUTPUT_DIR=${OUTPUT_DIR:-/data/wangk/dreamzero/video_compare_agibot_vs_g1ft_vs_g2ft/$CHECKPOINT_NAME} + +GPU_IDS=${GPU_IDS:-0,1} +EPISODE_INDICES=${EPISODE_INDICES:-0} +FRAME_INDEX=${FRAME_INDEX:-30} +FUTURE_FRAMES=${FUTURE_FRAMES:-33} +NUM_INFERENCE_TIMESTEPS=${NUM_INFERENCE_TIMESTEPS:-4} +SEED=${SEED:-42} +PREFLIGHT_ONLY=${PREFLIGHT_ONLY:-0} +PROMPT_OVERRIDE=${PROMPT_OVERRIDE:-} + +fail() { + echo "[ERROR] $*" >&2 + exit 1 +} + +[[ -d "$PROJECT_ROOT" ]] || fail "Missing project root: $PROJECT_ROOT" +[[ -f "$EVAL_SCRIPT" ]] || fail "Missing evaluator: $EVAL_SCRIPT" +[[ -d "$TEST_DATA_ROOT" ]] || fail "Missing G2 test split: $TEST_DATA_ROOT" +[[ -d "$MODEL_PATH" ]] || fail "Missing checkpoint: $MODEL_PATH" +[[ -f "$MODEL_PATH/experiment_cfg/conf.yaml" ]] || fail "Missing checkpoint conf.yaml" +[[ -f "$MODEL_PATH/experiment_cfg/metadata.json" ]] || fail "Missing checkpoint metadata.json" +[[ -d "$WAN_CKPT_DIR" ]] || fail "Missing Wan checkpoint directory: $WAN_CKPT_DIR" +[[ -d "$TOKENIZER_PATH" ]] || fail "Missing tokenizer directory: $TOKENIZER_PATH" + +mkdir -p "$OUTPUT_DIR" +cd "$PROJECT_ROOT" + +ARGS=( + --model-path "$MODEL_PATH" + --wan-ckpt-dir "$WAN_CKPT_DIR" + --tokenizer-path "$TOKENIZER_PATH" + --test-data-root "$TEST_DATA_ROOT" + --output-dir "$OUTPUT_DIR" + --episode-indices "$EPISODE_INDICES" + --frame-index "$FRAME_INDEX" + --future-frames "$FUTURE_FRAMES" + --num-inference-timesteps "$NUM_INFERENCE_TIMESTEPS" + --seed "$SEED" + --embodiment-tag agibot +) + +if [[ -n "$PROMPT_OVERRIDE" ]]; then + ARGS+=(--prompt-override "$PROMPT_OVERRIDE") +fi + +if [[ "$PREFLIGHT_ONLY" == "1" ]]; then + echo "[INFO] Cheap preflight only; model will not be loaded." + python "$EVAL_SCRIPT" "${ARGS[@]}" --preflight-only + exit 0 +fi + +IFS=',' read -r -a GPU_ARRAY <<< "$GPU_IDS" +NPROC_PER_NODE=${#GPU_ARRAY[@]} +(( NPROC_PER_NODE == 1 || NPROC_PER_NODE == 2 )) \ + || fail "Expected one or two GPUs; got GPU_IDS=$GPU_IDS" + +export CUDA_VISIBLE_DEVICES="$GPU_IDS" +export TOKENIZERS_PARALLELISM=false + +echo "[INFO] checkpoint embodiment: agibot" +echo "[INFO] visual source: G2 held-out test videos" +echo "[INFO] checkpoint: $MODEL_PATH" +echo "[INFO] test split: $TEST_DATA_ROOT" +echo "[INFO] episodes: $EPISODE_INDICES" +echo "[INFO] frame: $FRAME_INDEX" +echo "[INFO] output: $OUTPUT_DIR" +echo "[INFO] GPUs: $GPU_IDS" + +python -m torch.distributed.run \ + --standalone \ + --nproc_per_node="$NPROC_PER_NODE" \ + "$EVAL_SCRIPT" \ + "${ARGS[@]}" diff --git a/run_agibot_fruit_testset_video_eval.sh b/run_agibot_fruit_testset_video_eval.sh new file mode 100755 index 00000000..751df5fa --- /dev/null +++ b/run_agibot_fruit_testset_video_eval.sh @@ -0,0 +1,86 @@ +#!/usr/bin/env bash +set -Eeuo pipefail + +PROJECT_ROOT=${PROJECT_ROOT:-/home/ubuntu/projects/wangk/dreamzero} +EVAL_SCRIPT=${EVAL_SCRIPT:-$PROJECT_ROOT/eval_agibot_checkpoint_on_fruit_testset.py} + +MODEL_PATH=${MODEL_PATH:-/data/wangk/checkpoints/dreamzero_agibot_fruit_lora_20k/checkpoint-4000} +TEST_DATA_ROOT=${TEST_DATA_ROOT:-} +WAN_CKPT_DIR=${WAN_CKPT_DIR:-/data/wangk/checkpoints/Wan2.1-I2V-14B-480P} +TOKENIZER_PATH=${TOKENIZER_PATH:-/data/wangk/checkpoints/umt5-xxl} + +CHECKPOINT_NAME=$(basename "$MODEL_PATH") +OUTPUT_DIR=${OUTPUT_DIR:-/data/wangk/dreamzero/agibot_fruit_testset_video_eval/$CHECKPOINT_NAME} + +GPU_IDS=${GPU_IDS:-0,1} +EPISODE_INDICES=${EPISODE_INDICES:-0} +FRAME_INDEX=${FRAME_INDEX:--1} +FUTURE_FRAMES=${FUTURE_FRAMES:-33} +NUM_INFERENCE_TIMESTEPS=${NUM_INFERENCE_TIMESTEPS:-4} +SEED=${SEED:-42} +PREFLIGHT_ONLY=${PREFLIGHT_ONLY:-0} +PROMPT_OVERRIDE=${PROMPT_OVERRIDE:-} + +fail() { + echo "[ERROR] $*" >&2 + exit 1 +} + +[[ -n "$TEST_DATA_ROOT" ]] || fail \ + "TEST_DATA_ROOT is empty. Set it to the packaged AgiBot fruit test split." +[[ -d "$PROJECT_ROOT" ]] || fail "Missing project root: $PROJECT_ROOT" +[[ -f "$EVAL_SCRIPT" ]] || fail "Missing evaluator: $EVAL_SCRIPT" +[[ -d "$TEST_DATA_ROOT" ]] || fail "Missing fruit test split: $TEST_DATA_ROOT" +[[ -d "$MODEL_PATH" ]] || fail "Missing checkpoint: $MODEL_PATH" +[[ -f "$MODEL_PATH/experiment_cfg/conf.yaml" ]] || fail "Missing checkpoint conf.yaml" +[[ -f "$MODEL_PATH/experiment_cfg/metadata.json" ]] || fail "Missing checkpoint metadata.json" +[[ -d "$WAN_CKPT_DIR" ]] || fail "Missing Wan checkpoint directory: $WAN_CKPT_DIR" +[[ -d "$TOKENIZER_PATH" ]] || fail "Missing tokenizer directory: $TOKENIZER_PATH" + +mkdir -p "$OUTPUT_DIR" +cd "$PROJECT_ROOT" + +ARGS=( + --model-path "$MODEL_PATH" + --wan-ckpt-dir "$WAN_CKPT_DIR" + --tokenizer-path "$TOKENIZER_PATH" + --test-data-root "$TEST_DATA_ROOT" + --output-dir "$OUTPUT_DIR" + --episode-indices "$EPISODE_INDICES" + --frame-index "$FRAME_INDEX" + --future-frames "$FUTURE_FRAMES" + --num-inference-timesteps "$NUM_INFERENCE_TIMESTEPS" + --seed "$SEED" + --embodiment-tag agibot +) + +if [[ -n "$PROMPT_OVERRIDE" ]]; then + ARGS+=(--prompt-override "$PROMPT_OVERRIDE") +fi + +if [[ "$PREFLIGHT_ONLY" == "1" ]]; then + echo "[INFO] Cheap preflight only; model will not be loaded." + python "$EVAL_SCRIPT" "${ARGS[@]}" --preflight-only + exit 0 +fi + +IFS=',' read -r -a GPU_ARRAY <<< "$GPU_IDS" +NPROC_PER_NODE=${#GPU_ARRAY[@]} +(( NPROC_PER_NODE == 1 || NPROC_PER_NODE == 2 )) \ + || fail "Expected one or two GPUs; got GPU_IDS=$GPU_IDS" + +export CUDA_VISIBLE_DEVICES="$GPU_IDS" +export TOKENIZERS_PARALLELISM=false + +echo "[INFO] checkpoint: $MODEL_PATH" +echo "[INFO] fruit test split: $TEST_DATA_ROOT" +echo "[INFO] episodes: $EPISODE_INDICES" +echo "[INFO] frame: $FRAME_INDEX" +echo "[INFO] output: $OUTPUT_DIR" +echo "[INFO] GPUs: $GPU_IDS" + +python -m torch.distributed.run \ + --standalone \ + --nproc_per_node="$NPROC_PER_NODE" \ + "$EVAL_SCRIPT" \ + "${ARGS[@]}" diff --git a/run_g2_testset_video_eval.sh b/run_g2_testset_video_eval.sh new file mode 100644 index 00000000..112d5e9a --- /dev/null +++ b/run_g2_testset_video_eval.sh @@ -0,0 +1,118 @@ +#!/usr/bin/env bash +set -Eeuo pipefail + +# Offline G2 held-out test-set video prediction. +# No robot, WebSocket client, JPEG transport, or SDK action execution. + +PROJECT_ROOT=${PROJECT_ROOT:-/home/ubuntu/projects/wangk/dreamzero} +PYTHON_BIN=${PYTHON_BIN:-/data/wangk/conda/envs/dreamzero/bin/python} +EVAL_SCRIPT=${EVAL_SCRIPT:-$PROJECT_ROOT/eval_g2_checkpoint_on_testset.py} +TEST_DATA_ROOT=${TEST_DATA_ROOT:-/data/training_data/teleop/g2/g2_tasks_g1_g7_joint_gear_subtask_v2/test} + +# Change this path when sweeping joint/video-only checkpoints. +MODEL_PATH=${MODEL_PATH:-/data/wangk/checkpoints/dreamzero_g2_video_only_lora_1k/checkpoint-1000} + +WAN_CKPT_DIR=${WAN_CKPT_DIR:-/data/wangk/checkpoints/Wan2.1-I2V-14B-480P} +TOKENIZER_PATH=${TOKENIZER_PATH:-/data/wangk/checkpoints/umt5-xxl} + +RUN_NAME="$(basename "$(dirname "$MODEL_PATH")")_$(basename "$MODEL_PATH")" +OUTPUT_DIR=${OUTPUT_DIR:-/data/wangk/dreamzero/g2_testset_video_eval/$RUN_NAME} + +GPU_IDS=${GPU_IDS:-0,1} +EPISODE_INDICES=${EPISODE_INDICES:-0} +FRAME_INDEX=${FRAME_INDEX:-30} +FUTURE_FRAMES=${FUTURE_FRAMES:-33} +NUM_INFERENCE_TIMESTEPS=${NUM_INFERENCE_TIMESTEPS:-16} +NUM_DIT_STEPS=${NUM_DIT_STEPS:-8} +SEED=${SEED:-42} +PREFLIGHT_ONLY=${PREFLIGHT_ONLY:-0} +PROMPT_OVERRIDE=${PROMPT_OVERRIDE:-} +WINDOWED=${WINDOWED:-0} +WINDOW_HISTORY=${WINDOW_HISTORY:-4} +WINDOW_STRIDE=${WINDOW_STRIDE:-4} +WINDOW_STARTS=${WINDOW_STARTS:-0,4,8,12,16,20,24,28} +ROLLOUT_FUTURE_FRAMES=${ROLLOUT_FUTURE_FRAMES:-33} +ROLLOUT_BLOCKS=${ROLLOUT_BLOCKS:-4} +SAVE_WINDOW_ARTIFACTS=${SAVE_WINDOW_ARTIFACTS:-0} + +fail() { + echo "[ERROR] $*" >&2 + exit 1 +} + +[[ -x "$PYTHON_BIN" ]] || fail "Python is not executable: $PYTHON_BIN" +[[ -d "$PROJECT_ROOT" ]] || fail "Missing project root: $PROJECT_ROOT" +[[ -f "$EVAL_SCRIPT" ]] || fail "Missing evaluator: $EVAL_SCRIPT" +[[ -d "$TEST_DATA_ROOT" ]] || fail "Missing G2 test split: $TEST_DATA_ROOT" +[[ -f "$TEST_DATA_ROOT/meta/info.json" ]] || fail "Missing test info.json" +[[ -f "$TEST_DATA_ROOT/meta/modality.json" ]] || fail "Missing test modality.json" +[[ -f "$TEST_DATA_ROOT/meta/embodiment.json" ]] || fail "Missing test embodiment.json" +[[ -d "$MODEL_PATH" ]] || fail "Missing checkpoint: $MODEL_PATH" +[[ -f "$MODEL_PATH/experiment_cfg/conf.yaml" ]] || fail "Missing checkpoint conf.yaml" +[[ -f "$MODEL_PATH/experiment_cfg/metadata.json" ]] || fail "Missing checkpoint metadata.json" +[[ -d "$WAN_CKPT_DIR" ]] || fail "Missing Wan checkpoint directory: $WAN_CKPT_DIR" +[[ -d "$TOKENIZER_PATH" ]] || fail "Missing tokenizer directory: $TOKENIZER_PATH" + +mkdir -p "$OUTPUT_DIR" +cd "$PROJECT_ROOT" + +ARGS=( + --model-path "$MODEL_PATH" + --wan-ckpt-dir "$WAN_CKPT_DIR" + --tokenizer-path "$TOKENIZER_PATH" + --test-data-root "$TEST_DATA_ROOT" + --output-dir "$OUTPUT_DIR" + --episode-indices "$EPISODE_INDICES" + --frame-index "$FRAME_INDEX" + --future-frames "$FUTURE_FRAMES" + --num-inference-timesteps "$NUM_INFERENCE_TIMESTEPS" + --seed "$SEED" + --embodiment-tag g2 +) + +if [[ -n "$PROMPT_OVERRIDE" ]]; then + ARGS+=(--prompt-override "$PROMPT_OVERRIDE") +fi +if [[ -n "$NUM_DIT_STEPS" ]]; then + ARGS+=(--num-dit-steps "$NUM_DIT_STEPS") +fi +if [[ "$WINDOWED" == "1" ]]; then + ARGS+=( + --windowed + --window-history "$WINDOW_HISTORY" + --window-stride "$WINDOW_STRIDE" + --window-starts "$WINDOW_STARTS" + --rollout-future-frames "$ROLLOUT_FUTURE_FRAMES" + --rollout-blocks "$ROLLOUT_BLOCKS" + ) + if [[ "$SAVE_WINDOW_ARTIFACTS" == "1" ]]; then + ARGS+=(--save-window-artifacts) + fi +fi + +if [[ "$PREFLIGHT_ONLY" == "1" ]]; then + echo "[INFO] Running CPU preflight only; the 14B model will not be loaded." + "$PYTHON_BIN" "$EVAL_SCRIPT" "${ARGS[@]}" --preflight-only + exit 0 +fi + +IFS=',' read -r -a GPU_ARRAY <<< "$GPU_IDS" +NPROC_PER_NODE=${#GPU_ARRAY[@]} +(( NPROC_PER_NODE == 1 || NPROC_PER_NODE == 2 )) \ + || fail "DreamZero inference supports 1 or 2 GPUs here; got GPU_IDS=$GPU_IDS" + +export CUDA_VISIBLE_DEVICES="$GPU_IDS" +export TOKENIZERS_PARALLELISM=false + +echo "[INFO] checkpoint: $MODEL_PATH" +echo "[INFO] test split: $TEST_DATA_ROOT" +echo "[INFO] episodes: $EPISODE_INDICES" +echo "[INFO] frame: $FRAME_INDEX" +echo "[INFO] output: $OUTPUT_DIR" +echo "[INFO] GPUs: $GPU_IDS" + +"$PYTHON_BIN" -m torch.distributed.run \ + --standalone \ + --nproc_per_node="$NPROC_PER_NODE" \ + "$EVAL_SCRIPT" \ + "${ARGS[@]}" diff --git a/scripts/audit_g2_checkpoint.py b/scripts/audit_g2_checkpoint.py new file mode 100644 index 00000000..0904c939 --- /dev/null +++ b/scripts/audit_g2_checkpoint.py @@ -0,0 +1,189 @@ +#!/usr/bin/env python3 +"""CPU-only validation of a DreamZero G2 LoRA deployment checkpoint.""" + +from __future__ import annotations + +import argparse +import json +import re +import sys +from collections import Counter +from pathlib import Path + + +EXPECTED_ACTION_HORIZON = 24 +EXPECTED_NUM_FRAMES = 33 +EXPECTED_ACTION_DIM = 32 +EXPECTED_OUTPUT_DIM = 16 + + +def _read_safetensors_header(path: Path) -> dict: + with path.open("rb") as stream: + header_size_bytes = stream.read(8) + if len(header_size_bytes) != 8: + raise ValueError(f"{path} has an incomplete safetensors header") + header_size = int.from_bytes(header_size_bytes, "little") + if header_size <= 0 or header_size > path.stat().st_size - 8: + raise ValueError( + f"{path} has an invalid safetensors header size: {header_size}" + ) + return json.loads(stream.read(header_size)) + + +def _nested(config: dict, *keys: str): + value = config + for key in keys: + if not isinstance(value, dict) or key not in value: + return None + value = value[key] + return value + + +def audit(checkpoint: Path) -> list[str]: + errors: list[str] = [] + required = [ + checkpoint / "config.json", + checkpoint / "model.safetensors", + checkpoint / "experiment_cfg" / "conf.yaml", + checkpoint / "experiment_cfg" / "metadata.json", + ] + for path in required: + if not path.is_file(): + errors.append(f"missing required file: {path}") + if errors: + return errors + + config = json.loads((checkpoint / "config.json").read_text()) + inner = _nested(config, "action_head_cfg", "config") or {} + diffusion = inner.get("diffusion_model_cfg", {}) + checks = { + "config.action_horizon": ( + config.get("action_horizon"), + EXPECTED_ACTION_HORIZON, + ), + "action_head.action_horizon": ( + inner.get("action_horizon"), + EXPECTED_ACTION_HORIZON, + ), + "action_head.num_frames": ( + inner.get("num_frames"), + EXPECTED_NUM_FRAMES, + ), + "action_head.action_dim": ( + inner.get("action_dim"), + EXPECTED_ACTION_DIM, + ), + "diffusion.num_action_per_block": ( + diffusion.get("num_action_per_block"), + EXPECTED_ACTION_HORIZON, + ), + "diffusion.out_dim": ( + diffusion.get("out_dim"), + EXPECTED_OUTPUT_DIM, + ), + } + for name, (actual, expected) in checks.items(): + if actual != expected: + errors.append(f"{name}: expected {expected}, got {actual!r}") + + conf_text = ( + checkpoint / "experiment_cfg" / "conf.yaml" + ).read_text(errors="replace") + # Support both old LoRA-only checkpoints and new action-adapter-only + # checkpoints. The latter intentionally freezes shared DiT/LoRA and only + # stores state/action adapter parameters. + pretrained_match = re.search( + r"(?m)^pretrained_model_path:\s*(\S+)\s*$", + conf_text, + ) + if not pretrained_match: + errors.append("experiment config has no pretrained_model_path") + else: + base_path = Path(pretrained_match.group(1)) + if not base_path.is_dir(): + errors.append(f"pretrained base directory is missing: {base_path}") + elif not ( + (base_path / "model.safetensors").is_file() + or (base_path / "model.safetensors.index.json").is_file() + ): + errors.append( + f"pretrained base has no safetensors weights: {base_path}" + ) + + header = _read_safetensors_header(checkpoint / "model.safetensors") + keys = [key for key in header if key != "__metadata__"] + buckets = Counter() + for key in keys: + lowered = key.lower() + if "lora_a" in lowered: + buckets["lora_A"] += 1 + elif "lora_b" in lowered: + buckets["lora_B"] += 1 + elif "action_encoder" in lowered: + buckets["action_encoder"] += 1 + elif "action_decoder" in lowered: + buckets["action_decoder"] += 1 + elif "state_encoder" in lowered: + buckets["state_encoder"] += 1 + else: + buckets["other"] += 1 + + has_lora = buckets["lora_A"] > 0 or buckets["lora_B"] > 0 + has_action_adapter = ( + buckets["action_encoder"] > 0 + or buckets["action_decoder"] > 0 + or buckets["state_encoder"] > 0 + ) + + if has_lora: + print("checkpoint_type: LoRA") + if buckets["lora_A"] == 0: + errors.append("checkpoint contains no LoRA A weights") + if buckets["lora_A"] != buckets["lora_B"]: + errors.append( + "unbalanced LoRA weights: " + f"A={buckets['lora_A']} B={buckets['lora_B']}" + ) + elif has_action_adapter: + print("checkpoint_type: action_adapter") + else: + errors.append( + "checkpoint contains neither LoRA weights nor action adapter weights" + ) + for name, minimum in ( + ("action_encoder", 6), + ("action_decoder", 4), + ("state_encoder", 4), + ): + if buckets[name] < minimum: + errors.append( + f"incomplete {name}: expected at least {minimum}, " + f"got {buckets[name]}" + ) + + print(f"checkpoint: {checkpoint}") + print(f"pretrained_model_path: {pretrained_match.group(1) if pretrained_match else 'MISSING'}") + print(f"tensor_keys: {len(keys)}") + print(f"tensor_buckets: {dict(buckets)}") + for name, (actual, expected) in checks.items(): + print(f"{name}: {actual!r} (expected {expected})") + return errors + + +def main() -> int: + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint", type=Path) + args = parser.parse_args() + + errors = audit(args.checkpoint.expanduser().resolve()) + if errors: + print("G2 CHECKPOINT AUDIT FAILED:", file=sys.stderr) + for error in errors: + print(f" - {error}", file=sys.stderr) + return 1 + print("G2 CHECKPOINT AUDIT PASSED") + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/scripts/audit_g2_checkpoint_action_adapter.py b/scripts/audit_g2_checkpoint_action_adapter.py new file mode 100644 index 00000000..42193b5c --- /dev/null +++ b/scripts/audit_g2_checkpoint_action_adapter.py @@ -0,0 +1,173 @@ +#!/usr/bin/env python3 +"""CPU-only validation of a DreamZero G2 LoRA deployment checkpoint.""" + +from __future__ import annotations + +import argparse +import json +import re +import sys +from collections import Counter +from pathlib import Path + + +EXPECTED_ACTION_HORIZON = 24 +EXPECTED_NUM_FRAMES = 33 +EXPECTED_ACTION_DIM = 32 +EXPECTED_OUTPUT_DIM = 16 + + +def _read_safetensors_header(path: Path) -> dict: + with path.open("rb") as stream: + header_size_bytes = stream.read(8) + if len(header_size_bytes) != 8: + raise ValueError(f"{path} has an incomplete safetensors header") + header_size = int.from_bytes(header_size_bytes, "little") + if header_size <= 0 or header_size > path.stat().st_size - 8: + raise ValueError( + f"{path} has an invalid safetensors header size: {header_size}" + ) + return json.loads(stream.read(header_size)) + + +def _nested(config: dict, *keys: str): + value = config + for key in keys: + if not isinstance(value, dict) or key not in value: + return None + value = value[key] + return value + + +def audit(checkpoint: Path) -> list[str]: + errors: list[str] = [] + required = [ + checkpoint / "config.json", + checkpoint / "model.safetensors", + checkpoint / "experiment_cfg" / "conf.yaml", + checkpoint / "experiment_cfg" / "metadata.json", + ] + for path in required: + if not path.is_file(): + errors.append(f"missing required file: {path}") + if errors: + return errors + + config = json.loads((checkpoint / "config.json").read_text()) + inner = _nested(config, "action_head_cfg", "config") or {} + diffusion = inner.get("diffusion_model_cfg", {}) + checks = { + "config.action_horizon": ( + config.get("action_horizon"), + EXPECTED_ACTION_HORIZON, + ), + "action_head.action_horizon": ( + inner.get("action_horizon"), + EXPECTED_ACTION_HORIZON, + ), + "action_head.num_frames": ( + inner.get("num_frames"), + EXPECTED_NUM_FRAMES, + ), + "action_head.action_dim": ( + inner.get("action_dim"), + EXPECTED_ACTION_DIM, + ), + "diffusion.num_action_per_block": ( + diffusion.get("num_action_per_block"), + EXPECTED_ACTION_HORIZON, + ), + "diffusion.out_dim": ( + diffusion.get("out_dim"), + EXPECTED_OUTPUT_DIM, + ), + } + for name, (actual, expected) in checks.items(): + if actual != expected: + errors.append(f"{name}: expected {expected}, got {actual!r}") + + conf_text = ( + checkpoint / "experiment_cfg" / "conf.yaml" + ).read_text(errors="replace") + if not re.search(r"(?m)^save_lora_only:\s*true\s*$", conf_text): + errors.append("experiment config does not declare save_lora_only: true") + pretrained_match = re.search( + r"(?m)^pretrained_model_path:\s*(\S+)\s*$", + conf_text, + ) + if not pretrained_match: + errors.append("experiment config has no pretrained_model_path") + else: + base_path = Path(pretrained_match.group(1)) + if not base_path.is_dir(): + errors.append(f"pretrained base directory is missing: {base_path}") + elif not ( + (base_path / "model.safetensors").is_file() + or (base_path / "model.safetensors.index.json").is_file() + ): + errors.append( + f"pretrained base has no safetensors weights: {base_path}" + ) + + header = _read_safetensors_header(checkpoint / "model.safetensors") + keys = [key for key in header if key != "__metadata__"] + buckets = Counter() + for key in keys: + lowered = key.lower() + if "lora_a" in lowered: + buckets["lora_A"] += 1 + elif "lora_b" in lowered: + buckets["lora_B"] += 1 + elif "action_encoder" in lowered: + buckets["action_encoder"] += 1 + elif "action_decoder" in lowered: + buckets["action_decoder"] += 1 + elif "state_encoder" in lowered: + buckets["state_encoder"] += 1 + else: + buckets["other"] += 1 + + if buckets["lora_A"] == 0: + errors.append("checkpoint contains no LoRA A weights") + if buckets["lora_A"] != buckets["lora_B"]: + errors.append( + "unbalanced LoRA weights: " + f"A={buckets['lora_A']} B={buckets['lora_B']}" + ) + for name, minimum in ( + ("action_encoder", 6), + ("action_decoder", 4), + ("state_encoder", 4), + ): + if buckets[name] < minimum: + errors.append( + f"incomplete {name}: expected at least {minimum}, " + f"got {buckets[name]}" + ) + + print(f"checkpoint: {checkpoint}") + print(f"pretrained_model_path: {pretrained_match.group(1) if pretrained_match else 'MISSING'}") + print(f"tensor_keys: {len(keys)}") + print(f"tensor_buckets: {dict(buckets)}") + for name, (actual, expected) in checks.items(): + print(f"{name}: {actual!r} (expected {expected})") + return errors + + +def main() -> int: + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint", type=Path) + args = parser.parse_args() + + errors = audit(args.checkpoint.expanduser().resolve()) + if errors: + print("G2 CHECKPOINT AUDIT FAILED:", file=sys.stderr) + for error in errors: + print(f" - {error}", file=sys.stderr) + return 1 + print("G2 CHECKPOINT AUDIT PASSED") + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) \ No newline at end of file diff --git a/scripts/data/build_g2_active_hold_windows.py b/scripts/data/build_g2_active_hold_windows.py new file mode 100644 index 00000000..0bda8320 --- /dev/null +++ b/scripts/data/build_g2_active_hold_windows.py @@ -0,0 +1,80 @@ +#!/usr/bin/env python3 +"""Build the active/hold sampling index for an existing policy-space G2 split.""" + +from __future__ import annotations + +import argparse +import json +from pathlib import Path + +import numpy as np +import pandas as pd + + +ARM_DIMS = (0, 1, 2, 3, 4, 5, 6, 8, 9, 10, 11, 12, 13, 14) +GRIP_DIMS = (7, 15) + + +def main() -> None: + parser = argparse.ArgumentParser() + parser.add_argument("split", type=Path) + parser.add_argument("--action-horizon", type=int, default=24) + parser.add_argument("--overwrite", action="store_true") + args = parser.parse_args() + + output = args.split / "meta/g2_active_hold_windows.json" + if output.exists() and not args.overwrite: + raise FileExistsError(f"Refusing to overwrite {output}; pass --overwrite") + + metrics: list[tuple[int, int, float, bool]] = [] + for path in sorted(args.split.glob("data/chunk-*/*.parquet")): + frame = pd.read_parquet(path, columns=["observation.state", "action", "episode_index"]) + state = np.stack(frame["observation.state"]).astype(np.float32) + action = np.stack(frame["action"]).astype(np.float32) + if state.shape[1:] != (16,) or action.shape[1:] != (16,): + raise ValueError(f"{path}: expected 16D state/action") + if np.any(state[:, GRIP_DIMS] < 0) or np.any(state[:, GRIP_DIMS] > 1): + raise ValueError(f"{path}: state gripper is not in policy space [0,1]") + if np.any(action[:, GRIP_DIMS] < 0) or np.any(action[:, GRIP_DIMS] > 1): + raise ValueError(f"{path}: action gripper is not in policy space [0,1]") + episode = int(frame["episode_index"].iloc[0]) + for step in range(max(0, len(frame) - args.action_horizon + 1)): + future = action[step : step + args.action_horizon] + arm_motion = float( + np.linalg.norm(future[:, ARM_DIMS] - state[step, ARM_DIMS], axis=1).max() + ) + grip_transition = bool( + np.any(np.abs(future[:, GRIP_DIMS] - state[step, GRIP_DIMS]) > 0.25) + ) + metrics.append((episode, step, arm_motion, grip_transition)) + + nonzero = np.asarray([motion for _, _, motion, _ in metrics if motion > 1e-8]) + if not nonzero.size: + raise RuntimeError("No nonzero G2 arm motion found") + threshold = float(np.quantile(nonzero, 0.30)) + active: dict[str, list[int]] = {} + hold: dict[str, list[int]] = {} + for episode, step, motion, grip_transition in metrics: + target = active if motion >= threshold or grip_transition else hold + target.setdefault(str(episode), []).append(step) + + payload = { + "schema_version": 1, + "action_horizon": args.action_horizon, + "arm_motion_threshold_rule": "p30_nonzero_24_step_l2", + "arm_motion_threshold": threshold, + "gripper_transition_threshold": 0.25, + "active_count": sum(map(len, active.values())), + "hold_count": sum(map(len, hold.values())), + "active": active, + "hold": hold, + } + output.write_text(json.dumps(payload, indent=2), encoding="utf-8") + print( + f"Wrote {output}: threshold={threshold:.8f} " + f"active={payload['active_count']} hold={payload['hold_count']}" + ) + + +if __name__ == "__main__": + main() diff --git a/scripts/data/convert_lerobot_g2_ee.py b/scripts/data/convert_lerobot_g2_ee.py new file mode 100755 index 00000000..47aab7bb --- /dev/null +++ b/scripts/data/convert_lerobot_g2_ee.py @@ -0,0 +1,891 @@ +#!/usr/bin/env python3 +"""Convert the FruitPackaging LeRobot v3 dataset to DreamZero's LeRobot v2 layout. + +The source dataset stores many episodes in shared parquet and video shards. DreamZero's +current loader expects one parquet and one video per episode, plus GEAR metadata files. +This converter preserves all samples, converts source xyzw quaternions to rotation-6d, +materializes the task text, splits videos on episode boundaries, and computes stats. +""" + +from __future__ import annotations + +import argparse +from concurrent.futures import ThreadPoolExecutor, as_completed +import json +import logging +from pathlib import Path +import shutil +import subprocess +import sys +from typing import Any + +import numpy as np +import pyarrow as pa +import pyarrow.compute as pc +import pyarrow.parquet as pq + + +LOG = logging.getLogger("fruit-v3-converter") + +SOURCE_VECTOR_NAMES = [ + "l.ee.x", "l.ee.y", "l.ee.z", + "l.ee.qx", "l.ee.qy", "l.ee.qz", "l.ee.qw", + "l.ee.gripper.pos", + "r.ee.x", "r.ee.y", "r.ee.z", + "r.ee.qx", "r.ee.qy", "r.ee.qz", "r.ee.qw", + "r.ee.gripper.pos", +] + +# PyTorch3D-compatible rotation-6d convention: flatten the first two matrix rows. +# Conversion is implemented in NumPy so DreamZero training does not need PyTorch3D. +OUTPUT_VECTOR_NAMES = [ + "l.ee.x", "l.ee.y", "l.ee.z", + "l.ee.rot6d.r0c0", "l.ee.rot6d.r0c1", "l.ee.rot6d.r0c2", + "l.ee.rot6d.r1c0", "l.ee.rot6d.r1c1", "l.ee.rot6d.r1c2", + "l.ee.gripper.pos", + "r.ee.x", "r.ee.y", "r.ee.z", + "r.ee.rot6d.r0c0", "r.ee.rot6d.r0c1", "r.ee.rot6d.r0c2", + "r.ee.rot6d.r1c0", "r.ee.rot6d.r1c1", "r.ee.rot6d.r1c2", + "r.ee.gripper.pos", +] + +OUTPUT_VECTOR_DIM = 20 +ROTATION_OUTPUT_SLICES = ((3, 9), (13, 19)) + +CAMERAS = { + "observation.images.head_color": "observation.images.top_head", + "observation.images.hand_left_color": "observation.images.hand_left", + "observation.images.hand_right_color": "observation.images.hand_right", +} + +STATE_ACTION_FIELDS = { + "left_effector_position": { + "bounds": (0, 3), + "rotation_type": None, + }, + "left_effector_rotation": { + "bounds": (3, 9), + "rotation_type": "rotation_6d", + }, + "left_gripper_position": { + "bounds": (9, 10), + "rotation_type": None, + }, + "right_effector_position": { + "bounds": (10, 13), + "rotation_type": None, + }, + "right_effector_rotation": { + "bounds": (13, 19), + "rotation_type": "rotation_6d", + }, + "right_gripper_position": { + "bounds": (19, 20), + "rotation_type": None, + }, +} + +RELATIVE_ACTION_KEYS = ( + "left_effector_position", + "left_gripper_position", + "right_effector_position", + "right_gripper_position", +) + + +def nonempty_path(value: str) -> Path: + """Reject an unset shell variable such as --output "$TEST_DATA".""" + if not value.strip(): + raise argparse.ArgumentTypeError("path cannot be empty") + return Path(value) + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--source", type=nonempty_path, required=True) + parser.add_argument("--output", type=nonempty_path, required=True) + parser.add_argument("--workers", type=int, default=4) + parser.add_argument("--action-horizon", type=int, default=24) + parser.add_argument("--video-preset", default="veryfast") + parser.add_argument("--video-crf", type=int, default=18) + parser.add_argument( + "--video-width", + type=int, + default=None, + help="Resize every output camera to this width (use with --video-height).", + ) + parser.add_argument( + "--video-height", + type=int, + default=None, + help="Resize every output camera to this height (use with --video-width).", + ) + parser.add_argument( + "--max-episodes", + type=int, + default=None, + help="Convert only the first N episodes (intended for converter tests).", + ) + parser.add_argument("--overwrite", action="store_true") + return parser.parse_args() + + +def load_json(path: Path) -> dict[str, Any]: + with path.open() as handle: + return json.load(handle) + + +def write_json(path: Path, value: Any) -> None: + path.parent.mkdir(parents=True, exist_ok=True) + with path.open("w") as handle: + json.dump(value, handle, indent=2, ensure_ascii=True) + handle.write("\n") + + +def check_tools() -> None: + for executable in ("ffmpeg", "ffprobe"): + if shutil.which(executable) is None: + raise RuntimeError(f"Required executable not found: {executable}") + + +def validate_output_path(source: Path, output: Path) -> None: + """Refuse paths where --overwrite could delete source code or source data.""" + cwd = Path.cwd().resolve() + home = Path.home().resolve() + filesystem_root = Path(output.anchor).resolve() + + if output in {filesystem_root, home}: + raise ValueError(f"Refusing dangerous output path: {output}") + + if output == cwd or output in cwd.parents: + raise ValueError( + f"Refusing output path {output}: it is the current directory or its parent" + ) + + if output == source or output in source.parents or source in output.parents: + raise ValueError( + f"Refusing output path {output}: it overlaps source dataset {source}" + ) + + if (output / ".git").exists(): + raise ValueError(f"Refusing to overwrite Git repository: {output}") + + +def source_table_paths(source: Path) -> tuple[Path, Path, Path]: + data_files = sorted((source / "data").glob("chunk-*/*.parquet")) + episode_files = sorted((source / "meta/episodes").glob("chunk-*/*.parquet")) + tasks_path = source / "meta/tasks.parquet" + if len(data_files) != 1: + raise ValueError(f"Expected one source data parquet, found {len(data_files)}") + if len(episode_files) != 1: + raise ValueError(f"Expected one source episode parquet, found {len(episode_files)}") + if not tasks_path.is_file(): + raise FileNotFoundError(tasks_path) + return data_files[0], episode_files[0], tasks_path + + +def validate_source( + info: dict[str, Any], data_file: Path, episode_rows: list[dict[str, Any]], max_episodes: int | None +) -> int: + if info.get("codebase_version") != "v3.0": + raise ValueError(f"Expected LeRobot v3.0, got {info.get('codebase_version')!r}") + for key in ("observation.state", "action", *CAMERAS): + if key not in info["features"]: + raise ValueError(f"Missing source feature: {key}") + if info["features"]["observation.state"]["shape"] != [16]: + raise ValueError("Fruit converter expects a 16-dimensional state") + if info["features"]["action"]["shape"] != [16]: + raise ValueError("Fruit converter expects a 16-dimensional action") + for key in ("observation.state", "action"): + names = info["features"][key].get("names") + if names != SOURCE_VECTOR_NAMES: + raise ValueError( + f"{key} names do not match the expected G2 EE xyzw layout: " + f"expected={SOURCE_VECTOR_NAMES}, actual={names}" + ) + + total = int(info["total_episodes"]) + if len(episode_rows) != total: + raise ValueError(f"Episode metadata has {len(episode_rows)} rows, expected {total}") + parquet = pq.ParquetFile(data_file) + if parquet.metadata.num_rows != int(info["total_frames"]): + raise ValueError("Data parquet frame count does not match info.json") + metadata_frames = sum(int(row["length"]) for row in episode_rows) + if metadata_frames != parquet.metadata.num_rows: + raise ValueError( + "Episode metadata lengths do not add up to the data parquet row count: " + f"{metadata_frames} != {parquet.metadata.num_rows}" + ) + + episode_index_table = pq.read_table(data_file, columns=["episode_index"]) + episode_ids = np.asarray( + episode_index_table["episode_index"].combine_chunks().to_numpy(), + dtype=np.int64, + ) + expected_ids = np.asarray( + [int(row["episode_index"]) for row in episode_rows], + dtype=np.int64, + ) + actual_ids, actual_counts = np.unique(episode_ids, return_counts=True) + if not np.array_equal(actual_ids, expected_ids): + raise ValueError( + "Episode ids in the data parquet do not match meta/episodes: " + f"data={actual_ids.tolist()}, metadata={expected_ids.tolist()}" + ) + expected_counts = np.asarray( + [int(row["length"]) for row in episode_rows], + dtype=np.int64, + ) + if not np.array_equal(actual_counts, expected_counts): + raise ValueError( + "Per-episode frame counts in the data parquet do not match meta/episodes: " + f"data={actual_counts.tolist()}, metadata={expected_counts.tolist()}" + ) + if max_episodes is not None: + if max_episodes < 1 or max_episodes > total: + raise ValueError(f"--max-episodes must be in [1, {total}]") + return max_episodes + return total + + +def feature_video_info( + source_feature: dict[str, Any], video_width: int | None, video_height: int | None +) -> dict[str, Any]: + shape = ( + [video_height, video_width, 3] + if video_width is not None and video_height is not None + else list(source_feature["shape"]) + ) + source_info = source_feature.get("info", source_feature.get("video_info", {})) + return { + "dtype": "video", + "shape": shape, + "names": ["height", "width", "channel"], + "video_info": { + "video.fps": float(source_info.get("video.fps", 30.0)), + "video.codec": "h264", + "video.pix_fmt": "yuv420p", + "video.is_depth_map": False, + "has_audio": False, + }, + } + + +def build_info( + source_info: dict[str, Any], + episode_rows: list[dict[str, Any]], + episode_count: int, + video_width: int | None, + video_height: int | None, +) -> dict[str, Any]: + total_frames = sum(int(row["length"]) for row in episode_rows[:episode_count]) + task_ids = { + int(row["episode_index"]): tuple(row["tasks"]) + for row in episode_rows[:episode_count] + } + features: dict[str, Any] = { + "observation.state": { + "dtype": "float32", + "shape": [OUTPUT_VECTOR_DIM], + "names": OUTPUT_VECTOR_NAMES, + }, + "action": { + "dtype": "float32", + "shape": [OUTPUT_VECTOR_DIM], + "names": OUTPUT_VECTOR_NAMES, + }, + "annotation.task": {"dtype": "string", "shape": [1]}, + "timestamp": {"dtype": "float32", "shape": [1]}, + "frame_index": {"dtype": "int64", "shape": [1]}, + "episode_index": {"dtype": "int64", "shape": [1]}, + "index": {"dtype": "int64", "shape": [1]}, + "task_index": {"dtype": "int64", "shape": [1]}, + } + for source_key, output_key in CAMERAS.items(): + features[output_key] = feature_video_info( + source_info["features"][source_key], video_width, video_height + ) + + return { + "codebase_version": "v2.1", + "robot_type": "A2D_FRUIT_AGIBOT", + "total_episodes": episode_count, + "total_frames": total_frames, + "total_tasks": len({tasks for tasks in task_ids.values()}), + "chunks_size": 1000, + "fps": float(source_info["fps"]), + "splits": {"train": f"0:{episode_count}"}, + "data_path": "data/chunk-{episode_chunk:03d}/episode_{episode_index:06d}.parquet", + "video_path": ( + "videos/chunk-{episode_chunk:03d}/{video_key}/episode_{episode_index:06d}.mp4" + ), + "features": features, + } + + +def build_modality() -> dict[str, Any]: + def packed_field( + bounds: tuple[int, int], original_key: str, rotation_type: str | None + ) -> dict[str, Any]: + return { + "original_key": original_key, + "start": bounds[0], + "end": bounds[1], + "rotation_type": rotation_type, + "absolute": True, + "dtype": "float32", + "range": None, + } + + return { + "state": { + key: packed_field(field["bounds"], "observation.state", field["rotation_type"]) + for key, field in STATE_ACTION_FIELDS.items() + }, + "action": { + key: packed_field(field["bounds"], "action", field["rotation_type"]) + for key, field in STATE_ACTION_FIELDS.items() + }, + "video": { + "top_head": {"original_key": "observation.images.top_head"}, + "hand_left": {"original_key": "observation.images.hand_left"}, + "hand_right": {"original_key": "observation.images.hand_right"}, + }, + "annotation": { + "language.action_text": {"original_key": "annotation.task"}, + }, + } + + +def compute_output_stats( + output_parts: dict[str, list[np.ndarray]], source_stats: dict[str, Any] +) -> dict[str, Any]: + required = ("min", "max", "mean", "std", "q01", "q99") + + result: dict[str, Any] = {} + for key in ("observation.state", "action"): + if not output_parts[key]: + raise ValueError(f"No converted values available for {key} stats") + values = np.concatenate(output_parts[key], axis=0).astype(np.float64) + result[key] = { + "min": np.min(values, axis=0).tolist(), + "max": np.max(values, axis=0).tolist(), + "mean": np.mean(values, axis=0).tolist(), + "std": np.std(values, axis=0).tolist(), + "q01": np.quantile(values, 0.01, axis=0).tolist(), + "q99": np.quantile(values, 0.99, axis=0).tolist(), + } + # rotation-6d is analytically bounded by [-1, 1]. Using fixed bounds + # avoids unstable min-max scaling when a small smoke subset has little + # rotational variation. + for start, end in ROTATION_OUTPUT_SLICES: + result[key]["min"][start:end] = [-1.0] * (end - start) + result[key]["max"][start:end] = [1.0] * (end - start) + # LeRobot v3 datasets do not always store timestamp statistics. DreamZero + # normalizes state/action, so retain timestamp stats only when supplied. + if "timestamp" in source_stats: + result["timestamp"] = { + name: source_stats["timestamp"][name] for name in required + } + return result + + +def quaternion_xyzw_to_rotation_6d(quaternion: np.ndarray) -> np.ndarray: + """Convert xyzw quaternions to PyTorch3D-compatible matrix rotation-6d.""" + norms = np.linalg.norm(quaternion, axis=1) + if np.any(norms < 1e-6): + raise ValueError("Zero quaternion found") + max_error = float(np.max(np.abs(norms - 1.0))) + if max_error > 5e-2: + raise ValueError(f"Quaternion max |norm-1| is too large: {max_error:.6f}") + normalized = quaternion / norms[:, None] + x, y, z, w = normalized.T + matrix = np.stack( + ( + 1.0 - 2.0 * (y * y + z * z), + 2.0 * (x * y - z * w), + 2.0 * (x * z + y * w), + 2.0 * (x * y + z * w), + 1.0 - 2.0 * (x * x + z * z), + 2.0 * (y * z - x * w), + 2.0 * (x * z - y * w), + 2.0 * (y * z + x * w), + 1.0 - 2.0 * (x * x + y * y), + ), + axis=1, + ).reshape(-1, 3, 3) + return matrix[:, :2, :].reshape(-1, 6).astype(np.float32, copy=False) + + +def convert_xyzw_to_rotation_6d(values: np.ndarray) -> np.ndarray: + """Convert both G2 EE quaternions and repack 16 source values as 20.""" + if values.ndim != 2 or values.shape[1] != 16: + raise ValueError(f"Expected [N, 16] G2 vector, got {values.shape}") + if not np.isfinite(values).all(): + raise ValueError("State/action contains NaN or Inf") + output = np.empty((values.shape[0], OUTPUT_VECTOR_DIM), dtype=np.float32) + output[:, 0:3] = values[:, 0:3] + output[:, 3:9] = quaternion_xyzw_to_rotation_6d(values[:, 3:7]) + output[:, 9] = values[:, 7] + output[:, 10:13] = values[:, 8:11] + output[:, 13:19] = quaternion_xyzw_to_rotation_6d(values[:, 11:15]) + output[:, 19] = values[:, 15] + return output + + +def replace_packed_column( + table: pa.Table, column_name: str, values: np.ndarray +) -> pa.Table: + index = table.schema.get_field_index(column_name) + if index < 0: + raise KeyError(column_name) + packed = pa.array(values.tolist(), type=pa.list_(pa.float32(), OUTPUT_VECTOR_DIM)) + return table.set_column(index, column_name, packed) + + +def write_parquets_and_collect_relative( + data_file: Path, + output: Path, + episode_rows: list[dict[str, Any]], + task_by_index: dict[int, str], + episode_count: int, + action_horizon: int, +) -> tuple[dict[str, list[np.ndarray]], dict[str, list[np.ndarray]]]: + relative_parts: dict[str, list[np.ndarray]] = { + key: [] for key in RELATIVE_ACTION_KEYS + } + output_parts: dict[str, list[np.ndarray]] = { + "observation.state": [], + "action": [], + } + columns = [ + "observation.state", + "action", + "timestamp", + "frame_index", + "episode_index", + "index", + "task_index", + ] + source_table = pq.read_table(data_file, columns=columns) + + for converted_index, metadata in enumerate(episode_rows[:episode_count], start=1): + episode_index = int(metadata["episode_index"]) + table = source_table.filter( + pc.equal(source_table["episode_index"], episode_index) + ) + expected = int(metadata["length"]) + if table.num_rows != expected: + raise ValueError( + f"Episode {episode_index}: parquet has {table.num_rows} rows, expected {expected}" + ) + episode_ids = np.asarray( + table["episode_index"].combine_chunks().to_numpy(), dtype=np.int64 + ) + if not np.all(episode_ids == episode_index): + raise ValueError(f"Episode {episode_index}: row group contains another episode id") + task_indices = np.asarray( + table["task_index"].combine_chunks().to_numpy(), dtype=np.int64 + ) + if len(np.unique(task_indices)) != 1: + raise ValueError(f"Episode {episode_index}: multiple task indices in one episode") + task = task_by_index[int(task_indices[0])] + declared_tasks = list(metadata["tasks"]) + if declared_tasks != [task]: + raise ValueError( + f"Episode {episode_index}: task mismatch: metadata={declared_tasks}, parquet={task!r}" + ) + + source_state = np.asarray( + table["observation.state"].to_pylist(), dtype=np.float32 + ) + source_action = np.asarray(table["action"].to_pylist(), dtype=np.float32) + state = convert_xyzw_to_rotation_6d(source_state) + action = convert_xyzw_to_rotation_6d(source_action) + output_parts["observation.state"].append(state) + output_parts["action"].append(action) + table = replace_packed_column(table, "observation.state", state) + table = replace_packed_column(table, "action", action) + + annotation = pa.array([task] * expected, type=pa.large_string()) + table = table.append_column("annotation.task", annotation) + table = table.select( + [ + "observation.state", + "action", + "annotation.task", + "timestamp", + "frame_index", + "episode_index", + "index", + "task_index", + ] + ) + chunk = episode_index // 1000 + destination = ( + output / f"data/chunk-{chunk:03d}/episode_{episode_index:06d}.parquet" + ) + destination.parent.mkdir(parents=True, exist_ok=True) + pq.write_table(table, destination, compression="zstd", row_group_size=expected) + + usable = expected - action_horizon + if usable > 0: + for key in RELATIVE_ACTION_KEYS: + start, end = STATE_ACTION_FIELDS[key]["bounds"] + episode_relative = np.empty( + (usable, action_horizon, end - start), dtype=np.float32 + ) + reference = state[:usable, None, start:end] + for horizon_index in range(action_horizon): + episode_relative[:, horizon_index] = ( + action[horizon_index : horizon_index + usable, start:end] - reference[:, 0] + ) + relative_parts[key].append(episode_relative.reshape(-1, end - start)) + + if converted_index % 50 == 0 or converted_index == episode_count: + LOG.info("Wrote parquet episodes: %d/%d", converted_index, episode_count) + + return relative_parts, output_parts + + +def compute_relative_stats(parts: dict[str, list[np.ndarray]]) -> dict[str, Any]: + result: dict[str, Any] = {} + for key, arrays in parts.items(): + if not arrays: + raise ValueError(f"No relative-action samples for {key}") + values = np.concatenate(arrays, axis=0).astype(np.float64) + result[key] = { + "max": np.max(values, axis=0).tolist(), + "min": np.min(values, axis=0).tolist(), + "mean": np.mean(values, axis=0).tolist(), + "std": np.std(values, axis=0).tolist(), + "q01": np.quantile(values, 0.01, axis=0).tolist(), + "q99": np.quantile(values, 0.99, axis=0).tolist(), + } + return result + + +def output_video_path(output: Path, episode_index: int, output_key: str) -> Path: + return ( + output + / f"videos/chunk-{episode_index // 1000:03d}" + / output_key + / f"episode_{episode_index:06d}.mp4" + ) + + +def probe_frames(path: Path) -> int: + completed = subprocess.run( + [ + "ffprobe", + "-v", + "error", + "-select_streams", + "v:0", + "-count_frames", + "-show_entries", + "stream=nb_read_frames", + "-of", + "default=noprint_wrappers=1:nokey=1", + str(path), + ], + check=True, + capture_output=True, + text=True, + ) + return int(completed.stdout.strip()) + + +def convert_video_job( + source: Path, + output: Path, + source_info: dict[str, Any], + metadata: dict[str, Any], + episode_index: int, + source_key: str, + output_key: str, + fps: float, + preset: str, + crf: int, + video_width: int | None, + video_height: int | None, +) -> dict[str, Any]: + prefix = f"videos/{source_key}" + file_index = int(metadata[f"{prefix}/file_index"]) + chunk_index = int(metadata[f"{prefix}/chunk_index"]) + start = float(metadata[f"{prefix}/from_timestamp"]) + end = float(metadata[f"{prefix}/to_timestamp"]) + expected_frames = int(metadata["length"]) + available_frames = int(round((end - start) * fps)) + if available_frames < 1 or available_frames > expected_frames: + raise ValueError( + f"Episode {episode_index} {source_key}: invalid available frame count " + f"{available_frames}, expected {expected_frames}" + ) + pad_frames = expected_frames - available_frames + source_path = source / source_info["video_path"].format( + video_key=source_key, + chunk_index=chunk_index, + file_index=file_index, + ) + destination = output_video_path(output, episode_index, output_key) + destination.parent.mkdir(parents=True, exist_ok=True) + video_filter = f"trim=end_frame={available_frames}" + if pad_frames: + video_filter += f",tpad=stop_mode=clone:stop={pad_frames}" + video_filter += ",setpts=PTS-STARTPTS" + if video_width is not None and video_height is not None: + video_filter += f",scale={video_width}:{video_height}:flags=lanczos,setsar=1" + command = [ + "ffmpeg", + "-nostdin", + "-v", + "error", + "-y", + "-ss", + f"{start:.9f}", + "-i", + str(source_path), + "-an", + "-vf", + video_filter, + "-frames:v", + str(expected_frames), + "-r", + f"{fps:g}", + "-vsync", + "cfr", + "-c:v", + "libx264", + "-preset", + preset, + "-crf", + str(crf), + "-pix_fmt", + "yuv420p", + "-g", + "2", + "-keyint_min", + "2", + "-sc_threshold", + "0", + "-movflags", + "+faststart", + str(destination), + ] + subprocess.run(command, check=True) + actual_frames = probe_frames(destination) + if actual_frames != expected_frames: + raise ValueError( + f"Episode {episode_index} {output_key}: encoded {actual_frames} frames, " + f"expected {expected_frames}" + ) + return { + "episode_index": episode_index, + "camera": output_key, + "source_file": str(source_path), + "source_start": start, + "source_end": end, + "expected_frames": expected_frames, + "source_interval_frames": available_frames, + "padded_frames": pad_frames, + "output": str(destination), + } + + +def convert_videos( + source: Path, + output: Path, + source_info: dict[str, Any], + episode_rows: list[dict[str, Any]], + episode_count: int, + workers: int, + preset: str, + crf: int, + video_width: int | None, + video_height: int | None, +) -> list[dict[str, Any]]: + fps = float(source_info["fps"]) + jobs = [] + for metadata in episode_rows[:episode_count]: + episode_index = int(metadata["episode_index"]) + for source_key, output_key in CAMERAS.items(): + jobs.append((episode_index, metadata, source_key, output_key)) + + reports = [] + with ThreadPoolExecutor(max_workers=workers) as executor: + futures = [ + executor.submit( + convert_video_job, + source, + output, + source_info, + metadata, + episode_index, + source_key, + output_key, + fps, + preset, + crf, + video_width, + video_height, + ) + for episode_index, metadata, source_key, output_key in jobs + ] + for completed_count, future in enumerate(as_completed(futures), start=1): + reports.append(future.result()) + if completed_count % 50 == 0 or completed_count == len(jobs): + LOG.info("Wrote episode videos: %d/%d", completed_count, len(jobs)) + return sorted(reports, key=lambda item: (item["episode_index"], item["camera"])) + + +def write_metadata_files( + output: Path, + source_info: dict[str, Any], + output_stats: dict[str, Any], + tasks: list[dict[str, Any]], + episode_rows: list[dict[str, Any]], + episode_count: int, + relative_stats: dict[str, Any], + video_report: list[dict[str, Any]], + action_horizon: int, + video_width: int | None, + video_height: int | None, +) -> None: + meta = output / "meta" + meta.mkdir(parents=True, exist_ok=True) + write_json( + meta / "info.json", + build_info( + source_info, episode_rows, episode_count, video_width, video_height + ), + ) + write_json(meta / "modality.json", build_modality()) + write_json(meta / "embodiment.json", {"embodiment_tag": "agibot"}) + write_json(meta / "stats.json", output_stats) + write_json(meta / "relative_stats_dreamzero.json", relative_stats) + + used_task_indices = { + int(row["task_index"]) + for row in tasks + if any(row["task"] in episode["tasks"] for episode in episode_rows[:episode_count]) + } + with (meta / "tasks.jsonl").open("w") as handle: + for row in tasks: + if int(row["task_index"]) in used_task_indices: + handle.write(json.dumps(row, ensure_ascii=True) + "\n") + with (meta / "episodes.jsonl").open("w") as handle: + for row in episode_rows[:episode_count]: + handle.write( + json.dumps( + { + "episode_index": int(row["episode_index"]), + "tasks": list(row["tasks"]), + "length": int(row["length"]), + }, + ensure_ascii=True, + ) + + "\n" + ) + + repairs = [item for item in video_report if item["padded_frames"]] + write_json( + meta / "conversion_report.json", + { + "source": str(source_info.get("source_path", "")), + "episode_count": episode_count, + "video_count": len(video_report), + "action_horizon": action_horizon, + "state_dim": OUTPUT_VECTOR_DIM, + "action_dim": OUTPUT_VECTOR_DIM, + "source_quaternion_order": "xyzw", + "output_rotation_representation": "rotation_6d", + "relative_action_keys": list(RELATIVE_ACTION_KEYS), + "repairs": repairs, + }, + ) + + +def main() -> int: + args = parse_args() + logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s %(message)s") + source = args.source.resolve() + output = args.output.resolve() + if (args.video_width is None) != (args.video_height is None): + raise ValueError("--video-width and --video-height must be provided together") + if args.workers < 1: + raise ValueError("--workers must be a positive integer") + if args.action_horizon < 1: + raise ValueError("--action-horizon must be a positive integer") + if args.video_width is not None and ( + args.video_width < 2 + or args.video_height < 2 + or args.video_width % 2 + or args.video_height % 2 + ): + raise ValueError("Output video width and height must be positive even integers") + if not source.is_dir(): + raise FileNotFoundError(source) + validate_output_path(source, output) + if output.exists(): + if not args.overwrite: + raise FileExistsError(f"Output already exists: {output}") + shutil.rmtree(output) + output.mkdir(parents=True) + check_tools() + + source_info = load_json(source / "meta/info.json") + source_info["source_path"] = str(source) + source_stats = load_json(source / "meta/stats.json") + data_file, episodes_file, tasks_file = source_table_paths(source) + episode_rows = pq.read_table(episodes_file).to_pylist() + tasks = pq.read_table(tasks_file).to_pylist() + episode_count = validate_source(source_info, data_file, episode_rows, args.max_episodes) + task_by_index = {int(row["task_index"]): str(row["task"]) for row in tasks} + + LOG.info("Converting %d episodes from %s", episode_count, source) + relative_parts, output_parts = write_parquets_and_collect_relative( + data_file, + output, + episode_rows, + task_by_index, + episode_count, + args.action_horizon, + ) + relative_stats = compute_relative_stats(relative_parts) + output_stats = compute_output_stats(output_parts, source_stats) + video_report = convert_videos( + source, + output, + source_info, + episode_rows, + episode_count, + args.workers, + args.video_preset, + args.video_crf, + args.video_width, + args.video_height, + ) + write_metadata_files( + output, + source_info, + output_stats, + tasks, + episode_rows, + episode_count, + relative_stats, + video_report, + args.action_horizon, + args.video_width, + args.video_height, + ) + LOG.info("Conversion completed: %s", output) + return 0 + + +if __name__ == "__main__": + try: + sys.exit(main()) + except Exception: + LOG.exception("Conversion failed") + sys.exit(1) diff --git a/scripts/data/convert_lerobot_g2_to_gear.py b/scripts/data/convert_lerobot_g2_to_gear.py new file mode 100644 index 00000000..6a88abe6 --- /dev/null +++ b/scripts/data/convert_lerobot_g2_to_gear.py @@ -0,0 +1,953 @@ +#!/usr/bin/env python3 +r"""Convert a dual-arm G2 LeRobot v3 dataset into DreamZero/GEAR datasets. + +The converter is intentionally strict. It validates the 16-D joint layout, +materializes one parquet and one video per episode, forces all three camera +streams to the same CFR/resolution, writes GEAR metadata, and can physically +separate a held-out test set from the training set. + +Expected packed state/action layout (from create_g2_dataset_using_lerobot.py): + [left_joint_1..7, left_gripper, right_joint_1..7, right_gripper] + +Example (120 episodes -> 110 train + 10 held-out test): + python convert_lerobot_g2_to_gear.py \ + --source /data/.../g2_mock_light_module_joint_streaming \ + --output /data/.../g2_mock_light_module_gear \ + --test-episodes 10 --split-mode tail \ + --video-width 320 --video-height 176 \ + --video-codec libx264 --workers 8 + +Output: + /train/ # pass this directory to DreamZero training + /test/ # never read by the training job + /split_manifest.json + +For a quick smoke test, add ``--max-source-episodes 2 --test-episodes 0``. +""" + +from __future__ import annotations + +import argparse +from concurrent.futures import ThreadPoolExecutor, as_completed +from dataclasses import dataclass +import json +import logging +from pathlib import Path +import random +import shutil +import subprocess +import sys +from typing import Any, Iterable + +import numpy as np +import pyarrow as pa +import pyarrow.dataset as pads +import pyarrow.parquet as pq + + +LOG = logging.getLogger("g2-v3-to-gear") + +SOURCE_CAMERAS = { + "observation.images.head_color": "observation.images.top_head", + "observation.images.hand_left_color": "observation.images.hand_left", + "observation.images.hand_right_color": "observation.images.hand_right", +} +VIDEO_MODALITY_ORDER = ("top_head", "hand_left", "hand_right") + +JOINT_SLICES: dict[str, tuple[int, int]] = { + "left_joint_position": (0, 7), + "left_gripper_position": (7, 8), + "right_joint_position": (8, 15), + "right_gripper_position": (15, 16), +} +EXPECTED_NAMES = [ + *(f"l.joint{i}.pos" for i in range(1, 8)), + "l.gripper.pos", + *(f"r.joint{i}.pos" for i in range(1, 8)), + "r.gripper.pos", +] + + +@dataclass(frozen=True) +class SourceEpisode: + source_index: int + length: int + task: str + metadata: dict[str, Any] + + +@dataclass(frozen=True) +class OutputEpisode: + source: SourceEpisode + output_index: int + + +def nonempty_path(value: str) -> Path: + if not value.strip(): + raise argparse.ArgumentTypeError("path cannot be empty") + return Path(value) + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser( + description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter + ) + parser.add_argument("--source", type=nonempty_path, required=True) + parser.add_argument("--output", type=nonempty_path, required=True) + parser.add_argument("--embodiment-tag", default="g2") + parser.add_argument("--workers", type=int, default=8) + parser.add_argument("--action-horizon", type=int, default=24) + parser.add_argument("--video-width", type=int, default=320) + parser.add_argument("--video-height", type=int, default=176) + parser.add_argument( + "--resize-mode", + choices=("stretch", "pad", "crop"), + default="stretch", + help="How to make all camera streams the same size.", + ) + parser.add_argument( + "--video-codec", + choices=("libx264", "h264_nvenc"), + default="libx264", + ) + parser.add_argument("--video-preset", default="veryfast") + parser.add_argument("--video-crf", type=int, default=18) + parser.add_argument( + "--test-episodes", + type=int, + default=10, + help="Number of physically isolated held-out episodes.", + ) + parser.add_argument( + "--split-mode", + choices=("tail", "random"), + default="tail", + help="tail reserves the final N source episodes; random uses --split-seed.", + ) + parser.add_argument("--split-seed", type=int, default=42) + parser.add_argument( + "--max-source-episodes", + type=int, + default=None, + help="Use only the first N source episodes (smoke tests only).", + ) + parser.add_argument( + "--min-episode-frames", + type=int, + default=None, + help="Default is action_horizon + 1; shorter episodes fail validation.", + ) + parser.add_argument("--overwrite", action="store_true") + return parser.parse_args() + + +def read_json(path: Path) -> dict[str, Any]: + with path.open(encoding="utf-8") as handle: + value = json.load(handle) + if not isinstance(value, dict): + raise ValueError(f"Expected a JSON object: {path}") + return value + + +def write_json(path: Path, value: Any) -> None: + path.parent.mkdir(parents=True, exist_ok=True) + with path.open("w", encoding="utf-8") as handle: + json.dump(value, handle, indent=2, ensure_ascii=False) + handle.write("\n") + + +def write_jsonl(path: Path, rows: Iterable[dict[str, Any]]) -> None: + path.parent.mkdir(parents=True, exist_ok=True) + with path.open("w", encoding="utf-8") as handle: + for row in rows: + handle.write(json.dumps(row, ensure_ascii=False) + "\n") + + +def require_tools(video_codec: str) -> None: + for executable in ("ffmpeg", "ffprobe"): + if shutil.which(executable) is None: + raise RuntimeError(f"Required executable not found: {executable}") + if video_codec == "h264_nvenc": + completed = subprocess.run( + ["ffmpeg", "-hide_banner", "-encoders"], + check=True, + capture_output=True, + text=True, + ) + if "h264_nvenc" not in completed.stdout: + raise RuntimeError("ffmpeg does not provide the h264_nvenc encoder") + + +def validate_paths(source: Path, output: Path, overwrite: bool) -> None: + source = source.resolve() + output = output.resolve() + if not source.is_dir(): + raise FileNotFoundError(source) + if output == source or output in source.parents or source in output.parents: + raise ValueError(f"Source and output must not overlap: {source} / {output}") + if output == Path(output.anchor) or output == Path.home().resolve(): + raise ValueError(f"Refusing dangerous output path: {output}") + if (output / ".git").exists(): + raise ValueError(f"Refusing to overwrite a Git repository: {output}") + if output.exists(): + if not overwrite: + raise FileExistsError(f"Output exists; pass --overwrite to rebuild: {output}") + shutil.rmtree(output) + output.mkdir(parents=True) + + +def parquet_files(path: Path) -> list[Path]: + files = sorted(path.glob("chunk-*/*.parquet")) + if not files: + raise FileNotFoundError(f"No parquet files below {path}") + return files + + +def read_many_parquets(files: list[Path]) -> pa.Table: + return pads.dataset([str(path) for path in files], format="parquet").to_table() + + +def validate_source_info(info: dict[str, Any]) -> None: + if info.get("codebase_version") != "v3.0": + raise ValueError( + f"Expected LeRobot v3.0, got {info.get('codebase_version')!r}" + ) + features = info.get("features", {}) + for key in ("observation.state", "action", *SOURCE_CAMERAS): + if key not in features: + raise ValueError(f"Missing source feature: {key}") + for key in ("observation.state", "action"): + shape = list(features[key].get("shape", [])) + if shape != [16]: + raise ValueError(f"{key} must be 16-D joint data, got shape {shape}") + names = features[key].get("names") + if names and list(names) != EXPECTED_NAMES: + raise ValueError( + f"{key} names do not match the G2 joint layout.\n" + f"Expected: {EXPECTED_NAMES}\nGot: {names}" + ) + fps = float(info.get("fps", 0)) + if fps <= 0: + raise ValueError(f"Invalid source fps: {fps}") + + +def load_source_episodes( + source: Path, info: dict[str, Any], max_source_episodes: int | None +) -> tuple[list[SourceEpisode], list[Path], pads.Dataset]: + data_files = parquet_files(source / "data") + episode_files = parquet_files(source / "meta/episodes") + tasks_path = source / "meta/tasks.parquet" + if not tasks_path.is_file(): + raise FileNotFoundError(tasks_path) + + episode_rows = read_many_parquets(episode_files).to_pylist() + task_rows = pq.read_table(tasks_path).to_pylist() + task_by_index = {int(row["task_index"]): str(row["task"]).strip() for row in task_rows} + episodes: list[SourceEpisode] = [] + for row in sorted(episode_rows, key=lambda item: int(item["episode_index"])): + source_index = int(row["episode_index"]) + declared_tasks = [str(task).strip() for task in row.get("tasks", []) if str(task).strip()] + task = declared_tasks[0] if declared_tasks else "" + if not task and "task_index" in row: + task = task_by_index.get(int(row["task_index"]), "") + if not task: + raise ValueError(f"Episode {source_index} has no language task") + if len(set(declared_tasks)) > 1: + raise ValueError(f"Episode {source_index} declares multiple tasks: {declared_tasks}") + episodes.append( + SourceEpisode( + source_index=source_index, + length=int(row["length"]), + task=task, + metadata=row, + ) + ) + + total = int(info["total_episodes"]) + if len(episodes) != total: + raise ValueError(f"Episode metadata has {len(episodes)} rows; info.json says {total}") + if max_source_episodes is not None: + if max_source_episodes < 1 or max_source_episodes > len(episodes): + raise ValueError( + f"--max-source-episodes must be in [1, {len(episodes)}]" + ) + episodes = episodes[:max_source_episodes] + + data_dataset = pads.dataset([str(path) for path in data_files], format="parquet") + required_columns = { + "observation.state", + "action", + "episode_index", + } + missing = required_columns - set(data_dataset.schema.names) + if missing: + raise ValueError(f"Source data is missing columns: {sorted(missing)}") + return episodes, data_files, data_dataset + + +def split_episodes( + episodes: list[SourceEpisode], test_count: int, mode: str, seed: int +) -> tuple[list[SourceEpisode], list[SourceEpisode]]: + if test_count < 0 or test_count >= len(episodes): + if test_count == 0: + return episodes, [] + raise ValueError(f"--test-episodes must be in [0, {len(episodes) - 1}]") + if test_count == 0: + return episodes, [] + if mode == "tail": + return episodes[:-test_count], episodes[-test_count:] + rng = random.Random(seed) + test_ids = set(rng.sample([item.source_index for item in episodes], test_count)) + train = [item for item in episodes if item.source_index not in test_ids] + test = [item for item in episodes if item.source_index in test_ids] + return train, test + + +def validate_episode_lengths( + episodes: list[SourceEpisode], min_frames: int, split_name: str +) -> None: + short = [(item.source_index, item.length) for item in episodes if item.length < min_frames] + if short: + raise ValueError( + f"{split_name} has episodes shorter than {min_frames} frames: {short[:20]}" + ) + + +def table_for_episode(data_dataset: pads.Dataset, episode: SourceEpisode) -> pa.Table: + columns = ["observation.state", "action", "episode_index"] + for optional in ("timestamp", "frame_index", "index", "task_index"): + if optional in data_dataset.schema.names: + columns.append(optional) + table = data_dataset.to_table( + filter=pads.field("episode_index") == episode.source_index, + columns=columns, + ) + if table.num_rows != episode.length: + raise ValueError( + f"Episode {episode.source_index}: data has {table.num_rows} rows; " + f"metadata says {episode.length}" + ) + return table + + +def fixed_size_vectors(column: pa.ChunkedArray, name: str) -> np.ndarray: + values = np.asarray(column.to_pylist(), dtype=np.float32) + if values.ndim != 2 or values.shape[1] != 16: + raise ValueError(f"{name} must have shape [N, 16], got {values.shape}") + if not np.isfinite(values).all(): + raise ValueError(f"{name} contains NaN or infinity") + return values + + +def output_table( + source_table: pa.Table, + episode: OutputEpisode, + fps: float, + task_index: int, + global_start: int, +) -> tuple[pa.Table, np.ndarray, np.ndarray]: + state = fixed_size_vectors(source_table["observation.state"], "observation.state") + action = fixed_size_vectors(source_table["action"], "action") + length = episode.source.length + table = pa.table( + { + "observation.state": pa.array(state.tolist(), type=pa.list_(pa.float32(), 16)), + "action": pa.array(action.tolist(), type=pa.list_(pa.float32(), 16)), + "annotation.language.action_text": pa.array( + [episode.source.task] * length, type=pa.large_string() + ), + "timestamp": pa.array( + np.arange(length, dtype=np.float32) / np.float32(fps) + ), + "frame_index": pa.array(np.arange(length, dtype=np.int64)), + "episode_index": pa.array( + np.full(length, episode.output_index, dtype=np.int64) + ), + "index": pa.array( + np.arange(global_start, global_start + length, dtype=np.int64) + ), + "task_index": pa.array(np.full(length, task_index, dtype=np.int64)), + } + ) + return table, state, action + + +class StatsAccumulator: + def __init__(self) -> None: + self.state: list[np.ndarray] = [] + self.action: list[np.ndarray] = [] + self.relative: dict[str, list[np.ndarray]] = {key: [] for key in JOINT_SLICES} + + def add(self, state: np.ndarray, action: np.ndarray, horizon: int) -> None: + self.state.append(state) + self.action.append(action) + usable = len(state) - horizon + 1 + if usable <= 0: + return + for key, (start, end) in JOINT_SLICES.items(): + reference = state[:usable, start:end] + chunks = np.stack( + [action[offset : offset + usable, start:end] - reference for offset in range(horizon)], + axis=1, + ) + self.relative[key].append(chunks.reshape(-1, end - start)) + + +def numeric_stats(values: np.ndarray) -> dict[str, list[float]]: + values64 = np.asarray(values, dtype=np.float64) + return { + "min": np.min(values64, axis=0).tolist(), + "max": np.max(values64, axis=0).tolist(), + "mean": np.mean(values64, axis=0).tolist(), + "std": np.std(values64, axis=0).tolist(), + "q01": np.quantile(values64, 0.01, axis=0).tolist(), + "q99": np.quantile(values64, 0.99, axis=0).tolist(), + } + + +def finish_stats(accumulator: StatsAccumulator) -> tuple[dict[str, Any], dict[str, Any]]: + if not accumulator.state or not accumulator.action: + raise ValueError("Cannot compute statistics for an empty split") + stats = { + "observation.state": numeric_stats(np.concatenate(accumulator.state, axis=0)), + "action": numeric_stats(np.concatenate(accumulator.action, axis=0)), + } + relative: dict[str, Any] = {} + for key, arrays in accumulator.relative.items(): + if not arrays: + raise ValueError(f"No relative-action samples for {key}") + relative[key] = numeric_stats(np.concatenate(arrays, axis=0)) + return stats, relative + + +def output_video_path(root: Path, episode_index: int, output_key: str) -> Path: + return ( + root + / f"videos/chunk-{episode_index // 1000:03d}" + / output_key + / f"episode_{episode_index:06d}.mp4" + ) + + +def probe_video(path: Path) -> tuple[int, int, int, float]: + completed = subprocess.run( + [ + "ffprobe", + "-v", + "error", + "-count_frames", + "-select_streams", + "v:0", + "-show_entries", + "stream=nb_read_frames,width,height,avg_frame_rate", + "-of", + "json", + str(path), + ], + check=True, + capture_output=True, + text=True, + ) + stream = json.loads(completed.stdout)["streams"][0] + numerator, denominator = (int(part) for part in stream["avg_frame_rate"].split("/")) + rate = numerator / denominator if denominator else 0.0 + return int(stream["nb_read_frames"]), int(stream["width"]), int(stream["height"]), rate + + +def resize_filter(width: int, height: int, mode: str) -> str: + if mode == "stretch": + return f"scale={width}:{height}:flags=lanczos,setsar=1" + if mode == "pad": + return ( + f"scale={width}:{height}:force_original_aspect_ratio=decrease:flags=lanczos," + f"pad={width}:{height}:(ow-iw)/2:(oh-ih)/2:black,setsar=1" + ) + return ( + f"scale={width}:{height}:force_original_aspect_ratio=increase:flags=lanczos," + f"crop={width}:{height},setsar=1" + ) + + +def source_video_path( + source: Path, + source_info: dict[str, Any], + metadata: dict[str, Any], + source_key: str, +) -> tuple[Path, float]: + prefix = f"videos/{source_key}" + required = ( + f"{prefix}/file_index", + f"{prefix}/chunk_index", + f"{prefix}/from_timestamp", + ) + missing = [key for key in required if key not in metadata] + if missing: + raise ValueError(f"Episode video metadata is missing: {missing}") + path = source / source_info["video_path"].format( + video_key=source_key, + chunk_index=int(metadata[f"{prefix}/chunk_index"]), + file_index=int(metadata[f"{prefix}/file_index"]), + ) + if not path.is_file(): + raise FileNotFoundError(path) + return path, float(metadata[f"{prefix}/from_timestamp"]) + + +def convert_video( + source: Path, + destination_root: Path, + source_info: dict[str, Any], + episode: OutputEpisode, + source_key: str, + output_key: str, + width: int, + height: int, + mode: str, + codec: str, + preset: str, + crf: int, +) -> dict[str, Any]: + source_path, start = source_video_path( + source, source_info, episode.source.metadata, source_key + ) + destination = output_video_path(destination_root, episode.output_index, output_key) + destination.parent.mkdir(parents=True, exist_ok=True) + fps = float(source_info["fps"]) + frame_count = episode.source.length + filters = ( + f"fps={fps:g},{resize_filter(width, height, mode)}," + f"tpad=stop_mode=clone:stop_duration=2,trim=end_frame={frame_count}," + "setpts=N/FRAME_RATE/TB" + ) + command = [ + "ffmpeg", + "-nostdin", + "-hide_banner", + "-loglevel", + "error", + "-y", + "-ss", + f"{start:.9f}", + "-i", + str(source_path), + "-an", + "-vf", + filters, + "-frames:v", + str(frame_count), + "-r", + f"{fps:g}", + "-c:v", + codec, + ] + if codec == "libx264": + command += ["-preset", preset, "-crf", str(crf)] + else: + command += ["-preset", "p4", "-cq", str(crf), "-b:v", "0"] + command += [ + "-pix_fmt", + "yuv420p", + "-g", + "2", + "-keyint_min", + "2", + "-sc_threshold", + "0", + "-movflags", + "+faststart", + str(destination), + ] + subprocess.run(command, check=True) + actual_frames, actual_width, actual_height, actual_fps = probe_video(destination) + if actual_frames != frame_count: + raise ValueError( + f"{destination}: {actual_frames} frames; expected {frame_count}" + ) + if (actual_width, actual_height) != (width, height): + raise ValueError( + f"{destination}: resolution {(actual_width, actual_height)}; " + f"expected {(width, height)}" + ) + if abs(actual_fps - fps) > 1e-3: + raise ValueError(f"{destination}: fps {actual_fps}; expected {fps}") + return { + "source_episode_index": episode.source.source_index, + "output_episode_index": episode.output_index, + "camera": output_key, + "frames": actual_frames, + "width": actual_width, + "height": actual_height, + "fps": actual_fps, + "source": str(source_path), + "output": str(destination), + } + + +def build_modality() -> dict[str, Any]: + def field(original_key: str, start: int, end: int) -> dict[str, Any]: + return { + "original_key": original_key, + "start": start, + "end": end, + "rotation_type": None, + "absolute": True, + "dtype": "float32", + "range": None, + } + + return { + "state": { + key: field("observation.state", start, end) + for key, (start, end) in JOINT_SLICES.items() + }, + "action": { + key: field("action", start, end) + for key, (start, end) in JOINT_SLICES.items() + }, + "video": { + "top_head": {"original_key": "observation.images.top_head"}, + "hand_left": {"original_key": "observation.images.hand_left"}, + "hand_right": {"original_key": "observation.images.hand_right"}, + }, + "annotation": { + "language.action_text": { + "original_key": "annotation.language.action_text" + } + }, + } + + +def build_info( + source_info: dict[str, Any], + episodes: list[OutputEpisode], + tasks: list[str], + width: int, + height: int, + split_name: str, +) -> dict[str, Any]: + fps = float(source_info["fps"]) + vector_features = { + "observation.state": { + "dtype": "float32", + "shape": [16], + "names": EXPECTED_NAMES, + }, + "action": {"dtype": "float32", "shape": [16], "names": EXPECTED_NAMES}, + } + features: dict[str, Any] = { + **vector_features, + "annotation.language.action_text": {"dtype": "string", "shape": [1]}, + "timestamp": {"dtype": "float32", "shape": [1]}, + "frame_index": {"dtype": "int64", "shape": [1]}, + "episode_index": {"dtype": "int64", "shape": [1]}, + "index": {"dtype": "int64", "shape": [1]}, + "task_index": {"dtype": "int64", "shape": [1]}, + } + for output_key in SOURCE_CAMERAS.values(): + features[output_key] = { + "dtype": "video", + "shape": [height, width, 3], + "names": ["height", "width", "channels"], + "video_info": { + "video.fps": fps, + "video.codec": "h264", + "video.pix_fmt": "yuv420p", + "video.is_depth_map": False, + "has_audio": False, + }, + } + return { + "codebase_version": "v2.1", + "robot_type": "g2", + "total_episodes": len(episodes), + "total_frames": sum(item.source.length for item in episodes), + "total_tasks": len(tasks), + "chunks_size": 1000, + "fps": fps, + "splits": {split_name: f"0:{len(episodes)}"}, + "data_path": "data/chunk-{episode_chunk:03d}/episode_{episode_index:06d}.parquet", + "video_path": "videos/chunk-{episode_chunk:03d}/{video_key}/episode_{episode_index:06d}.mp4", + "features": features, + } + + +def write_split_metadata( + root: Path, + source_info: dict[str, Any], + split_name: str, + episodes: list[OutputEpisode], + tasks: list[str], + stats: dict[str, Any], + relative_stats: dict[str, Any], + embodiment_tag: str, + width: int, + height: int, + action_horizon: int, + video_reports: list[dict[str, Any]], + stats_source: str, +) -> None: + meta = root / "meta" + task_to_index = {task: index for index, task in enumerate(tasks)} + write_json( + meta / "info.json", + build_info(source_info, episodes, tasks, width, height, split_name), + ) + write_json(meta / "modality.json", build_modality()) + write_json(meta / "embodiment.json", {"embodiment_tag": embodiment_tag}) + write_json(meta / "stats.json", stats) + write_json(meta / "relative_stats_dreamzero.json", relative_stats) + write_jsonl( + meta / "tasks.jsonl", + ({"task_index": index, "task": task} for index, task in enumerate(tasks)), + ) + write_jsonl( + meta / "episodes.jsonl", + ( + { + "episode_index": item.output_index, + "tasks": [item.source.task], + "length": item.source.length, + } + for item in episodes + ), + ) + write_json( + meta / "conversion_report.json", + { + "split": split_name, + "episode_count": len(episodes), + "frame_count": sum(item.source.length for item in episodes), + "source_episode_indices": [item.source.source_index for item in episodes], + "task_indices": { + str(item.output_index): task_to_index[item.source.task] for item in episodes + }, + "state_dim": 16, + "action_dim": 16, + "action_horizon": action_horizon, + "camera_order": list(VIDEO_MODALITY_ORDER), + "video_resolution": [width, height], + "video_count": len(video_reports), + "stats_source": stats_source, + }, + ) + + +def convert_split_data( + root: Path, + source_info: dict[str, Any], + data_dataset: pads.Dataset, + source_episodes: list[SourceEpisode], + horizon: int, +) -> tuple[list[OutputEpisode], list[str], StatsAccumulator]: + episodes = [OutputEpisode(source=item, output_index=index) for index, item in enumerate(source_episodes)] + tasks = list(dict.fromkeys(item.source.task for item in episodes)) + task_to_index = {task: index for index, task in enumerate(tasks)} + accumulator = StatsAccumulator() + global_index = 0 + for count, item in enumerate(episodes, start=1): + source_table = table_for_episode(data_dataset, item.source) + table, state, action = output_table( + source_table, + item, + float(source_info["fps"]), + task_to_index[item.source.task], + global_index, + ) + destination = ( + root + / f"data/chunk-{item.output_index // 1000:03d}" + / f"episode_{item.output_index:06d}.parquet" + ) + destination.parent.mkdir(parents=True, exist_ok=True) + pq.write_table(table, destination, compression="zstd", row_group_size=item.source.length) + accumulator.add(state, action, horizon) + global_index += item.source.length + if count % 20 == 0 or count == len(episodes): + LOG.info("%s parquet: %d/%d", root.name, count, len(episodes)) + return episodes, tasks, accumulator + + +def convert_split_videos( + source: Path, + root: Path, + source_info: dict[str, Any], + episodes: list[OutputEpisode], + args: argparse.Namespace, +) -> list[dict[str, Any]]: + jobs = [ + (episode, source_key, output_key) + for episode in episodes + for source_key, output_key in SOURCE_CAMERAS.items() + ] + reports: list[dict[str, Any]] = [] + with ThreadPoolExecutor(max_workers=args.workers) as executor: + futures = [ + executor.submit( + convert_video, + source, + root, + source_info, + episode, + source_key, + output_key, + args.video_width, + args.video_height, + args.resize_mode, + args.video_codec, + args.video_preset, + args.video_crf, + ) + for episode, source_key, output_key in jobs + ] + for count, future in enumerate(as_completed(futures), start=1): + reports.append(future.result()) + if count % 30 == 0 or count == len(jobs): + LOG.info("%s videos: %d/%d", root.name, count, len(jobs)) + return sorted( + reports, + key=lambda item: (item["output_episode_index"], item["camera"]), + ) + + +def validate_output(root: Path, split_name: str) -> None: + meta = root / "meta" + required = ( + "info.json", + "modality.json", + "embodiment.json", + "stats.json", + "relative_stats_dreamzero.json", + "tasks.jsonl", + "episodes.jsonl", + "conversion_report.json", + ) + missing = [name for name in required if not (meta / name).is_file()] + if missing: + raise ValueError(f"{split_name}: missing metadata files: {missing}") + info = read_json(meta / "info.json") + episode_count = int(info["total_episodes"]) + data_count = len(list((root / "data").glob("chunk-*/*.parquet"))) + video_count = len(list((root / "videos").glob("chunk-*/*/*.mp4"))) + if data_count != episode_count: + raise ValueError(f"{split_name}: {data_count} parquets for {episode_count} episodes") + if video_count != episode_count * 3: + raise ValueError(f"{split_name}: {video_count} videos; expected {episode_count * 3}") + + +def main() -> int: + args = parse_args() + logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s %(message)s") + if args.workers < 1: + raise ValueError("--workers must be positive") + if args.action_horizon < 1: + raise ValueError("--action-horizon must be positive") + if min(args.video_width, args.video_height) < 2 or args.video_width % 2 or args.video_height % 2: + raise ValueError("Video width and height must be positive even integers") + + source = args.source.resolve() + output = args.output.resolve() + require_tools(args.video_codec) + validate_paths(source, output, args.overwrite) + + source_info = read_json(source / "meta/info.json") + validate_source_info(source_info) + source_episodes, data_files, data_dataset = load_source_episodes( + source, source_info, args.max_source_episodes + ) + train_source, test_source = split_episodes( + source_episodes, args.test_episodes, args.split_mode, args.split_seed + ) + min_frames = args.min_episode_frames or (args.action_horizon + 1) + validate_episode_lengths(train_source, min_frames, "train") + if test_source: + validate_episode_lengths(test_source, min_frames, "test") + + LOG.info( + "Source: %d episodes, %d parquet shards; split into %d train / %d test", + len(source_episodes), + len(data_files), + len(train_source), + len(test_source), + ) + write_json( + output / "split_manifest.json", + { + "source": str(source), + "split_mode": args.split_mode, + "split_seed": args.split_seed if args.split_mode == "random" else None, + "train_source_episode_indices": [item.source_index for item in train_source], + "test_source_episode_indices": [item.source_index for item in test_source], + }, + ) + + train_root = output / "train" + train_episodes, train_tasks, train_acc = convert_split_data( + train_root, + source_info, + data_dataset, + train_source, + args.action_horizon, + ) + train_stats, train_relative_stats = finish_stats(train_acc) + train_video_reports = convert_split_videos( + source, train_root, source_info, train_episodes, args + ) + write_split_metadata( + train_root, + source_info, + "train", + train_episodes, + train_tasks, + train_stats, + train_relative_stats, + args.embodiment_tag, + args.video_width, + args.video_height, + args.action_horizon, + train_video_reports, + "train", + ) + validate_output(train_root, "train") + + if test_source: + test_root = output / "test" + test_episodes, test_tasks, _ = convert_split_data( + test_root, + source_info, + data_dataset, + test_source, + args.action_horizon, + ) + test_video_reports = convert_split_videos( + source, test_root, source_info, test_episodes, args + ) + # Test data deliberately uses training-set normalization statistics. + write_split_metadata( + test_root, + source_info, + "test", + test_episodes, + test_tasks, + train_stats, + train_relative_stats, + args.embodiment_tag, + args.video_width, + args.video_height, + args.action_horizon, + test_video_reports, + "train", + ) + validate_output(test_root, "test") + + LOG.info("Conversion complete. DreamZero train root: %s", train_root) + if test_source: + LOG.info("Held-out test root: %s", output / "test") + return 0 + + +if __name__ == "__main__": + try: + sys.exit(main()) + except Exception: + LOG.exception("Conversion failed") + sys.exit(1) diff --git a/scripts/data/convert_subtask.py b/scripts/data/convert_subtask.py new file mode 100644 index 00000000..c188a5c0 --- /dev/null +++ b/scripts/data/convert_subtask.py @@ -0,0 +1,1196 @@ +#!/usr/bin/env python3 +r"""Convert a dual-arm G2 LeRobot v3 dataset into DreamZero/GEAR datasets. + +The converter is intentionally strict. It validates the 16-D joint layout, +requires exactly one ``subtask_index`` per source episode, resolves that index +through ``meta/subtasks.parquet``, and writes the resolved subtask text to +``annotation.language.action_text``. It materializes one parquet and one video +per existing source episode, forces all three camera streams to the same +CFR/resolution, writes GEAR metadata, and can physically separate a held-out +test set from the training set. + +Important: the source episodes are already subtask clips. This converter does +not split frames or videos again; it fixes the language supervision used by +DreamZero. + +Expected packed state/action layout (from create_g2_dataset_using_lerobot.py): + [left_joint_1..7, left_gripper, right_joint_1..7, right_gripper] + +Example (120 episodes -> 110 train + 10 held-out test): + python convert_lerobot_g2_to_gear.py \ + --source /data/.../g2_mock_light_module_joint_streaming \ + --output /data/.../g2_mock_light_module_gear \ + --test-episodes 10 --split-mode tail \ + --video-width 320 --video-height 176 \ + --video-codec libx264 --workers 8 + +Output: + /train/ # pass this directory to DreamZero training + /test/ # never read by the training job + /split_manifest.json + +For a quick smoke test, add ``--max-source-episodes 2 --test-episodes 0``. +""" + +from __future__ import annotations + +import argparse +from concurrent.futures import ThreadPoolExecutor, as_completed +from dataclasses import dataclass +import json +import logging +from pathlib import Path +import random +import shutil +import subprocess +import sys +from typing import Any, Iterable + +import numpy as np +import pyarrow as pa +import pyarrow.dataset as pads +import pyarrow.parquet as pq + + +LOG = logging.getLogger("g2-v3-to-gear") + +SOURCE_CAMERAS = { + "observation.images.head_color": "observation.images.top_head", + "observation.images.hand_left_color": "observation.images.hand_left", + "observation.images.hand_right_color": "observation.images.hand_right", +} +VIDEO_MODALITY_ORDER = ("top_head", "hand_left", "hand_right") + +JOINT_SLICES: dict[str, tuple[int, int]] = { + "left_joint_position": (0, 7), + "left_gripper_position": (7, 8), + "right_joint_position": (8, 15), + "right_gripper_position": (15, 16), +} +EXPECTED_NAMES = [ + *(f"l.joint{i}.pos" for i in range(1, 8)), + "l.gripper.pos", + *(f"r.joint{i}.pos" for i in range(1, 8)), + "r.gripper.pos", +] + + +@dataclass(frozen=True) +class SourceEpisode: + source_index: int + length: int + task: str + source_task: str + subtask_index: int + metadata: dict[str, Any] + + +@dataclass(frozen=True) +class OutputEpisode: + source: SourceEpisode + output_index: int + + +def nonempty_path(value: str) -> Path: + if not value.strip(): + raise argparse.ArgumentTypeError("path cannot be empty") + return Path(value) + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser( + description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter + ) + parser.add_argument("--source", type=nonempty_path, required=True) + parser.add_argument("--output", type=nonempty_path, required=True) + parser.add_argument("--embodiment-tag", default="g2") + parser.add_argument("--workers", type=int, default=8) + parser.add_argument("--action-horizon", type=int, default=24) + parser.add_argument("--video-width", type=int, default=320) + parser.add_argument("--video-height", type=int, default=176) + parser.add_argument( + "--resize-mode", + choices=("stretch", "pad", "crop"), + default="stretch", + help="How to make all camera streams the same size.", + ) + parser.add_argument( + "--video-codec", + choices=("libx264", "h264_nvenc"), + default="libx264", + ) + parser.add_argument("--video-preset", default="veryfast") + parser.add_argument("--video-crf", type=int, default=18) + parser.add_argument( + "--test-episodes", + type=int, + default=10, + help="Number of physically isolated held-out episodes.", + ) + parser.add_argument( + "--split-mode", + choices=("tail", "random", "stratified"), + default="tail", + help=( + "tail reserves the final N source episodes; random samples globally; " + "stratified samples evenly from every subtask." + ), + ) + parser.add_argument("--split-seed", type=int, default=42) + parser.add_argument( + "--max-source-episodes", + type=int, + default=None, + help="Use only the first N source episodes (smoke tests only).", + ) + parser.add_argument( + "--min-episode-frames", + type=int, + default=None, + help="Default is action_horizon + 1; shorter episodes fail validation.", + ) + parser.add_argument("--overwrite", action="store_true") + return parser.parse_args() + + +def read_json(path: Path) -> dict[str, Any]: + with path.open(encoding="utf-8") as handle: + value = json.load(handle) + if not isinstance(value, dict): + raise ValueError(f"Expected a JSON object: {path}") + return value + + +def write_json(path: Path, value: Any) -> None: + path.parent.mkdir(parents=True, exist_ok=True) + with path.open("w", encoding="utf-8") as handle: + json.dump(value, handle, indent=2, ensure_ascii=False) + handle.write("\n") + + +def write_jsonl(path: Path, rows: Iterable[dict[str, Any]]) -> None: + path.parent.mkdir(parents=True, exist_ok=True) + with path.open("w", encoding="utf-8") as handle: + for row in rows: + handle.write(json.dumps(row, ensure_ascii=False) + "\n") + + +def require_tools(video_codec: str) -> None: + for executable in ("ffmpeg", "ffprobe"): + if shutil.which(executable) is None: + raise RuntimeError(f"Required executable not found: {executable}") + if video_codec == "h264_nvenc": + completed = subprocess.run( + ["ffmpeg", "-hide_banner", "-encoders"], + check=True, + capture_output=True, + text=True, + ) + if "h264_nvenc" not in completed.stdout: + raise RuntimeError("ffmpeg does not provide the h264_nvenc encoder") + + +def validate_paths(source: Path, output: Path, overwrite: bool) -> None: + source = source.resolve() + output = output.resolve() + if not source.is_dir(): + raise FileNotFoundError(source) + if output == source or output in source.parents or source in output.parents: + raise ValueError(f"Source and output must not overlap: {source} / {output}") + if output == Path(output.anchor) or output == Path.home().resolve(): + raise ValueError(f"Refusing dangerous output path: {output}") + if (output / ".git").exists(): + raise ValueError(f"Refusing to overwrite a Git repository: {output}") + if output.exists(): + if not overwrite: + raise FileExistsError(f"Output exists; pass --overwrite to rebuild: {output}") + shutil.rmtree(output) + output.mkdir(parents=True) + + +def parquet_files(path: Path) -> list[Path]: + files = sorted(path.glob("chunk-*/*.parquet")) + if not files: + raise FileNotFoundError(f"No parquet files below {path}") + return files + + +def read_many_parquets(files: list[Path]) -> pa.Table: + return pads.dataset([str(path) for path in files], format="parquet").to_table() + + +def indexed_text_map( + path: Path, + index_column: str, + text_columns: tuple[str, ...], +) -> dict[int, str]: + if not path.is_file(): + raise FileNotFoundError(path) + table = pq.read_table(path) + if index_column not in table.column_names: + raise ValueError( + f"{path} is missing {index_column!r}; columns are {table.column_names}" + ) + text_column = next( + (name for name in text_columns if name in table.column_names), + None, + ) + if text_column is None: + raise ValueError( + f"{path} has no text column among {text_columns}; " + f"columns are {table.column_names}" + ) + + result: dict[int, str] = {} + for row in table.select([index_column, text_column]).to_pylist(): + index = int(row[index_column]) + text = str(row[text_column]).strip() + if not text: + raise ValueError(f"{path}: empty text for {index_column}={index}") + previous = result.get(index) + if previous is not None and previous != text: + raise ValueError( + f"{path}: conflicting text for {index_column}={index}: " + f"{previous!r} / {text!r}" + ) + result[index] = text + if not result: + raise ValueError(f"{path} contains no mappings") + return result + + +def episode_subtask_indices(data_dataset: pads.Dataset) -> dict[int, int]: + table = data_dataset.to_table(columns=["episode_index", "subtask_index"]) + episode_values = table["episode_index"].to_numpy(zero_copy_only=False) + subtask_values = table["subtask_index"].to_numpy(zero_copy_only=False) + seen: dict[int, int] = {} + conflicts: dict[int, set[int]] = {} + + for episode_value, subtask_value in zip( + episode_values, subtask_values, strict=True + ): + episode_index = int(episode_value) + subtask_index = int(subtask_value) + previous = seen.setdefault(episode_index, subtask_index) + if previous != subtask_index: + conflicts.setdefault(episode_index, {previous}).add(subtask_index) + + if conflicts: + examples = { + episode: sorted(indices) + for episode, indices in list(sorted(conflicts.items()))[:20] + } + raise ValueError( + "Every source episode must contain exactly one subtask_index; " + f"conflicts: {examples}" + ) + return seen + + +def validate_source_info(info: dict[str, Any]) -> None: + if info.get("codebase_version") != "v3.0": + raise ValueError( + f"Expected LeRobot v3.0, got {info.get('codebase_version')!r}" + ) + features = info.get("features", {}) + for key in ("observation.state", "action", *SOURCE_CAMERAS): + if key not in features: + raise ValueError(f"Missing source feature: {key}") + for key in ("observation.state", "action"): + shape = list(features[key].get("shape", [])) + if shape != [16]: + raise ValueError(f"{key} must be 16-D joint data, got shape {shape}") + names = features[key].get("names") + if names and list(names) != EXPECTED_NAMES: + raise ValueError( + f"{key} names do not match the G2 joint layout.\n" + f"Expected: {EXPECTED_NAMES}\nGot: {names}" + ) + fps = float(info.get("fps", 0)) + if fps <= 0: + raise ValueError(f"Invalid source fps: {fps}") + + +def load_source_episodes( + source: Path, info: dict[str, Any], max_source_episodes: int | None +) -> tuple[list[SourceEpisode], list[Path], pads.Dataset]: + data_files = parquet_files(source / "data") + episode_files = parquet_files(source / "meta/episodes") + tasks_path = source / "meta/tasks.parquet" + subtasks_path = source / "meta/subtasks.parquet" + + data_dataset = pads.dataset([str(path) for path in data_files], format="parquet") + required_columns = { + "observation.state", + "action", + "episode_index", + "subtask_index", + } + missing = required_columns - set(data_dataset.schema.names) + if missing: + raise ValueError(f"Source data is missing columns: {sorted(missing)}") + + episode_rows = read_many_parquets(episode_files).to_pylist() + task_by_index = indexed_text_map(tasks_path, "task_index", ("task",)) + subtask_by_index = indexed_text_map( + subtasks_path, + "subtask_index", + ("subtask", "task", "text"), + ) + subtask_index_by_episode = episode_subtask_indices(data_dataset) + + episodes: list[SourceEpisode] = [] + for row in sorted(episode_rows, key=lambda item: int(item["episode_index"])): + source_index = int(row["episode_index"]) + declared_tasks = [ + str(task).strip() + for task in row.get("tasks", []) + if str(task).strip() + ] + if len(set(declared_tasks)) > 1: + raise ValueError( + f"Episode {source_index} declares multiple high-level tasks: " + f"{declared_tasks}" + ) + source_task = declared_tasks[0] if declared_tasks else "" + if not source_task and "task_index" in row: + source_task = task_by_index.get(int(row["task_index"]), "") + if not source_task: + raise ValueError(f"Episode {source_index} has no high-level task") + + if source_index not in subtask_index_by_episode: + raise ValueError(f"Episode {source_index} has no frame-level subtask_index") + subtask_index = subtask_index_by_episode[source_index] + task = subtask_by_index.get(subtask_index, "") + if not task: + raise ValueError( + f"Episode {source_index}: subtask_index={subtask_index} " + f"is absent from {subtasks_path}" + ) + + episodes.append( + SourceEpisode( + source_index=source_index, + length=int(row["length"]), + task=task, + source_task=source_task, + subtask_index=subtask_index, + metadata=row, + ) + ) + + total = int(info["total_episodes"]) + if len(episodes) != total: + raise ValueError(f"Episode metadata has {len(episodes)} rows; info.json says {total}") + if max_source_episodes is not None: + if max_source_episodes < 1 or max_source_episodes > len(episodes): + raise ValueError( + f"--max-source-episodes must be in [1, {len(episodes)}]" + ) + episodes = episodes[:max_source_episodes] + + used_subtask_indices = {item.subtask_index for item in episodes} + LOG.info( + "Resolved %d source episodes to %d subtask labels " + "(%d high-level task labels retained for provenance)", + len(episodes), + len(used_subtask_indices), + len({item.source_task for item in episodes}), + ) + return episodes, data_files, data_dataset + + +def split_episodes( + episodes: list[SourceEpisode], test_count: int, mode: str, seed: int +) -> tuple[list[SourceEpisode], list[SourceEpisode]]: + if test_count < 0 or test_count >= len(episodes): + if test_count == 0: + return episodes, [] + raise ValueError(f"--test-episodes must be in [0, {len(episodes) - 1}]") + if test_count == 0: + return episodes, [] + if mode == "tail": + return episodes[:-test_count], episodes[-test_count:] + rng = random.Random(seed) + if mode == "stratified": + by_subtask: dict[int, list[SourceEpisode]] = {} + for item in episodes: + by_subtask.setdefault(item.subtask_index, []).append(item) + subtask_ids = sorted(by_subtask) + if test_count < len(subtask_ids): + raise ValueError( + "Stratified split needs at least one test episode per subtask: " + f"test_count={test_count}, subtasks={len(subtask_ids)}" + ) + base, remainder = divmod(test_count, len(subtask_ids)) + test_ids: set[int] = set() + for position, subtask_id in enumerate(subtask_ids): + count = base + int(position < remainder) + group = by_subtask[subtask_id] + if count >= len(group): + raise ValueError( + f"Subtask {subtask_id} has {len(group)} episodes; " + f"cannot reserve {count} and retain training data" + ) + test_ids.update( + item.source_index for item in rng.sample(group, count) + ) + else: + test_ids = set(rng.sample([item.source_index for item in episodes], test_count)) + train = [item for item in episodes if item.source_index not in test_ids] + test = [item for item in episodes if item.source_index in test_ids] + return train, test + + +def validate_episode_lengths( + episodes: list[SourceEpisode], min_frames: int, split_name: str +) -> None: + short = [(item.source_index, item.length) for item in episodes if item.length < min_frames] + if short: + raise ValueError( + f"{split_name} has episodes shorter than {min_frames} frames: {short[:20]}" + ) + + +def table_for_episode(data_dataset: pads.Dataset, episode: SourceEpisode) -> pa.Table: + columns = ["observation.state", "action", "episode_index"] + for optional in ( + "timestamp", + "frame_index", + "index", + "task_index", + "subtask_index", + ): + if optional in data_dataset.schema.names: + columns.append(optional) + table = data_dataset.to_table( + filter=pads.field("episode_index") == episode.source_index, + columns=columns, + ) + if table.num_rows != episode.length: + raise ValueError( + f"Episode {episode.source_index}: data has {table.num_rows} rows; " + f"metadata says {episode.length}" + ) + actual_subtasks = { + int(value) + for value in table["subtask_index"].to_pylist() + if value is not None + } + if actual_subtasks != {episode.subtask_index}: + raise ValueError( + f"Episode {episode.source_index}: expected subtask_index " + f"{episode.subtask_index}, got {sorted(actual_subtasks)}" + ) + return table + + +def fixed_size_vectors(column: pa.ChunkedArray, name: str) -> np.ndarray: + values = np.asarray(column.to_pylist(), dtype=np.float32) + if values.ndim != 2 or values.shape[1] != 16: + raise ValueError(f"{name} must have shape [N, 16], got {values.shape}") + if not np.isfinite(values).all(): + raise ValueError(f"{name} contains NaN or infinity") + return values + + +def output_table( + source_table: pa.Table, + episode: OutputEpisode, + fps: float, + task_index: int, + global_start: int, +) -> tuple[pa.Table, np.ndarray, np.ndarray]: + state = fixed_size_vectors(source_table["observation.state"], "observation.state") + action = fixed_size_vectors(source_table["action"], "action") + length = episode.source.length + table = pa.table( + { + "observation.state": pa.array(state.tolist(), type=pa.list_(pa.float32(), 16)), + "action": pa.array(action.tolist(), type=pa.list_(pa.float32(), 16)), + "annotation.language.action_text": pa.array( + [episode.source.task] * length, type=pa.large_string() + ), + "timestamp": pa.array( + np.arange(length, dtype=np.float32) / np.float32(fps) + ), + "frame_index": pa.array(np.arange(length, dtype=np.int64)), + "episode_index": pa.array( + np.full(length, episode.output_index, dtype=np.int64) + ), + "index": pa.array( + np.arange(global_start, global_start + length, dtype=np.int64) + ), + "task_index": pa.array(np.full(length, task_index, dtype=np.int64)), + } + ) + return table, state, action + + +class StatsAccumulator: + def __init__(self) -> None: + self.state: list[np.ndarray] = [] + self.action: list[np.ndarray] = [] + self.relative: dict[str, list[np.ndarray]] = {key: [] for key in JOINT_SLICES} + + def add(self, state: np.ndarray, action: np.ndarray, horizon: int) -> None: + self.state.append(state) + self.action.append(action) + usable = len(state) - horizon + 1 + if usable <= 0: + return + for key, (start, end) in JOINT_SLICES.items(): + reference = state[:usable, start:end] + chunks = np.stack( + [action[offset : offset + usable, start:end] - reference for offset in range(horizon)], + axis=1, + ) + self.relative[key].append(chunks.reshape(-1, end - start)) + + +def numeric_stats(values: np.ndarray) -> dict[str, list[float]]: + values64 = np.asarray(values, dtype=np.float64) + return { + "min": np.min(values64, axis=0).tolist(), + "max": np.max(values64, axis=0).tolist(), + "mean": np.mean(values64, axis=0).tolist(), + "std": np.std(values64, axis=0).tolist(), + "q01": np.quantile(values64, 0.01, axis=0).tolist(), + "q99": np.quantile(values64, 0.99, axis=0).tolist(), + } + + +def finish_stats(accumulator: StatsAccumulator) -> tuple[dict[str, Any], dict[str, Any]]: + if not accumulator.state or not accumulator.action: + raise ValueError("Cannot compute statistics for an empty split") + stats = { + "observation.state": numeric_stats(np.concatenate(accumulator.state, axis=0)), + "action": numeric_stats(np.concatenate(accumulator.action, axis=0)), + } + relative: dict[str, Any] = {} + for key, arrays in accumulator.relative.items(): + if not arrays: + raise ValueError(f"No relative-action samples for {key}") + relative[key] = numeric_stats(np.concatenate(arrays, axis=0)) + return stats, relative + + +def output_video_path(root: Path, episode_index: int, output_key: str) -> Path: + return ( + root + / f"videos/chunk-{episode_index // 1000:03d}" + / output_key + / f"episode_{episode_index:06d}.mp4" + ) + + +def probe_video(path: Path) -> tuple[int, int, int, float]: + completed = subprocess.run( + [ + "ffprobe", + "-v", + "error", + "-count_frames", + "-select_streams", + "v:0", + "-show_entries", + "stream=nb_read_frames,width,height,avg_frame_rate", + "-of", + "json", + str(path), + ], + check=True, + capture_output=True, + text=True, + ) + stream = json.loads(completed.stdout)["streams"][0] + numerator, denominator = (int(part) for part in stream["avg_frame_rate"].split("/")) + rate = numerator / denominator if denominator else 0.0 + return int(stream["nb_read_frames"]), int(stream["width"]), int(stream["height"]), rate + + +def resize_filter(width: int, height: int, mode: str) -> str: + if mode == "stretch": + return f"scale={width}:{height}:flags=lanczos,setsar=1" + if mode == "pad": + return ( + f"scale={width}:{height}:force_original_aspect_ratio=decrease:flags=lanczos," + f"pad={width}:{height}:(ow-iw)/2:(oh-ih)/2:black,setsar=1" + ) + return ( + f"scale={width}:{height}:force_original_aspect_ratio=increase:flags=lanczos," + f"crop={width}:{height},setsar=1" + ) + + +def source_video_path( + source: Path, + source_info: dict[str, Any], + metadata: dict[str, Any], + source_key: str, +) -> tuple[Path, float]: + prefix = f"videos/{source_key}" + required = ( + f"{prefix}/file_index", + f"{prefix}/chunk_index", + f"{prefix}/from_timestamp", + ) + missing = [key for key in required if key not in metadata] + if missing: + raise ValueError(f"Episode video metadata is missing: {missing}") + path = source / source_info["video_path"].format( + video_key=source_key, + chunk_index=int(metadata[f"{prefix}/chunk_index"]), + file_index=int(metadata[f"{prefix}/file_index"]), + ) + if not path.is_file(): + raise FileNotFoundError(path) + return path, float(metadata[f"{prefix}/from_timestamp"]) + + +def convert_video( + source: Path, + destination_root: Path, + source_info: dict[str, Any], + episode: OutputEpisode, + source_key: str, + output_key: str, + width: int, + height: int, + mode: str, + codec: str, + preset: str, + crf: int, +) -> dict[str, Any]: + source_path, start = source_video_path( + source, source_info, episode.source.metadata, source_key + ) + destination = output_video_path(destination_root, episode.output_index, output_key) + destination.parent.mkdir(parents=True, exist_ok=True) + fps = float(source_info["fps"]) + frame_count = episode.source.length + filters = ( + f"fps={fps:g},{resize_filter(width, height, mode)}," + f"tpad=stop_mode=clone:stop_duration=2,trim=end_frame={frame_count}," + "setpts=N/FRAME_RATE/TB" + ) + command = [ + "ffmpeg", + "-nostdin", + "-hide_banner", + "-loglevel", + "error", + "-y", + "-ss", + f"{start:.9f}", + "-i", + str(source_path), + "-an", + "-vf", + filters, + "-frames:v", + str(frame_count), + "-r", + f"{fps:g}", + "-c:v", + codec, + ] + if codec == "libx264": + command += ["-preset", preset, "-crf", str(crf)] + else: + command += ["-preset", "p4", "-cq", str(crf), "-b:v", "0"] + command += [ + "-pix_fmt", + "yuv420p", + "-g", + "2", + "-keyint_min", + "2", + "-sc_threshold", + "0", + "-movflags", + "+faststart", + str(destination), + ] + subprocess.run(command, check=True) + actual_frames, actual_width, actual_height, actual_fps = probe_video(destination) + if actual_frames != frame_count: + raise ValueError( + f"{destination}: {actual_frames} frames; expected {frame_count}" + ) + if (actual_width, actual_height) != (width, height): + raise ValueError( + f"{destination}: resolution {(actual_width, actual_height)}; " + f"expected {(width, height)}" + ) + if abs(actual_fps - fps) > 1e-3: + raise ValueError(f"{destination}: fps {actual_fps}; expected {fps}") + return { + "source_episode_index": episode.source.source_index, + "output_episode_index": episode.output_index, + "camera": output_key, + "frames": actual_frames, + "width": actual_width, + "height": actual_height, + "fps": actual_fps, + "source": str(source_path), + "output": str(destination), + } + + +def build_modality() -> dict[str, Any]: + def field(original_key: str, start: int, end: int) -> dict[str, Any]: + return { + "original_key": original_key, + "start": start, + "end": end, + "rotation_type": None, + "absolute": True, + "dtype": "float32", + "range": None, + } + + return { + "state": { + key: field("observation.state", start, end) + for key, (start, end) in JOINT_SLICES.items() + }, + "action": { + key: field("action", start, end) + for key, (start, end) in JOINT_SLICES.items() + }, + "video": { + "top_head": {"original_key": "observation.images.top_head"}, + "hand_left": {"original_key": "observation.images.hand_left"}, + "hand_right": {"original_key": "observation.images.hand_right"}, + }, + "annotation": { + "language.action_text": { + "original_key": "annotation.language.action_text" + } + }, + } + + +def build_info( + source_info: dict[str, Any], + episodes: list[OutputEpisode], + tasks: list[str], + width: int, + height: int, + split_name: str, +) -> dict[str, Any]: + fps = float(source_info["fps"]) + vector_features = { + "observation.state": { + "dtype": "float32", + "shape": [16], + "names": EXPECTED_NAMES, + }, + "action": {"dtype": "float32", "shape": [16], "names": EXPECTED_NAMES}, + } + features: dict[str, Any] = { + **vector_features, + "annotation.language.action_text": {"dtype": "string", "shape": [1]}, + "timestamp": {"dtype": "float32", "shape": [1]}, + "frame_index": {"dtype": "int64", "shape": [1]}, + "episode_index": {"dtype": "int64", "shape": [1]}, + "index": {"dtype": "int64", "shape": [1]}, + "task_index": {"dtype": "int64", "shape": [1]}, + } + for output_key in SOURCE_CAMERAS.values(): + features[output_key] = { + "dtype": "video", + "shape": [height, width, 3], + "names": ["height", "width", "channels"], + "video_info": { + "video.fps": fps, + "video.codec": "h264", + "video.pix_fmt": "yuv420p", + "video.is_depth_map": False, + "has_audio": False, + }, + } + return { + "codebase_version": "v2.1", + "robot_type": "g2", + "total_episodes": len(episodes), + "total_frames": sum(item.source.length for item in episodes), + "total_tasks": len(tasks), + "chunks_size": 1000, + "fps": fps, + "splits": {split_name: f"0:{len(episodes)}"}, + "data_path": "data/chunk-{episode_chunk:03d}/episode_{episode_index:06d}.parquet", + "video_path": "videos/chunk-{episode_chunk:03d}/{video_key}/episode_{episode_index:06d}.mp4", + "features": features, + } + + +def write_split_metadata( + root: Path, + source_info: dict[str, Any], + split_name: str, + episodes: list[OutputEpisode], + tasks: list[str], + stats: dict[str, Any], + relative_stats: dict[str, Any], + embodiment_tag: str, + width: int, + height: int, + action_horizon: int, + video_reports: list[dict[str, Any]], + stats_source: str, +) -> None: + meta = root / "meta" + task_to_index = {task: index for index, task in enumerate(tasks)} + write_json( + meta / "info.json", + build_info(source_info, episodes, tasks, width, height, split_name), + ) + write_json(meta / "modality.json", build_modality()) + write_json(meta / "embodiment.json", {"embodiment_tag": embodiment_tag}) + write_json(meta / "stats.json", stats) + write_json(meta / "relative_stats_dreamzero.json", relative_stats) + write_jsonl( + meta / "tasks.jsonl", + ({"task_index": index, "task": task} for index, task in enumerate(tasks)), + ) + write_jsonl( + meta / "episodes.jsonl", + ( + { + "episode_index": item.output_index, + "tasks": [item.source.task], + "length": item.source.length, + } + for item in episodes + ), + ) + write_json( + meta / "conversion_report.json", + { + "split": split_name, + "episode_count": len(episodes), + "frame_count": sum(item.source.length for item in episodes), + "source_episode_indices": [item.source.source_index for item in episodes], + "task_indices": { + str(item.output_index): task_to_index[item.source.task] for item in episodes + }, + "source_subtask_indices": { + str(item.output_index): item.source.subtask_index for item in episodes + }, + "source_high_level_tasks": { + str(item.output_index): item.source.source_task for item in episodes + }, + "language_supervision": "subtask_index -> meta/subtasks.parquet", + "language_task_count": len(tasks), + "language_tasks": tasks, + "state_dim": 16, + "action_dim": 16, + "action_horizon": action_horizon, + "camera_order": list(VIDEO_MODALITY_ORDER), + "video_resolution": [width, height], + "video_count": len(video_reports), + "stats_source": stats_source, + }, + ) + + +def convert_split_data( + root: Path, + source_info: dict[str, Any], + data_dataset: pads.Dataset, + source_episodes: list[SourceEpisode], + horizon: int, +) -> tuple[list[OutputEpisode], list[str], StatsAccumulator]: + episodes = [OutputEpisode(source=item, output_index=index) for index, item in enumerate(source_episodes)] + tasks = list(dict.fromkeys(item.source.task for item in episodes)) + task_to_index = {task: index for index, task in enumerate(tasks)} + accumulator = StatsAccumulator() + global_index = 0 + for count, item in enumerate(episodes, start=1): + source_table = table_for_episode(data_dataset, item.source) + table, state, action = output_table( + source_table, + item, + float(source_info["fps"]), + task_to_index[item.source.task], + global_index, + ) + destination = ( + root + / f"data/chunk-{item.output_index // 1000:03d}" + / f"episode_{item.output_index:06d}.parquet" + ) + destination.parent.mkdir(parents=True, exist_ok=True) + pq.write_table(table, destination, compression="zstd", row_group_size=item.source.length) + accumulator.add(state, action, horizon) + global_index += item.source.length + if count % 20 == 0 or count == len(episodes): + LOG.info("%s parquet: %d/%d", root.name, count, len(episodes)) + return episodes, tasks, accumulator + + +def convert_split_videos( + source: Path, + root: Path, + source_info: dict[str, Any], + episodes: list[OutputEpisode], + args: argparse.Namespace, +) -> list[dict[str, Any]]: + jobs = [ + (episode, source_key, output_key) + for episode in episodes + for source_key, output_key in SOURCE_CAMERAS.items() + ] + reports: list[dict[str, Any]] = [] + with ThreadPoolExecutor(max_workers=args.workers) as executor: + futures = [ + executor.submit( + convert_video, + source, + root, + source_info, + episode, + source_key, + output_key, + args.video_width, + args.video_height, + args.resize_mode, + args.video_codec, + args.video_preset, + args.video_crf, + ) + for episode, source_key, output_key in jobs + ] + for count, future in enumerate(as_completed(futures), start=1): + reports.append(future.result()) + if count % 30 == 0 or count == len(jobs): + LOG.info("%s videos: %d/%d", root.name, count, len(jobs)) + return sorted( + reports, + key=lambda item: (item["output_episode_index"], item["camera"]), + ) + + +def validate_output(root: Path, split_name: str) -> None: + meta = root / "meta" + required = ( + "info.json", + "modality.json", + "embodiment.json", + "stats.json", + "relative_stats_dreamzero.json", + "tasks.jsonl", + "episodes.jsonl", + "conversion_report.json", + ) + missing = [name for name in required if not (meta / name).is_file()] + if missing: + raise ValueError(f"{split_name}: missing metadata files: {missing}") + info = read_json(meta / "info.json") + episode_count = int(info["total_episodes"]) + data_count = len(list((root / "data").glob("chunk-*/*.parquet"))) + video_count = len(list((root / "videos").glob("chunk-*/*/*.mp4"))) + if data_count != episode_count: + raise ValueError(f"{split_name}: {data_count} parquets for {episode_count} episodes") + if video_count != episode_count * 3: + raise ValueError(f"{split_name}: {video_count} videos; expected {episode_count * 3}") + + output_data = pads.dataset( + [str(path) for path in parquet_files(root / "data")], + format="parquet", + ).to_table( + columns=[ + "episode_index", + "annotation.language.action_text", + "task_index", + ] + ) + output_episode_values = output_data["episode_index"].to_numpy( + zero_copy_only=False + ) + output_task_values = output_data["task_index"].to_numpy(zero_copy_only=False) + output_text_values = output_data[ + "annotation.language.action_text" + ].to_pylist() + + episode_labels: dict[int, tuple[int, str]] = {} + for episode_value, task_value, text_value in zip( + output_episode_values, + output_task_values, + output_text_values, + strict=True, + ): + episode_index = int(episode_value) + label = (int(task_value), str(text_value).strip()) + previous = episode_labels.setdefault(episode_index, label) + if previous != label: + raise ValueError( + f"{split_name}: episode {episode_index} has multiple language labels: " + f"{previous!r} / {label!r}" + ) + + if len(episode_labels) != episode_count: + raise ValueError( + f"{split_name}: found language labels for {len(episode_labels)} " + f"of {episode_count} episodes" + ) + + task_rows: list[dict[str, Any]] = [] + with (meta / "tasks.jsonl").open(encoding="utf-8") as handle: + for line in handle: + if line.strip(): + task_rows.append(json.loads(line)) + metadata_tasks = { + int(row["task_index"]): str(row["task"]).strip() for row in task_rows + } + observed_tasks: dict[int, str] = {} + for task_index, text in episode_labels.values(): + previous = observed_tasks.setdefault(task_index, text) + if previous != text: + raise ValueError( + f"{split_name}: task_index {task_index} maps to multiple texts: " + f"{previous!r} / {text!r}" + ) + if observed_tasks != metadata_tasks: + raise ValueError( + f"{split_name}: parquet language labels do not match tasks.jsonl" + ) + if len(metadata_tasks) != int(info["total_tasks"]): + raise ValueError( + f"{split_name}: tasks.jsonl has {len(metadata_tasks)} tasks; " + f"info.json says {info['total_tasks']}" + ) + LOG.info( + "%s language audit: %d episodes, %d unique subtask texts", + split_name, + episode_count, + len(metadata_tasks), + ) + + +def main() -> int: + args = parse_args() + logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s %(message)s") + if args.workers < 1: + raise ValueError("--workers must be positive") + if args.action_horizon < 1: + raise ValueError("--action-horizon must be positive") + if min(args.video_width, args.video_height) < 2 or args.video_width % 2 or args.video_height % 2: + raise ValueError("Video width and height must be positive even integers") + + source = args.source.resolve() + output = args.output.resolve() + require_tools(args.video_codec) + validate_paths(source, output, args.overwrite) + + source_info = read_json(source / "meta/info.json") + validate_source_info(source_info) + source_episodes, data_files, data_dataset = load_source_episodes( + source, source_info, args.max_source_episodes + ) + train_source, test_source = split_episodes( + source_episodes, args.test_episodes, args.split_mode, args.split_seed + ) + min_frames = args.min_episode_frames or (args.action_horizon + 1) + validate_episode_lengths(train_source, min_frames, "train") + if test_source: + validate_episode_lengths(test_source, min_frames, "test") + + LOG.info( + "Source: %d episodes, %d parquet shards; split into %d train / %d test", + len(source_episodes), + len(data_files), + len(train_source), + len(test_source), + ) + write_json( + output / "split_manifest.json", + { + "source": str(source), + "split_mode": args.split_mode, + "split_seed": args.split_seed if args.split_mode != "tail" else None, + "language_supervision": "subtask_index -> meta/subtasks.parquet", + "train_source_episode_indices": [item.source_index for item in train_source], + "test_source_episode_indices": [item.source_index for item in test_source], + "train_subtask_indices": sorted( + {item.subtask_index for item in train_source} + ), + "test_subtask_indices": sorted( + {item.subtask_index for item in test_source} + ), + }, + ) + + train_root = output / "train" + train_episodes, train_tasks, train_acc = convert_split_data( + train_root, + source_info, + data_dataset, + train_source, + args.action_horizon, + ) + train_stats, train_relative_stats = finish_stats(train_acc) + train_video_reports = convert_split_videos( + source, train_root, source_info, train_episodes, args + ) + write_split_metadata( + train_root, + source_info, + "train", + train_episodes, + train_tasks, + train_stats, + train_relative_stats, + args.embodiment_tag, + args.video_width, + args.video_height, + args.action_horizon, + train_video_reports, + "train", + ) + validate_output(train_root, "train") + + if test_source: + test_root = output / "test" + test_episodes, test_tasks, _ = convert_split_data( + test_root, + source_info, + data_dataset, + test_source, + args.action_horizon, + ) + test_video_reports = convert_split_videos( + source, test_root, source_info, test_episodes, args + ) + # Test data deliberately uses training-set normalization statistics. + write_split_metadata( + test_root, + source_info, + "test", + test_episodes, + test_tasks, + train_stats, + train_relative_stats, + args.embodiment_tag, + args.video_width, + args.video_height, + args.action_horizon, + test_video_reports, + "train", + ) + validate_output(test_root, "test") + + LOG.info("Conversion complete. DreamZero train root: %s", train_root) + if test_source: + LOG.info("Held-out test root: %s", output / "test") + return 0 + + +if __name__ == "__main__": + try: + sys.exit(main()) + except Exception: + LOG.exception("Conversion failed") + sys.exit(1) diff --git a/scripts/data/prepare_g2_action_adapter_dataset.py b/scripts/data/prepare_g2_action_adapter_dataset.py new file mode 100755 index 00000000..e58afb0d --- /dev/null +++ b/scripts/data/prepare_g2_action_adapter_dataset.py @@ -0,0 +1,181 @@ +#!/usr/bin/env python3 +"""Prepare a non-destructive G2 dataset for action-adapter-only training. + +The source GEAR dataset stores G2 grippers in SDK space [0, -0.785]. This +script rewrites state/action gripper dimensions 7 and 15 into policy space +0=closed, 1=open, recomputes normalization metadata, and records active/hold +24-step windows. Videos are symlinked; source data is never modified. +""" + +from __future__ import annotations + +import argparse +import json +import shutil +import sys +from pathlib import Path + +import numpy as np +import pandas as pd + +REPO_ROOT = Path(__file__).resolve().parents[2] +if str(REPO_ROOT) not in sys.path: + sys.path.insert(0, str(REPO_ROOT)) + +from scripts.data.convert_lerobot_to_gear import ( + compute_relative_stats, + compute_stats, +) + + +GRIPPER_DIMS = (7, 15) +ARM_DIMS = (0, 1, 2, 3, 4, 5, 6, 8, 9, 10, 11, 12, 13, 14) +G2_OPEN_POSITION = -0.785 + + +def read_json(path: Path) -> dict: + with path.open() as stream: + return json.load(stream) + + +def write_json(path: Path, value: dict) -> None: + with path.open("w") as stream: + json.dump(value, stream, indent=2, ensure_ascii=False) + + +def sdk_to_policy(vector: object) -> np.ndarray: + result = np.asarray(vector, dtype=np.float32).copy() + if result.shape != (16,): + raise ValueError(f"Expected 16D G2 vector, got {result.shape}") + result[list(GRIPPER_DIMS)] = np.clip( + result[list(GRIPPER_DIMS)] / G2_OPEN_POSITION, + 0.0, + 1.0, + ) + return result + + +def prepare_split(source: Path, output: Path, action_horizon: int) -> None: + if output.exists(): + raise FileExistsError( + f"Refusing to overwrite derived dataset: {output}" + ) + (output / "data").mkdir(parents=True) + shutil.copytree(source / "meta", output / "meta") + write_json( + output / "meta" / "embodiment.json", + {"robot_type": "g2", "embodiment_tag": "g2"}, + ) + + source_videos = source / "videos" + if source_videos.exists(): + (output / "videos").symlink_to( + source_videos.resolve(), + target_is_directory=True, + ) + + output_parquets: list[Path] = [] + window_metrics: list[tuple[int, int, float, bool]] = [] + for source_parquet in sorted(source.glob("data/chunk-*/*.parquet")): + relative = source_parquet.relative_to(source) + output_parquet = output / relative + output_parquet.parent.mkdir(parents=True, exist_ok=True) + frame = pd.read_parquet(source_parquet) + for column in ("observation.state", "action"): + if column not in frame: + raise KeyError(f"{source_parquet}: missing {column}") + frame[column] = [sdk_to_policy(value) for value in frame[column]] + frame.to_parquet(output_parquet, index=False) + output_parquets.append(output_parquet) + + episode = int(frame["episode_index"].iloc[0]) + states = np.stack(frame["observation.state"]).astype(np.float32) + actions = np.stack(frame["action"]).astype(np.float32) + usable = max(0, len(frame) - action_horizon + 1) + for step in range(usable): + future = actions[step : step + action_horizon] + arm_delta = future[:, ARM_DIMS] - states[step, ARM_DIMS] + arm_motion = float(np.linalg.norm(arm_delta, axis=1).max()) + grip_transition = bool( + np.any( + np.abs( + future[:, GRIPPER_DIMS] + - states[step, GRIPPER_DIMS] + ) + > 0.25 + ) + ) + window_metrics.append( + (episode, step, arm_motion, grip_transition) + ) + + if not output_parquets: + raise FileNotFoundError(f"No GEAR parquet files below {source}") + + modality = read_json(output / "meta" / "modality.json") + info = read_json(output / "meta" / "info.json") + numeric_columns = list(read_json(source / "meta" / "stats.json")) + write_json( + output / "meta" / "stats.json", + compute_stats(output_parquets, numeric_columns), + ) + write_json( + output / "meta" / "relative_stats_dreamzero.json", + compute_relative_stats( + output_parquets, + modality, + ["left_joint_position", "right_joint_position"], + action_horizon=action_horizon, + ), + ) + + nonzero_motion = np.asarray( + [metric[2] for metric in window_metrics if metric[2] > 1e-8], + dtype=np.float64, + ) + if nonzero_motion.size == 0: + raise RuntimeError(f"{source}: no nonzero arm-motion windows") + threshold = float(np.quantile(nonzero_motion, 0.30)) + active: dict[str, list[int]] = {} + hold: dict[str, list[int]] = {} + for episode, step, arm_motion, grip_transition in window_metrics: + target = active if arm_motion >= threshold or grip_transition else hold + target.setdefault(str(episode), []).append(step) + write_json( + output / "meta" / "g2_active_hold_windows.json", + { + "schema_version": 1, + "action_horizon": action_horizon, + "arm_motion_threshold_rule": "p30_nonzero_24_step_l2", + "arm_motion_threshold": threshold, + "gripper_transition_threshold": 0.25, + "active_count": sum(map(len, active.values())), + "hold_count": sum(map(len, hold.values())), + "active": active, + "hold": hold, + }, + ) + print( + f"[G2 DATASET] {source} -> {output} " + f"episodes={info['total_episodes']} threshold={threshold:.8f} " + f"active={sum(map(len, active.values()))} " + f"hold={sum(map(len, hold.values()))}" + ) + + +def main() -> None: + parser = argparse.ArgumentParser() + parser.add_argument("--source-root", type=Path, required=True) + parser.add_argument("--output-root", type=Path, required=True) + parser.add_argument("--action-horizon", type=int, default=24) + args = parser.parse_args() + for split in ("train", "test"): + prepare_split( + args.source_root / split, + args.output_root / split, + args.action_horizon, + ) + + +if __name__ == "__main__": + main() diff --git a/scripts/data/repartition_g2_gear_stratified.py b/scripts/data/repartition_g2_gear_stratified.py new file mode 100644 index 00000000..a0800d30 --- /dev/null +++ b/scripts/data/repartition_g2_gear_stratified.py @@ -0,0 +1,239 @@ +#!/usr/bin/env python3 +"""Repartition existing GEAR train/test splits into a task-stratified dataset.""" + +from __future__ import annotations + +import argparse +import json +import os +import random +import shutil +import sys +from dataclasses import dataclass +from pathlib import Path + +import pandas as pd +import numpy as np + +REPO_ROOT = Path(__file__).resolve().parents[2] +if str(REPO_ROOT) not in sys.path: + sys.path.insert(0, str(REPO_ROOT)) + +from scripts.data.convert_lerobot_to_gear import compute_relative_stats, compute_stats + + +@dataclass(frozen=True) +class Episode: + root: Path + source_split: str + source_index: int + task: str + length: int + + +def read_json(path: Path) -> dict: + return json.loads(path.read_text(encoding="utf-8")) + + +def write_json(path: Path, value: object) -> None: + path.write_text( + json.dumps(value, ensure_ascii=False, indent=2), encoding="utf-8" + ) + + +def load_episodes(source: Path) -> tuple[list[Episode], list[str]]: + episodes: list[Episode] = [] + canonical_tasks: list[str] = [] + for split in ("train", "test"): + root = source / split + task_rows = [json.loads(line) for line in (root / "meta/tasks.jsonl").read_text().splitlines()] + tasks = {int(row["task_index"]): row["task"] for row in task_rows} + for task in tasks.values(): + if task not in canonical_tasks: + canonical_tasks.append(task) + for line in (root / "meta/episodes.jsonl").read_text().splitlines(): + row = json.loads(line) + task = row["tasks"][0] + episodes.append( + Episode(root, split, int(row["episode_index"]), task, int(row["length"])) + ) + return episodes, canonical_tasks + + +def source_parquet(item: Episode) -> Path: + return item.root / f"data/chunk-{item.source_index // 1000:03d}/episode_{item.source_index:06d}.parquet" + + +def materialize_split( + output: Path, + split: str, + episodes: list[Episode], + canonical_tasks: list[str], + info_template: dict, + modality: dict, + embodiment: dict, + invert_policy_gripper: bool, +) -> tuple[list[Path], list[dict], list[str]]: + root = output / split + (root / "meta").mkdir(parents=True) + used_tasks = [task for task in canonical_tasks if any(e.task == task for e in episodes)] + task_ids = {task: index for index, task in enumerate(used_tasks)} + parquet_paths: list[Path] = [] + episode_rows: list[dict] = [] + global_index = 0 + video_keys = tuple(modality["video"]) + + for output_index, item in enumerate(episodes): + chunk = output_index // 1000 + destination = root / f"data/chunk-{chunk:03d}/episode_{output_index:06d}.parquet" + destination.parent.mkdir(parents=True, exist_ok=True) + frame = pd.read_parquet(source_parquet(item)) + if len(frame) != item.length: + raise RuntimeError(f"Length mismatch for {source_parquet(item)}") + if invert_policy_gripper: + for column in ("observation.state", "action"): + values = np.stack(frame[column]).astype(np.float32) + if values.shape[1:] != (16,): + raise ValueError(f"{source_parquet(item)}: {column} is not 16D") + if np.any(values[:, (7, 15)] < 0) or np.any(values[:, (7, 15)] > 1): + raise ValueError( + f"{source_parquet(item)}: {column} gripper is outside [0,1]" + ) + values[:, (7, 15)] = 1.0 - values[:, (7, 15)] + frame[column] = list(values) + frame["episode_index"] = output_index + frame["task_index"] = task_ids[item.task] + frame["index"] = range(global_index, global_index + len(frame)) + frame.to_parquet(destination, index=False) + parquet_paths.append(destination) + global_index += len(frame) + episode_rows.append( + {"episode_index": output_index, "tasks": [item.task], "length": len(frame)} + ) + + source_chunk = item.source_index // 1000 + for key in video_keys: + source_video = item.root / f"videos/chunk-{source_chunk:03d}/observation.images.{key}/episode_{item.source_index:06d}.mp4" + destination_video = root / f"videos/chunk-{chunk:03d}/observation.images.{key}/episode_{output_index:06d}.mp4" + destination_video.parent.mkdir(parents=True, exist_ok=True) + if not source_video.is_file(): + raise FileNotFoundError(source_video) + os.link(source_video, destination_video) + + info = dict(info_template) + info["total_episodes"] = len(episodes) + info["total_frames"] = global_index + info["total_tasks"] = len(used_tasks) + info["splits"] = {split: f"0:{len(episodes)}"} + write_json(root / "meta/info.json", info) + write_json(root / "meta/modality.json", modality) + write_json(root / "meta/embodiment.json", embodiment) + (root / "meta/tasks.jsonl").write_text( + "".join( + json.dumps({"task_index": index, "task": task}, ensure_ascii=False) + "\n" + for index, task in enumerate(used_tasks) + ), + encoding="utf-8", + ) + (root / "meta/episodes.jsonl").write_text( + "".join(json.dumps(row, ensure_ascii=False) + "\n" for row in episode_rows), + encoding="utf-8", + ) + return parquet_paths, episode_rows, used_tasks + + +def main() -> None: + parser = argparse.ArgumentParser() + parser.add_argument("--source", type=Path, required=True) + parser.add_argument("--output", type=Path, required=True) + parser.add_argument("--test-per-task", type=int, default=2) + parser.add_argument("--seed", type=int, default=42) + parser.add_argument( + "--invert-policy-gripper", + action="store_true", + help="Convert source 0=open,1=closed to canonical 0=closed,1=open.", + ) + args = parser.parse_args() + if args.output.exists(): + raise FileExistsError(f"Refusing to overwrite {args.output}") + + episodes, canonical_tasks = load_episodes(args.source) + by_task: dict[str, list[Episode]] = {task: [] for task in canonical_tasks} + for item in episodes: + by_task[item.task].append(item) + rng = random.Random(args.seed) + test_set: set[Episode] = set() + for task, group in by_task.items(): + if len(group) <= args.test_per_task: + raise RuntimeError(f"Task {task!r} has only {len(group)} episodes") + test_set.update(rng.sample(group, args.test_per_task)) + train = [item for item in episodes if item not in test_set] + test = [item for item in episodes if item in test_set] + + template_root = args.source / "train/meta" + info = read_json(template_root / "info.json") + modality = read_json(template_root / "modality.json") + embodiment = read_json(template_root / "embodiment.json") + try: + train_paths, train_rows, train_tasks = materialize_split( + args.output, "train", train, canonical_tasks, info, modality, + embodiment, args.invert_policy_gripper, + ) + test_paths, test_rows, test_tasks = materialize_split( + args.output, "test", test, canonical_tasks, info, modality, + embodiment, args.invert_policy_gripper, + ) + numeric_columns = list(read_json(template_root / "stats.json")) + train_stats = compute_stats(train_paths, numeric_columns) + relative_stats = compute_relative_stats( + train_paths, + modality, + ["left_joint_position", "right_joint_position"], + action_horizon=24, + ) + for split in ("train", "test"): + write_json(args.output / split / "meta/stats.json", train_stats) + write_json( + args.output / split / "meta/relative_stats_dreamzero.json", + relative_stats, + ) + write_json( + args.output / split / "meta/conversion_report.json", + { + "source": str(args.source), + "split_mode": "stratified_by_subtask", + "test_per_task": args.test_per_task, + "seed": args.seed, + "invert_policy_gripper": args.invert_policy_gripper, + }, + ) + write_json( + args.output / "split_manifest.json", + { + "source": str(args.source), + "split_mode": "stratified_by_subtask", + "seed": args.seed, + "test_per_task": args.test_per_task, + "invert_policy_gripper": args.invert_policy_gripper, + "train_episodes": len(train_rows), + "test_episodes": len(test_rows), + "train_tasks": train_tasks, + "test_tasks": test_tasks, + "test_source_episodes": [ + {"split": item.source_split, "episode_index": item.source_index, "task": item.task} + for item in test + ], + }, + ) + except Exception: + shutil.rmtree(args.output, ignore_errors=True) + raise + + print( + f"Created {args.output}: train={len(train)} test={len(test)} " + f"tasks={len(canonical_tasks)} test_per_task={args.test_per_task}" + ) + + +if __name__ == "__main__": + main() diff --git a/scripts/eval/action_input_sensitivity.py b/scripts/eval/action_input_sensitivity.py new file mode 100644 index 00000000..f8b3ccb9 --- /dev/null +++ b/scripts/eval/action_input_sensitivity.py @@ -0,0 +1,1532 @@ +#!/usr/bin/env python3 +"""Counterfactual G2 action/video sensitivity evaluation against a live server. + +For every target training episode this script keeps the target video and the +target action trajectory fixed, then compares: + +* ``baseline``: target state + chronological four-frame history; +* ``state_swap``: a time-aligned state trajectory and four-frame history from + another episode, both aligned to the target timeline; +* ``history_reverse``: target state + the same four target frames in reverse + order; +* ``state_swap_history_reverse``: the same swapped state/history input as + ``state_swap``, but with those four swapped frames in reverse order. +* ``temporal_shift``: target state and chronological four-frame history from + another target-episode time, while the GT/action remains at the baseline + time. + +The server still receives exactly the production G2 websocket observation +contract. Action chunks are recorded as (24, 16) arrays and compared with +the target episode's GT action horizon. Every decoded frame from the +server's raw generated video is retained and written at a configurable slow +display FPS; GT frames are aligned to each continuous inference round for an +intuitive prediction-vs-GT video. +""" + +from __future__ import annotations + +import argparse +import csv +import json +import logging +import shutil +import sys +import time +import uuid +from pathlib import Path + +import cv2 +import imageio.v2 as imageio +import numpy as np +import pyarrow.parquet as pq + +ROOT = Path(__file__).resolve().parents[2] +sys.path.insert(0, str(ROOT)) + +from eval_utils.policy_client import WebsocketClientPolicy # noqa: E402 + + +ACTION_NAMES = tuple( + [f"left_joint_{i}" for i in range(7)] + + ["left_gripper"] + + [f"right_joint_{i}" for i in range(7)] + + ["right_gripper"] +) +ARM_DIMS = np.asarray([*range(7), *range(8, 15)], dtype=np.int64) +GRIPPER_DIMS = np.asarray([7, 15], dtype=np.int64) +ACTION_HORIZON = 24 +HISTORY_LEN = 4 +STAGE_NAMES = ( + "right_grasp_ethernet_dock", + "left_grasp_cable_1", + "insert_cable_1_into_dock", + "left_grasp_cable_2", + "insert_cable_2_into_dock", + "insert_dock_hold_5s", + "unplug_ethernet_dock", + "return_ethernet_dock", +) + +CONDITION_ORDER = ( + "baseline", + "state_only", + "history_only", + "state_history_swap", + "state_swap", + "history_reverse", + "state_swap_history_reverse", + "temporal_shift", +) +CONDITION_COLORS = { + "baseline": "tab:blue", + "state_only": "tab:red", + "history_only": "tab:green", + "state_history_swap": "tab:purple", + "state_swap": "tab:red", + "history_reverse": "tab:orange", + "state_swap_history_reverse": "tab:brown", + "temporal_shift": "tab:cyan", +} + + +def _stage_info(anchor: int, num_rows: int) -> dict[str, object]: + """Return the same eight-bin trajectory-stage approximation used for selection.""" + stage_index = min(len(STAGE_NAMES) - 1, int(anchor * len(STAGE_NAMES) / num_rows)) + return { + "index": stage_index, + "name": STAGE_NAMES[stage_index], + "source": "normalized episode progress split into eight bins", + } + + +def _condition_label( + condition: str, + target_stage: dict[str, object], + donor_stage: dict[str, object], +) -> str: + target_name = str(target_stage["name"]) + donor_name = str(donor_stage["name"]) + return { + "baseline": f"BASELINE target={target_name}", + "state_only": f"STATE ONLY swap-source={donor_name}", + "history_only": f"HISTORY ONLY swap-source={donor_name}", + "state_history_swap": f"STATE+HISTORY swap-source={donor_name}", + }.get(condition, condition.upper()) + + +def _episode_path( + root: Path, + template: str, + episode: int, + chunks_size: int, + **kwargs: object, +) -> Path: + return root / template.format( + episode_chunk=episode // chunks_size, + episode_index=episode, + **kwargs, + ) + + +def _read_video(path: Path, expected_rows: int) -> np.ndarray: + """Read a G2 training video as BGR, matching the live client capture path.""" + capture = cv2.VideoCapture(str(path)) + if not capture.isOpened(): + raise RuntimeError(f"Failed to open G2 video: {path}") + frames: list[np.ndarray] = [] + try: + while len(frames) < expected_rows: + ok, frame = capture.read() + if not ok: + break + frames.append(np.ascontiguousarray(frame)) + finally: + capture.release() + if len(frames) < expected_rows: + raise RuntimeError( + f"{path} contains {len(frames)} frames, expected {expected_rows}" + ) + return np.stack(frames, axis=0) + + +def _encode_video_observation( + frames: np.ndarray, + quality: int, +) -> dict[str, object]: + """Encode BGR frames exactly as the production G2 JPEG transport does.""" + array = np.asarray(frames) + if array.ndim != 4 or array.shape[-1] != 3: + raise ValueError(f"Expected (T,H,W,3) video frames, got {array.shape}") + encoded: list[bytes] = [] + jpeg_quality = int(np.clip(quality, 1, 100)) + for frame in array: + ok, payload = cv2.imencode( + ".jpg", + np.ascontiguousarray(frame), + [int(cv2.IMWRITE_JPEG_QUALITY), jpeg_quality], + ) + if not ok: + raise RuntimeError("Failed to JPEG-encode an evaluation frame") + encoded.append(payload.tobytes()) + return { + "__dreamzero_image_encoding__": "jpeg_sequence", + "shape": tuple(int(dim) for dim in array.shape), + "dtype": str(array.dtype), + "quality": jpeg_quality, + "frames": encoded, + } + + +def _grid_rgb( + top_bgr: np.ndarray, + left_bgr: np.ndarray, + right_bgr: np.ndarray, +) -> np.ndarray: + top = cv2.cvtColor(top_bgr, cv2.COLOR_BGR2RGB) + left = cv2.cvtColor(left_bgr, cv2.COLOR_BGR2RGB) + right = cv2.cvtColor(right_bgr, cv2.COLOR_BGR2RGB) + black = np.zeros_like(top) + return np.concatenate( + [np.concatenate([top, left], axis=1), np.concatenate([right, black], axis=1)], + axis=0, + ) + + +def _label(frame_rgb: np.ndarray, text: str) -> np.ndarray: + frame_bgr = cv2.cvtColor(frame_rgb, cv2.COLOR_RGB2BGR) + cv2.rectangle(frame_bgr, (0, 0), (max(430, frame_bgr.shape[1] // 2), 32), (0, 0, 0), -1) + cv2.putText( + frame_bgr, + text, + (8, 22), + cv2.FONT_HERSHEY_SIMPLEX, + 0.55, + (255, 255, 255), + 2, + cv2.LINE_AA, + ) + return cv2.cvtColor(frame_bgr, cv2.COLOR_BGR2RGB) + + +def _save_video(path: Path, frames: list[np.ndarray], fps: int) -> None: + if not frames: + raise RuntimeError(f"No frames to save: {path}") + path.parent.mkdir(parents=True, exist_ok=True) + imageio.mimsave(path, frames, fps=int(fps), codec="libx264", macro_block_size=None) + + +def _parse_pairs(spec: str) -> list[tuple[int, int]]: + pairs: list[tuple[int, int]] = [] + for token in spec.split(","): + token = token.strip() + if not token: + continue + parts = token.split(":") + if len(parts) != 2: + raise ValueError( + f"Invalid pair {token!r}; expected target_episode:donor_episode" + ) + target, donor = (int(part.strip()) for part in parts) + if target < 0 or donor < 0: + raise ValueError(f"Episode indices must be non-negative: {token!r}") + pairs.append((target, donor)) + if not pairs: + raise ValueError("--pairs must contain at least one target:donor pair") + return pairs + + +def _parse_anchor_pairs(spec: str) -> list[tuple[int, int, int, int]]: + """Parse target:donor:target_anchor:donor_anchor specifications.""" + pairs: list[tuple[int, int, int, int]] = [] + for token in spec.split(","): + token = token.strip() + if not token: + continue + parts = token.split(":") + if len(parts) != 4: + raise ValueError( + f"Invalid anchor pair {token!r}; expected " + "target_episode:donor_episode:target_anchor:donor_anchor" + ) + target, donor, target_anchor, donor_anchor = ( + int(part.strip()) for part in parts + ) + if min(target, donor, target_anchor, donor_anchor) < 0: + raise ValueError(f"Anchor pair values must be non-negative: {token!r}") + pairs.append((target, donor, target_anchor, donor_anchor)) + if not pairs: + raise ValueError("--anchor-pairs must contain at least one item") + return pairs + + +def _parse_anchor_offsets(spec: str) -> list[int]: + offsets = [int(token.strip()) for token in spec.split(",") if token.strip()] + if not offsets: + raise ValueError("--anchor-offsets must contain at least one integer") + if min(offsets) < 0: + raise ValueError("--anchor-offsets values must be non-negative") + return offsets + + +def _parse_conditions(spec: str) -> list[str]: + conditions = [token.strip() for token in spec.split(",") if token.strip()] + unknown = [token for token in conditions if token not in CONDITION_ORDER] + if unknown: + raise ValueError( + f"Unknown condition(s) {unknown}; choose from {CONDITION_ORDER}" + ) + if "baseline" not in conditions: + conditions.insert(0, "baseline") + # Stable order makes all plots and summary files directly comparable. + return [condition for condition in CONDITION_ORDER if condition in conditions] + + +def _parse_start_frames(spec: str | None) -> dict[int, int]: + if not spec: + return {} + mapping: dict[int, int] = {} + for token in spec.split(","): + token = token.strip() + if not token: + continue + parts = token.split(":") + if len(parts) != 2: + raise ValueError( + f"Invalid --start-frames item {token!r}; expected episode:start_frame" + ) + episode, start_frame = (int(part.strip()) for part in parts) + if episode < 0 or start_frame < 0: + raise ValueError(f"Episode and start frame must be non-negative: {token!r}") + mapping[episode] = start_frame + return mapping + + +def _load_episode( + root: Path, + info: dict[str, object], + episode: int, +) -> dict[str, object]: + chunks_size = int(info.get("chunks_size", 1000)) + data_template = str( + info.get( + "data_path", + "data/chunk-{episode_chunk:03d}/episode_{episode_index:06d}.parquet", + ) + ) + parquet_path = _episode_path(root, data_template, episode, chunks_size) + if not parquet_path.exists(): + raise FileNotFoundError(f"Episode {episode} parquet not found: {parquet_path}") + table = pq.read_table(parquet_path) + num_rows = int(table.num_rows) + state = np.asarray(table["observation.state"].to_pylist(), dtype=np.float32).reshape(num_rows, 16) + action = np.asarray(table["action"].to_pylist(), dtype=np.float32).reshape(num_rows, 16) + + video_template = str( + info.get( + "video_path", + "videos/chunk-{episode_chunk:03d}/{video_key}/episode_{episode_index:06d}.mp4", + ) + ) + videos = { + name: _read_video( + _episode_path(root, video_template, episode, chunks_size, video_key=key), + num_rows, + ) + for name, key in ( + ("top_head", "observation.images.top_head"), + ("hand_left", "observation.images.hand_left"), + ("hand_right", "observation.images.hand_right"), + ) + } + prompt_key = "annotation.language.action_text" + prompt = str(table[prompt_key][0].as_py()) if prompt_key in table.column_names else "" + return { + "episode": episode, + "num_rows": num_rows, + "state": state, + "action": action, + "videos": videos, + "prompt": prompt, + } + + +def _choose_starts( + num_rows: int, + block_stride: int, + max_blocks: int | None, + start_frame: int = 0, + contiguous: bool = False, +) -> list[int]: + # Packet [s:s+4] is anchored at s+3. The GT horizon is action[s+3:s+27]. + max_start = num_rows - (HISTORY_LEN + ACTION_HORIZON) + if max_start < 0: + return [] + if start_frame < 0 or start_frame > max_start: + return [] + starts = list(range(start_frame, max_start + 1, block_stride)) + if max_blocks is not None and len(starts) > max_blocks: + if contiguous: + # Presentation videos should be a continuous piece of one + # trajectory. The old evenly-spaced selection is useful for a + # broad metric sweep but makes a visually confusing jump-cut + # video, so it is now opt-in through the default=False path. + starts = starts[:max_blocks] + else: + # Cover the whole trajectory for the metric-oriented report. + positions = np.rint(np.linspace(0, len(starts) - 1, max_blocks)).astype(int) + starts = [starts[int(pos)] for pos in positions] + return starts + + +def _state_alignment_indices(target_rows: int, donor_rows: int) -> np.ndarray: + if target_rows <= 1 or donor_rows <= 1: + return np.zeros(target_rows, dtype=np.int64) + return np.rint(np.linspace(0, donor_rows - 1, target_rows)).astype(np.int64) + + +def _list_video_names(server_video_dir: Path) -> set[str]: + try: + return {path.name for path in server_video_dir.glob("*.mp4")} + except OSError as exc: + raise RuntimeError(f"Cannot scan server video directory {server_video_dir}: {exc}") from exc + + +def _wait_for_new_video( + server_video_dir: Path, + before: set[str], + timeout_seconds: float, + min_frames: int = 1, +) -> Path: + deadline = time.time() + timeout_seconds + observed: dict[str, tuple[int, float]] = {} + while time.time() < deadline: + candidates = [ + path + for path in server_video_dir.glob("*.mp4") + if path.name not in before and path.stat().st_size > 0 + ] + for candidate in sorted(candidates, key=lambda path: path.stat().st_mtime_ns, reverse=True): + try: + size = int(candidate.stat().st_size) + except OSError: + continue + now = time.time() + previous = observed.get(candidate.name) + if previous is None or previous[0] != size: + observed[candidate.name] = (size, now) + continue + if now - previous[1] < 1.0: + continue + try: + frame_count = len(_read_mp4(candidate)) + except Exception: + continue + if frame_count < min_frames: + # A newly-created but incomplete MP4 can be readable while + # only containing the first inference chunk. Keep waiting + # for the same file to grow instead of silently copying it. + continue + return candidate + time.sleep(0.25) + raise RuntimeError( + f"Server reset did not produce a new video in {server_video_dir}; " + f"existing tail={sorted(before)[-3:]}; " + f"expected at least {min_frames} decoded frames" + ) + + +def _read_mp4(path: Path) -> list[np.ndarray]: + reader = imageio.get_reader(path) + try: + frames = [np.ascontiguousarray(frame) for frame in reader] + finally: + reader.close() + if not frames: + raise RuntimeError(f"Server video contains no frames: {path}") + return frames + + +def _resample_frames(frames: list[np.ndarray], target_count: int) -> list[np.ndarray]: + if target_count <= 0: + return [] + if len(frames) == target_count: + return frames + indices = np.rint(np.linspace(0, len(frames) - 1, target_count)).astype(np.int64) + return [frames[int(index)] for index in indices] + + +def _metric(error: np.ndarray) -> dict[str, float | int]: + return { + "overall_mae": float(error.mean()), + "arm_mae": float(error[..., ARM_DIMS].mean()), + "gripper_mae": float(error[..., GRIPPER_DIMS].mean()), + "left_arm_mae": float(error[..., :7].mean()), + "right_arm_mae": float(error[..., 8:15].mean()), + "num_action_rows": int(np.prod(error.shape[:-1])), + } + + +def _write_action_tables( + output_dir: Path, + target_episode: int, + donor_episode: int, + starts: list[int], + anchors: list[int], + gt: np.ndarray, + predictions: dict[str, np.ndarray], +) -> dict[str, object]: + baseline = predictions["baseline"] + detail_path = output_dir / ( + f"episode_{target_episode:06d}_donor_{donor_episode:06d}_action_detail.csv" + ) + fields = [ + "condition", + "block_index", + "packet_start", + "anchor_frame", + "horizon_step", + "action_dim", + "action_name", + "gt", + "pred", + "abs_error_vs_gt", + "pred_delta_vs_baseline", + "abs_pred_delta_vs_baseline", + ] + with detail_path.open("w", newline="", encoding="utf-8") as stream: + writer = csv.DictWriter(stream, fieldnames=fields) + writer.writeheader() + for condition, predicted in predictions.items(): + for block_index, (start, anchor) in enumerate(zip(starts, anchors)): + for horizon_step in range(ACTION_HORIZON): + for dim, name in enumerate(ACTION_NAMES): + pred_value = float(predicted[block_index, horizon_step, dim]) + base_value = float(baseline[block_index, horizon_step, dim]) + gt_value = float(gt[block_index, horizon_step, dim]) + signed_delta = pred_value - base_value + writer.writerow( + { + "condition": condition, + "block_index": block_index, + "packet_start": start, + "anchor_frame": anchor, + "horizon_step": horizon_step + 1, + "action_dim": dim, + "action_name": name, + "gt": gt_value, + "pred": pred_value, + "abs_error_vs_gt": abs(pred_value - gt_value), + "pred_delta_vs_baseline": signed_delta, + "abs_pred_delta_vs_baseline": abs(signed_delta), + } + ) + + metrics_path = output_dir / ( + f"episode_{target_episode:06d}_donor_{donor_episode:06d}_metrics.csv" + ) + metric_fields = [ + "condition", + "scope", + "block_index", + "anchor_frame", + "horizon_step", + "overall_mae", + "arm_mae", + "gripper_mae", + "left_arm_mae", + "right_arm_mae", + "pred_delta_vs_baseline_mae", + "state_input_delta_mae", + "num_action_rows", + ] + metric_rows: list[dict[str, object]] = [] + for condition, predicted in predictions.items(): + error = np.abs(predicted - gt) + delta = np.abs(predicted - baseline) + overall = _metric(error) + metric_rows.append( + { + "condition": condition, + "scope": "overall", + "block_index": "", + "anchor_frame": "", + "horizon_step": "", + **overall, + "pred_delta_vs_baseline_mae": float(delta.mean()), + "state_input_delta_mae": "", + } + ) + for block_index, anchor in enumerate(anchors): + block_metric = _metric(error[block_index : block_index + 1]) + metric_rows.append( + { + "condition": condition, + "scope": "block", + "block_index": block_index, + "anchor_frame": anchor, + "horizon_step": "", + **block_metric, + "pred_delta_vs_baseline_mae": float(delta[block_index].mean()), + "state_input_delta_mae": "", + } + ) + horizon_mae = error.mean(axis=(0, 2)) + horizon_delta = delta.mean(axis=(0, 2)) + for horizon_step, (mae, delta_mae) in enumerate(zip(horizon_mae, horizon_delta), start=1): + metric_rows.append( + { + "condition": condition, + "scope": "horizon", + "block_index": "", + "anchor_frame": "", + "horizon_step": horizon_step, + "overall_mae": float(mae), + "arm_mae": float(error[:, horizon_step - 1, ARM_DIMS].mean()), + "gripper_mae": float(error[:, horizon_step - 1, GRIPPER_DIMS].mean()), + "left_arm_mae": float(error[:, horizon_step - 1, :7].mean()), + "right_arm_mae": float(error[:, horizon_step - 1, 8:15].mean()), + "pred_delta_vs_baseline_mae": float(delta_mae), + "state_input_delta_mae": "", + "num_action_rows": int(error.shape[0]), + } + ) + with metrics_path.open("w", newline="", encoding="utf-8") as stream: + writer = csv.DictWriter(stream, fieldnames=metric_fields) + writer.writeheader() + writer.writerows(metric_rows) + + arrays_path = output_dir / ( + f"episode_{target_episode:06d}_donor_{donor_episode:06d}_action_arrays.npz" + ) + np.savez_compressed( + arrays_path, + ground_truth=gt, + starts=np.asarray(starts, dtype=np.int32), + anchors=np.asarray(anchors, dtype=np.int32), + **{f"pred_{condition}": predicted for condition, predicted in predictions.items()}, + ) + return { + "action_detail_csv": str(detail_path), + "metrics_csv": str(metrics_path), + "action_arrays_npz": str(arrays_path), + "overall_metrics": { + condition: { + **_metric(np.abs(predicted - gt)), + "pred_delta_vs_baseline_mae": float(np.abs(predicted - baseline).mean()), + } + for condition, predicted in predictions.items() + }, + } + + +def _write_action_plots( + output_dir: Path, + target_episode: int, + donor_episode: int, + anchors: list[int], + gt: np.ndarray, + predictions: dict[str, np.ndarray], +) -> dict[str, object]: + try: + import matplotlib + + matplotlib.use("Agg") + import matplotlib.pyplot as plt + except Exception as exc: # pragma: no cover - report still has CSV/NPZ + logging.warning("Could not import matplotlib: %s", exc) + return { + "first_action_plot": None, + "all16_first_action_plot": None, + "horizon_mae_plot": None, + "all16_horizon_plots": [], + } + + x = np.arange(1, len(anchors) + 1, dtype=np.int32) + representative_dims = [0, 7, 8, 15] + first_plot = output_dir / ( + f"episode_{target_episode:06d}_donor_{donor_episode:06d}_first_action_curves.png" + ) + fig, axes = plt.subplots(2, 3, figsize=(18, 9), sharex=True) + for axis, dim in zip(axes.flat[:4], representative_dims): + axis.plot(x, gt[:, 0, dim], color="black", linewidth=2.0, label="GT") + for condition, predicted in predictions.items(): + axis.plot( + x, + predicted[:, 0, dim], + color=CONDITION_COLORS[condition], + linewidth=1.5, + marker="o", + markersize=3.5, + label=condition, + ) + axis.set_title(f"{dim}: {ACTION_NAMES[dim]}") + axis.grid(alpha=0.22) + axis.set_ylabel("joint position") + axes[0, 0].legend(fontsize=8, loc="best") + round_error = { + condition: np.abs(predicted - gt).mean(axis=(1, 2)) + for condition, predicted in predictions.items() + } + for condition, values in round_error.items(): + axes[1, 1].plot( + x, + values, + color=CONDITION_COLORS[condition], + linewidth=1.8, + marker="o", + label=condition, + ) + axes[1, 1].set_title("action-chunk MAE vs GT") + axes[1, 1].set_ylabel("MAE") + axes[1, 1].grid(alpha=0.22) + baseline = predictions["baseline"] + for condition, predicted in predictions.items(): + axes[1, 2].plot( + x, + np.abs(predicted - baseline).mean(axis=(1, 2)), + color=CONDITION_COLORS[condition], + linewidth=1.8, + marker="o", + label=condition, + ) + axes[1, 2].set_title("change from baseline") + axes[1, 2].set_ylabel("MAE") + axes[1, 2].grid(alpha=0.22) + for axis in axes[-1, :]: + axis.set_xlabel("continuous inference round") + fig.suptitle( + f"G2 action curves: target episode {target_episode}, state donor {donor_episode}" + ) + fig.tight_layout() + fig.savefig(first_plot, dpi=140) + plt.close(fig) + + all16_first_plot = output_dir / ( + f"episode_{target_episode:06d}_donor_{donor_episode:06d}_all16_first_action_curves.png" + ) + fig, axes = plt.subplots(4, 4, figsize=(22, 16), sharex=True) + for axis, dim in zip(axes.flat, range(16)): + axis.plot(x, gt[:, 0, dim], color="black", linewidth=2.0, label="GT") + for condition, predicted in predictions.items(): + axis.plot( + x, + predicted[:, 0, dim], + color=CONDITION_COLORS[condition], + linewidth=1.5, + marker="o", + markersize=3.5, + label=condition, + ) + axis.set_title(f"{dim}: {ACTION_NAMES[dim]}") + axis.set_xlabel("continuous inference round") + axis.set_ylabel("position") + axis.grid(alpha=0.22) + axes[0, 0].legend(fontsize=8, loc="best") + fig.suptitle( + f"G2 all 16 action dimensions: target episode {target_episode}, " + f"donor episode {donor_episode}", + fontsize=16, + ) + fig.tight_layout() + fig.savefig(all16_first_plot, dpi=140) + plt.close(fig) + + mae_plot = output_dir / ( + f"episode_{target_episode:06d}_donor_{donor_episode:06d}_horizon_mae.png" + ) + selected_rounds = sorted(set([0, len(anchors) // 2, len(anchors) - 1])) + fig, axes = plt.subplots( + len(selected_rounds), len(representative_dims), + figsize=(18, 4.2 * len(selected_rounds)), + squeeze=False, + sharex=True, + ) + horizon = np.arange(1, ACTION_HORIZON + 1) + for row, round_index in enumerate(selected_rounds): + for col, dim in enumerate(representative_dims): + axis = axes[row, col] + axis.plot(horizon, gt[round_index, :, dim], color="black", linewidth=2.0, label="GT") + for condition, predicted in predictions.items(): + axis.plot( + horizon, + predicted[round_index, :, dim], + color=CONDITION_COLORS[condition], + linewidth=1.3, + label=condition, + ) + axis.set_title( + f"round {round_index + 1}, {ACTION_NAMES[dim]} " + f"(anchor {anchors[round_index]})" + ) + axis.grid(alpha=0.22) + axis.set_xlabel("action chunk horizon step") + axis.set_ylabel("position") + if row == 0 and col == 0: + axis.legend(fontsize=8, loc="best") + fig.tight_layout() + fig.savefig(mae_plot, dpi=140) + plt.close(fig) + + all16_horizon_plots: list[str] = [] + for round_index, anchor in enumerate(anchors): + full_horizon_plot = output_dir / ( + f"episode_{target_episode:06d}_donor_{donor_episode:06d}_" + f"all16_horizon_round_{round_index + 1:02d}.png" + ) + fig, axes = plt.subplots(4, 4, figsize=(22, 16), sharex=True) + for axis, dim in zip(axes.flat, range(16)): + axis.plot( + horizon, + gt[round_index, :, dim], + color="black", + linewidth=2.0, + label="GT", + ) + for condition, predicted in predictions.items(): + axis.plot( + horizon, + predicted[round_index, :, dim], + color=CONDITION_COLORS[condition], + linewidth=1.3, + label=condition, + ) + axis.set_title(f"{dim}: {ACTION_NAMES[dim]}") + axis.set_xlabel("action horizon step") + axis.set_ylabel("position") + axis.grid(alpha=0.22) + axes[0, 0].legend(fontsize=8, loc="best") + fig.suptitle( + f"G2 all 16 action dimensions, round {round_index + 1}, " + f"target anchor {anchor}: target episode {target_episode}, " + f"donor episode {donor_episode}", + fontsize=16, + ) + fig.tight_layout() + fig.savefig(full_horizon_plot, dpi=140) + plt.close(fig) + all16_horizon_plots.append(str(full_horizon_plot)) + return { + "first_action_plot": str(first_plot), + "all16_first_action_plot": str(all16_first_plot), + "horizon_mae_plot": str(mae_plot), + "all16_horizon_plots": all16_horizon_plots, + } + + +def _round_boundaries(total_frames: int, num_rounds: int) -> np.ndarray: + if total_frames < 1 or num_rounds < 1: + return np.zeros(num_rounds + 1, dtype=np.int64) + boundaries = np.rint(np.linspace(0, total_frames, num_rounds + 1)).astype(np.int64) + boundaries[0] = 0 + boundaries[-1] = total_frames + return boundaries + + +def _make_aligned_gt_video_frames( + target: dict[str, object], + starts: list[int], + total_frames: int, +) -> tuple[list[np.ndarray], list[int]]: + """Build a GT sequence with the same number of frames as server output. + + Each server inference round contributes one contiguous video chunk. The + matching GT chunk starts immediately after the four-frame observation + packet (the predicted future starts at anchor+1). We preserve every + decoded server frame and only sample the GT timeline to that same length. + """ + videos = target["videos"] + assert isinstance(videos, dict) + frames: list[np.ndarray] = [] + gt_indices: list[int] = [] + boundaries = _round_boundaries(total_frames, len(starts)) + num_rows = int(target["num_rows"]) + for block_index, start in enumerate(starts): + anchor = start + HISTORY_LEN - 1 + chunk_count = max(1, int(boundaries[block_index + 1] - boundaries[block_index])) + gt_start = min(anchor + 1, num_rows - 1) + gt_end = min(gt_start + chunk_count - 1, num_rows - 1) + indices = np.rint(np.linspace(gt_start, gt_end, chunk_count)).astype(np.int64) + for frame_index in indices: + frame = _grid_rgb( + videos["top_head"][frame_index], + videos["hand_left"][frame_index], + videos["hand_right"][frame_index], + ) + frames.append(_label(frame, f"GROUND TRUTH round={block_index + 1} anchor={anchor}")) + gt_indices.append(int(frame_index)) + return frames, gt_indices + + +def _make_condition_video( + output_dir: Path, + target_episode: int, + donor_episode: int, + condition: str, + raw_video: Path, + target: dict[str, object], + starts: list[int], + video_fps: int, + condition_label: str | None = None, +) -> tuple[Path, list[np.ndarray], list[np.ndarray], list[int], int]: + raw_frames = _read_mp4(raw_video) + gt_frames, gt_indices = _make_aligned_gt_video_frames(target, starts, len(raw_frames)) + side_by_side: list[np.ndarray] = [] + for block_frame, (predicted, gt) in enumerate(zip(raw_frames, gt_frames)): + if predicted.shape != gt.shape: + predicted = cv2.resize(predicted, (gt.shape[1], gt.shape[0]), interpolation=cv2.INTER_AREA) + predicted = _label( + predicted, + f"{condition_label or condition.upper()} server-frame={block_frame + 1}/{len(raw_frames)}", + ) + side_by_side.append(np.concatenate([predicted, gt], axis=1)) + output_path = output_dir / ( + f"episode_{target_episode:06d}_donor_{donor_episode:06d}_{condition}_slow.mp4" + ) + _save_video(output_path, side_by_side, video_fps) + return output_path, raw_frames, gt_frames, gt_indices, len(raw_frames) + + +def _make_montage_video( + output_path: Path, + predictions: dict[str, list[np.ndarray]], + gt_frames: list[np.ndarray], + video_fps: int, + condition_labels: dict[str, str] | None = None, +) -> None: + conditions = list(predictions) + if not conditions: + raise RuntimeError("Cannot make a montage without predictions") + common_count = min([len(gt_frames)] + [len(predictions[name]) for name in conditions]) + if common_count < 1: + raise RuntimeError("Cannot make an empty montage") + normalized_predictions = { + condition: _resample_frames(predictions[condition], common_count) + for condition in conditions + } + normalized_gt = _resample_frames(gt_frames, common_count) + panels: list[np.ndarray] = [] + # The strong comparison uses four counterfactual conditions plus GT. A + # fixed 2x2 layout (used by the old three-condition report) silently + # dropped the fourth condition, so build a compact grid for any number of + # conditions and keep GT as the final panel. + panel_count = len(conditions) + 1 + columns = 3 if panel_count > 4 else 2 + rows = int(np.ceil(panel_count / columns)) + for index, gt in enumerate(normalized_gt): + predicted_panels: list[np.ndarray] = [] + for condition in conditions: + frame = normalized_predictions[condition][index] + if frame.shape != gt.shape: + frame = cv2.resize(frame, (gt.shape[1], gt.shape[0]), interpolation=cv2.INTER_AREA) + predicted_panels.append( + _label( + frame, + (condition_labels or {}).get(condition, condition.upper()), + ) + ) + predicted_panels.append(_label(gt, "GROUND TRUTH")) + while len(predicted_panels) < rows * columns: + predicted_panels.append(np.zeros_like(gt)) + row_images = [ + np.concatenate(predicted_panels[row * columns : (row + 1) * columns], axis=1) + for row in range(rows) + ] + panels.append(np.concatenate(row_images, axis=0)) + _save_video(output_path, panels, video_fps) + + +def _write_round_contact_sheet( + output_path: Path, + target: dict[str, object], + starts: list[int], + predictions: dict[str, list[np.ndarray]], + gt_frames: list[np.ndarray], + condition_labels: dict[str, str] | None = None, +) -> None: + """Save three representative continuous rounds as a simple image sheet.""" + conditions = list(predictions) + common_count = min([len(gt_frames)] + [len(predictions[name]) for name in conditions]) + normalized_predictions = { + condition: _resample_frames(predictions[condition], common_count) + for condition in conditions + } + normalized_gt = _resample_frames(gt_frames, common_count) + selected_rounds = sorted(set([0, len(starts) // 2, len(starts) - 1])) + boundaries = _round_boundaries(common_count, len(starts)) + rows: list[np.ndarray] = [] + for round_index in selected_rounds: + center = int((boundaries[round_index] + boundaries[round_index + 1] - 1) // 2) + gt_panel = _label( + normalized_gt[center], + f"GT round={round_index + 1} anchor={starts[round_index] + HISTORY_LEN - 1}", + ) + row_panels = [gt_panel] + for condition in conditions: + frame = normalized_predictions[condition][center] + if frame.shape != gt_panel.shape: + frame = cv2.resize(frame, (gt_panel.shape[1], gt_panel.shape[0]), interpolation=cv2.INTER_AREA) + row_panels.append( + _label( + frame, + (condition_labels or {}).get(condition, condition.upper()), + ) + ) + rows.append(np.concatenate(row_panels, axis=1)) + sheet = np.concatenate(rows, axis=0) + output_path.parent.mkdir(parents=True, exist_ok=True) + imageio.imwrite(output_path, sheet) + + +def _parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--host", default="127.0.0.1") + parser.add_argument("--port", type=int, default=9443) + parser.add_argument("--test-data-root", type=Path, required=True) + parser.add_argument("--output-dir", type=Path, required=True) + parser.add_argument( + "--server-video-dir", + type=Path, + required=True, + help="The live server's VIDEO_SAVE_MODE=full output directory.", + ) + parser.add_argument( + "--pairs", + default="0:935,769:1129,1044:1221", + help="Comma-separated target_episode:donor_episode pairs.", + ) + parser.add_argument( + "--anchor-pairs", + default=None, + help=( + "Explicit cross-stage anchors, comma-separated as " + "target_episode:donor_episode:target_anchor:donor_anchor. " + "When supplied, these replace --pairs and use --anchor-offsets." + ), + ) + parser.add_argument( + "--anchor-offsets", + default="0,4,8,12", + help="Offsets added to each explicit target/donor anchor (default: 0,4,8,12).", + ) + parser.add_argument( + "--conditions", + default="baseline,state_only,history_only,state_history_swap", + help="Conditions to run; combined condition is optional.", + ) + parser.add_argument("--block-stride", type=int, default=4) + parser.add_argument( + "--max-blocks", + type=int, + default=12, + help="Maximum action anchors per target episode.", + ) + parser.add_argument( + "--start-frame", + type=int, + default=0, + help="First target frame used when selecting the presentation window.", + ) + parser.add_argument( + "--start-frames", + default=None, + help="Optional per-target overrides, e.g. 0:740,30:1332.", + ) + parser.add_argument( + "--temporal-offset", + type=int, + default=4, + help="For temporal_shift, use the target episode packet at start+offset while evaluating GT at start.", + ) + parser.add_argument( + "--contiguous-blocks", + action="store_true", + help="Keep selected inference rounds contiguous for readable videos.", + ) + parser.add_argument("--image-jpeg-quality", type=int, default=80) + parser.add_argument( + "--video-fps", + type=int, + default=5, + help="Display FPS for slow diagnostic videos; source server videos remain 30 FPS.", + ) + parser.add_argument("--video-wait-seconds", type=float, default=120.0) + parser.add_argument("--prompt", default=None) + return parser.parse_args() + + +def main() -> None: + logging.basicConfig(level=logging.INFO, force=True) + args = _parse_args() + if args.block_stride < 1: + raise ValueError("--block-stride must be positive") + if args.max_blocks is not None and args.max_blocks < 1: + raise ValueError("--max-blocks must be positive") + if args.video_fps < 1: + raise ValueError("--video-fps must be positive") + + root = args.test_data_root.resolve() + info = json.loads((root / "meta/info.json").read_text(encoding="utf-8")) + pairs = _parse_pairs(args.pairs) + explicit_anchor_pairs = _parse_anchor_pairs(args.anchor_pairs) if args.anchor_pairs else None + anchor_offsets = _parse_anchor_offsets(args.anchor_offsets) + if explicit_anchor_pairs is not None: + pairs = [(target, donor) for target, donor, _, _ in explicit_anchor_pairs] + conditions = _parse_conditions(args.conditions) + start_frame_overrides = _parse_start_frames(args.start_frames) + args.output_dir.mkdir(parents=True, exist_ok=True) + server_video_dir = args.server_video_dir.resolve() + server_video_dir.mkdir(parents=True, exist_ok=True) + + client = WebsocketClientPolicy(host=args.host, port=args.port) + metadata = client.get_server_metadata() + logging.info("Connected to DreamZero server %s:%s metadata=%s", args.host, args.port, metadata) + all_pair_reports: list[dict[str, object]] = [] + all_summary_rows: list[dict[str, object]] = [] + started = time.time() + + try: + for pair_index, (target_episode, donor_episode) in enumerate(pairs): + target = _load_episode(root, info, target_episode) + donor = _load_episode(root, info, donor_episode) + target_rows = int(target["num_rows"]) + donor_rows = int(donor["num_rows"]) + target_state = np.asarray(target["state"], dtype=np.float32) + target_action = np.asarray(target["action"], dtype=np.float32) + donor_state = np.asarray(donor["state"], dtype=np.float32) + donor_indices = _state_alignment_indices(target_rows, donor_rows) + explicit_anchor = ( + explicit_anchor_pairs[pair_index] + if explicit_anchor_pairs is not None + else None + ) + if explicit_anchor is not None: + _, _, target_anchor_base, donor_anchor_base = explicit_anchor + starts = [ + target_anchor_base - HISTORY_LEN + 1 + offset + for offset in anchor_offsets + ] + donor_starts = [ + donor_anchor_base - HISTORY_LEN + 1 + offset + for offset in anchor_offsets + ] + if any( + start < 0 or start + HISTORY_LEN + ACTION_HORIZON > target_rows + for start in starts + ): + raise RuntimeError( + f"Explicit target anchors {starts} do not leave a complete " + f"{HISTORY_LEN}-frame packet and {ACTION_HORIZON}-step horizon " + f"in episode {target_episode} ({target_rows} rows)" + ) + if any( + start < 0 or start + HISTORY_LEN > donor_rows + for start in donor_starts + ): + raise RuntimeError( + f"Explicit donor anchors {donor_starts} do not leave a complete " + f"{HISTORY_LEN}-frame packet in episode {donor_episode} " + f"({donor_rows} rows)" + ) + pair_start_frame = starts[0] + anchors = [start + HISTORY_LEN - 1 for start in starts] + donor_anchors = [start + HISTORY_LEN - 1 for start in donor_starts] + else: + pair_start_frame = start_frame_overrides.get(target_episode, args.start_frame) + starts = _choose_starts( + target_rows, + args.block_stride, + args.max_blocks, + start_frame=pair_start_frame, + contiguous=args.contiguous_blocks, + ) + if not starts: + raise RuntimeError( + f"Target episode {target_episode} has only {target_rows} rows; " + f"no complete {ACTION_HORIZON}-step horizons" + ) + anchors = [start + HISTORY_LEN - 1 for start in starts] + # Legacy state_swap alignment maps the target timeline onto the + # donor timeline. Explicit strong-comparison runs instead use + # donor anchors supplied by --anchor-pairs above. + donor_starts = [ + int(donor_indices[start]) - HISTORY_LEN + 1 for start in starts + ] + donor_anchors = [int(donor_indices[anchor]) for anchor in anchors] + target_stage = _stage_info(anchors[0], target_rows) + donor_stage = _stage_info(donor_anchors[0], donor_rows) + condition_labels = { + condition: _condition_label(condition, target_stage, donor_stage) + for condition in conditions + } + gt = np.stack([target_action[anchor : anchor + ACTION_HORIZON] for anchor in anchors]) + pair_output = args.output_dir / f"target_{target_episode:06d}_donor_{donor_episode:06d}" + pair_output.mkdir(parents=True, exist_ok=True) + predictions: dict[str, np.ndarray] = {} + condition_video_frames: dict[str, list[np.ndarray]] = {} + condition_gt_video_frames: dict[str, list[np.ndarray]] = {} + condition_video_reports: dict[str, dict[str, object]] = {} + session_ids: dict[str, str] = {} + + logging.info( + "Pair %d/%d target=%d (%d rows, task=%s) donor=%d (%d rows, task=%s) blocks=%d anchors=%s", + pair_index + 1, + len(pairs), + target_episode, + target_rows, + target["prompt"], + donor_episode, + donor_rows, + donor["prompt"], + len(starts), + anchors, + ) + + for condition in conditions: + session_id = ( + f"action-sensitivity-{target_episode:06d}-{donor_episode:06d}-" + f"{condition}-{uuid.uuid4()}" + ) + session_ids[condition] = session_id + # Flush any previous condition and establish an empty server + # video buffer before this condition's first inference. + client.reset({"session_id": session_id}) + before_inferences = _list_video_names(server_video_dir) + condition_predictions: list[np.ndarray] = [] + condition_started = time.time() + for block_index, (start, anchor) in enumerate(zip(starts, anchors)): + reverse_history = condition in { + "history_reverse", + "state_swap_history_reverse", + } + if condition == "temporal_shift": + source_start = start + int(args.temporal_offset) + if source_start < 0 or source_start + HISTORY_LEN > target_rows: + raise RuntimeError( + f"temporal_shift source window [{source_start}, " + f"{source_start + HISTORY_LEN}) is outside target episode {target_episode}" + ) + source_anchor = source_start + HISTORY_LEN - 1 + input_state = target_state[source_anchor] + history_indices = np.arange( + source_start, + source_start + HISTORY_LEN, + dtype=np.int64, + ) + source_videos = target["videos"] + else: + use_donor_state = condition in { + "state_only", + "state_history_swap", + "state_swap", + "state_swap_history_reverse", + } + use_donor_history = condition in { + "history_only", + "state_history_swap", + "state_swap", + "state_swap_history_reverse", + } + if use_donor_state: + state_index = donor_anchors[block_index] + input_state = donor_state[state_index] + else: + input_state = target_state[anchor] + if use_donor_history: + if explicit_anchor is not None: + history_indices = np.arange( + donor_starts[block_index], + donor_starts[block_index] + HISTORY_LEN, + dtype=np.int64, + ) + else: + # Preserve the legacy linearly time-aligned + # behavior when --anchor-pairs is not used. + history_indices = donor_indices[start : start + HISTORY_LEN] + source_videos = donor["videos"] + else: + history_indices = np.arange(start, start + HISTORY_LEN, dtype=np.int64) + source_videos = target["videos"] + camera_frames = { + key: source_videos[key][history_indices] + for key in ("top_head", "hand_left", "hand_right") + } + if reverse_history: + camera_frames = {key: frames[::-1] for key, frames in camera_frames.items()} + prompt = args.prompt if args.prompt is not None else str(target["prompt"]) + observation = { + "observation/top_head": _encode_video_observation(camera_frames["top_head"], args.image_jpeg_quality), + "observation/hand_left": _encode_video_observation(camera_frames["hand_left"], args.image_jpeg_quality), + "observation/hand_right": _encode_video_observation(camera_frames["hand_right"], args.image_jpeg_quality), + "observation/state": input_state, + "prompt": prompt, + # Keep both aliases explicit. The live G2 adapter + # accepts either one, and this makes the language + # payload auditable independently of adapter defaults. + "annotation.language.action_text": prompt, + "session_id": session_id, + } + result = np.asarray(client.infer(observation), dtype=np.float32) + if result.shape != (ACTION_HORIZON, 16): + raise RuntimeError( + f"target={target_episode} condition={condition} block={block_index} " + f"returned {result.shape}, expected {(ACTION_HORIZON, 16)}" + ) + condition_predictions.append(result) + if block_index == 0 or block_index + 1 == len(starts): + logging.info( + "target=%d condition=%s block=%d/%d anchor=%d elapsed=%.1fs", + target_episode, + condition, + block_index + 1, + len(starts), + anchor, + time.time() - condition_started, + ) + + # reset flushes the accumulated generated video for exactly + # this condition while keeping the loaded checkpoint alive. + before_flush = _list_video_names(server_video_dir) + client.reset({"session_id": session_id}) + raw_video = _wait_for_new_video( + server_video_dir, + before_flush, + args.video_wait_seconds, + min_frames=max(1, len(starts) * 5), + ) + # If a stale reset-created file appeared before the first + # inference, it is intentionally not treated as this run's + # video. Keep the variable for a useful audit trail. + stale_reset_files = sorted(_list_video_names(server_video_dir) - before_inferences) + raw_copy = pair_output / f"{condition}_server_raw.mp4" + shutil.copy2(raw_video, raw_copy) + predicted = np.stack(condition_predictions) + predictions[condition] = predicted + condition_video_path, predicted_video_frames, aligned_gt_frames, gt_indices, raw_frame_count = _make_condition_video( + pair_output, + target_episode, + donor_episode, + condition, + raw_video, + target, + starts, + args.video_fps, + condition_labels[condition], + ) + condition_video_frames[condition] = predicted_video_frames + condition_gt_video_frames[condition] = aligned_gt_frames + condition_video_reports[condition] = { + "server_raw_video": str(raw_copy), + "server_raw_frames": raw_frame_count, + "slow_video": str(condition_video_path), + "slow_video_fps": args.video_fps, + "display_frames": raw_frame_count, + "source_server_fps": 30, + "gt_frame_indices": gt_indices, + "video_frames_preserved": True, + "stale_reset_files_seen": stale_reset_files, + } + + table_report = _write_action_tables( + pair_output, + target_episode, + donor_episode, + starts, + anchors, + gt, + predictions, + ) + plot_report = _write_action_plots( + pair_output, + target_episode, + donor_episode, + anchors, + gt, + predictions, + ) + montage_path = pair_output / ( + f"episode_{target_episode:06d}_donor_{donor_episode:06d}_conditions_montage_slow.mp4" + ) + reference_condition = conditions[0] + montage_gt_frames = condition_gt_video_frames[reference_condition] + _make_montage_video( + montage_path, + condition_video_frames, + montage_gt_frames, + args.video_fps, + condition_labels, + ) + contact_sheet_path = pair_output / ( + f"episode_{target_episode:06d}_donor_{donor_episode:06d}_rounds_contact_sheet.png" + ) + _write_round_contact_sheet( + contact_sheet_path, + target, + starts, + condition_video_frames, + montage_gt_frames, + condition_labels, + ) + + state_delta = np.stack( + [ + np.abs(target_state[anchor] - donor_state[donor_anchor]) + for anchor, donor_anchor in zip(anchors, donor_anchors) + ] + ) + pair_report = { + "target_episode": target_episode, + "donor_episode": donor_episode, + "target_rows": target_rows, + "donor_rows": donor_rows, + "target_prompt": target["prompt"], + "donor_prompt": donor["prompt"], + "target_stage": target_stage, + "donor_stage": donor_stage, + "server": {"host": args.host, "port": args.port, "metadata": metadata}, + "protocol": { + "history_length": HISTORY_LEN, + "history_normal": "[t-3,t-2,t-1,t]", + "history_reverse": "[t,t-1,t-2,t-3]", + "state_only": "donor state at donor_anchor + target chronological four-frame history", + "history_only": "target state at target_anchor + donor chronological four-frame history", + "state_history_swap": "donor state and donor chronological four-frame history", + "state_swap_alignment": ( + "explicit target/donor anchors and shared offsets" + if explicit_anchor is not None + else "donor state and donor four-frame history linearly time-aligned to target episode row count" + ), + "history_reorder_control": "state_swap and state_swap_history_reverse share the same donor state and four frames; only frame order differs", + "temporal_shift": "input state and chronological four-frame history are read from target start+temporal_offset; GT/action stays at target start", + "temporal_offset": args.temporal_offset, + "gt_action": "target episode action[anchor:anchor+24]", + "action_shape": [ACTION_HORIZON, 16], + "prompt_sent": target["prompt"], + "prompt_sent_length": len(str(target["prompt"])), + "prompt_input_aliases": ["prompt", "annotation.language.action_text"], + }, + "starts": starts, + "anchors": anchors, + "donor_starts": donor_starts, + "donor_anchors": donor_anchors, + "anchor_offsets": anchor_offsets if explicit_anchor is not None else None, + "start_frame": pair_start_frame, + "state_input_delta": { + "mean_abs": float(state_delta.mean()), + "max_abs": float(state_delta.max()), + "by_dim_mean_abs": state_delta.mean(axis=0).tolist(), + }, + "conditions": conditions, + "session_ids": session_ids, + "video_montage_slow": str(montage_path), + "rounds_contact_sheet": str(contact_sheet_path), + "display_video_fps": args.video_fps, + "actual_target_fps": 30, + "display_slowdown_factor": 30.0 / args.video_fps, + "video_note": "Every decoded server output frame is kept; GT is aligned per continuous inference round.", + **table_report, + **plot_report, + "condition_videos": condition_video_reports, + } + report_path = pair_output / ( + f"episode_{target_episode:06d}_donor_{donor_episode:06d}_report.json" + ) + report_path.write_text(json.dumps(pair_report, ensure_ascii=False, indent=2), encoding="utf-8") + pair_report["report_json"] = str(report_path) + all_pair_reports.append(pair_report) + + for condition, metrics in table_report["overall_metrics"].items(): + all_summary_rows.append( + { + "target_episode": target_episode, + "donor_episode": donor_episode, + "target_stage_index": target_stage["index"], + "target_stage": target_stage["name"], + "donor_stage_index": donor_stage["index"], + "donor_stage": donor_stage["name"], + "condition": condition, + **metrics, + "state_input_delta_mae": float(state_delta.mean()), + } + ) + + finally: + try: + client.reset({"session_id": f"action-sensitivity-final-{uuid.uuid4()}"}) + except Exception: + logging.exception("Final server reset failed") + + summary_csv = args.output_dir / "action_sensitivity_summary.csv" + summary_fields = [ + "target_episode", + "donor_episode", + "target_stage_index", + "target_stage", + "donor_stage_index", + "donor_stage", + "condition", + "overall_mae", + "arm_mae", + "gripper_mae", + "left_arm_mae", + "right_arm_mae", + "num_action_rows", + "pred_delta_vs_baseline_mae", + "state_input_delta_mae", + ] + with summary_csv.open("w", newline="", encoding="utf-8") as stream: + writer = csv.DictWriter(stream, fieldnames=summary_fields) + writer.writeheader() + writer.writerows(all_summary_rows) + + aggregate: dict[str, dict[str, float | int]] = {} + for condition in conditions: + rows = [row for row in all_summary_rows if row["condition"] == condition] + if not rows: + continue + aggregate[condition] = { + key: float(np.mean([float(row[key]) for row in rows])) + for key in ( + "overall_mae", + "arm_mae", + "gripper_mae", + "left_arm_mae", + "right_arm_mae", + "pred_delta_vs_baseline_mae", + "state_input_delta_mae", + ) + } + aggregate[condition]["num_action_rows"] = int( + sum(int(row["num_action_rows"]) for row in rows) + ) + summary_json = args.output_dir / "action_sensitivity_summary.json" + summary_json.write_text( + json.dumps( + { + "server": {"host": args.host, "port": args.port, "metadata": metadata}, + "dataset_root": str(root), + "pairs": pairs, + "anchor_pairs": explicit_anchor_pairs, + "anchor_offsets": anchor_offsets, + "conditions": conditions, + "block_stride": args.block_stride, + "max_blocks": args.max_blocks, + "start_frame": args.start_frame, + "start_frame_overrides": start_frame_overrides, + "contiguous_blocks": args.contiguous_blocks, + "display_video_fps": args.video_fps, + "pair_reports": all_pair_reports, + "aggregate_mean_over_pairs": aggregate, + "elapsed_seconds": time.time() - started, + "summary_csv": str(summary_csv), + }, + ensure_ascii=False, + indent=2, + ), + encoding="utf-8", + ) + logging.info("Saved sensitivity summary CSV: %s", summary_csv) + logging.info("Saved sensitivity summary JSON: %s", summary_json) + + +if __name__ == "__main__": + main() diff --git a/scripts/eval/render_existing_action_sensitivity.py b/scripts/eval/render_existing_action_sensitivity.py new file mode 100644 index 00000000..3dd2b5ad --- /dev/null +++ b/scripts/eval/render_existing_action_sensitivity.py @@ -0,0 +1,109 @@ +#!/usr/bin/env python3 +"""Re-render plots and presentation sheets from saved sensitivity outputs.""" + +from __future__ import annotations + +import argparse +import json +from pathlib import Path + +import numpy as np + +from action_input_sensitivity import ( + _condition_label, + _label, + _make_montage_video, + _read_mp4, + _save_video, + _stage_info, + _write_action_plots, + _write_round_contact_sheet, +) + + +def main() -> None: + parser = argparse.ArgumentParser() + parser.add_argument("--root", type=Path, required=True) + args = parser.parse_args() + + for pair_dir in sorted(args.root.glob("target_*_donor_*")): + if not pair_dir.is_dir(): + continue + report_path = next(pair_dir.glob("*_report.json")) + arrays_path = next(pair_dir.glob("*_action_arrays.npz")) + report = json.loads(report_path.read_text(encoding="utf-8")) + arrays = np.load(arrays_path) + predictions = { + key.removeprefix("pred_"): arrays[key] + for key in arrays.files + if key.startswith("pred_") + } + target_episode = int(report["target_episode"]) + donor_episode = int(report["donor_episode"]) + anchors = [int(value) for value in report["anchors"]] + donor_anchors = [int(value) for value in report["donor_anchors"]] + target_stage = _stage_info(anchors[0], int(report["target_rows"])) + donor_stage = _stage_info(donor_anchors[0], int(report["donor_rows"])) + labels = { + condition: _condition_label(condition, target_stage, donor_stage) + for condition in predictions + } + + plot_report = _write_action_plots( + pair_dir, + target_episode, + donor_episode, + anchors, + arrays["ground_truth"], + predictions, + ) + + condition_frames: dict[str, list[np.ndarray]] = {} + gt_frames: list[np.ndarray] | None = None + for condition in predictions: + slow_path = Path(report["condition_videos"][condition]["slow_video"]) + combined = _read_mp4(slow_path) + half = combined[0].shape[1] // 2 + predicted = [frame[:, :half].copy() for frame in combined] + ground_truth = [frame[:, half:].copy() for frame in combined] + relabeled = [ + np.concatenate([_label(frame, labels[condition]), gt], axis=1) + for frame, gt in zip(predicted, ground_truth) + ] + _save_video(slow_path, relabeled, int(report["display_video_fps"])) + condition_frames[condition] = predicted + if gt_frames is None: + gt_frames = ground_truth + + assert gt_frames is not None + montage_path = Path(report["video_montage_slow"]) + _make_montage_video( + montage_path, + condition_frames, + gt_frames, + int(report["display_video_fps"]), + labels, + ) + contact_sheet_path = Path(report["rounds_contact_sheet"]) + _write_round_contact_sheet( + contact_sheet_path, + {}, + [int(value) for value in report["starts"]], + condition_frames, + gt_frames, + labels, + ) + + report["target_stage"] = target_stage + report["donor_stage"] = donor_stage + report["presentation_labels"] = labels + report.update(plot_report) + report_path.write_text( + json.dumps(report, ensure_ascii=False, indent=2), + encoding="utf-8", + ) + print(pair_dir) + + +if __name__ == "__main__": + main() diff --git a/scripts/eval/render_four_action_chunks.py b/scripts/eval/render_four_action_chunks.py new file mode 100644 index 00000000..39c8b99d --- /dev/null +++ b/scripts/eval/render_four_action_chunks.py @@ -0,0 +1,111 @@ +#!/usr/bin/env python3 +"""Render one compact all-16D plot containing every saved action chunk.""" + +from __future__ import annotations + +import argparse +import json +from pathlib import Path + +import numpy as np + +ACTION_NAMES = ( + *[f"left_joint_{index}" for index in range(7)], + "left_gripper", + *[f"right_joint_{index}" for index in range(7)], + "right_gripper", +) +COLORS = { + "baseline": "tab:blue", + "state_only": "tab:red", + "history_only": "tab:green", + "state_history_swap": "tab:purple", +} +LINESTYLES = { + "baseline": "-", + "state_only": "--", + "history_only": ":", + "state_history_swap": "-.", +} + + +def render(pair_dir: Path) -> Path: + import matplotlib + + matplotlib.use("Agg") + import matplotlib.pyplot as plt + + arrays_path = next(pair_dir.glob("*_action_arrays.npz")) + report_path = next(pair_dir.glob("*_report.json")) + arrays = np.load(arrays_path) + report = json.loads(report_path.read_text(encoding="utf-8")) + ground_truth = arrays["ground_truth"] + conditions = [ + condition + for condition in COLORS + if f"pred_{condition}" in arrays.files + ] + chunks, horizon, dimensions = ground_truth.shape + if dimensions != 16: + raise RuntimeError(f"Expected 16 action dimensions, got {ground_truth.shape}") + total_rows = chunks * horizon + x = np.arange(1, total_rows + 1) + gt_unrolled = ground_truth.reshape(total_rows, dimensions) + predictions = { + condition: arrays[f"pred_{condition}"].reshape(total_rows, dimensions) + for condition in conditions + } + + target_episode = int(report["target_episode"]) + source_episode = int(report["donor_episode"]) + output_path = pair_dir / ( + f"episode_{target_episode:06d}_source_{source_episode:06d}_" + f"all16_{chunks}chunks_x_{horizon}steps.png" + ) + fig, axes = plt.subplots(4, 4, figsize=(24, 16), sharex=True) + for axis, dim in zip(axes.flat, range(dimensions)): + axis.plot(x, gt_unrolled[:, dim], color="black", linewidth=1.8, label="GT") + for condition in conditions: + axis.plot( + x, + predictions[condition][:, dim], + color=COLORS[condition], + linestyle=LINESTYLES[condition], + linewidth=1.65, + alpha=0.82, + label=condition, + ) + for chunk_index in range(1, chunks): + axis.axvline( + chunk_index * horizon + 0.5, + color="0.35", + linestyle="--", + linewidth=0.9, + ) + axis.set_title(f"{dim}: {ACTION_NAMES[dim]}") + axis.set_xlabel(f"all action rows ({chunks} chunks x {horizon} steps)") + axis.set_ylabel("joint position") + axis.grid(alpha=0.20) + axes[0, 0].legend(fontsize=8, loc="best") + fig.suptitle( + f"G2 all 16 action dimensions | target episode {target_episode} | " + f"swap-source episode {source_episode} | {chunks} chunks x {horizon} steps", + fontsize=17, + ) + fig.tight_layout() + fig.savefig(output_path, dpi=160) + plt.close(fig) + return output_path + + +def main() -> None: + parser = argparse.ArgumentParser() + parser.add_argument("--root", type=Path, required=True) + args = parser.parse_args() + for pair_dir in sorted(args.root.glob("target_*_donor_*")): + if pair_dir.is_dir(): + print(render(pair_dir)) + + +if __name__ == "__main__": + main() diff --git a/scripts/eval/rollout_g2_action_full.py b/scripts/eval/rollout_g2_action_full.py new file mode 100644 index 00000000..2cdcc47b --- /dev/null +++ b/scripts/eval/rollout_g2_action_full.py @@ -0,0 +1,165 @@ +#!/usr/bin/env python3 +"""Full-episode G2 action-only evaluation against every 4-frame anchor. + +This is deliberately separate from robot deployment. It sends the same +4-frame teacher-forced observation packets used by the live client, keeps four +causal packets per server cache window, and records every returned 24x16 +action chunk against the training episode GT. Video files are ignored; the +result is the complete 16-D action table/curve requested for the task. +""" + +from __future__ import annotations + +import argparse +import csv +import json +import logging +import sys +import time +import uuid +from pathlib import Path + +import cv2 +import numpy as np +import pyarrow.parquet as pq + +ROOT = Path(__file__).resolve().parents[2] +sys.path.insert(0, str(ROOT)) + +from eval_utils.policy_client import WebsocketClientPolicy # noqa: E402 +from scripts.eval.rollout_g2_server_windowed import ( # noqa: E402 + ACTION_NAMES, + ARM_DIMS, + GRIPPER_DIMS, + _episode_path, + _encode_video_observation, + _read_video, +) + + +def _parse_args() -> argparse.Namespace: + p = argparse.ArgumentParser() + p.add_argument("--host", default="127.0.0.1") + p.add_argument("--port", type=int, default=30002) + p.add_argument("--test-data-root", type=Path, required=True) + p.add_argument("--episode-index", type=int, default=0) + p.add_argument("--output-dir", type=Path, required=True) + p.add_argument( + "--server-video-dir", + type=Path, + default=None, + help="Temporary server video directory; remove files flushed by this action-only run.", + ) + p.add_argument("--window-blocks", type=int, default=4) + p.add_argument("--max-blocks", type=int, default=None) + p.add_argument("--prompt", default=None) + p.add_argument("--image-jpeg-quality", type=int, default=80) + return p.parse_args() + + +def _remove_new_server_videos(server_dir: Path | None, before: set[str]) -> None: + """Drop per-reset video files; this pass evaluates actions only.""" + if server_dir is None: + return + try: + for path in server_dir.glob("*.mp4"): + if path.name not in before: + try: + path.unlink() + except OSError as exc: + logging.warning("Could not remove temporary server video %s: %s", path, exc) + except OSError as exc: + logging.warning("Could not scan temporary server video directory %s: %s", server_dir, exc) + + +def _write_reports(out: Path, episode: int, starts: list[int], anchors: list[int], pred: np.ndarray, gt: np.ndarray) -> None: + err = np.abs(pred - gt) + np.savez_compressed(out / f"episode_{episode:06d}_full_action_arrays.npz", predicted=pred, ground_truth=gt, packet_starts=np.asarray(starts), anchors=np.asarray(anchors)) + detail = out / f"episode_{episode:06d}_full_action_pred_vs_gt.csv" + fields = ["block_index", "packet_start", "anchor_frame", "horizon_step"] + [f"gt_{x}" for x in ACTION_NAMES] + [f"pred_{x}" for x in ACTION_NAMES] + with detail.open("w", newline="", encoding="utf-8") as f: + w = csv.DictWriter(f, fieldnames=fields); w.writeheader() + for bi, (start, anchor) in enumerate(zip(starts, anchors)): + for hi in range(24): + row = {"block_index": bi, "packet_start": start, "anchor_frame": anchor, "horizon_step": hi + 1} + row.update({f"gt_{n}": float(v) for n, v in zip(ACTION_NAMES, gt[bi, hi])}) + row.update({f"pred_{n}": float(v) for n, v in zip(ACTION_NAMES, pred[bi, hi])}) + w.writerow(row) + + def metric(scope: str, e: np.ndarray) -> dict[str, object]: + return {"scope": scope, "overall_mae": float(e.mean()), "arm_mae": float(e[..., ARM_DIMS].mean()), "gripper_mae": float(e[..., GRIPPER_DIMS].mean()), "left_arm_mae": float(e[..., :7].mean()), "right_arm_mae": float(e[..., 8:15].mean()), "num_action_rows": int(np.prod(e.shape[:-1]))} + + metrics = out / f"episode_{episode:06d}_full_action_metrics.csv" + rows = [metric(f"block_{i:04d}", err[i : i + 1]) | {"block_index": i, "packet_start": starts[i], "anchor_frame": anchors[i]} for i in range(len(pred))] + overall = metric("overall", err) | {"block_index": "", "packet_start": "", "anchor_frame": ""} + rows.append(overall) + mf = ["scope", "block_index", "packet_start", "anchor_frame", "overall_mae", "arm_mae", "gripper_mae", "left_arm_mae", "right_arm_mae", "num_action_rows"] + with metrics.open("w", newline="", encoding="utf-8") as f: + w = csv.DictWriter(f, fieldnames=mf); w.writeheader(); w.writerows(rows) + + plot = out / f"episode_{episode:06d}_full_action_pred_vs_gt.png" + mae_plot = out / f"episode_{episode:06d}_full_action_mae_by_horizon.png" + try: + import matplotlib; matplotlib.use("Agg") + import matplotlib.pyplot as plt + x = np.arange(len(pred) * 24); pf = pred.reshape(-1, 16); gf = gt.reshape(-1, 16) + fig, axes = plt.subplots(4, 4, figsize=(22, 13), sharex=True) + for d, ax in enumerate(axes.flat): + ax.plot(x, gf[:, d], color="tab:blue", lw=.8, label="GT") + ax.plot(x, pf[:, d], color="tab:orange", lw=.8, label="PRED") + ax.set_title(f"{d}: {ACTION_NAMES[d]}"); ax.grid(alpha=.2) + axes[0, 0].legend(); fig.suptitle(f"G2 full 4-frame-anchor action vs GT, episode {episode}"); fig.tight_layout(); fig.savefig(plot, dpi=140); plt.close(fig) + h = np.arange(1, 25); fig, ax = plt.subplots(figsize=(12, 5)); ax.plot(h, err.mean((0, 2)), label="all 16"); ax.plot(h, err[..., ARM_DIMS].mean((0, 2)), label="arm 14"); ax.plot(h, err[..., GRIPPER_DIMS].mean((0, 2)), label="gripper 2"); ax.set_xlabel("horizon step"); ax.set_ylabel("MAE"); ax.grid(alpha=.25); ax.legend(); fig.tight_layout(); fig.savefig(mae_plot, dpi=140); plt.close(fig) + except Exception as exc: + logging.warning("plotting failed: %s", exc); plot = None; mae_plot = None + report = {"episode_index": episode, "num_action_blocks": len(pred), "packet_stride": 4, "window_blocks": 4, "video_not_evaluated": True, "action_detail_csv": str(detail), "action_metrics_csv": str(metrics), "action_plot": str(plot) if plot else None, "action_mae_by_horizon": str(mae_plot) if mae_plot else None, "overall_action_metrics": overall} + (out / f"episode_{episode:06d}_full_action_report.json").write_text(json.dumps(report, ensure_ascii=False, indent=2), encoding="utf-8") + logging.info("Saved full action report: %s", report) + + +def main() -> None: + logging.basicConfig(level=logging.INFO, force=True) + a = _parse_args(); a.output_dir.mkdir(parents=True, exist_ok=True) + if a.window_blocks < 1: raise ValueError("--window-blocks must be positive") + root = a.test_data_root.resolve(); info = json.loads((root / "meta/info.json").read_text(encoding="utf-8")); episode = int(a.episode_index); chunks = int(info.get("chunks_size", 1000)) + parquet = _episode_path(root, info.get("data_path", "data/chunk-{episode_chunk:03d}/episode_{episode_index:06d}.parquet"), episode, chunks) + table = pq.read_table(parquet); n = table.num_rows + state = np.asarray(table["observation.state"].to_pylist(), dtype=np.float32).reshape(n, 16); action = np.asarray(table["action"].to_pylist(), dtype=np.float32).reshape(n, 16) + tmpl = info.get("video_path", "videos/chunk-{episode_chunk:03d}/{video_key}/episode_{episode_index:06d}.mp4") + videos = {name: _read_video(_episode_path(root, tmpl, episode, chunks, video_key=key), n) for name, key in (("top_head", "observation.images.top_head"), ("hand_left", "observation.images.hand_left"), ("hand_right", "observation.images.hand_right"))} + max_start = n - 27; starts = list(range(0, max_start + 1, 4)); + if a.max_blocks is not None: starts = starts[: max(0, int(a.max_blocks))] + if not starts: raise RuntimeError("No valid 4-frame action anchors") + prompt_key = "annotation.language.action_text"; prompt = a.prompt if a.prompt is not None else (str(table[prompt_key][0].as_py()) if prompt_key in table.column_names else "") + sid = f"g2-full-action-episode-{episode:06d}-{uuid.uuid4()}"; client = WebsocketClientPolicy(host=a.host, port=a.port); logging.info("Full action pass episode=%d rows=%d blocks=%d server=%s:%d", episode, n, len(starts), a.host, a.port) + pred: list[np.ndarray] = []; gt: list[np.ndarray] = []; anchors: list[int] = []; started = time.time() + try: + before_videos = {p.name for p in a.server_video_dir.glob("*.mp4")} if a.server_video_dir else set() + client.reset({"session_id": sid}) + _remove_new_server_videos(a.server_video_dir, before_videos) + for i, start in enumerate(starts): + idx = list(range(start, start + 4)); anchor = start + 3 + result = np.asarray(client.infer({ + "observation/top_head": _encode_video_observation(videos["top_head"][idx], a.image_jpeg_quality), + "observation/hand_left": _encode_video_observation(videos["hand_left"][idx], a.image_jpeg_quality), + "observation/hand_right": _encode_video_observation(videos["hand_right"][idx], a.image_jpeg_quality), + "observation/state": state[anchor], "prompt": prompt, "session_id": sid, + }), dtype=np.float32) + if result.shape != (24, 16): raise RuntimeError(f"block {i} returned {result.shape}") + pred.append(result); gt.append(action[anchor : anchor + 24]); anchors.append(anchor) + if (i + 1) % a.window_blocks == 0 or i + 1 == len(starts): + before_videos = {p.name for p in a.server_video_dir.glob("*.mp4")} if a.server_video_dir else set() + client.reset({"session_id": sid}) + _remove_new_server_videos(a.server_video_dir, before_videos) + if i == 0 or (i + 1) % 25 == 0 or i + 1 == len(starts): logging.info("action block %d/%d anchor=%d elapsed=%.1fs", i + 1, len(starts), anchor, time.time() - started) + finally: + try: + before_videos = {p.name for p in a.server_video_dir.glob("*.mp4")} if a.server_video_dir else set() + client.reset({"session_id": sid}) + _remove_new_server_videos(a.server_video_dir, before_videos) + except Exception: logging.exception("final reset failed") + _write_reports(a.output_dir, episode, starts, anchors, np.stack(pred), np.stack(gt)) + + +if __name__ == "__main__": + main() diff --git a/scripts/eval/rollout_g2_g1g7_action_dit_windowed.py b/scripts/eval/rollout_g2_g1g7_action_dit_windowed.py new file mode 100644 index 00000000..4d9caa07 --- /dev/null +++ b/scripts/eval/rollout_g2_g1g7_action_dit_windowed.py @@ -0,0 +1,446 @@ +#!/usr/bin/env python3 +"""Isolated G1-G7 Action-DiT LoRA G2 video/action window evaluation. + +The persistent G1-G7 server must be started with +``--no-reset-cache-each-request``. +Each evaluation window sends four consecutive 4-frame observations through the +same causal cache. DreamZero then decodes one 33-frame video window. The +matching GT segment is taken from the same episode timeline and is stitched +side-by-side at 30 Hz; no global frame-rate resampling is performed. +""" + +from __future__ import annotations + +import argparse +import csv +import json +import logging +import sys +import time +import uuid +from pathlib import Path + +# Keep evaluation code isolated under scripts/eval while importing the shared +# websocket client from the repository root. This file never touches the +# robot deployment client. +sys.path.insert(0, str(Path(__file__).resolve().parents[2])) + +import cv2 +import imageio.v2 as imageio +import numpy as np +import pyarrow.parquet as pq + +from eval_utils.policy_client import WebsocketClientPolicy + + +ACTION_NAMES = ( + [f"left_joint_{i}" for i in range(7)] + + ["left_gripper"] + + [f"right_joint_{i}" for i in range(7)] + + ["right_gripper"] +) +ARM_DIMS = [*range(7), *range(8, 15)] +GRIPPER_DIMS = [7, 15] + + +def _episode_path(root: Path, template: str, episode: int, chunks_size: int, **kwargs: object) -> Path: + return root / template.format( + episode_chunk=episode // chunks_size, + episode_index=episode, + **kwargs, + ) + + +def _read_video(path: Path, expected_rows: int) -> np.ndarray: + capture = cv2.VideoCapture(str(path)) + if not capture.isOpened(): + raise RuntimeError(f"Failed to open G2 video: {path}") + frames: list[np.ndarray] = [] + try: + while len(frames) < expected_rows: + ok, frame = capture.read() + if not ok: + break + frames.append(np.ascontiguousarray(frame)) + finally: + capture.release() + if len(frames) < expected_rows: + raise RuntimeError( + f"{path} contains {len(frames)} frames, expected {expected_rows}" + ) + return np.stack(frames, axis=0) + + +def _grid_rgb(top_bgr: np.ndarray, left_bgr: np.ndarray, right_bgr: np.ndarray) -> np.ndarray: + top = cv2.cvtColor(top_bgr, cv2.COLOR_BGR2RGB) + left = cv2.cvtColor(left_bgr, cv2.COLOR_BGR2RGB) + right = cv2.cvtColor(right_bgr, cv2.COLOR_BGR2RGB) + black = np.zeros_like(top) + return np.concatenate( + [np.concatenate([top, left], axis=1), np.concatenate([right, black], axis=1)], + axis=0, + ) + + +def _label(frame_rgb: np.ndarray, text: str) -> np.ndarray: + frame_bgr = cv2.cvtColor(frame_rgb, cv2.COLOR_RGB2BGR) + cv2.rectangle(frame_bgr, (0, 0), (370, 34), (0, 0, 0), -1) + cv2.putText( + frame_bgr, + text, + (8, 24), + cv2.FONT_HERSHEY_SIMPLEX, + 0.65, + (255, 255, 255), + 2, + cv2.LINE_AA, + ) + return cv2.cvtColor(frame_bgr, cv2.COLOR_BGR2RGB) + + +def _save_video(path: Path, frames: list[np.ndarray], fps: int = 30) -> None: + if not frames: + raise RuntimeError(f"No frames to save: {path}") + path.parent.mkdir(parents=True, exist_ok=True) + imageio.mimsave(path, frames, fps=fps, codec="libx264", macro_block_size=None) + + +def _find_new_server_video(server_dir: Path, before: set[str], timeout: float = 60.0) -> Path: + deadline = time.time() + timeout + while time.time() < deadline: + candidates = [ + p for p in server_dir.glob("*.mp4") + if p.name not in before and p.stat().st_size > 0 + ] + if candidates: + return max(candidates, key=lambda p: p.stat().st_mtime_ns) + time.sleep(0.25) + raise RuntimeError( + f"Server reset did not create a new video in {server_dir}; " + f"existing={sorted(before)[-3:]}" + ) + + +def _read_server_video(path: Path, expected_frames: int) -> list[np.ndarray]: + reader = imageio.get_reader(path) + try: + frames = [np.asarray(frame) for frame in reader] + finally: + reader.close() + if len(frames) != expected_frames: + raise RuntimeError( + f"Expected exactly {expected_frames} decoded prediction frames from " + f"the causal window, got {len(frames)} ({path.name}). " + "Do not resample this diagnostic: check server cache mode and VAE decode." + ) + return [np.ascontiguousarray(frame) for frame in frames] + + +def _write_action_reports( + output_dir: Path, + episode: int, + window_starts: list[int], + anchors: list[int], + predicted: np.ndarray, + ground_truth: np.ndarray, +) -> dict[str, object]: + # Shapes: [window, causal_block, horizon, action_dim]. + if predicted.shape != ground_truth.shape or predicted.ndim != 4: + raise ValueError(f"Action arrays must match 4-D shape, got {predicted.shape} and {ground_truth.shape}") + error = np.abs(predicted - ground_truth) + detail_path = output_dir / f"episode_{episode:06d}_windowed_action_pred_vs_gt.csv" + fields = ["window_index", "window_start", "block_index", "anchor_frame", "horizon_step"] + fields += [f"gt_{name}" for name in ACTION_NAMES] + fields += [f"pred_{name}" for name in ACTION_NAMES] + with detail_path.open("w", newline="", encoding="utf-8") as stream: + writer = csv.DictWriter(stream, fieldnames=fields) + writer.writeheader() + for wi, start in enumerate(window_starts): + for bi in range(predicted.shape[1]): + anchor = anchors[wi * predicted.shape[1] + bi] + for hi in range(predicted.shape[2]): + row: dict[str, object] = { + "window_index": wi, + "window_start": start, + "block_index": bi, + "anchor_frame": anchor, + "horizon_step": hi + 1, + } + row.update({f"gt_{name}": float(v) for name, v in zip(ACTION_NAMES, ground_truth[wi, bi, hi])}) + row.update({f"pred_{name}": float(v) for name, v in zip(ACTION_NAMES, predicted[wi, bi, hi])}) + writer.writerow(row) + + metrics_path = output_dir / f"episode_{episode:06d}_windowed_action_metrics.csv" + metric_fields = ["scope", "window_index", "window_start", "block_index", "anchor_frame", "overall_mae", "arm_mae", "gripper_mae", "left_arm_mae", "right_arm_mae", "num_action_rows"] + + def metric_row(scope: str, e: np.ndarray, **extra: object) -> dict[str, object]: + row: dict[str, object] = { + "scope": scope, + "overall_mae": float(e.mean()), + "arm_mae": float(e[..., ARM_DIMS].mean()), + "gripper_mae": float(e[..., GRIPPER_DIMS].mean()), + "left_arm_mae": float(e[..., :7].mean()), + "right_arm_mae": float(e[..., 8:15].mean()), + "num_action_rows": int(np.prod(e.shape[:-1])), + } + row.update(extra) + return row + + metric_rows: list[dict[str, object]] = [] + num_blocks = predicted.shape[1] + for wi, start in enumerate(window_starts): + for bi in range(num_blocks): + metric_rows.append( + metric_row( + f"window_{wi:04d}_block_{bi:02d}", + error[wi, bi], + window_index=wi, + window_start=start, + block_index=bi, + anchor_frame=anchors[wi * num_blocks + bi], + ) + ) + overall = metric_row("overall", error, window_index="", window_start="", block_index="", anchor_frame="") + metric_rows.append(overall) + with metrics_path.open("w", newline="", encoding="utf-8") as stream: + writer = csv.DictWriter(stream, fieldnames=metric_fields) + writer.writeheader() + writer.writerows(metric_rows) + + # One 16-dimension plot over the complete sampled task. Each action row + # is a predicted 24-step horizon at one causal observation anchor. + pred_flat = predicted.reshape(-1, predicted.shape[-1]) + gt_flat = ground_truth.reshape(-1, ground_truth.shape[-1]) + plot_path = output_dir / f"episode_{episode:06d}_windowed_action_pred_vs_gt.png" + try: + import matplotlib + + matplotlib.use("Agg") + import matplotlib.pyplot as plt + + fig, axes = plt.subplots(4, 4, figsize=(22, 13), sharex=True) + x = np.arange(len(pred_flat)) + for dim, axis in enumerate(axes.flat): + axis.plot(x, gt_flat[:, dim], color="tab:blue", linewidth=0.8, label="GT") + axis.plot(x, pred_flat[:, dim], color="tab:orange", linewidth=0.8, alpha=0.9, label="PRED") + axis.set_title(f"{dim}: {ACTION_NAMES[dim]}") + axis.grid(alpha=0.2) + axes[0, 0].legend(loc="upper right") + axes[-1, 0].set_xlabel("window/block/horizon action row") + axes[-1, 1].set_xlabel("GT and predicted action") + fig.suptitle(f"G2 checkpoint-3000: full windowed action vs GT, episode {episode}") + fig.tight_layout() + fig.savefig(plot_path, dpi=140) + plt.close(fig) + + mae_path = output_dir / f"episode_{episode:06d}_windowed_action_mae_by_horizon.png" + horizon_mae = error.mean(axis=(0, 1, 3)) + arm_mae = error[..., ARM_DIMS].mean(axis=(0, 1, 3)) + grip_mae = error[..., GRIPPER_DIMS].mean(axis=(0, 1, 3)) + fig, axis = plt.subplots(figsize=(12, 5)) + h = np.arange(1, error.shape[2] + 1) + axis.plot(h, horizon_mae, label="all 16 dims", linewidth=2) + axis.plot(h, arm_mae, label="arm 14 dims", linewidth=2) + axis.plot(h, grip_mae, label="gripper 2 dims", linewidth=2) + axis.set_xlabel("action horizon step") + axis.set_ylabel("mean absolute error") + axis.set_title("Predicted action vs training GT: error by horizon step") + axis.grid(alpha=0.25) + axis.legend() + fig.tight_layout() + fig.savefig(mae_path, dpi=140) + plt.close(fig) + except Exception as exc: # diagnostic output should not hide CSV results + logging.warning("Could not create action plots: %s", exc) + plot_path = None + mae_path = None + + return { + "action_detail_csv": str(detail_path), + "action_metrics_csv": str(metrics_path), + "action_plot": str(plot_path) if plot_path else None, + "action_horizon_mae_plot": str(mae_path) if mae_path else None, + "overall_action_metrics": overall, + } + + +def _parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser() + parser.add_argument("--host", default="127.0.0.1") + parser.add_argument("--port", type=int, default=30002) + parser.add_argument("--test-data-root", type=Path, required=True) + parser.add_argument("--episode-index", type=int, default=0) + parser.add_argument("--output-dir", type=Path, required=True) + parser.add_argument("--server-video-dir", type=Path, required=True) + parser.add_argument("--window-future-frames", type=int, default=33) + parser.add_argument("--window-history", type=int, default=4) + parser.add_argument("--causal-blocks", type=int, default=4) + parser.add_argument("--window-stride", type=int, default=33) + parser.add_argument("--max-windows", type=int, default=None) + parser.add_argument("--prompt", default=None) + return parser.parse_args() + + +def main() -> None: + logging.basicConfig(level=logging.INFO, force=True) + args = _parse_args() + if args.window_history != 4 or args.causal_blocks != 4: + raise ValueError("This protocol requires four consecutive 4-frame causal packets") + if args.window_future_frames != 33: + raise ValueError("This protocol requires exactly a 33-frame decoded video window") + if args.window_stride < 1: + raise ValueError("--window-stride must be positive") + + args.output_dir.mkdir(parents=True, exist_ok=True) + root = args.test_data_root.resolve() + info = json.loads((root / "meta/info.json").read_text(encoding="utf-8")) + chunks_size = int(info.get("chunks_size", 1000)) + episode = int(args.episode_index) + parquet_path = _episode_path( + root, + info.get("data_path", "data/chunk-{episode_chunk:03d}/episode_{episode_index:06d}.parquet"), + episode, + chunks_size, + ) + table = pq.read_table(parquet_path) + state = np.asarray(table["observation.state"].to_pylist(), dtype=np.float32).reshape(table.num_rows, 16) + action = np.asarray(table["action"].to_pylist(), dtype=np.float32).reshape(table.num_rows, 16) + num_rows = int(table.num_rows) + video_template = info.get("video_path", "videos/chunk-{episode_chunk:03d}/{video_key}/episode_{episode_index:06d}.mp4") + videos: dict[str, np.ndarray] = {} + for name, key in (("top_head", "observation.images.top_head"), ("hand_left", "observation.images.hand_left"), ("hand_right", "observation.images.hand_right")): + videos[name] = _read_video(_episode_path(root, video_template, episode, chunks_size, video_key=key), num_rows) + + # A window starts with the four-frame packet [start:start+4]. Its last + # causal packet is [start+12:start+16], anchored at start+15; the decoded + # GT/prediction segment is anchor:anchor+33, i.e. rows start+3..start+35. + max_start = num_rows - (args.window_history * args.causal_blocks + args.window_future_frames - 1) + starts = list(range(0, max_start + 1, args.window_stride)) + if args.max_windows is not None: + starts = starts[: max(0, int(args.max_windows))] + if not starts: + raise RuntimeError(f"No valid 33-frame windows in episode with {num_rows} rows") + prompt_key = "annotation.language.action_text" + prompt = args.prompt if args.prompt is not None else (str(table[prompt_key][0].as_py()) if prompt_key in table.column_names else "") + session_id = f"g2-windowed-episode-{episode:06d}-{uuid.uuid4()}" + server_dir = args.server_video_dir.resolve() + server_dir.mkdir(parents=True, exist_ok=True) + + logging.info( + "Windowed rollout: episode=%d rows=%d windows=%d stride=%d history=%d causal_blocks=%d future=%d", + episode, num_rows, len(starts), args.window_stride, args.window_history, args.causal_blocks, args.window_future_frames, + ) + client = WebsocketClientPolicy(host=args.host, port=args.port) + logging.info("Server metadata: %s", client.get_server_metadata()) + # Clear any stale server-side cache/video from an earlier client session. + try: + client.reset({"session_id": session_id}) + except Exception: + logging.exception("Initial server reset failed") + + all_pred: list[np.ndarray] = [] + all_gt: list[np.ndarray] = [] + all_anchors: list[int] = [] + all_comparison_frames: list[np.ndarray] = [] + window_reports: list[dict[str, object]] = [] + started = time.time() + for wi, window_start in enumerate(starts): + before = {p.name for p in server_dir.glob("*.mp4")} + pred_chunks: list[np.ndarray] = [] + gt_chunks: list[np.ndarray] = [] + anchors_window: list[int] = [] + for block_index in range(args.causal_blocks): + packet_start = window_start + block_index * args.window_history + indices = list(range(packet_start, packet_start + args.window_history)) + anchor = indices[-1] + result = np.asarray( + client.infer( + { + "observation/top_head": videos["top_head"][indices], + "observation/hand_left": videos["hand_left"][indices], + "observation/hand_right": videos["hand_right"][indices], + "observation/state": state[anchor], + "prompt": prompt, + "session_id": session_id, + } + ), + dtype=np.float32, + ) + if result.shape != (24, 16): + raise RuntimeError(f"window {wi} block {block_index}: server returned {result.shape}, expected (24,16)") + pred_chunks.append(result) + gt_chunks.append(action[anchor : anchor + 24]) + anchors_window.append(anchor) + # Reset flushes exactly this four-forward causal video window; it does + # not unload the server/checkpoint. + client.reset({"session_id": session_id}) + predicted_video_path = _find_new_server_video(server_dir, before) + predicted_frames = _read_server_video(predicted_video_path, args.window_future_frames) + gt_indices = range(window_start + 3, window_start + 3 + args.window_future_frames) + gt_frames = [ + _grid_rgb(videos["top_head"][idx], videos["hand_left"][idx], videos["hand_right"][idx]) + for idx in gt_indices + ] + comparison_frames = [ + np.concatenate([_label(pred, "PREDICTED"), _label(gt, "G2 GROUND TRUTH")], axis=1) + for pred, gt in zip(predicted_frames, gt_frames) + ] + all_comparison_frames.extend(comparison_frames) + pred_array = np.stack(pred_chunks) + gt_array = np.stack(gt_chunks) + all_pred.append(pred_array) + all_gt.append(gt_array) + all_anchors.extend(anchors_window) + window_reports.append( + { + "window_index": wi, + "window_start": window_start, + "packet_starts": [window_start + args.window_history * i for i in range(args.causal_blocks)], + "anchor_indices": anchors_window, + "gt_video_indices": [window_start + 3, window_start + 3 + args.window_future_frames - 1], + "predicted_video": str(predicted_video_path), + "predicted_frames": len(predicted_frames), + "ground_truth_frames": len(gt_frames), + "action_mae": float(np.abs(pred_array - gt_array).mean()), + } + ) + logging.info( + "window %d/%d start=%d video=%d frames action_mae=%.6f elapsed=%.1fs", + wi + 1, len(starts), window_start, len(predicted_frames), window_reports[-1]["action_mae"], time.time() - started, + ) + + predicted = np.stack(all_pred) + ground_truth = np.stack(all_gt) + np.savez_compressed( + args.output_dir / f"episode_{episode:06d}_windowed_action_arrays.npz", + predicted=predicted, + ground_truth=ground_truth, + anchors=np.asarray(all_anchors, dtype=np.int32), + window_starts=np.asarray(starts, dtype=np.int32), + ) + comparison_path = args.output_dir / f"episode_{episode:06d}_windowed_predicted_vs_gt_f{len(all_comparison_frames)}.mp4" + _save_video(comparison_path, all_comparison_frames, fps=30) + reports = _write_action_reports(args.output_dir, episode, starts, all_anchors, predicted, ground_truth) + report = { + "episode_index": episode, + "checkpoint_server_protocol": "four 4-frame history packets -> one causal 33-frame decoded window; reset only between windows", + "server_cache_required": "DREAMZERO_RESET_AR_EACH_REQUEST=false", + "num_rows": num_rows, + "window_starts": starts, + "window_stride": args.window_stride, + "window_history": args.window_history, + "causal_blocks": args.causal_blocks, + "future_frames": args.window_future_frames, + "comparison_fps": 30, + "comparison_video": str(comparison_path), + "window_reports": window_reports, + **reports, + } + report_path = args.output_dir / f"episode_{episode:06d}_windowed_report.json" + report_path.write_text(json.dumps(report, ensure_ascii=False, indent=2), encoding="utf-8") + logging.info("Saved windowed comparison: %s", comparison_path) + logging.info("Saved windowed report: %s", report_path) + + +if __name__ == "__main__": + main() diff --git a/scripts/eval/rollout_g2_server_windowed.py b/scripts/eval/rollout_g2_server_windowed.py new file mode 100644 index 00000000..ac2b24f7 --- /dev/null +++ b/scripts/eval/rollout_g2_server_windowed.py @@ -0,0 +1,515 @@ +#!/usr/bin/env python3 +"""Teacher-forced DreamZero G2 video/action window evaluation. + +The persistent server must be started with ``--no-reset-cache-each-request``. +Each evaluation window sends four consecutive 4-frame observations through the +same causal cache. DreamZero then decodes one 33-frame video window. The +matching GT segment is taken from the same episode timeline and is stitched +side-by-side at 30 Hz; no global frame-rate resampling is performed. +""" + +from __future__ import annotations + +import argparse +import csv +import json +import logging +import sys +import time +import uuid +from pathlib import Path + +# Keep evaluation code isolated under scripts/eval while importing the shared +# websocket client from the repository root. This file never touches the +# robot deployment client. +sys.path.insert(0, str(Path(__file__).resolve().parents[2])) + +import cv2 +import imageio.v2 as imageio +import numpy as np +import pyarrow.parquet as pq + +from eval_utils.policy_client import WebsocketClientPolicy + + +ACTION_NAMES = ( + [f"left_joint_{i}" for i in range(7)] + + ["left_gripper"] + + [f"right_joint_{i}" for i in range(7)] + + ["right_gripper"] +) +ARM_DIMS = [*range(7), *range(8, 15)] +GRIPPER_DIMS = [7, 15] + + +def _episode_path(root: Path, template: str, episode: int, chunks_size: int, **kwargs: object) -> Path: + return root / template.format( + episode_chunk=episode // chunks_size, + episode_index=episode, + **kwargs, + ) + + +def _read_video(path: Path, expected_rows: int) -> np.ndarray: + capture = cv2.VideoCapture(str(path)) + if not capture.isOpened(): + raise RuntimeError(f"Failed to open G2 video: {path}") + frames: list[np.ndarray] = [] + try: + while len(frames) < expected_rows: + ok, frame = capture.read() + if not ok: + break + frames.append(np.ascontiguousarray(frame)) + finally: + capture.release() + if len(frames) < expected_rows: + raise RuntimeError( + f"{path} contains {len(frames)} frames, expected {expected_rows}" + ) + return np.stack(frames, axis=0) + + +def _encode_video_observation(frames: np.ndarray, quality: int = 80) -> dict[str, object]: + """Encode BGR frames exactly like the live G2 JPEG transport. + + The live client captures BGR, JPEG-encodes it, and the G2 server decodes + then converts BGR to RGB before DreamZero sees it. Sending cv2's raw BGR + array would silently create a different offline evaluation condition. + """ + array = np.asarray(frames) + if array.ndim != 4 or array.shape[-1] != 3: + raise ValueError(f"Expected video frames with shape (T,H,W,3), got {array.shape}") + encoded: list[bytes] = [] + jpeg_quality = int(np.clip(quality, 1, 100)) + for frame in array: + ok, payload = cv2.imencode( + ".jpg", np.ascontiguousarray(frame), + [int(cv2.IMWRITE_JPEG_QUALITY), jpeg_quality], + ) + if not ok: + raise RuntimeError("Failed to JPEG-encode evaluation frame") + encoded.append(payload.tobytes()) + return { + "__dreamzero_image_encoding__": "jpeg_sequence", + "shape": tuple(int(dim) for dim in array.shape), + "dtype": str(array.dtype), + "quality": jpeg_quality, + "frames": encoded, + } + + +def _grid_rgb(top_bgr: np.ndarray, left_bgr: np.ndarray, right_bgr: np.ndarray) -> np.ndarray: + top = cv2.cvtColor(top_bgr, cv2.COLOR_BGR2RGB) + left = cv2.cvtColor(left_bgr, cv2.COLOR_BGR2RGB) + right = cv2.cvtColor(right_bgr, cv2.COLOR_BGR2RGB) + black = np.zeros_like(top) + return np.concatenate( + [np.concatenate([top, left], axis=1), np.concatenate([right, black], axis=1)], + axis=0, + ) + + +def _label(frame_rgb: np.ndarray, text: str) -> np.ndarray: + frame_bgr = cv2.cvtColor(frame_rgb, cv2.COLOR_RGB2BGR) + cv2.rectangle(frame_bgr, (0, 0), (370, 34), (0, 0, 0), -1) + cv2.putText( + frame_bgr, + text, + (8, 24), + cv2.FONT_HERSHEY_SIMPLEX, + 0.65, + (255, 255, 255), + 2, + cv2.LINE_AA, + ) + return cv2.cvtColor(frame_bgr, cv2.COLOR_BGR2RGB) + + +def _save_video(path: Path, frames: list[np.ndarray], fps: int = 30) -> None: + if not frames: + raise RuntimeError(f"No frames to save: {path}") + path.parent.mkdir(parents=True, exist_ok=True) + imageio.mimsave(path, frames, fps=fps, codec="libx264", macro_block_size=None) + + +def _find_new_server_video(server_dir: Path, before: set[str], timeout: float = 60.0) -> Path: + deadline = time.time() + timeout + while time.time() < deadline: + candidates = [ + p for p in server_dir.glob("*.mp4") + if p.name not in before and p.stat().st_size > 0 + ] + if candidates: + return max(candidates, key=lambda p: p.stat().st_mtime_ns) + time.sleep(0.25) + raise RuntimeError( + f"Server reset did not create a new video in {server_dir}; " + f"existing={sorted(before)[-3:]}" + ) + + +def _read_server_video(path: Path, expected_frames: int) -> list[np.ndarray]: + reader = imageio.get_reader(path) + try: + frames = [np.asarray(frame) for frame in reader] + finally: + reader.close() + if len(frames) != expected_frames: + raise RuntimeError( + f"Expected exactly {expected_frames} decoded prediction frames from " + f"the causal window, got {len(frames)} ({path.name}). " + "Do not resample this diagnostic: check server cache mode and VAE decode." + ) + return [np.ascontiguousarray(frame) for frame in frames] + + +def _write_action_reports( + output_dir: Path, + episode: int, + window_starts: list[int], + anchors: list[int], + predicted: np.ndarray, + ground_truth: np.ndarray, +) -> dict[str, object]: + # Shapes: [window, causal_block, horizon, action_dim]. + if predicted.shape != ground_truth.shape or predicted.ndim != 4: + raise ValueError(f"Action arrays must match 4-D shape, got {predicted.shape} and {ground_truth.shape}") + error = np.abs(predicted - ground_truth) + detail_path = output_dir / f"episode_{episode:06d}_windowed_action_pred_vs_gt.csv" + fields = ["window_index", "window_start", "block_index", "anchor_frame", "horizon_step"] + fields += [f"gt_{name}" for name in ACTION_NAMES] + fields += [f"pred_{name}" for name in ACTION_NAMES] + with detail_path.open("w", newline="", encoding="utf-8") as stream: + writer = csv.DictWriter(stream, fieldnames=fields) + writer.writeheader() + for wi, start in enumerate(window_starts): + for bi in range(predicted.shape[1]): + anchor = anchors[wi * predicted.shape[1] + bi] + for hi in range(predicted.shape[2]): + row: dict[str, object] = { + "window_index": wi, + "window_start": start, + "block_index": bi, + "anchor_frame": anchor, + "horizon_step": hi + 1, + } + row.update({f"gt_{name}": float(v) for name, v in zip(ACTION_NAMES, ground_truth[wi, bi, hi])}) + row.update({f"pred_{name}": float(v) for name, v in zip(ACTION_NAMES, predicted[wi, bi, hi])}) + writer.writerow(row) + + metrics_path = output_dir / f"episode_{episode:06d}_windowed_action_metrics.csv" + metric_fields = ["scope", "window_index", "window_start", "block_index", "anchor_frame", "overall_mae", "arm_mae", "gripper_mae", "left_arm_mae", "right_arm_mae", "num_action_rows"] + + def metric_row(scope: str, e: np.ndarray, **extra: object) -> dict[str, object]: + row: dict[str, object] = { + "scope": scope, + "overall_mae": float(e.mean()), + "arm_mae": float(e[..., ARM_DIMS].mean()), + "gripper_mae": float(e[..., GRIPPER_DIMS].mean()), + "left_arm_mae": float(e[..., :7].mean()), + "right_arm_mae": float(e[..., 8:15].mean()), + "num_action_rows": int(np.prod(e.shape[:-1])), + } + row.update(extra) + return row + + metric_rows: list[dict[str, object]] = [] + num_blocks = predicted.shape[1] + for wi, start in enumerate(window_starts): + for bi in range(num_blocks): + metric_rows.append( + metric_row( + f"window_{wi:04d}_block_{bi:02d}", + error[wi, bi], + window_index=wi, + window_start=start, + block_index=bi, + anchor_frame=anchors[wi * num_blocks + bi], + ) + ) + overall = metric_row("overall", error, window_index="", window_start="", block_index="", anchor_frame="") + metric_rows.append(overall) + with metrics_path.open("w", newline="", encoding="utf-8") as stream: + writer = csv.DictWriter(stream, fieldnames=metric_fields) + writer.writeheader() + writer.writerows(metric_rows) + + # One 16-dimension plot over the complete sampled task. Each action row + # is a predicted 24-step horizon at one causal observation anchor. + pred_flat = predicted.reshape(-1, predicted.shape[-1]) + gt_flat = ground_truth.reshape(-1, ground_truth.shape[-1]) + plot_path = output_dir / f"episode_{episode:06d}_windowed_action_pred_vs_gt.png" + try: + import matplotlib + + matplotlib.use("Agg") + import matplotlib.pyplot as plt + + fig, axes = plt.subplots(4, 4, figsize=(22, 13), sharex=True) + x = np.arange(len(pred_flat)) + for dim, axis in enumerate(axes.flat): + axis.plot(x, gt_flat[:, dim], color="tab:blue", linewidth=0.8, label="GT") + axis.plot(x, pred_flat[:, dim], color="tab:orange", linewidth=0.8, alpha=0.9, label="PRED") + axis.set_title(f"{dim}: {ACTION_NAMES[dim]}") + axis.grid(alpha=0.2) + axes[0, 0].legend(loc="upper right") + axes[-1, 0].set_xlabel("window/block/horizon action row") + axes[-1, 1].set_xlabel("GT and predicted action") + fig.suptitle(f"G2 checkpoint-3000: full windowed action vs GT, episode {episode}") + fig.tight_layout() + fig.savefig(plot_path, dpi=140) + plt.close(fig) + + mae_path = output_dir / f"episode_{episode:06d}_windowed_action_mae_by_horizon.png" + horizon_mae = error.mean(axis=(0, 1, 3)) + arm_mae = error[..., ARM_DIMS].mean(axis=(0, 1, 3)) + grip_mae = error[..., GRIPPER_DIMS].mean(axis=(0, 1, 3)) + fig, axis = plt.subplots(figsize=(12, 5)) + h = np.arange(1, error.shape[2] + 1) + axis.plot(h, horizon_mae, label="all 16 dims", linewidth=2) + axis.plot(h, arm_mae, label="arm 14 dims", linewidth=2) + axis.plot(h, grip_mae, label="gripper 2 dims", linewidth=2) + axis.set_xlabel("action horizon step") + axis.set_ylabel("mean absolute error") + axis.set_title("Predicted action vs training GT: error by horizon step") + axis.grid(alpha=0.25) + axis.legend() + fig.tight_layout() + fig.savefig(mae_path, dpi=140) + plt.close(fig) + except Exception as exc: # diagnostic output should not hide CSV results + logging.warning("Could not create action plots: %s", exc) + plot_path = None + mae_path = None + + return { + "action_detail_csv": str(detail_path), + "action_metrics_csv": str(metrics_path), + "action_plot": str(plot_path) if plot_path else None, + "action_horizon_mae_plot": str(mae_path) if mae_path else None, + "overall_action_metrics": overall, + } + + +def _parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser() + parser.add_argument("--host", default="127.0.0.1") + parser.add_argument("--port", type=int, default=30002) + parser.add_argument("--test-data-root", type=Path, required=True) + parser.add_argument("--episode-index", type=int, default=0) + parser.add_argument("--output-dir", type=Path, required=True) + parser.add_argument("--server-video-dir", type=Path, required=True) + parser.add_argument("--window-future-frames", type=int, default=33) + parser.add_argument("--window-history", type=int, default=4) + parser.add_argument("--causal-blocks", type=int, default=4) + parser.add_argument("--window-stride", type=int, default=33) + parser.add_argument("--image-jpeg-quality", type=int, default=80) + parser.add_argument("--max-windows", type=int, default=None) + parser.add_argument("--prompt", default=None) + return parser.parse_args() + + +def main() -> None: + logging.basicConfig(level=logging.INFO, force=True) + args = _parse_args() + if args.window_history != 4 or args.causal_blocks != 4: + raise ValueError("This protocol requires four consecutive 4-frame causal packets") + if args.window_future_frames != 33: + raise ValueError("This protocol requires exactly a 33-frame decoded video window") + if args.window_stride < 1: + raise ValueError("--window-stride must be positive") + + args.output_dir.mkdir(parents=True, exist_ok=True) + root = args.test_data_root.resolve() + info = json.loads((root / "meta/info.json").read_text(encoding="utf-8")) + chunks_size = int(info.get("chunks_size", 1000)) + episode = int(args.episode_index) + parquet_path = _episode_path( + root, + info.get("data_path", "data/chunk-{episode_chunk:03d}/episode_{episode_index:06d}.parquet"), + episode, + chunks_size, + ) + table = pq.read_table(parquet_path) + state = np.asarray(table["observation.state"].to_pylist(), dtype=np.float32).reshape(table.num_rows, 16) + action = np.asarray(table["action"].to_pylist(), dtype=np.float32).reshape(table.num_rows, 16) + num_rows = int(table.num_rows) + video_template = info.get("video_path", "videos/chunk-{episode_chunk:03d}/{video_key}/episode_{episode_index:06d}.mp4") + videos: dict[str, np.ndarray] = {} + for name, key in (("top_head", "observation.images.top_head"), ("hand_left", "observation.images.hand_left"), ("hand_right", "observation.images.hand_right")): + videos[name] = _read_video(_episode_path(root, video_template, episode, chunks_size, video_key=key), num_rows) + + # A window starts with the four-frame packet [start:start+4]. The video + # timeline used for comparison begins at the first packet anchor + # (start+3) and contains 33 frames, i.e. rows start+3..start+35. The + # four causal packets are context for that same request; they are not + # appended to the future-video duration. The old bound added the 16 + # context frames to the 33 future frames and silently dropped the tail of + # every episode. + max_start = num_rows - (args.window_history + args.window_future_frames - 1) + starts = list(range(0, max_start + 1, args.window_stride)) + if starts[-1] < max_start: + starts.append(max_start) + if args.max_windows is not None: + starts = starts[: max(0, int(args.max_windows))] + if not starts: + raise RuntimeError(f"No valid 33-frame windows in episode with {num_rows} rows") + prompt_key = "annotation.language.action_text" + prompt = args.prompt if args.prompt is not None else (str(table[prompt_key][0].as_py()) if prompt_key in table.column_names else "") + session_id = f"g2-windowed-episode-{episode:06d}-{uuid.uuid4()}" + server_dir = args.server_video_dir.resolve() + server_dir.mkdir(parents=True, exist_ok=True) + + logging.info( + "Windowed rollout: episode=%d rows=%d windows=%d stride=%d history=%d causal_blocks=%d future=%d", + episode, num_rows, len(starts), args.window_stride, args.window_history, args.causal_blocks, args.window_future_frames, + ) + client = WebsocketClientPolicy(host=args.host, port=args.port) + logging.info("Server metadata: %s", client.get_server_metadata()) + # Clear any stale server-side cache/video from an earlier client session. + try: + client.reset({"session_id": session_id}) + except Exception: + logging.exception("Initial server reset failed") + + all_pred: list[np.ndarray] = [] + all_gt: list[np.ndarray] = [] + all_anchors: list[int] = [] + all_comparison_frames: list[np.ndarray] = [] + window_reports: list[dict[str, object]] = [] + action_window_starts: list[int] = [] + started = time.time() + for wi, window_start in enumerate(starts): + before = {p.name for p in server_dir.glob("*.mp4")} + pred_chunks: list[np.ndarray] = [] + gt_chunks: list[np.ndarray] = [] + anchors_window: list[int] = [] + for block_index in range(args.causal_blocks): + packet_start = window_start + block_index * args.window_history + indices = list(range(packet_start, packet_start + args.window_history)) + anchor = indices[-1] + result = np.asarray( + client.infer( + { + "observation/top_head": _encode_video_observation( + videos["top_head"][indices], args.image_jpeg_quality + ), + "observation/hand_left": _encode_video_observation( + videos["hand_left"][indices], args.image_jpeg_quality + ), + "observation/hand_right": _encode_video_observation( + videos["hand_right"][indices], args.image_jpeg_quality + ), + "observation/state": state[anchor], + "prompt": prompt, + "session_id": session_id, + } + ), + dtype=np.float32, + ) + if result.shape != (24, 16): + raise RuntimeError(f"window {wi} block {block_index}: server returned {result.shape}, expected (24,16)") + pred_chunks.append(result) + gt_chunks.append(action[anchor : anchor + 24]) + anchors_window.append(anchor) + # Reset flushes exactly this four-forward causal video window; it does + # not unload the server/checkpoint. + client.reset({"session_id": session_id}) + predicted_video_path = _find_new_server_video(server_dir, before) + predicted_frames = _read_server_video(predicted_video_path, args.window_future_frames) + # The server has to save this temporary causal decode so that the + # client can read it, but it is not an evaluation artifact. Keep only + # the single stitched episode video under args.output_dir instead of + # leaving one 33-frame mp4 per causal window in the server directory. + try: + predicted_video_path.unlink() + except OSError as exc: + logging.warning("Could not remove temporary server video %s: %s", predicted_video_path, exc) + gt_indices = range(window_start + 3, window_start + 3 + args.window_future_frames) + gt_frames = [ + _grid_rgb(videos["top_head"][idx], videos["hand_left"][idx], videos["hand_right"][idx]) + for idx in gt_indices + ] + comparison_frames = [ + np.concatenate([_label(pred, "PREDICTED"), _label(gt, "G2 GROUND TRUTH")], axis=1) + for pred, gt in zip(predicted_frames, gt_frames) + ] + # Avoid duplicate frames when a final tail window was appended to + # cover an episode whose length is not an exact multiple of the + # requested stride. The GT index is the authoritative timeline. + gt_start = window_start + 3 + trim = max(0, len(all_comparison_frames) + 3 - gt_start) + all_comparison_frames.extend(comparison_frames[trim:]) + # A video tail can be valid while its last causal packet has fewer + # than 24 GT action rows remaining. Keep that video window, but do + # not mix a short action slice into the fixed [24,16] action tensor; + # the dedicated full-action pass below evaluates every valid anchor. + action_valid = all(chunk.shape == (24, 16) for chunk in gt_chunks) + action_mae: float | None = None + if action_valid: + pred_array = np.stack(pred_chunks) + gt_array = np.stack(gt_chunks) + all_pred.append(pred_array) + all_gt.append(gt_array) + all_anchors.extend(anchors_window) + action_window_starts.append(window_start) + action_mae = float(np.abs(pred_array - gt_array).mean()) + window_reports.append( + { + "window_index": wi, + "window_start": window_start, + "packet_starts": [window_start + args.window_history * i for i in range(args.causal_blocks)], + "anchor_indices": anchors_window, + "gt_video_indices": [window_start + 3, window_start + 3 + args.window_future_frames - 1], + "predicted_video": str(predicted_video_path), + "predicted_frames": len(predicted_frames), + "ground_truth_frames": len(gt_frames), + "action_mae": action_mae, + } + ) + logging.info( + "window %d/%d start=%d video=%d frames action_mae=%s elapsed=%.1fs", + wi + 1, len(starts), window_start, len(predicted_frames), + f"{action_mae:.6f}" if action_mae is not None else "n/a", + time.time() - started, + ) + + if not all_pred: + raise RuntimeError("No complete 24-step action windows remained for action reporting") + predicted = np.stack(all_pred) + ground_truth = np.stack(all_gt) + np.savez_compressed( + args.output_dir / f"episode_{episode:06d}_windowed_action_arrays.npz", + predicted=predicted, + ground_truth=ground_truth, + anchors=np.asarray(all_anchors, dtype=np.int32), + window_starts=np.asarray(action_window_starts, dtype=np.int32), + ) + comparison_path = args.output_dir / f"episode_{episode:06d}_windowed_predicted_vs_gt_f{len(all_comparison_frames)}.mp4" + _save_video(comparison_path, all_comparison_frames, fps=30) + reports = _write_action_reports(args.output_dir, episode, action_window_starts, all_anchors, predicted, ground_truth) + report = { + "episode_index": episode, + "checkpoint_server_protocol": "four 4-frame history packets -> one causal 33-frame decoded window; reset only between windows", + "server_cache_required": "DREAMZERO_RESET_AR_EACH_REQUEST=false", + "num_rows": num_rows, + "video_window_starts": starts, + "action_window_starts": action_window_starts, + "window_stride": args.window_stride, + "window_history": args.window_history, + "causal_blocks": args.causal_blocks, + "future_frames": args.window_future_frames, + "comparison_fps": 30, + "comparison_video": str(comparison_path), + "window_reports": window_reports, + **reports, + } + report_path = args.output_dir / f"episode_{episode:06d}_windowed_report.json" + report_path.write_text(json.dumps(report, ensure_ascii=False, indent=2), encoding="utf-8") + logging.info("Saved windowed comparison: %s", comparison_path) + logging.info("Saved windowed report: %s", report_path) + + +if __name__ == "__main__": + main() diff --git a/scripts/eval/summarize_g2_action_timeline.py b/scripts/eval/summarize_g2_action_timeline.py new file mode 100644 index 00000000..7a965744 --- /dev/null +++ b/scripts/eval/summarize_g2_action_timeline.py @@ -0,0 +1,142 @@ +#!/usr/bin/env python3 +"""Turn overlapping 24-step action chunks into a chronological task curve. + +The full-anchor evaluator intentionally preserves every [anchor, horizon] +comparison. This helper also selects the latest prediction available for +each episode frame, which is the useful receding-horizon view for deployment; +it avoids the misleading plot where overlapping chunks are simply flattened +end-to-end. +""" + +from __future__ import annotations + +import argparse +import csv +import json +from pathlib import Path + +import numpy as np + +ACTION_NAMES = ( + [f"left_joint_{i}" for i in range(7)] + + ["left_gripper"] + + [f"right_joint_{i}" for i in range(7)] + + ["right_gripper"] +) +ARM_DIMS = [*range(7), *range(8, 15)] +GRIPPER_DIMS = [7, 15] + + +def _metric(name: str, error: np.ndarray) -> dict[str, object]: + return { + "scope": name, + "overall_mae": float(error.mean()), + "arm_mae": float(error[..., ARM_DIMS].mean()), + "gripper_mae": float(error[..., GRIPPER_DIMS].mean()), + "left_arm_mae": float(error[..., :7].mean()), + "right_arm_mae": float(error[..., 8:15].mean()), + "num_action_rows": int(error.shape[0]), + } + + +def main() -> None: + parser = argparse.ArgumentParser() + parser.add_argument("--input-dir", type=Path, required=True) + parser.add_argument("--episode-index", type=int, default=0) + parser.add_argument("--output-dir", type=Path, default=None) + args = parser.parse_args() + out = (args.output_dir or args.input_dir).resolve() + out.mkdir(parents=True, exist_ok=True) + source = out / f"episode_{args.episode_index:06d}_full_action_arrays.npz" + if not source.is_file(): + source = args.input_dir.resolve() / f"episode_{args.episode_index:06d}_full_action_arrays.npz" + data = np.load(source) + predicted = np.asarray(data["predicted"], dtype=np.float32) + ground_truth = np.asarray(data["ground_truth"], dtype=np.float32) + anchors = np.asarray(data["anchors"], dtype=np.int64) + if predicted.ndim != 3 or predicted.shape != ground_truth.shape or predicted.shape[-1] != 16: + raise ValueError(f"Expected matching [blocks,24,16] arrays, got {predicted.shape} and {ground_truth.shape}") + + # Later anchors overwrite earlier ones. This is the receding-horizon + # action that would be current at each frame when replanning every packet. + rows: dict[int, tuple[int, int, np.ndarray, np.ndarray]] = {} + coverage: dict[int, int] = {} + for block, anchor in enumerate(anchors.tolist()): + for horizon in range(predicted.shape[1]): + frame = int(anchor) + horizon + rows[frame] = (block, horizon, predicted[block, horizon], ground_truth[block, horizon]) + coverage[frame] = coverage.get(frame, 0) + 1 + frames = np.asarray(sorted(rows), dtype=np.int64) + block_used = np.asarray([rows[int(frame)][0] for frame in frames], dtype=np.int64) + horizon_used = np.asarray([rows[int(frame)][1] for frame in frames], dtype=np.int64) + pred = np.stack([rows[int(frame)][2] for frame in frames]).astype(np.float32) + gt = np.stack([rows[int(frame)][3] for frame in frames]).astype(np.float32) + counts = np.asarray([coverage[int(frame)] for frame in frames], dtype=np.int64) + err = np.abs(pred - gt) + + np.savez_compressed( + out / f"episode_{args.episode_index:06d}_full_action_timeline_arrays.npz", + frame_indices=frames, + predicted_latest=pred, + ground_truth=gt, + block_used=block_used, + horizon_used=horizon_used, + coverage_count=counts, + ) + csv_path = out / f"episode_{args.episode_index:06d}_full_action_timeline_pred_vs_gt.csv" + fields = ["frame_index", "block_used", "horizon_step", "coverage_count"] + fields += [f"gt_{name}" for name in ACTION_NAMES] + fields += [f"pred_{name}" for name in ACTION_NAMES] + with csv_path.open("w", newline="", encoding="utf-8") as stream: + writer = csv.DictWriter(stream, fieldnames=fields) + writer.writeheader() + for i, frame in enumerate(frames): + row: dict[str, object] = { + "frame_index": int(frame), + "block_used": int(block_used[i]), + "horizon_step": int(horizon_used[i]) + 1, + "coverage_count": int(counts[i]), + } + row.update({f"gt_{name}": float(value) for name, value in zip(ACTION_NAMES, gt[i])}) + row.update({f"pred_{name}": float(value) for name, value in zip(ACTION_NAMES, pred[i])}) + writer.writerow(row) + + plot_path = out / f"episode_{args.episode_index:06d}_full_action_timeline_pred_vs_gt.png" + try: + import matplotlib + + matplotlib.use("Agg") + import matplotlib.pyplot as plt + + fig, axes = plt.subplots(4, 4, figsize=(22, 13), sharex=True) + for dim, axis in enumerate(axes.flat): + axis.plot(frames, gt[:, dim], color="tab:blue", lw=0.8, label="GT") + axis.plot(frames, pred[:, dim], color="tab:orange", lw=0.8, label="PRED latest") + axis.set_title(f"{dim}: {ACTION_NAMES[dim]}") + axis.grid(alpha=0.2) + axes[0, 0].legend() + fig.suptitle(f"G2 chronological latest-anchor action vs GT, episode {args.episode_index}") + fig.tight_layout() + fig.savefig(plot_path, dpi=140) + plt.close(fig) + except Exception: + plot_path = None + + report = { + "episode_index": args.episode_index, + "source_arrays": str(source), + "timeline_frames": int(len(frames)), + "frame_start": int(frames[0]) if len(frames) else None, + "frame_end": int(frames[-1]) if len(frames) else None, + "selection": "latest prediction for each frame among overlapping 24-step chunks", + "timeline_metrics": _metric("latest_timeline", err), + "timeline_csv": str(csv_path), + "timeline_plot": str(plot_path) if plot_path else None, + } + report_path = out / f"episode_{args.episode_index:06d}_full_action_timeline_report.json" + report_path.write_text(json.dumps(report, ensure_ascii=False, indent=2), encoding="utf-8") + print(json.dumps(report, ensure_ascii=False, indent=2)) + + +if __name__ == "__main__": + main() diff --git a/scripts/infer/compare_g2_live_to_dataset.py b/scripts/infer/compare_g2_live_to_dataset.py new file mode 100644 index 00000000..8cd6e265 --- /dev/null +++ b/scripts/infer/compare_g2_live_to_dataset.py @@ -0,0 +1,139 @@ +"""Rank G2 Gear episode starts by similarity to one live camera snapshot.""" + +from __future__ import annotations + +import argparse +import json +from pathlib import Path + +import cv2 +import numpy as np +import pyarrow.parquet as pq + + +CAMERAS = ("top_head", "hand_left", "hand_right") + + +def _first_rgb(path: Path) -> np.ndarray: + capture = cv2.VideoCapture(str(path)) + try: + ok, frame = capture.read() + finally: + capture.release() + if not ok: + raise RuntimeError(f"Could not decode first frame: {path}") + return cv2.cvtColor(frame, cv2.COLOR_BGR2RGB) + + +def _image_score(reference: np.ndarray, candidate: np.ndarray) -> float: + if candidate.shape != reference.shape: + candidate = cv2.resize( + candidate, + (reference.shape[1], reference.shape[0]), + interpolation=cv2.INTER_AREA, + ) + # Low-frequency comparison emphasizes camera framing and object layout over + # JPEG noise and small illumination changes. + reference_small = cv2.resize(reference, (80, 44), interpolation=cv2.INTER_AREA) + candidate_small = cv2.resize(candidate, (80, 44), interpolation=cv2.INTER_AREA) + return float( + np.mean( + np.abs( + reference_small.astype(np.float32) + - candidate_small.astype(np.float32) + ) + ) + / 255.0 + ) + + +def main() -> None: + parser = argparse.ArgumentParser() + parser.add_argument("snapshot", type=Path) + parser.add_argument("dataset_root", type=Path) + parser.add_argument("--output-dir", type=Path, required=True) + parser.add_argument("--top-k", type=int, default=10) + args = parser.parse_args() + + snapshot = np.load(args.snapshot, allow_pickle=False) + live = { + camera: np.asarray( + snapshot[f"{camera}_jpeg_decoded"][-1], dtype=np.uint8 + ) + for camera in CAMERAS + } + live_state = np.asarray(snapshot["state_16"], dtype=np.float32) + + with (args.dataset_root / "meta/info.json").open( + "r", encoding="utf-8" + ) as stream: + info = json.load(stream) + total = int(info["total_episodes"]) + + results: list[dict[str, object]] = [] + cached_frames: dict[int, dict[str, np.ndarray]] = {} + for episode in range(total): + frames: dict[str, np.ndarray] = {} + view_scores: dict[str, float] = {} + for camera in CAMERAS: + path = ( + args.dataset_root + / "videos" + / f"chunk-{episode // 1000:03d}" + / f"observation.images.{camera}" + / f"episode_{episode:06d}.mp4" + ) + frames[camera] = _first_rgb(path) + view_scores[camera] = _image_score(live[camera], frames[camera]) + + parquet = ( + args.dataset_root + / "data" + / f"chunk-{episode // 1000:03d}" + / f"episode_{episode:06d}.parquet" + ) + state0 = np.asarray( + pq.read_table( + parquet, columns=["observation.state"] + )["observation.state"][0].as_py(), + dtype=np.float32, + ) + arm_indices = np.r_[0:7, 8:15] + state_arm_mae = float( + np.mean(np.abs(state0[arm_indices] - live_state[arm_indices])) + ) + results.append( + { + "episode": episode, + "image_score_mean": float(np.mean(list(view_scores.values()))), + "view_scores": view_scores, + "state_arm_mae_rad": state_arm_mae, + "dataset_grippers": state0[[7, 15]].tolist(), + "live_grippers": live_state[[7, 15]].tolist(), + } + ) + cached_frames[episode] = frames + + results.sort(key=lambda item: float(item["image_score_mean"])) + args.output_dir.mkdir(parents=True, exist_ok=True) + with (args.output_dir / "ranking.json").open("w", encoding="utf-8") as stream: + json.dump(results, stream, ensure_ascii=False, indent=2) + + # One readable sheet per top match: dataset row followed by live row. + for result in results[: args.top_k]: + episode = int(result["episode"]) + rows = [ + np.concatenate([cached_frames[episode][camera] for camera in CAMERAS], axis=1), + np.concatenate([live[camera] for camera in CAMERAS], axis=1), + ] + rgb = np.concatenate(rows, axis=0) + cv2.imwrite( + str(args.output_dir / f"episode_{episode:06d}_dataset_vs_live.png"), + cv2.cvtColor(rgb, cv2.COLOR_RGB2BGR), + ) + + print(json.dumps(results[: args.top_k], ensure_ascii=False, indent=2)) + + +if __name__ == "__main__": + main() diff --git a/scripts/make_zero_lora_control.py b/scripts/make_zero_lora_control.py new file mode 100644 index 00000000..9a8e9200 --- /dev/null +++ b/scripts/make_zero_lora_control.py @@ -0,0 +1,66 @@ +#!/usr/bin/env python3 +"""Create a zero-delta LoRA checkpoint for base-model inference controls.""" + +from __future__ import annotations + +import argparse +import json +import shutil +from pathlib import Path + +import torch +from safetensors.torch import load_file, save_file + + +def main() -> int: + parser = argparse.ArgumentParser() + parser.add_argument("source", type=Path) + parser.add_argument("output", type=Path) + args = parser.parse_args() + + source = args.source.resolve() + output = args.output.resolve() + if output.exists(): + raise FileExistsError(f"Refusing to overwrite existing path: {output}") + + source_weights = source / "model.safetensors" + source_config = source / "config.json" + source_experiment = source / "experiment_cfg" + for path in (source_weights, source_config, source_experiment): + if not path.exists(): + raise FileNotFoundError(path) + + state = load_file(source_weights) + zero_lora = { + key: torch.zeros_like(value) + for key, value in state.items() + if "lora_A" in key or "lora_B" in key + } + if not zero_lora: + raise RuntimeError(f"No LoRA tensors found in {source_weights}") + + output.mkdir(parents=True) + save_file(zero_lora, output / "model.safetensors") + shutil.copy2(source_config, output / "config.json") + shutil.copytree(source_experiment, output / "experiment_cfg") + (output / "CONTROL_INFO.json").write_text( + json.dumps( + { + "purpose": "DreamZero-AgiBot base control in G2 architecture", + "source_checkpoint": str(source), + "zeroed_lora_tensors": len(zero_lora), + "non_lora_checkpoint_tensors_loaded": 0, + }, + indent=2, + ) + + "\n", + encoding="utf-8", + ) + print( + f"Created {output} with {len(zero_lora)} zero LoRA tensors" + ) + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/scripts/open_loop_g2.py b/scripts/open_loop_g2.py new file mode 100644 index 00000000..75cb99c7 --- /dev/null +++ b/scripts/open_loop_g2.py @@ -0,0 +1,508 @@ +#!/usr/bin/env python3 +"""Simple offline open-loop evaluation for DreamZero G2. + +This is the G2 counterpart of ``scripts/open_loop_yam.py``. It evaluates +the action head directly from a checkpoint; no websocket server, server-side +cache, temporary MP4 files, or causal-window stitching is involved. + +At each sampled anchor ``t`` the evaluator feeds the four ground-truth +frames ``[t-3, t]``, the state at ``t``, and the dataset's task prompt. The +returned 24-step action chunk is compared with ``action[t:t+24]``. Anchors +are independent: the action-head inference cache is reset before every +anchor, so a result cannot accidentally depend on the previous sample. + +Example (on the training server): + + python scripts/open_loop_g2.py \ + --model-path /data/wangk/checkpoints/.../checkpoint-500 \ + --dataset-path /data/training_data/teleop/g2/g2_mock_light_module_joint_gear_policy_gripper/train \ + --episodes 0,1,2 \ + --max-anchors 8 \ + --output-dir /tmp/g2_open_loop_checkpoint_500 + +The script intentionally reports all 16 action dimensions: 14 arm joints +and two grippers. Long generated-video comparison remains a separate visual +spot check because ``open_loop_yam.py`` itself is an action-fit evaluator, +not a video-rollout evaluator. +""" + +from __future__ import annotations + +import argparse +import csv +import json +import os +import time +from dataclasses import dataclass +from pathlib import Path + +import torch._dynamo + +torch._dynamo.config.disable = True + +import cv2 +import matplotlib + +matplotlib.use("Agg") +import matplotlib.pyplot as plt +import numpy as np +import pyarrow.parquet as pq +import torch +import torch.distributed as dist +from tianshou.data import Batch + +from groot.vla.data.schema import EmbodimentTag +from groot.vla.model.n1_5.sim_policy import GrootSimPolicy + + +VIDEO_KEYS = { + "top_head": "observation.images.top_head", + "hand_left": "observation.images.hand_left", + "hand_right": "observation.images.hand_right", +} + +STATE_KEYS = { + "state.left_joint_position": (0, 7), + "state.left_gripper_position": (7, 8), + "state.right_joint_position": (8, 15), + "state.right_gripper_position": (15, 16), +} + +ACTION_KEYS = ( + "action.left_joint_position", + "action.left_gripper_position", + "action.right_joint_position", + "action.right_gripper_position", +) + +ACTION_NAMES = ( + *(f"left_joint_{i}" for i in range(7)), + "left_gripper", + *(f"right_joint_{i}" for i in range(7)), + "right_gripper", +) + +ARM_DIMS = np.asarray([*range(7), *range(8, 15)], dtype=np.int64) +GRIPPER_DIMS = np.asarray([7, 15], dtype=np.int64) + + +@dataclass +class Episode: + index: int + state: np.ndarray + action: np.ndarray + videos: dict[str, np.ndarray] + prompt: str + + +def _episode_path( + root: Path, + template: str, + episode: int, + chunks_size: int, + **kwargs: object, +) -> Path: + return root / template.format( + episode_chunk=episode // chunks_size, + episode_index=episode, + **kwargs, + ) + + +def _read_video(path: Path, expected_rows: int) -> np.ndarray: + """Read an MP4 as RGB frames with shape ``[T,H,W,3]``.""" + + capture = cv2.VideoCapture(str(path)) + if not capture.isOpened(): + raise RuntimeError(f"Failed to open G2 video: {path}") + frames: list[np.ndarray] = [] + try: + while len(frames) < expected_rows: + ok, frame = capture.read() + if not ok: + break + frames.append(cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)) + finally: + capture.release() + if len(frames) < expected_rows: + raise RuntimeError( + f"{path} contains {len(frames)} frames, expected {expected_rows}" + ) + return np.stack(frames, axis=0) + + +def _scalar(value: object) -> str: + if isinstance(value, np.ndarray): + if value.size == 0: + return "" + return _scalar(value.reshape(-1)[0]) + if isinstance(value, (list, tuple)): + return _scalar(value[0]) if value else "" + return str(value) + + +def load_episode(root: Path, episode: int) -> Episode: + info = json.loads((root / "meta/info.json").read_text(encoding="utf-8")) + chunks_size = int(info.get("chunks_size", 1000)) + parquet = _episode_path( + root, + info.get( + "data_path", + "data/chunk-{episode_chunk:03d}/episode_{episode_index:06d}.parquet", + ), + episode, + chunks_size, + ) + table = pq.read_table(parquet) + n = int(table.num_rows) + state = np.asarray(table["observation.state"].to_pylist(), dtype=np.float32).reshape(n, 16) + action = np.asarray(table["action"].to_pylist(), dtype=np.float32).reshape(n, 16) + + video_template = info.get( + "video_path", + "videos/chunk-{episode_chunk:03d}/{video_key}/episode_{episode_index:06d}.mp4", + ) + videos = { + name: _read_video( + _episode_path( + root, + video_template, + episode, + chunks_size, + video_key=key, + ), + n, + ) + for name, key in VIDEO_KEYS.items() + } + + prompt_key = "annotation.language.action_text" + if prompt_key in table.column_names: + prompt = _scalar(table[prompt_key][0].as_py()) + else: + prompt = "" + return Episode(index=episode, state=state, action=action, videos=videos, prompt=prompt) + + +def _build_obs(episode: Episode, anchor: int, prompt: str) -> dict[str, object]: + """Build the raw G2 observation expected by the G2 transform pipeline.""" + + if anchor < 3: + raise ValueError(f"G2 requires anchor >= 3 for four-frame history, got {anchor}") + obs: dict[str, object] = { + f"video.{name}": frames[anchor - 3 : anchor + 1] + for name, frames in episode.videos.items() + } + state = episode.state[anchor] + for key, (start, end) in STATE_KEYS.items(): + # The leading singleton is the state-time dimension. GrootSimPolicy + # adds the batch dimension itself, just as open_loop_yam.py does. + obs[key] = state[start:end].reshape(1, -1).astype(np.float32) + obs["annotation.language.action_text"] = prompt + return obs + + +def _as_2d(value: object, width: int) -> np.ndarray: + array = value.detach().cpu().numpy() if isinstance(value, torch.Tensor) else np.asarray(value) + if array.ndim == 3 and array.shape[0] == 1: + array = array[0] + if array.ndim == 1: + if width == 1: + array = array[:, None] + else: + array = array[None, :] + if array.ndim != 2 or array.shape[1] != width: + raise RuntimeError(f"Expected action array [24,{width}], got {array.shape}") + return array.astype(np.float32, copy=False) + + +def _extract_action(result: Batch) -> np.ndarray: + if not hasattr(result, "act"): + raise RuntimeError("Policy result has no .act field") + action = result.act + pieces: list[np.ndarray] = [] + widths = (7, 1, 7, 1) + for key, width in zip(ACTION_KEYS, widths): + if key not in action: + break + pieces.append(_as_2d(action[key], width)) + if len(pieces) == len(ACTION_KEYS): + output = np.concatenate(pieces, axis=-1) + elif "action" in action: + output = _as_2d(action["action"], 16) + else: + raise RuntimeError(f"Unexpected policy action keys: {list(action.keys())}") + if output.shape != (24, 16): + raise RuntimeError(f"Expected complete G2 action chunk [24,16], got {output.shape}") + return output + + +def _reset_action_cache(policy: GrootSimPolicy) -> None: + head = getattr(policy.trained_model, "action_head", None) + reset = getattr(head, "reset_inference_cache", None) + if reset is None: + raise RuntimeError( + "The loaded action head has no reset_inference_cache(); " + "independent open-loop anchors would not be safe." + ) + reset() + + +def _select_anchors(num_rows: int, stride: int, max_anchors: int | None) -> list[int]: + last = num_rows - 24 + if last < 3: + return [] + anchors = list(range(3, last + 1, max(1, stride))) + if max_anchors is not None: + anchors = anchors[: max(0, max_anchors)] + return anchors + + +def _plot_action_chunk(path: Path, predicted: np.ndarray, ground_truth: np.ndarray, anchor: int) -> None: + fig, axes = plt.subplots(1, 3, figsize=(19, 5), sharex=True) + groups = (("left arm (7)", range(0, 7)), ("right arm (7)", range(8, 15)), ("grippers (2)", (7, 15))) + horizon = np.arange(predicted.shape[0]) + for ax, (title, dims) in zip(axes, groups): + for dim in dims: + ax.plot(horizon, ground_truth[:, dim], "--", linewidth=1.4, alpha=0.75, label=f"GT {ACTION_NAMES[dim]}") + ax.plot(horizon, predicted[:, dim], linewidth=1.4, alpha=0.85, label=f"Pred {ACTION_NAMES[dim]}") + ax.set_title(title) + ax.set_xlabel("chunk horizon (24 steps)") + ax.grid(alpha=0.25) + axes[0].set_ylabel("absolute action") + axes[-1].legend(fontsize=7, ncol=2, loc="best") + fig.suptitle(f"G2 open-loop action chunk | anchor={anchor} | solid=Pred, dashed=GT") + fig.tight_layout() + fig.savefig(path, dpi=160) + plt.close(fig) + + +def _plot_anchor_first_step(path: Path, predicted: np.ndarray, ground_truth: np.ndarray, anchors: list[int]) -> None: + fig, axes = plt.subplots(1, 3, figsize=(19, 5), sharex=True) + groups = (("left arm (7)", range(0, 7)), ("right arm (7)", range(8, 15)), ("grippers (2)", (7, 15))) + x = np.asarray(anchors) + for ax, (title, dims) in zip(axes, groups): + for dim in dims: + ax.plot(x, ground_truth[:, 0, dim], "--", linewidth=1.2, alpha=0.75, label=f"GT {ACTION_NAMES[dim]}") + ax.plot(x, predicted[:, 0, dim], linewidth=1.2, alpha=0.85, label=f"Pred {ACTION_NAMES[dim]}") + ax.set_title(title) + ax.set_xlabel("GT timeline anchor") + ax.grid(alpha=0.25) + axes[0].set_ylabel("first action in predicted chunk") + axes[-1].legend(fontsize=7, ncol=2, loc="best") + fig.suptitle("G2 open-loop first-action timeline | solid=Pred, dashed=GT") + fig.tight_layout() + fig.savefig(path, dpi=160) + plt.close(fig) + + +def _plot_horizon_error(path: Path, error: np.ndarray) -> None: + fig, ax = plt.subplots(figsize=(10, 5)) + ax.plot(np.mean(error, axis=(0, 2)), label="all 16", linewidth=2) + ax.plot(np.mean(error[:, :, ARM_DIMS], axis=(0, 2)), label="arm 14", linewidth=2) + ax.plot(np.mean(error[:, :, GRIPPER_DIMS], axis=(0, 2)), label="gripper 2", linewidth=2) + ax.set_xlabel("action horizon step") + ax.set_ylabel("mean absolute error") + ax.set_title("G2 open-loop error by action horizon") + ax.grid(alpha=0.25) + ax.legend() + fig.tight_layout() + fig.savefig(path, dpi=160) + plt.close(fig) + + +def _write_episode_report( + output_dir: Path, + episode: Episode, + anchors: list[int], + predicted: np.ndarray, + ground_truth: np.ndarray, + inference_seconds: list[float], +) -> dict[str, object]: + error = np.abs(predicted - ground_truth) + metrics = { + "episode": episode.index, + "prompt": episode.prompt, + "num_rows": int(episode.action.shape[0]), + "num_anchors": len(anchors), + "anchors": anchors, + "action_shape": list(predicted.shape), + "mae_all_16": float(error.mean()), + "mae_arm_14": float(error[:, :, ARM_DIMS].mean()), + "mae_gripper_2": float(error[:, :, GRIPPER_DIMS].mean()), + "mae_first_action_all_16": float(error[:, 0].mean()), + "mae_by_horizon_all_16": error.mean(axis=(0, 2)).tolist(), + "mae_by_dim": error.mean(axis=(0, 1)).tolist(), + "avg_inference_seconds": float(np.mean(inference_seconds)), + "median_inference_seconds": float(np.median(inference_seconds)), + } + + np.savez_compressed( + output_dir / f"episode_{episode.index:06d}_open_loop_actions.npz", + anchors=np.asarray(anchors, dtype=np.int32), + predicted=predicted, + ground_truth=ground_truth, + ) + with (output_dir / f"episode_{episode.index:06d}_anchor_metrics.csv").open("w", newline="", encoding="utf-8") as handle: + writer = csv.writer(handle) + writer.writerow(("anchor", "mae_all_16", "mae_arm_14", "mae_gripper_2", "inference_seconds")) + for i, anchor in enumerate(anchors): + row_error = error[i] + writer.writerow(( + anchor, + float(row_error.mean()), + float(row_error[:, ARM_DIMS].mean()), + float(row_error[:, GRIPPER_DIMS].mean()), + inference_seconds[i], + )) + + _plot_action_chunk( + output_dir / f"episode_{episode.index:06d}_action_chunk_anchor_{anchors[0]:06d}.png", + predicted[0], + ground_truth[0], + anchors[0], + ) + _plot_anchor_first_step( + output_dir / f"episode_{episode.index:06d}_action_first_step_timeline.png", + predicted, + ground_truth, + anchors, + ) + _plot_horizon_error( + output_dir / f"episode_{episode.index:06d}_action_horizon_mae.png", + error, + ) + report_path = output_dir / f"episode_{episode.index:06d}_open_loop_report.json" + report_path.write_text(json.dumps(metrics, ensure_ascii=False, indent=2), encoding="utf-8") + return metrics + + +def _init_single_process_dist(port: int) -> None: + if dist.is_initialized(): + return + os.environ.setdefault("MASTER_ADDR", "127.0.0.1") + os.environ.setdefault("MASTER_PORT", str(port)) + dist.init_process_group(backend="gloo", world_size=1, rank=0) + + +def _set_cuda_device(device: str) -> None: + if not device.startswith("cuda") or not torch.cuda.is_available(): + return + if ":" in device: + torch.cuda.set_device(int(device.rsplit(":", 1)[1])) + + +def evaluate(args: argparse.Namespace) -> None: + _init_single_process_dist(args.dist_port) + _set_cuda_device(args.device) + args.output_dir.mkdir(parents=True, exist_ok=True) + + print(f"Loading G2 checkpoint directly from {args.model_path} ...") + policy = GrootSimPolicy( + embodiment_tag=EmbodimentTag.G2, + model_path=str(args.model_path), + device=args.device, + ) + print("Model loaded.") + + episode_indices = [int(item) for item in args.episodes.split(",") if item.strip()] + all_reports: list[dict[str, object]] = [] + for episode_index in episode_indices: + episode = load_episode(args.dataset_path, episode_index) + prompt = args.prompt if args.prompt is not None else episode.prompt + if not prompt.strip(): + raise RuntimeError( + f"Episode {episode_index} has an empty prompt. Pass --prompt explicitly " + "or fix annotation.language.action_text in the dataset." + ) + anchors = _select_anchors(len(episode.action), args.anchor_stride, args.max_anchors) + if not anchors: + raise RuntimeError( + f"Episode {episode_index} has no valid anchors: rows={len(episode.action)}" + ) + predicted_chunks: list[np.ndarray] = [] + ground_truth_chunks: list[np.ndarray] = [] + inference_seconds: list[float] = [] + print( + f"\nEpisode {episode_index}: rows={len(episode.action)} " + f"anchors={len(anchors)} prompt={prompt!r}" + ) + for i, anchor in enumerate(anchors): + _reset_action_cache(policy) + obs = _build_obs(episode, anchor, prompt) + started = time.perf_counter() + with torch.inference_mode(): + result, _ = policy.lazy_joint_forward_causal(Batch(obs=obs)) + elapsed = time.perf_counter() - started + pred = _extract_action(result) + gt = episode.action[anchor : anchor + 24].astype(np.float32, copy=False) + if gt.shape != (24, 16): + raise RuntimeError(f"Episode {episode_index} anchor {anchor}: GT shape={gt.shape}") + predicted_chunks.append(pred) + ground_truth_chunks.append(gt) + inference_seconds.append(elapsed) + if i == 0 or (i + 1) % args.log_every == 0 or i + 1 == len(anchors): + chunk_mae = float(np.abs(pred - gt).mean()) + print( + f" [{i + 1:>3d}/{len(anchors)}] anchor={anchor} " + f"mae={chunk_mae:.6f} infer={elapsed:.3f}s" + ) + + predicted = np.stack(predicted_chunks) + ground_truth = np.stack(ground_truth_chunks) + metrics = _write_episode_report( + args.output_dir, + episode, + anchors, + predicted, + ground_truth, + inference_seconds, + ) + metrics["prompt_used"] = prompt + all_reports.append(metrics) + print( + f" episode metrics: all16={metrics['mae_all_16']:.6f} " + f"arm14={metrics['mae_arm_14']:.6f} " + f"gripper2={metrics['mae_gripper_2']:.6f}" + ) + + summary = { + "protocol": "direct checkpoint open-loop; 4 GT history frames; independent cache reset per anchor", + "checkpoint": str(args.model_path), + "dataset": str(args.dataset_path), + "episodes": episode_indices, + "anchor_stride": args.anchor_stride, + "max_anchors": args.max_anchors, + "reports": all_reports, + "mean_mae_all_16": float(np.mean([item["mae_all_16"] for item in all_reports])), + "mean_mae_arm_14": float(np.mean([item["mae_arm_14"] for item in all_reports])), + "mean_mae_gripper_2": float(np.mean([item["mae_gripper_2"] for item in all_reports])), + } + (args.output_dir / "summary.json").write_text( + json.dumps(summary, ensure_ascii=False, indent=2), encoding="utf-8" + ) + print( + f"\nSummary: all16={summary['mean_mae_all_16']:.6f} " + f"arm14={summary['mean_mae_arm_14']:.6f} " + f"gripper2={summary['mean_mae_gripper_2']:.6f}" + ) + print(f"Results saved to {args.output_dir.resolve()}") + + +def _parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(formatter_class=argparse.ArgumentDefaultsHelpFormatter) + parser.add_argument("--model-path", type=Path, required=True) + parser.add_argument("--dataset-path", type=Path, required=True) + parser.add_argument("--device", default="cuda:0") + parser.add_argument("--episodes", default="0,1,2", help="Comma-separated episode indices") + parser.add_argument("--prompt", default=None, help="Override the dataset prompt") + parser.add_argument("--anchor-stride", type=int, default=24) + parser.add_argument("--max-anchors", type=int, default=8) + parser.add_argument("--output-dir", type=Path, default=Path("results_g2_open_loop")) + parser.add_argument("--log-every", type=int, default=1) + parser.add_argument("--dist-port", type=int, default=29501) + return parser.parse_args() + + +if __name__ == "__main__": + evaluate(_parse_args()) diff --git a/scripts/plot_train_vs_infer_actions.py b/scripts/plot_train_vs_infer_actions.py new file mode 100644 index 00000000..a94fb2a4 --- /dev/null +++ b/scripts/plot_train_vs_infer_actions.py @@ -0,0 +1,445 @@ +#!/usr/bin/env python3 +# -*- coding: utf-8 -*- + +import os +import json +import glob +import math +import argparse +from pathlib import Path + +import numpy as np +import pandas as pd +import matplotlib.pyplot as plt + + +DIM_NAMES = [ + "l.joint1", + "l.joint2", + "l.joint3", + "l.joint4", + "l.joint5", + "l.joint6", + "l.joint7", + "l.gripper", + "r.joint1", + "r.joint2", + "r.joint3", + "r.joint4", + "r.joint5", + "r.joint6", + "r.joint7", + "r.gripper", +] + + +def load_train_actions(dataset_root: str, max_files: int | None = None): + """ + 读取训练集 parquet: + - absolute action: df["action"] + - relative action: df["action"] - df["observation.state"] + + 返回: + train_abs: (N, 16) + train_rel: (N, 16) + """ + pattern = os.path.join(dataset_root, "data", "chunk-*", "*.parquet") + parquet_files = sorted(glob.glob(pattern)) + + if not parquet_files: + raise FileNotFoundError(f"No parquet files found under: {pattern}") + + if max_files is not None: + parquet_files = parquet_files[:max_files] + + abs_list = [] + rel_list = [] + + total_rows = 0 + for i, path in enumerate(parquet_files, 1): + df = pd.read_parquet(path, columns=["observation.state", "action"]) + + # 列里每个元素都是 list[16] + states = np.stack(df["observation.state"].to_numpy()).astype(np.float32) + actions = np.stack(df["action"].to_numpy()).astype(np.float32) + + if states.shape[1] != 16 or actions.shape[1] != 16: + raise ValueError( + f"Expected 16-dim state/action, got state={states.shape}, action={actions.shape}, file={path}" + ) + + rel = actions - states + + abs_list.append(actions) + rel_list.append(rel) + + total_rows += len(df) + print(f"[train] loaded {i}/{len(parquet_files)}: {path} rows={len(df)} total_rows={total_rows}") + + train_abs = np.concatenate(abs_list, axis=0) + train_rel = np.concatenate(rel_list, axis=0) + + print(f"[train] absolute actions shape: {train_abs.shape}") + print(f"[train] relative actions shape: {train_rel.shape}") + return train_abs, train_rel + + +def load_infer_json(json_path: str): + """ + 读取 robot_json / agibot_actions_xxx.json + 返回: + infer_raw : 模型原始输出动作 (T_all, 16) -> 优先用 raw_actions + infer_exec : 最终执行动作 (T_all, 16) -> 用 actions + records : 原始 step 记录 + """ + with open(json_path, "r", encoding="utf-8") as f: + records = json.load(f) + + if not isinstance(records, list) or len(records) == 0: + raise ValueError(f"Expected a non-empty list in {json_path}") + + raw_list = [] + exec_list = [] + + for idx, rec in enumerate(records): + if "raw_actions" in rec and rec["raw_actions"] is not None: + arr = np.asarray(rec["raw_actions"], dtype=np.float32) + if arr.ndim == 1: + arr = arr.reshape(1, -1) + if arr.shape[1] == 16: + raw_list.append(arr) + + if "actions" in rec and rec["actions"] is not None: + arr = np.asarray(rec["actions"], dtype=np.float32) + if arr.ndim == 1: + arr = arr.reshape(1, -1) + if arr.shape[1] == 16: + exec_list.append(arr) + + if len(raw_list) == 0: + raise ValueError(f"No 16-D raw_actions found in {json_path}") + + if len(exec_list) == 0: + raise ValueError(f"No 16-D actions found in {json_path}") + + infer_raw = np.concatenate(raw_list, axis=0) + infer_exec = np.concatenate(exec_list, axis=0) + + print(f"[infer] raw_actions shape : {infer_raw.shape}") + print(f"[infer] actions shape : {infer_exec.shape}") + print(f"[infer] num step records : {len(records)}") + return infer_raw, infer_exec, records + + +def summarize_array(name: str, arr: np.ndarray): + print("\n" + "=" * 80) + print(name) + print("=" * 80) + print(f"shape: {arr.shape}") + for d in range(arr.shape[1]): + col = arr[:, d] + print( + f"{d:2d} {DIM_NAMES[d]:>10s} | " + f"mean={col.mean(): .5f} std={col.std(): .5f} " + f"min={col.min(): .5f} p01={np.percentile(col,1): .5f} " + f"p99={np.percentile(col,99): .5f} max={col.max(): .5f}" + ) + + +def make_output_dir(out_dir: str): + Path(out_dir).mkdir(parents=True, exist_ok=True) + return out_dir + + +def _auto_grid(n: int): + cols = 4 + rows = math.ceil(n / cols) + return rows, cols + + +def plot_hist_compare( + ref_arr: np.ndarray, + cmp_arr: np.ndarray, + ref_name: str, + cmp_name: str, + out_path: str, + bins: int = 80, +): + """ + 每个维度一张子图,画两个直方图轮廓 + """ + rows, cols = _auto_grid(ref_arr.shape[1]) + fig = plt.figure(figsize=(20, 4 * rows)) + + for d in range(ref_arr.shape[1]): + ax = fig.add_subplot(rows, cols, d + 1) + ref_col = ref_arr[:, d] + cmp_col = cmp_arr[:, d] + + all_min = min(ref_col.min(), cmp_col.min()) + all_max = max(ref_col.max(), cmp_col.max()) + + if all_min == all_max: + all_min -= 1e-3 + all_max += 1e-3 + + edges = np.linspace(all_min, all_max, bins + 1) + + ax.hist( + ref_col, + bins=edges, + histtype="step", + density=True, + label=ref_name, + linewidth=1.5, + ) + ax.hist( + cmp_col, + bins=edges, + histtype="step", + density=True, + label=cmp_name, + linewidth=1.5, + ) + + ax.set_title(f"{d}: {DIM_NAMES[d]}") + ax.grid(True, alpha=0.3) + if d == 0: + ax.legend() + + fig.suptitle(f"{ref_name} vs {cmp_name}", fontsize=16) + fig.tight_layout(rect=[0, 0, 1, 0.97]) + fig.savefig(out_path, dpi=180) + plt.close(fig) + print(f"[saved] {out_path}") + + +def plot_percentile_compare( + ref_arr: np.ndarray, + cmp_arr: np.ndarray, + ref_name: str, + cmp_name: str, + out_path: str, +): + """ + 每个维度画 p01/p50/p99 对比柱状图 + """ + dims = np.arange(ref_arr.shape[1]) + + ref_p01 = np.percentile(ref_arr, 1, axis=0) + ref_p50 = np.percentile(ref_arr, 50, axis=0) + ref_p99 = np.percentile(ref_arr, 99, axis=0) + + cmp_p01 = np.percentile(cmp_arr, 1, axis=0) + cmp_p50 = np.percentile(cmp_arr, 50, axis=0) + cmp_p99 = np.percentile(cmp_arr, 99, axis=0) + + fig = plt.figure(figsize=(18, 8)) + ax = fig.add_subplot(111) + + width = 0.12 + ax.bar(dims - 2.5 * width, ref_p01, width=width, label=f"{ref_name} p01") + ax.bar(dims - 1.5 * width, ref_p50, width=width, label=f"{ref_name} p50") + ax.bar(dims - 0.5 * width, ref_p99, width=width, label=f"{ref_name} p99") + + ax.bar(dims + 0.5 * width, cmp_p01, width=width, label=f"{cmp_name} p01") + ax.bar(dims + 1.5 * width, cmp_p50, width=width, label=f"{cmp_name} p50") + ax.bar(dims + 2.5 * width, cmp_p99, width=width, label=f"{cmp_name} p99") + + ax.set_xticks(dims) + ax.set_xticklabels(DIM_NAMES, rotation=45, ha="right") + ax.set_title(f"Percentile Comparison: {ref_name} vs {cmp_name}") + ax.grid(True, axis="y", alpha=0.3) + ax.legend(ncol=2) + fig.tight_layout() + fig.savefig(out_path, dpi=180) + plt.close(fig) + print(f"[saved] {out_path}") + + +def plot_single_rollout_timeseries( + records: list, + out_path: str, + use_key: str = "actions", +): + """ + 画单个推理 JSON 中所有 step 的 action 时间序列。 + 这里把所有 step 的 action 直接串起来。 + """ + arrs = [] + boundaries = [] + + cur = 0 + for rec in records: + if use_key not in rec or rec[use_key] is None: + continue + arr = np.asarray(rec[use_key], dtype=np.float32) + if arr.ndim == 1: + arr = arr.reshape(1, -1) + if arr.shape[1] != 16: + continue + arrs.append(arr) + cur += arr.shape[0] + boundaries.append(cur) + + if not arrs: + raise ValueError(f"No valid {use_key} found for time-series plot.") + + seq = np.concatenate(arrs, axis=0) + rows, cols = _auto_grid(seq.shape[1]) + + fig = plt.figure(figsize=(20, 4 * rows)) + x = np.arange(seq.shape[0]) + + for d in range(seq.shape[1]): + ax = fig.add_subplot(rows, cols, d + 1) + ax.plot(x, seq[:, d], linewidth=1.2) + for b in boundaries[:-1]: + ax.axvline(b, linestyle="--", linewidth=0.8, alpha=0.5) + ax.set_title(f"{d}: {DIM_NAMES[d]}") + ax.grid(True, alpha=0.3) + + fig.suptitle(f"Inference {use_key} Time Series", fontsize=16) + fig.tight_layout(rect=[0, 0, 1, 0.97]) + fig.savefig(out_path, dpi=180) + plt.close(fig) + print(f"[saved] {out_path}") + + +def save_summary_csv( + ref_arr: np.ndarray, + cmp_arr: np.ndarray, + ref_name: str, + cmp_name: str, + out_csv: str, +): + rows = [] + for d in range(ref_arr.shape[1]): + r = ref_arr[:, d] + c = cmp_arr[:, d] + rows.append({ + "dim": d, + "name": DIM_NAMES[d], + f"{ref_name}_mean": float(r.mean()), + f"{ref_name}_std": float(r.std()), + f"{ref_name}_p01": float(np.percentile(r, 1)), + f"{ref_name}_p50": float(np.percentile(r, 50)), + f"{ref_name}_p99": float(np.percentile(r, 99)), + f"{cmp_name}_mean": float(c.mean()), + f"{cmp_name}_std": float(c.std()), + f"{cmp_name}_p01": float(np.percentile(c, 1)), + f"{cmp_name}_p50": float(np.percentile(c, 50)), + f"{cmp_name}_p99": float(np.percentile(c, 99)), + "mean_abs_gap": float(abs(r.mean() - c.mean())), + "p99_abs_gap": float(abs(np.percentile(r, 99) - np.percentile(c, 99))), + }) + + df = pd.DataFrame(rows) + df.to_csv(out_csv, index=False) + print(f"[saved] {out_csv}") + + +def main(): + parser = argparse.ArgumentParser() + parser.add_argument( + "--train-root", + required=True, + help="GEAR train root, e.g. /data/training_data/teleop/g2/g2_mock_light_module_joint_gear/train", + ) + parser.add_argument( + "--infer-json", + required=True, + help="robot_json/agibot_actions_xxx.json", + ) + parser.add_argument( + "--out-dir", + default="./action_compare_outputs", + help="output directory", + ) + parser.add_argument( + "--max-train-files", + type=int, + default=None, + help="only load first N parquet files for quick debug", + ) + args = parser.parse_args() + + out_dir = make_output_dir(args.out_dir) + + # 1) load + train_abs, train_rel = load_train_actions(args.train_root, args.max_train_files) + infer_raw, infer_exec, records = load_infer_json(args.infer_json) + + # 2) summary + summarize_array("TRAIN ABSOLUTE ACTION", train_abs) + summarize_array("TRAIN RELATIVE ACTION", train_rel) + summarize_array("INFER RAW ACTION", infer_raw) + summarize_array("INFER EXEC ACTION", infer_exec) + + # 3) save csv + save_summary_csv( + train_abs, + infer_exec, + "train_abs", + "infer_exec", + os.path.join(out_dir, "summary_train_abs_vs_infer_exec.csv"), + ) + + save_summary_csv( + train_rel, + infer_raw, + "train_rel", + "infer_raw", + os.path.join(out_dir, "summary_train_rel_vs_infer_raw.csv"), + ) + + # 4) plots + # A. 最重要:训练 absolute vs 推理最终执行 + plot_hist_compare( + train_abs, + infer_exec, + "train_abs", + "infer_exec", + os.path.join(out_dir, "hist_train_abs_vs_infer_exec.png"), + ) + plot_percentile_compare( + train_abs, + infer_exec, + "train_abs", + "infer_exec", + os.path.join(out_dir, "percentile_train_abs_vs_infer_exec.png"), + ) + + # B. 用来查 relative/absolute 是否错位:训练 relative vs 推理原始输出 + plot_hist_compare( + train_rel, + infer_raw, + "train_rel", + "infer_raw", + os.path.join(out_dir, "hist_train_rel_vs_infer_raw.png"), + ) + plot_percentile_compare( + train_rel, + infer_raw, + "train_rel", + "infer_raw", + os.path.join(out_dir, "percentile_train_rel_vs_infer_raw.png"), + ) + + # C. 单次 rollout 时间序列 + plot_single_rollout_timeseries( + records, + os.path.join(out_dir, "timeseries_infer_raw_actions.png"), + use_key="raw_actions", + ) + plot_single_rollout_timeseries( + records, + os.path.join(out_dir, "timeseries_infer_exec_actions.png"), + use_key="actions", + ) + + print("\nDone.") + + +if __name__ == "__main__": + main() diff --git a/scripts/run_agibot_fruit_eval.sh b/scripts/run_agibot_fruit_eval.sh new file mode 100644 index 00000000..d28fecbb --- /dev/null +++ b/scripts/run_agibot_fruit_eval.sh @@ -0,0 +1,109 @@ +#!/bin/bash +set -euo pipefail + +ROOT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")/.." && pwd)" +cd "$ROOT_DIR" + +unset http_proxy https_proxy HTTP_PROXY HTTPS_PROXY all_proxy ALL_PROXY +unset ws_proxy wss_proxy WS_PROXY WSS_PROXY + +export HF_HUB_OFFLINE="${HF_HUB_OFFLINE:-1}" +export PYTORCH_NVML_BASED_CUDA_CHECK="${PYTORCH_NVML_BASED_CUDA_CHECK:-1}" +export NO_ALBUMENTATIONS_UPDATE="${NO_ALBUMENTATIONS_UPDATE:-1}" +export ATTENTION_BACKEND="${ATTENTION_BACKEND:-FA2}" + +PYTHON_BIN="${PYTHON_BIN:-/data/wangk/conda/envs/dreamzero/bin/python}" +PORT="${PORT:-6006}" +NUM_GPUS="${NUM_GPUS:-2}" +CUDA_VISIBLE_DEVICES_VALUE="${CUDA_VISIBLE_DEVICES:-0,1}" + +SERVER_MODULE="${SERVER_MODULE:-socket_test_optimized_AR_agibot_fruit_video.py}" +EMBODIMENT_TAG="agibot" + +# Select one experiment explicitly. +MODEL_PATH="${MODEL_PATH:-/data/wangk/checkpoints/DreamZero-AgiBot}" +RUN_NAME="${RUN_NAME:-official_agibot}" + +WAN_CKPT_DIR="${WAN_CKPT_DIR:-/data/wangk/checkpoints/Wan2.1-I2V-14B-480P}" +TOKENIZER_PATH="${TOKENIZER_PATH:-/data/wangk/checkpoints/umt5-xxl}" + +# Generated-model video and real first-person video are separate. +VIDEO_SAVE_MODE="${VIDEO_SAVE_MODE:-full}" +INPUT_VIDEO_SAVE_MODE="${INPUT_VIDEO_SAVE_MODE:-top_head}" +INPUT_VIDEO_FPS="${INPUT_VIDEO_FPS:-30}" + +NUM_INFERENCE_TIMESTEPS="${NUM_INFERENCE_TIMESTEPS:-0}" +OUTPUT_ROOT="${OUTPUT_ROOT:-/data/wangk/dreamzero/fruit_eval}" +OUTPUT_DIR="${OUTPUT_DIR:-$OUTPUT_ROOT/$RUN_NAME}" + +if [ ! -x "$PYTHON_BIN" ]; then + echo "ERROR: PYTHON_BIN is not executable: $PYTHON_BIN" + exit 1 +fi + +if [ ! -f "$SERVER_MODULE" ]; then + echo "ERROR: SERVER_MODULE does not exist: $ROOT_DIR/$SERVER_MODULE" + exit 1 +fi + +if [ ! -d "$MODEL_PATH" ]; then + echo "ERROR: MODEL_PATH does not exist: $MODEL_PATH" + exit 1 +fi + +if [ ! -d "$WAN_CKPT_DIR" ]; then + echo "ERROR: WAN_CKPT_DIR does not exist: $WAN_CKPT_DIR" + exit 1 +fi + +if [ ! -e "$TOKENIZER_PATH" ]; then + echo "ERROR: TOKENIZER_PATH does not exist: $TOKENIZER_PATH" + exit 1 +fi + +IFS=',' read -r -a CUDA_DEVICE_ARRAY <<< "$CUDA_VISIBLE_DEVICES_VALUE" +if [ "${#CUDA_DEVICE_ARRAY[@]}" -ne "$NUM_GPUS" ]; then + echo "ERROR: NUM_GPUS=$NUM_GPUS but CUDA_VISIBLE_DEVICES has ${#CUDA_DEVICE_ARRAY[@]} device(s)" + exit 1 +fi + +mkdir -p "$OUTPUT_DIR" + +cat > "$OUTPUT_DIR/run_manifest.txt" <&2 + exit 1 +} + +[[ "$GPU_IDS" == "0,1,2,3" ]] || fail "This run is pinned to GPUs 0,1,2,3" +for path in "$PROJECT_ROOT" "$G2_DATA_ROOT" "$WAN_CKPT_DIR" \ + "$TOKENIZER_DIR" "$PRETRAINED_MODEL_PATH"; do + [[ -d "$path" ]] || fail "Missing directory: $path" +done +[[ -x "$PYTHON_BIN" ]] || fail "Missing Python interpreter: $PYTHON_BIN" + +"$PYTHON_BIN" - <<'PY' +import json +import os +from pathlib import Path + +root = Path(os.environ["G2_DATA_ROOT"]) +info = json.loads((root / "meta/info.json").read_text()) +expected = int(os.environ["EXPECTED_EPISODES"]) +if info["total_episodes"] != expected: + raise RuntimeError(f"Expected {expected} episodes, got {info['total_episodes']}") +if info["features"]["observation.state"]["shape"] != [16]: + raise RuntimeError("G2 state must be 16D") +if info["features"]["action"]["shape"] != [16]: + raise RuntimeError("G2 action must be 16D") +modality = json.loads((root / "meta/modality.json").read_text()) +expected_video = {"top_head", "hand_left", "hand_right"} +if set(modality["video"]) != expected_video: + raise RuntimeError(f"Unexpected video modalities: {sorted(modality['video'])}") +stats = json.loads((root / "meta/stats.json").read_text()) +for feature_name in ("observation.state", "action"): + values = stats[feature_name] + for index, key in ((7, "left_gripper_position"), (15, "right_gripper_position")): + if values["min"][index] < -1e-4 or values["max"][index] > 1.0001: + raise RuntimeError( + f"{feature_name} {key} is not policy-space [0,1]: " + f"min={values['min'][index]} max={values['max'][index]}" + ) +# All-absolute action contract: joint positions must be absolute (not near-zero +# relative deltas). An absolute joint position spans a wide [q01, q99] range, +# whereas a relative delta sits in a narrow band around 0. Check the SPREAD. +for feature_name in ("observation.state", "action"): + values = stats[feature_name] + arm_indices = [*range(0, 7), *range(8, 15)] + for i in arm_indices: + spread = values["q99"][i] - values["q01"][i] + if spread < 0.1: + raise RuntimeError( + f"{feature_name} dim {i} looks like a relative delta " + f"(q01={values['q01'][i]:.4f} q99={values['q99'][i]:.4f} " + f"spread={spread:.4f}); absolute joint-position stats expected." + ) +active_hold = json.loads((root / "meta/g2_active_hold_windows.json").read_text()) +if active_hold["action_horizon"] != 24: + raise RuntimeError("Active/hold index must use a 24-step horizon") +if active_hold["active_count"] <= 0 or active_hold["hold_count"] <= 0: + raise RuntimeError("Active/hold index is empty") +tasks_path = root / "meta/tasks.jsonl" +tasks = [json.loads(line) for line in tasks_path.read_text().splitlines() if line.strip()] +if len(tasks) != 1 or not str(tasks[0].get("task", "")).strip(): + raise RuntimeError("Expected exactly one non-empty G2 task prompt") +print(f"Prompt contract: {tasks[0]['task']}") +parquets = list(root.glob("data/chunk-*/*.parquet")) +videos = list(root.glob("videos/chunk-*/*/*.mp4")) +if len(parquets) != expected or len(videos) != expected * 3: + raise RuntimeError( + f"File count mismatch: parquets={len(parquets)} videos={len(videos)}" + ) +print( + "Validated G2 all-absolute dataset: " + f"episodes={expected} frames={info['total_frames']} tasks={info['total_tasks']} " + f"videos={len(videos)} active={active_hold['active_count']} " + f"hold={active_hold['hold_count']}" +) +PY + +if [[ -d "$OUTPUT_DIR" ]] && [[ -n "$(find "$OUTPUT_DIR" -mindepth 1 -maxdepth 1 -print -quit)" ]]; then + fail "Output directory is not empty: $OUTPUT_DIR" +fi +mkdir -p "$OUTPUT_DIR" +TRAIN_LOG="$OUTPUT_DIR/train.log" +cd "$PROJECT_ROOT" + +echo "G2 all-absolute joint training: Future-Video→Action bottleneck + state late gated residual" +echo "GPUs: $CUDA_VISIBLE_DEVICES" +echo "Base: $PRETRAINED_MODEL_PATH" +echo "Dataset: $G2_DATA_ROOT" +echo "Output: $OUTPUT_DIR" +echo "Objective: LoRA rank8 + flow action + decoded-action + endpoint (motion mask OFF)" + +"$PYTHON_BIN" -m torch.distributed.run \ + --nproc_per_node 4 \ + --standalone \ + groot/vla/experiment/experiment.py \ + report_to=wandb \ + data=dreamzero/g2_absolute \ + wandb_project=dreamzero \ + train_architecture=lora \ + num_frames=33 \ + action_horizon=24 \ + num_views=3 \ + model=dreamzero/vla \ + model/dreamzero/action_head=wan_flow_matching_action_tf \ + model/dreamzero/transform=dreamzero_cotrain \ + num_frame_per_block=2 \ + num_action_per_block=24 \ + num_state_per_block=1 \ + seed=42 \ + training_args.learning_rate=1e-5 \ + training_args.deepspeed="groot/vla/configs/deepspeed/zero2.json" \ + training_args.gradient_accumulation_steps=2 \ + save_steps="$SAVE_STEPS" \ + training_args.warmup_ratio=0.05 \ + output_dir="$OUTPUT_DIR" \ + per_device_train_batch_size=1 \ + max_steps="$MAX_STEPS" \ + weight_decay=1e-5 \ + save_total_limit=9 \ + upload_checkpoints=false \ + bf16=true \ + tf32=true \ + eval_bf16=true \ + dataloader_pin_memory=false \ + dataloader_num_workers=1 \ + image_resolution_width=320 \ + image_resolution_height=176 \ + save_lora_only=true \ + max_chunk_size=4 \ + frame_seqlen=880 \ + save_strategy=steps \ + active_window_ratio=0.9 \ + g2_data_root="$G2_DATA_ROOT" \ + dit_version="$WAN_CKPT_DIR" \ + text_encoder_pretrained_path="$WAN_CKPT_DIR/models_t5_umt5-xxl-enc-bf16.pth" \ + image_encoder_pretrained_path="$WAN_CKPT_DIR/models_clip_open-clip-xlm-roberta-large-vit-huge-14.pth" \ + vae_pretrained_path="$WAN_CKPT_DIR/Wan2.1_VAE.pth" \ + tokenizer_path="$TOKENIZER_DIR" \ + pretrained_model_path="$PRETRAINED_MODEL_PATH" \ + pretrained_lora_path=null \ + ++model_specific_transform.embodiment_tag_mapping.g2=33 \ + ++model_specific_transform.always_use_default_instruction=false \ + language_dropout_prob=0.0 \ + ++action_head_cfg.config.skip_component_loading=true \ + ++action_head_cfg.config.defer_lora_injection=true \ + action_head_cfg.config.action_only_training=false \ + action_head_cfg.config.tune_diffusion_model=true \ + ++action_head_cfg.config.lora_rank=8 \ + ++action_head_cfg.config.lora_alpha=8 \ + action_head_cfg.config.decouple_video_action_noise=false \ + action_head_cfg.config.dynamics_loss_weight=1.0 \ + action_head_cfg.config.action_loss_weight=1.0 \ + # x0 (decoded absolute-action reconstruction) is ~60-75x the flow loss in + # scale under absolute targets. Weight it down so flow matching stays the + # main action objective; x0 is only a small auxiliary. + ++action_head_cfg.config.action_x0_loss_weight=0.01 \ + ++action_head_cfg.config.action_endpoint_loss_weight=0.5 \ + ++action_head_cfg.config.action_velocity_loss_weight=0.5 \ + ++action_head_cfg.config.action_acc_loss_weight=0.1 \ + ++action_head_cfg.config.state_dropout=0.0 \ + ++action_head_cfg.config.cut_state_attention=true \ + ++action_head_cfg.config.motion_mask_enabled=false \ + 2>&1 | tee "$TRAIN_LOG" + +echo "Completed G2 all-absolute joint training" diff --git a/scripts/train/train_dreamzero_g2_action_only_stage2.sh b/scripts/train/train_dreamzero_g2_action_only_stage2.sh new file mode 100755 index 00000000..62c645a5 --- /dev/null +++ b/scripts/train/train_dreamzero_g2_action_only_stage2.sh @@ -0,0 +1,140 @@ +#!/usr/bin/env bash +set -Eeuo pipefail + +PROJECT_ROOT=${PROJECT_ROOT:-/home/ubuntu/projects/wangk/dreamzero} +G2_DATA_ROOT=${G2_DATA_ROOT:-/data/training_data/teleop/g2/g2_mock_light_module_joint_gear_policy_gripper/train} +OUTPUT_DIR=${OUTPUT_DIR:-/data/wangk/checkpoints/dreamzero_g2_action_adapter_clean_v1} +WAN_CKPT_DIR=${WAN_CKPT_DIR:-/data/wangk/checkpoints/Wan2.1-I2V-14B-480P} +TOKENIZER_DIR=${TOKENIZER_DIR:-/data/wangk/checkpoints/umt5-xxl} +PRETRAINED_MODEL_PATH=${PRETRAINED_MODEL_PATH:-/data/wangk/checkpoints/DreamZero-AgiBot} +PYTHON_BIN=${PYTHON_BIN:-/data/wangk/conda/envs/dreamzero/bin/python} + +GPU_IDS=${GPU_IDS:-0,1,2,3} +EXPECTED_EPISODES=${EXPECTED_EPISODES:-110} +MAX_STEPS=${MAX_STEPS:-4000} +WANDB_MODE=${WANDB_MODE:-offline} +HYDRA_FULL_ERROR=${HYDRA_FULL_ERROR:-1} + +export G2_DATA_ROOT OUTPUT_DIR WAN_CKPT_DIR TOKENIZER_DIR +export PRETRAINED_MODEL_PATH GPU_IDS +export EXPECTED_EPISODES MAX_STEPS WANDB_MODE HYDRA_FULL_ERROR + +fail() { + echo "[ERROR] $*" >&2 + exit 1 +} + +for path in \ + "$PROJECT_ROOT" \ + "$G2_DATA_ROOT" \ + "$WAN_CKPT_DIR" \ + "$TOKENIZER_DIR" \ + "$PRETRAINED_MODEL_PATH"; do + [[ -d "$path" ]] || fail "Missing directory: $path" +done +[[ -x "$PYTHON_BIN" ]] || fail "Missing Python interpreter: $PYTHON_BIN" + +IFS=',' read -r -a GPU_ARRAY <<< "$GPU_IDS" +[[ ${#GPU_ARRAY[@]} -eq 4 ]] || fail "Stage 2 requires exactly four GPUs: 0,1,2,3" +[[ "$GPU_IDS" == "0,1,2,3" ]] || fail "Refusing GPU selection other than 0,1,2,3" +export CUDA_VISIBLE_DEVICES="$GPU_IDS" + +"$PYTHON_BIN" - <<'PY' +import json +import os +from pathlib import Path + +root = Path(os.environ["G2_DATA_ROOT"]) +with (root / "meta/info.json").open() as stream: + info = json.load(stream) +expected = int(os.environ["EXPECTED_EPISODES"]) +if int(info["total_episodes"]) != expected: + raise RuntimeError( + f"Expected {expected} G2 episodes, got {info['total_episodes']}" + ) +if info["features"]["observation.state"]["shape"] != [16]: + raise RuntimeError("G2 state must be 16D") +if info["features"]["action"]["shape"] != [16]: + raise RuntimeError("G2 action must be 16D") +with (root / "meta/embodiment.json").open() as stream: + embodiment = json.load(stream) +if embodiment.get("embodiment_tag") != "g2": + raise RuntimeError(f"Expected embodiment_tag=g2, got {embodiment}") +with (root / "meta/g2_active_hold_windows.json").open() as stream: + windows = json.load(stream) +if windows["active_count"] <= 0 or windows["hold_count"] <= 0: + raise RuntimeError("Active/hold index must contain both groups") +print( + "Validated G2 action-adapter dataset: " + f"{expected} episodes, 16D state/action, " + f"active={windows['active_count']} hold={windows['hold_count']}" +) +PY + +mkdir -p "$OUTPUT_DIR" +TRAIN_LOG="$OUTPUT_DIR/train.log" +cd "$PROJECT_ROOT" + +echo "G2 Action-Adapter-Only training" +echo "GPUs: $CUDA_VISIBLE_DEVICES" +echo "DreamZero base: $PRETRAINED_MODEL_PATH" +echo "G2 LoRA/action initialization: none (clean official-base start)" +echo "Output: $OUTPUT_DIR" + +"$PYTHON_BIN" -m torch.distributed.run \ + --nproc_per_node 4 \ + --standalone \ + groot/vla/experiment/experiment.py \ + report_to=wandb \ + data=dreamzero/g2_relative \ + wandb_project=dreamzero \ + train_architecture=action_only \ + num_frames=33 \ + action_horizon=24 \ + num_views=3 \ + model=dreamzero/vla \ + model/dreamzero/action_head=wan_flow_matching_action_tf \ + model/dreamzero/transform=dreamzero_cotrain \ + num_frame_per_block=2 \ + num_action_per_block=24 \ + num_state_per_block=1 \ + seed=42 \ + training_args.learning_rate=1e-5 \ + training_args.deepspeed="groot/vla/configs/deepspeed/zero2.json" \ + milestone_save_steps="[100,250,500,1000,2000,3000,4000]" \ + training_args.warmup_ratio=0.05 \ + output_dir="$OUTPUT_DIR" \ + per_device_train_batch_size=1 \ + max_steps="$MAX_STEPS" \ + weight_decay=1e-5 \ + save_total_limit=7 \ + upload_checkpoints=false \ + bf16=true \ + tf32=true \ + eval_bf16=true \ + dataloader_pin_memory=false \ + dataloader_num_workers=1 \ + image_resolution_width=320 \ + image_resolution_height=176 \ + save_lora_only=false \ + max_chunk_size=4 \ + frame_seqlen=880 \ + save_strategy=no \ + g2_data_root="$G2_DATA_ROOT" \ + dit_version="$WAN_CKPT_DIR" \ + text_encoder_pretrained_path="$WAN_CKPT_DIR/models_t5_umt5-xxl-enc-bf16.pth" \ + image_encoder_pretrained_path="$WAN_CKPT_DIR/models_clip_open-clip-xlm-roberta-large-vit-huge-14.pth" \ + vae_pretrained_path="$WAN_CKPT_DIR/Wan2.1_VAE.pth" \ + tokenizer_path="$TOKENIZER_DIR" \ + pretrained_model_path="$PRETRAINED_MODEL_PATH" \ + pretrained_lora_path=null \ + ++model_specific_transform.embodiment_tag_mapping.g2=33 \ + ++action_head_cfg.config.skip_component_loading=true \ + ++action_head_cfg.config.defer_lora_injection=false \ + action_head_cfg.config.action_only_training=true \ + action_head_cfg.config.tune_diffusion_model=false \ + action_head_cfg.config.dynamics_loss_weight=1.0 \ + action_head_cfg.config.action_loss_weight=1.0 \ + 2>&1 | tee "$TRAIN_LOG" + +echo "Completed G2 Action-Adapter-Only training: $OUTPUT_DIR" diff --git a/scripts/train/train_dreamzero_g2_g1_g7_action_dit_lora.sh b/scripts/train/train_dreamzero_g2_g1_g7_action_dit_lora.sh new file mode 100755 index 00000000..6c9702a3 --- /dev/null +++ b/scripts/train/train_dreamzero_g2_g1_g7_action_dit_lora.sh @@ -0,0 +1,155 @@ +#!/usr/bin/env bash +set -Eeuo pipefail + +PROJECT_ROOT=${PROJECT_ROOT:-/home/ubuntu/projects/wangk/dreamzero} +G2_DATA_ROOT=${G2_DATA_ROOT:-/data/training_data/teleop/g2/g2_mock_light_module_joint_gear_policy_gripper/train} +OUTPUT_DIR=${OUTPUT_DIR:-/data/wangk/checkpoints/dreamzero_g2_mock_action10_video1_2k_v1} +WAN_CKPT_DIR=${WAN_CKPT_DIR:-/data/wangk/checkpoints/Wan2.1-I2V-14B-480P} +TOKENIZER_DIR=${TOKENIZER_DIR:-/data/wangk/checkpoints/umt5-xxl} +PRETRAINED_MODEL_PATH=${PRETRAINED_MODEL_PATH:-/data/wangk/checkpoints/DreamZero-AgiBot} +PYTHON_BIN=${PYTHON_BIN:-/data/wangk/conda/envs/dreamzero/bin/python} + +GPU_IDS=${GPU_IDS:-0,1,2,3} +EXPECTED_EPISODES=${EXPECTED_EPISODES:-110} +MAX_STEPS=${MAX_STEPS:-2000} +SAVE_STEPS=${SAVE_STEPS:-500} +WANDB_MODE=${WANDB_MODE:-offline} +HYDRA_FULL_ERROR=${HYDRA_FULL_ERROR:-1} + +export G2_DATA_ROOT OUTPUT_DIR WAN_CKPT_DIR TOKENIZER_DIR +export PRETRAINED_MODEL_PATH GPU_IDS EXPECTED_EPISODES MAX_STEPS SAVE_STEPS +export WANDB_MODE HYDRA_FULL_ERROR CUDA_VISIBLE_DEVICES="$GPU_IDS" + +fail() { + echo "[ERROR] $*" >&2 + exit 1 +} + +[[ "$GPU_IDS" == "0,1,2,3" ]] || fail "This run is pinned to GPUs 0,1,2,3" +for path in "$PROJECT_ROOT" "$G2_DATA_ROOT" "$WAN_CKPT_DIR" \ + "$TOKENIZER_DIR" "$PRETRAINED_MODEL_PATH"; do + [[ -d "$path" ]] || fail "Missing directory: $path" +done +[[ -x "$PYTHON_BIN" ]] || fail "Missing Python interpreter: $PYTHON_BIN" + +"$PYTHON_BIN" - <<'PY' +import json +import os +from pathlib import Path + +root = Path(os.environ["G2_DATA_ROOT"]) +info = json.loads((root / "meta/info.json").read_text()) +expected = int(os.environ["EXPECTED_EPISODES"]) +if info["total_episodes"] != expected: + raise RuntimeError(f"Expected {expected} episodes, got {info['total_episodes']}") +if info["features"]["observation.state"]["shape"] != [16]: + raise RuntimeError("G2 state must be 16D") +if info["features"]["action"]["shape"] != [16]: + raise RuntimeError("G2 action must be 16D") +modality = json.loads((root / "meta/modality.json").read_text()) +expected_video = {"top_head", "hand_left", "hand_right"} +if set(modality["video"]) != expected_video: + raise RuntimeError(f"Unexpected video modalities: {sorted(modality['video'])}") +stats = json.loads((root / "meta/stats.json").read_text()) +for feature_name in ("observation.state", "action"): + values = stats[feature_name] + for index, key in ((7, "left_gripper_position"), (15, "right_gripper_position")): + if values["min"][index] < -1e-4 or values["max"][index] > 1.0001: + raise RuntimeError( + f"{feature_name} {key} is not policy-space [0,1]: " + f"min={values['min'][index]} max={values['max'][index]}" + ) +relative = json.loads((root / "meta/relative_stats_dreamzero.json").read_text()) +for key in ("left_joint_position", "right_joint_position"): + if key not in relative: + raise RuntimeError(f"Missing relative stats for {key}") +active_hold = json.loads((root / "meta/g2_active_hold_windows.json").read_text()) +if active_hold["action_horizon"] != 24: + raise RuntimeError("Active/hold index must use a 24-step horizon") +if active_hold["active_count"] <= 0 or active_hold["hold_count"] <= 0: + raise RuntimeError("Active/hold index is empty") +parquets = list(root.glob("data/chunk-*/*.parquet")) +videos = list(root.glob("videos/chunk-*/*/*.mp4")) +if len(parquets) != expected or len(videos) != expected * 3: + raise RuntimeError( + f"File count mismatch: parquets={len(parquets)} videos={len(videos)}" + ) +print( + "Validated mock G2 dataset: " + f"episodes={expected} frames={info['total_frames']} tasks={info['total_tasks']} " + f"videos={len(videos)} active={active_hold['active_count']} " + f"hold={active_hold['hold_count']}" +) +PY + +if [[ -d "$OUTPUT_DIR" ]] && [[ -n "$(find "$OUTPUT_DIR" -mindepth 1 -maxdepth 1 -print -quit)" ]]; then + fail "Output directory is not empty: $OUTPUT_DIR" +fi +mkdir -p "$OUTPUT_DIR" +TRAIN_LOG="$OUTPUT_DIR/train.log" +cd "$PROJECT_ROOT" + +echo "G2 mock action-focused shared-DiT LoRA training" +echo "GPUs: $CUDA_VISIBLE_DEVICES" +echo "Base: $PRETRAINED_MODEL_PATH" +echo "Dataset: $G2_DATA_ROOT" +echo "Output: $OUTPUT_DIR" + +"$PYTHON_BIN" -m torch.distributed.run \ + --nproc_per_node 4 \ + --standalone \ + groot/vla/experiment/experiment.py \ + report_to=wandb \ + data=dreamzero/g2_relative \ + wandb_project=dreamzero \ + train_architecture=lora \ + num_frames=33 \ + action_horizon=24 \ + num_views=3 \ + model=dreamzero/vla \ + model/dreamzero/action_head=wan_flow_matching_action_tf \ + model/dreamzero/transform=dreamzero_cotrain \ + num_frame_per_block=2 \ + num_action_per_block=24 \ + num_state_per_block=1 \ + seed=42 \ + training_args.learning_rate=1e-5 \ + training_args.deepspeed="groot/vla/configs/deepspeed/zero2.json" \ + save_steps="$SAVE_STEPS" \ + training_args.warmup_ratio=0.05 \ + output_dir="$OUTPUT_DIR" \ + per_device_train_batch_size=1 \ + max_steps="$MAX_STEPS" \ + weight_decay=1e-5 \ + save_total_limit=7 \ + upload_checkpoints=false \ + bf16=true \ + tf32=true \ + eval_bf16=true \ + dataloader_pin_memory=false \ + dataloader_num_workers=1 \ + image_resolution_width=320 \ + image_resolution_height=176 \ + save_lora_only=true \ + max_chunk_size=4 \ + frame_seqlen=880 \ + save_strategy=steps \ + g2_data_root="$G2_DATA_ROOT" \ + dit_version="$WAN_CKPT_DIR" \ + text_encoder_pretrained_path="$WAN_CKPT_DIR/models_t5_umt5-xxl-enc-bf16.pth" \ + image_encoder_pretrained_path="$WAN_CKPT_DIR/models_clip_open-clip-xlm-roberta-large-vit-huge-14.pth" \ + vae_pretrained_path="$WAN_CKPT_DIR/Wan2.1_VAE.pth" \ + tokenizer_path="$TOKENIZER_DIR" \ + pretrained_model_path="$PRETRAINED_MODEL_PATH" \ + pretrained_lora_path=null \ + ++model_specific_transform.embodiment_tag_mapping.g2=33 \ + ++action_head_cfg.config.skip_component_loading=true \ + ++action_head_cfg.config.defer_lora_injection=true \ + action_head_cfg.config.action_only_training=false \ + action_head_cfg.config.tune_diffusion_model=true \ + action_head_cfg.config.decouple_video_action_noise=true \ + action_head_cfg.config.dynamics_loss_weight=1.0 \ + action_head_cfg.config.action_loss_weight=10.0 \ + 2>&1 | tee "$TRAIN_LOG" + +echo "Completed G2 mock action-focused shared-DiT LoRA training" diff --git a/scripts/train/train_dreamzero_g2_joint_lora.sh b/scripts/train/train_dreamzero_g2_joint_lora.sh new file mode 100644 index 00000000..d8e4ee0e --- /dev/null +++ b/scripts/train/train_dreamzero_g2_joint_lora.sh @@ -0,0 +1,244 @@ +#!/usr/bin/env bash +set -Eeuo pipefail + +# DreamZero G2 joint-space LoRA training. +# Point G2_DATA_ROOT at the GEAR train split, never at its parent/test folder. + +PROJECT_ROOT=${PROJECT_ROOT:-/home/ubuntu/projects/wangk/dreamzero} +G2_DATA_ROOT=${G2_DATA_ROOT:-/data/training_data/teleop/g2/g2_tasks_g1_g7_joint_gear_subtask_v2/train} + +OUTPUT_DIR=${OUTPUT_DIR:-/data/wangk/checkpoints/dreamzero_g2_joint_subtask_lora_v2} +WAN_CKPT_DIR=${WAN_CKPT_DIR:-/data/wangk/checkpoints/Wan2.1-I2V-14B-480P} +TOKENIZER_DIR=${TOKENIZER_DIR:-/data/wangk/checkpoints/umt5-xxl} +PRETRAINED_MODEL_PATH=${PRETRAINED_MODEL_PATH:-/data/wangk/checkpoints/DreamZero-AgiBot} + +# Only use physical GPUs 4 and above. GPU_IDS is preferred, while an existing +# CUDA_VISIBLE_DEVICES remains supported. The torchrun process count is always +# derived from the resulting list (for example, 6,7 means exactly 2 workers). +GPU_IDS=${GPU_IDS:-${CUDA_VISIBLE_DEVICES:-4,5,6,7}} +EXPECTED_EPISODES=${EXPECTED_EPISODES:-1346} +MAX_STEPS=${MAX_STEPS:-3000} +SAVE_STEPS=${SAVE_STEPS:-500} +WANDB_MODE=${WANDB_MODE:-offline} +HYDRA_FULL_ERROR=${HYDRA_FULL_ERROR:-1} + +export G2_DATA_ROOT OUTPUT_DIR WAN_CKPT_DIR TOKENIZER_DIR +export PRETRAINED_MODEL_PATH GPU_IDS +export EXPECTED_EPISODES MAX_STEPS SAVE_STEPS WANDB_MODE HYDRA_FULL_ERROR + +# Pin every Python entry point to the currently activated conda environment. +# This prevents ~/.local/bin/torchrun from launching /usr/bin/python3. +PYTHON_BIN="${PYTHON_BIN:-$(command -v python)}" +[[ -x "$PYTHON_BIN" ]] || { + echo "[ERROR] Python executable not found: $PYTHON_BIN" >&2 + exit 1 +} +export PYTHON_BIN + +fail() { + echo "[ERROR] $*" >&2 + exit 1 +} + +require_dir() { + [[ -d "$1" ]] || fail "Missing directory: $1" +} + +require_file() { + [[ -f "$1" ]] || fail "Missing file: $1" +} + +trap 'echo "[ERROR] Failed at line $LINENO" >&2' ERR + +require_dir "$PROJECT_ROOT" +require_dir "$G2_DATA_ROOT" +require_dir "$WAN_CKPT_DIR" +require_dir "$TOKENIZER_DIR" +require_dir "$PRETRAINED_MODEL_PATH" +require_file "$PROJECT_ROOT/groot/vla/experiment/experiment.py" +require_file "$PROJECT_ROOT/groot/vla/configs/data/dreamzero/g2_relative.yaml" +require_file "$PROJECT_ROOT/groot/vla/configs/data/dreamzero/base_48_wan_fine_aug_relative.yaml" +require_file "$WAN_CKPT_DIR/models_t5_umt5-xxl-enc-bf16.pth" +require_file "$WAN_CKPT_DIR/models_clip_open-clip-xlm-roberta-large-vit-huge-14.pth" +require_file "$WAN_CKPT_DIR/Wan2.1_VAE.pth" +require_file "$G2_DATA_ROOT/meta/info.json" +require_file "$G2_DATA_ROOT/meta/modality.json" +require_file "$G2_DATA_ROOT/meta/embodiment.json" +require_file "$G2_DATA_ROOT/meta/stats.json" +require_file "$G2_DATA_ROOT/meta/relative_stats_dreamzero.json" + +"$PYTHON_BIN" - <<'PY' +import sys + +required = { + "hydra": "hydra-core", + "torch": "torch", + "omegaconf": "omegaconf", +} + +missing = [] +for module_name, package_name in required.items(): + try: + __import__(module_name) + except Exception: + missing.append(package_name) + +print("Python executable:", sys.executable) +print("Python version:", sys.version.replace("\n", " ")) + +if missing: + raise SystemExit( + "Missing packages in the active Python environment: " + + ", ".join(missing) + ) +PY + +# Keep GPU visibility and torchrun world size under one source of truth. This +# also prevents an inherited NUM_GPUS from launching more workers than visible +# devices. +IFS=',' read -r -a GPU_ARRAY <<< "$GPU_IDS" +(( ${#GPU_ARRAY[@]} > 0 )) || fail "GPU_IDS must contain at least one GPU index" + +declare -A SEEN_GPUS=() +for gpu in "${GPU_ARRAY[@]}"; do + [[ "$gpu" =~ ^[0-9]+$ ]] || fail "GPU_IDS must be comma-separated physical GPU indices; got: $GPU_IDS" + (( gpu >= 4 )) || fail "Refusing to use physical GPU $gpu: GPUs 0-3 are reserved" + [[ -z "${SEEN_GPUS[$gpu]+x}" ]] || fail "Duplicate GPU index in GPU_IDS: $gpu" + SEEN_GPUS[$gpu]=1 +done + +CUDA_VISIBLE_DEVICES=$(IFS=,; echo "${GPU_ARRAY[*]}") +NPROC_PER_NODE=${#GPU_ARRAY[@]} +export CUDA_VISIBLE_DEVICES + +# Fail before torchrun if any selected physical index does not exist. +for gpu in "${GPU_ARRAY[@]}"; do + nvidia-smi -i "$gpu" --query-gpu=index --format=csv,noheader >/dev/null \ + || fail "Physical GPU $gpu is unavailable" +done + +cd "$PROJECT_ROOT" + +echo "[1/2] Validating the G2 joint subtask training split" +"$PYTHON_BIN" - <<'PY' +import json +import os +from pathlib import Path + +root = Path(os.environ["G2_DATA_ROOT"]) +with (root / "meta/info.json").open() as handle: + info = json.load(handle) +with (root / "meta/embodiment.json").open() as handle: + embodiment = json.load(handle) + +episodes = int(info["total_episodes"]) +expected_episodes = int(os.environ["EXPECTED_EPISODES"]) +parquets = list(root.glob("data/chunk-*/*.parquet")) +videos = list(root.glob("videos/chunk-*/*/*.mp4")) + +if episodes != expected_episodes: + raise RuntimeError( + f"Expected {expected_episodes} training episodes, got {episodes}" + ) +if len(parquets) != episodes: + raise RuntimeError(f"Expected {episodes} parquets, got {len(parquets)}") +if len(videos) != episodes * 3: + raise RuntimeError(f"Expected {episodes * 3} videos, got {len(videos)}") +if info["features"]["observation.state"]["shape"] != [16]: + raise RuntimeError("observation.state must be 16-dimensional G2 joint state") +if info["features"]["action"]["shape"] != [16]: + raise RuntimeError("action must be 16-dimensional G2 joint action") +if embodiment.get("embodiment_tag") != "g2": + raise RuntimeError(f"Expected embodiment_tag=g2, got {embodiment}") + +video_features = { + key: feature + for key, feature in info["features"].items() + if feature.get("dtype") == "video" +} +expected_video_keys = { + "observation.images.top_head", + "observation.images.hand_left", + "observation.images.hand_right", +} +if set(video_features) != expected_video_keys: + raise RuntimeError(f"Unexpected video features: {sorted(video_features)}") +for key, feature in video_features.items(): + if feature["shape"] != [176, 320, 3]: + raise RuntimeError(f"Unexpected shape for {key}: {feature['shape']}") + +print( + f"Validation OK: {episodes} train episodes, " + f"{len(parquets)} parquets, {len(videos)} videos, G2 joint 16D" +) +PY + +mkdir -p "$OUTPUT_DIR" +TRAIN_LOG="$OUTPUT_DIR/train.log" + +echo "Physical GPUs selected: $CUDA_VISIBLE_DEVICES" +echo "Training processes: $NPROC_PER_NODE (one per selected GPU)" +nvidia-smi -i "$CUDA_VISIBLE_DEVICES" \ + --query-gpu=index,name,memory.total,memory.used,utilization.gpu \ + --format=csv + +echo "W&B mode: $WANDB_MODE" +echo "Checkpoint interval: every $SAVE_STEPS steps" +echo "[2/2] Starting ${MAX_STEPS}-step DreamZero G2 joint LoRA training" +CUDA_VISIBLE_DEVICES="$CUDA_VISIBLE_DEVICES" "$PYTHON_BIN" -m torch.distributed.run \ + --nproc_per_node "$NPROC_PER_NODE" \ + --standalone \ + groot/vla/experiment/experiment.py \ + report_to=wandb \ + data=dreamzero/g2_relative \ + wandb_project=dreamzero \ + train_architecture=lora \ + num_frames=33 \ + action_horizon=24 \ + num_views=3 \ + model=dreamzero/vla \ + model/dreamzero/action_head=wan_flow_matching_action_tf \ + model/dreamzero/transform=dreamzero_cotrain \ + num_frame_per_block=2 \ + num_action_per_block=24 \ + num_state_per_block=1 \ + seed=42 \ + training_args.learning_rate=1e-5 \ + training_args.deepspeed="groot/vla/configs/deepspeed/zero2.json" \ + save_steps="$SAVE_STEPS" \ + training_args.warmup_ratio=0.05 \ + output_dir="$OUTPUT_DIR" \ + per_device_train_batch_size=1 \ + max_steps="$MAX_STEPS" \ + weight_decay=1e-5 \ + save_total_limit=5 \ + upload_checkpoints=false \ + bf16=true \ + tf32=true \ + eval_bf16=true \ + dataloader_pin_memory=false \ + dataloader_num_workers=1 \ + image_resolution_width=320 \ + image_resolution_height=176 \ + save_lora_only=true \ + max_chunk_size=4 \ + mixture_dataset_cls=groot.vla.data.dataset.lerobot_sharded.ShardedLeRobotMixtureDataset.from_mixture_spec \ + single_dataset_cls=groot.vla.data.dataset.lerobot_sharded.ShardedLeRobotSubLangSingleActionChunkDatasetDROID \ + frame_seqlen=880 \ + save_strategy=steps \ + g2_data_root="$G2_DATA_ROOT" \ + dit_version="$WAN_CKPT_DIR" \ + text_encoder_pretrained_path="$WAN_CKPT_DIR/models_t5_umt5-xxl-enc-bf16.pth" \ + image_encoder_pretrained_path="$WAN_CKPT_DIR/models_clip_open-clip-xlm-roberta-large-vit-huge-14.pth" \ + vae_pretrained_path="$WAN_CKPT_DIR/Wan2.1_VAE.pth" \ + tokenizer_path="$TOKENIZER_DIR" \ + pretrained_model_path="$PRETRAINED_MODEL_PATH" \ + ++model_specific_transform.embodiment_tag_mapping.g2=26 \ + ++action_head_cfg.config.skip_component_loading=true \ + ++action_head_cfg.config.defer_lora_injection=true \ + 2>&1 | tee "$TRAIN_LOG" + +echo "Completed successfully" +echo "Dataset: $G2_DATA_ROOT" +echo "Training output: $OUTPUT_DIR" +echo "Training log: $TRAIN_LOG" \ No newline at end of file diff --git a/scripts/train/train_dreamzero_g2_motion_mask_lora.sh b/scripts/train/train_dreamzero_g2_motion_mask_lora.sh new file mode 100755 index 00000000..6ca3f64c --- /dev/null +++ b/scripts/train/train_dreamzero_g2_motion_mask_lora.sh @@ -0,0 +1,173 @@ +#!/usr/bin/env bash +set -Eeuo pipefail + +PROJECT_ROOT=${PROJECT_ROOT:-/home/ubuntu/projects/wangk/dreamzero} +G2_DATA_ROOT=${G2_DATA_ROOT:-/data/training_data/teleop/g2/g2_mock_light_module_joint_gear_policy_gripper/train} +OUTPUT_DIR=${OUTPUT_DIR:-/data/wangk/checkpoints/dreamzero_g2_full_lora8_action_video_v1_4k} +WAN_CKPT_DIR=${WAN_CKPT_DIR:-/data/wangk/checkpoints/Wan2.1-I2V-14B-480P} +TOKENIZER_DIR=${TOKENIZER_DIR:-/data/wangk/checkpoints/umt5-xxl} +PRETRAINED_MODEL_PATH=${PRETRAINED_MODEL_PATH:-/data/wangk/checkpoints/DreamZero-AgiBot} +PYTHON_BIN=${PYTHON_BIN:-/data/wangk/conda/envs/dreamzero/bin/python} + +GPU_IDS=${GPU_IDS:-0,1,2,3} +EXPECTED_EPISODES=${EXPECTED_EPISODES:-110} +MAX_STEPS=${MAX_STEPS:-4000} +SAVE_STEPS=${SAVE_STEPS:-500} +WANDB_MODE=${WANDB_MODE:-offline} +HYDRA_FULL_ERROR=${HYDRA_FULL_ERROR:-1} + +export G2_DATA_ROOT OUTPUT_DIR WAN_CKPT_DIR TOKENIZER_DIR +export PRETRAINED_MODEL_PATH GPU_IDS EXPECTED_EPISODES MAX_STEPS SAVE_STEPS +export WANDB_MODE HYDRA_FULL_ERROR CUDA_VISIBLE_DEVICES="$GPU_IDS" + +fail() { + echo "[ERROR] $*" >&2 + exit 1 +} + +[[ "$GPU_IDS" == "0,1,2,3" ]] || fail "This run is pinned to GPUs 0,1,2,3" +for path in "$PROJECT_ROOT" "$G2_DATA_ROOT" "$WAN_CKPT_DIR" \ + "$TOKENIZER_DIR" "$PRETRAINED_MODEL_PATH"; do + [[ -d "$path" ]] || fail "Missing directory: $path" +done +[[ -x "$PYTHON_BIN" ]] || fail "Missing Python interpreter: $PYTHON_BIN" + +"$PYTHON_BIN" - <<'PY' +import json +import os +from pathlib import Path + +root = Path(os.environ["G2_DATA_ROOT"]) +info = json.loads((root / "meta/info.json").read_text()) +expected = int(os.environ["EXPECTED_EPISODES"]) +if info["total_episodes"] != expected: + raise RuntimeError(f"Expected {expected} episodes, got {info['total_episodes']}") +if info["features"]["observation.state"]["shape"] != [16]: + raise RuntimeError("G2 state must be 16D") +if info["features"]["action"]["shape"] != [16]: + raise RuntimeError("G2 action must be 16D") +modality = json.loads((root / "meta/modality.json").read_text()) +expected_video = {"top_head", "hand_left", "hand_right"} +if set(modality["video"]) != expected_video: + raise RuntimeError(f"Unexpected video modalities: {sorted(modality['video'])}") +stats = json.loads((root / "meta/stats.json").read_text()) +for feature_name in ("observation.state", "action"): + values = stats[feature_name] + for index, key in ((7, "left_gripper_position"), (15, "right_gripper_position")): + if values["min"][index] < -1e-4 or values["max"][index] > 1.0001: + raise RuntimeError( + f"{feature_name} {key} is not policy-space [0,1]: " + f"min={values['min'][index]} max={values['max'][index]}" + ) +relative = json.loads((root / "meta/relative_stats_dreamzero.json").read_text()) +for key in ("left_joint_position", "right_joint_position"): + if key not in relative: + raise RuntimeError(f"Missing relative stats for {key}") + if len(relative[key]["q01"]) != 7 or len(relative[key]["q99"]) != 7: + raise RuntimeError(f"Motion-mask stats for {key} must be 7D") +active_hold = json.loads((root / "meta/g2_active_hold_windows.json").read_text()) +if active_hold["action_horizon"] != 24: + raise RuntimeError("Active/hold index must use a 24-step horizon") +if active_hold["active_count"] <= 0 or active_hold["hold_count"] <= 0: + raise RuntimeError("Active/hold index is empty") +tasks_path = root / "meta/tasks.jsonl" +tasks = [json.loads(line) for line in tasks_path.read_text().splitlines() if line.strip()] +if len(tasks) != 1 or not str(tasks[0].get("task", "")).strip(): + raise RuntimeError("Expected exactly one non-empty G2 task prompt") +print(f"Prompt contract: {tasks[0]['task']}") +parquets = list(root.glob("data/chunk-*/*.parquet")) +videos = list(root.glob("videos/chunk-*/*/*.mp4")) +if len(parquets) != expected or len(videos) != expected * 3: + raise RuntimeError( + f"File count mismatch: parquets={len(parquets)} videos={len(videos)}" + ) +print( + "Validated mock G2 dataset: " + f"episodes={expected} frames={info['total_frames']} tasks={info['total_tasks']} " + f"videos={len(videos)} active={active_hold['active_count']} " + f"hold={active_hold['hold_count']}" +) +PY + +if [[ -d "$OUTPUT_DIR" ]] && [[ -n "$(find "$OUTPUT_DIR" -mindepth 1 -maxdepth 1 -print -quit)" ]]; then + fail "Output directory is not empty: $OUTPUT_DIR" +fi +mkdir -p "$OUTPUT_DIR" +TRAIN_LOG="$OUTPUT_DIR/train.log" +cd "$PROJECT_ROOT" + +echo "G2 mock shared-DiT LoRA training with full video+action objective" +echo "GPUs: $CUDA_VISIBLE_DEVICES" +echo "Base: $PRETRAINED_MODEL_PATH" +echo "Dataset: $G2_DATA_ROOT" +echo "Output: $OUTPUT_DIR" +echo "Objective: LoRA rank8 + flow action + decoded-action + endpoint + motion mask" + +"$PYTHON_BIN" -m torch.distributed.run \ + --nproc_per_node 4 \ + --standalone \ + groot/vla/experiment/experiment.py \ + report_to=wandb \ + data=dreamzero/g2_relative \ + wandb_project=dreamzero \ + train_architecture=lora \ + num_frames=33 \ + action_horizon=24 \ + num_views=3 \ + model=dreamzero/vla \ + model/dreamzero/action_head=wan_flow_matching_action_tf \ + model/dreamzero/transform=dreamzero_cotrain \ + num_frame_per_block=2 \ + num_action_per_block=24 \ + num_state_per_block=1 \ + seed=42 \ + training_args.learning_rate=1e-5 \ + training_args.deepspeed="groot/vla/configs/deepspeed/zero2.json" \ + save_steps="$SAVE_STEPS" \ + training_args.warmup_ratio=0.05 \ + output_dir="$OUTPUT_DIR" \ + per_device_train_batch_size=1 \ + max_steps="$MAX_STEPS" \ + weight_decay=1e-5 \ + save_total_limit=9 \ + upload_checkpoints=false \ + bf16=true \ + tf32=true \ + eval_bf16=true \ + dataloader_pin_memory=false \ + dataloader_num_workers=1 \ + image_resolution_width=320 \ + image_resolution_height=176 \ + save_lora_only=true \ + max_chunk_size=4 \ + frame_seqlen=880 \ + save_strategy=steps \ + active_window_ratio=0.9 \ + g2_data_root="$G2_DATA_ROOT" \ + dit_version="$WAN_CKPT_DIR" \ + text_encoder_pretrained_path="$WAN_CKPT_DIR/models_t5_umt5-xxl-enc-bf16.pth" \ + image_encoder_pretrained_path="$WAN_CKPT_DIR/models_clip_open-clip-xlm-roberta-large-vit-huge-14.pth" \ + vae_pretrained_path="$WAN_CKPT_DIR/Wan2.1_VAE.pth" \ + tokenizer_path="$TOKENIZER_DIR" \ + pretrained_model_path="$PRETRAINED_MODEL_PATH" \ + pretrained_lora_path=null \ + ++model_specific_transform.embodiment_tag_mapping.g2=33 \ + ++model_specific_transform.always_use_default_instruction=false \ + language_dropout_prob=0.0 \ + ++action_head_cfg.config.skip_component_loading=true \ + ++action_head_cfg.config.defer_lora_injection=true \ + action_head_cfg.config.action_only_training=false \ + action_head_cfg.config.tune_diffusion_model=true \ + ++action_head_cfg.config.lora_rank=8 \ + ++action_head_cfg.config.lora_alpha=8 \ + action_head_cfg.config.decouple_video_action_noise=true \ + action_head_cfg.config.dynamics_loss_weight=1.0 \ + action_head_cfg.config.action_loss_weight=10.0 \ + ++action_head_cfg.config.action_x0_loss_weight=1.0 \ + ++action_head_cfg.config.action_endpoint_loss_weight=0.5 \ + ++action_head_cfg.config.motion_mask_enabled=true \ + ++action_head_cfg.config.motion_mask_threshold_rad=0.03 \ + ++action_head_cfg.config.motion_mask_stats_path="$G2_DATA_ROOT/meta/relative_stats_dreamzero.json" \ + 2>&1 | tee "$TRAIN_LOG" + +echo "Completed G2 mock stationary-frame-masked LoRA training" diff --git a/scripts/train/train_dreamzero_g2_video_only_lora.sh b/scripts/train/train_dreamzero_g2_video_only_lora.sh new file mode 100755 index 00000000..944a9a62 --- /dev/null +++ b/scripts/train/train_dreamzero_g2_video_only_lora.sh @@ -0,0 +1,93 @@ +#!/usr/bin/env bash +set -Eeuo pipefail + +PROJECT_ROOT=/home/ubuntu/projects/wangk/dreamzero +G2_DATA_ROOT=/data/training_data/teleop/g2/g2_tasks_g1_g7_joint_gear_subtask_v2/train +WAN_CKPT_DIR=/data/wangk/checkpoints/Wan2.1-I2V-14B-480P +TOKENIZER_DIR=/data/wangk/checkpoints/umt5-xxl +PRETRAINED_MODEL_PATH=/data/wangk/checkpoints/DreamZero-AgiBot +OUTPUT_DIR=/data/wangk/checkpoints/dreamzero_g2_video_only_lora_1k +RUN_NAME=g2_video_only_wan_lora_1k + +MAX_STEPS=1000 +SAVE_STEPS=250 + +# 用户使用物理 GPU 0-3;物理 GPU 4-7 留给同事。 +export CUDA_VISIBLE_DEVICES=0,1,2,3 +export WANDB_MODE=offline +export HYDRA_FULL_ERROR=1 + +# 始终使用已激活的 dreamzero Conda 环境,避免命中 ~/.local/bin/torchrun。 +PYTHON_BIN="${CONDA_PREFIX}/bin/python" + +cd "$PROJECT_ROOT" +mkdir -p "$OUTPUT_DIR" +TRAIN_LOG="$OUTPUT_DIR/train.log" + +echo "Python: $PYTHON_BIN" +echo "Physical GPUs: $CUDA_VISIBLE_DEVICES" +echo "Training processes: 4" +echo "Dataset: $G2_DATA_ROOT" +echo "Run name: $RUN_NAME" +echo "Schedule: $MAX_STEPS steps, save every $SAVE_STEPS steps" +nvidia-smi -i "$CUDA_VISIBLE_DEVICES" \ + --query-gpu=index,name,memory.total,memory.used,utilization.gpu \ + --format=csv + +echo "Starting DreamZero G2 video-only Wan LoRA training" +CUDA_VISIBLE_DEVICES="$CUDA_VISIBLE_DEVICES" "$PYTHON_BIN" -m torch.distributed.run \ + --nproc_per_node 4 \ + --standalone \ + groot/vla/experiment/experiment.py \ + report_to=wandb \ + +run_name="$RUN_NAME" \ + data=dreamzero/g2_relative \ + wandb_project=dreamzero \ + train_architecture=lora \ + num_frames=33 \ + action_horizon=24 \ + num_views=3 \ + model=dreamzero/vla \ + model/dreamzero/action_head=wan_flow_matching_action_tf \ + model/dreamzero/transform=dreamzero_cotrain \ + num_frame_per_block=2 \ + num_action_per_block=24 \ + num_state_per_block=1 \ + seed=42 \ + training_args.learning_rate=1e-5 \ + training_args.deepspeed="groot/vla/configs/deepspeed/zero2.json" \ + save_steps="$SAVE_STEPS" \ + training_args.warmup_ratio=0.05 \ + output_dir="$OUTPUT_DIR" \ + per_device_train_batch_size=1 \ + max_steps="$MAX_STEPS" \ + weight_decay=1e-5 \ + save_total_limit=5 \ + upload_checkpoints=false \ + bf16=true \ + tf32=true \ + eval_bf16=true \ + dataloader_pin_memory=false \ + dataloader_num_workers=1 \ + image_resolution_width=320 \ + image_resolution_height=176 \ + save_lora_only=true \ + max_chunk_size=4 \ + frame_seqlen=880 \ + save_strategy=steps \ + g2_data_root="$G2_DATA_ROOT" \ + dit_version="$WAN_CKPT_DIR" \ + text_encoder_pretrained_path="$WAN_CKPT_DIR/models_t5_umt5-xxl-enc-bf16.pth" \ + image_encoder_pretrained_path="$WAN_CKPT_DIR/models_clip_open-clip-xlm-roberta-large-vit-huge-14.pth" \ + vae_pretrained_path="$WAN_CKPT_DIR/Wan2.1_VAE.pth" \ + tokenizer_path="$TOKENIZER_DIR" \ + pretrained_model_path="$PRETRAINED_MODEL_PATH" \ + ++model_specific_transform.embodiment_tag_mapping.g2=26 \ + ++action_head_cfg.config.skip_component_loading=true \ + ++action_head_cfg.config.defer_lora_injection=true \ + ++action_head_cfg.config.video_only_training=true \ + 2>&1 | tee "$TRAIN_LOG" + +echo "Completed successfully" +echo "Training output: $OUTPUT_DIR" +echo "Training log: $TRAIN_LOG" \ No newline at end of file diff --git a/socket_optimized_AR_g2.py b/socket_optimized_AR_g2.py new file mode 100644 index 00000000..a476d630 --- /dev/null +++ b/socket_optimized_AR_g2.py @@ -0,0 +1,1248 @@ +import dataclasses +import logging +import socket +import asyncio +import os +import http +import logging +import time +import traceback +import torch +import tyro +from einops import rearrange +import datetime +import cv2 + +from groot.vla.model.n1_5.sim_policy import GrootSimPolicy +from groot.vla.data.schema import EmbodimentTag +import imageio +import numpy as np + +from openpi_client import base_policy as _base_policy +from openpi_client import msgpack_numpy +import websockets.asyncio.server as _server +import websockets.frames +from tianshou.data import Batch +import torch.distributed as dist +from torch.distributed.device_mesh import DeviceMesh, init_device_mesh + +# Use roboarena policy server interface +from eval_utils.policy_server import WebsocketPolicyServer as RoboarenaServer +from eval_utils.policy_server import PolicyServerConfig + +logger = logging.getLogger(__name__) + +SIGNAL_INFER = 0 +SIGNAL_SHUTDOWN = 1 +SIGNAL_IDLE = 2 +SIGNAL_RESET_CACHE = 3 + + +def _run_policy_forward( + policy: GrootSimPolicy, + batch: Batch, +) -> tuple[object, torch.Tensor]: + """Select causal production or ordinary joint inference for A/B audits.""" + mode = os.environ.get( + "DREAMZERO_FORWARD_MODE", "causal" + ).strip().lower() + if mode == "causal": + return policy.lazy_joint_forward_causal(batch) + if mode == "joint": + return policy.lazy_joint_forward(batch) + raise ValueError( + "DREAMZERO_FORWARD_MODE must be 'causal' or 'joint', " + f"got {mode!r}" + ) + + +def _reset_policy_inference_cache(policy: object, reason: str) -> None: + trained_model = getattr(policy, "trained_model", None) + action_head = getattr(trained_model, "action_head", None) + reset_fn = getattr(action_head, "reset_inference_cache", None) + if callable(reset_fn): + reset_fn() + logger.info("Reset action-head inference cache on rank %s (%s)", dist.get_rank() if dist.is_initialized() else "?", reason) + else: + logger.warning("Policy action head does not expose reset_inference_cache(); cache reset skipped (%s)", reason) + +@dataclasses.dataclass +class Args: + port: int = 8000 + timeout_seconds: int = 50000 # 10 hours default, configurable + model_path: str = "./checkpoints/dreamzero" + wan_ckpt_dir: str | None = None # Override Wan2.1 component paths stored in config.json. + tokenizer_path: str | None = None # Override tokenizer_path stored in experiment_cfg/conf.yaml. + enable_dit_cache: bool = False # Backward-compatible alias for num_dit_steps=8. + num_dit_steps: int | None = None # Actual DiT compute steps. Supported fast masks: 5, 6, 7, 8. None keeps model default. + num_inference_timesteps: int | None = None # Positive values override diffusion steps; 0 keeps checkpoint default. + index: int = 0 + embodiment_tag: str = "oxe_droid" + max_chunk_size: int | None = None # If None, use config value. Otherwise override max_chunk_size for inference. + video_save_mode: str = "first" # one of: none, first, full. Controls generated video saved on reset/client close. + output_dir: str = "/data/wangk/dreamzero/video_rollout" + reset_cache_each_request: bool = True + + +class DistributedRoboarenaPolicyBase: + """Shared distributed inference plumbing for websocket policy wrappers.""" + + def __init__( + self, + groot_policy: GrootSimPolicy, + signal_group: dist.ProcessGroup, + output_dir: str | None = None, + video_save_mode: str = "first", + ) -> None: + self._policy = groot_policy + self._signal_group = signal_group + self._output_dir = output_dir + self._video_save_mode = video_save_mode + self._frame_buffers = self._init_frame_buffers() + self._current_session_id: str | None = None + self.video_across_time = [] + self._msg_index = 0 + + if self._output_dir: + os.makedirs(self._output_dir, exist_ok=True) + + def _init_frame_buffers(self) -> dict[str, list[np.ndarray]]: + return { + "video.top_head": [], + "video.hand_left": [], + "video.hand_right": [], + } + + def _reset_custom_state(self) -> None: + pass + + def _after_infer(self) -> None: + pass + + def _prepare_video_chunk(self, video_pred: torch.Tensor) -> torch.Tensor | None: + if self._video_save_mode == "none": + return None + return video_pred + + def _video_save_fps(self) -> int: + return 5 + + def _convert_observation(self, obs: dict) -> dict: + raise NotImplementedError + + def _convert_action(self, action_dict: dict) -> np.ndarray: + raise NotImplementedError + + def _broadcast_batch_to_workers(self, obs: dict) -> None: + import pickle + + serialized = pickle.dumps(obs) + data_size = len(serialized) + + size_tensor = torch.tensor([data_size], dtype=torch.int64, device='cuda') + dist.broadcast(size_tensor, src=0) + + data_tensor = torch.frombuffer(serialized, dtype=torch.uint8).clone().cuda() + dist.broadcast(data_tensor, src=0) + + def _extract_action_dict(self, action_chunk_dict: object) -> dict[str, object]: + action_dict: dict[str, object] = {} + for key in dir(action_chunk_dict): + if key.startswith('action.'): + action_dict[key] = getattr(action_chunk_dict, key) + return action_dict + + def _broadcast_signal_to_workers(self, signal: int) -> None: + signal_tensor = torch.tensor([signal], dtype=torch.int32, device='cpu') + dist.broadcast(signal_tensor, src=0, group=self._signal_group) + + def infer(self, obs: dict) -> np.ndarray: + session_id = obs.get('session_id') + if session_id is not None and session_id != self._current_session_id: + if self._current_session_id is not None: + logger.info("Session changed from '%s' to '%s', resetting state", self._current_session_id, session_id) + self._broadcast_signal_to_workers(SIGNAL_RESET_CACHE) + self._reset_state() + else: + logger.info("New session started: '%s'", session_id) + self._current_session_id = session_id + + if os.environ.get( + "DREAMZERO_RESET_AR_EACH_REQUEST", "true" + ).strip().lower() in {"1", "true", "yes", "on"}: + self._broadcast_signal_to_workers(SIGNAL_RESET_CACHE) + _reset_policy_inference_cache( + self._policy, + "fresh real observation before request", + ) + + self._msg_index += 1 + converted_obs = self._convert_observation(obs) + + self._broadcast_signal_to_workers(SIGNAL_INFER) + self._broadcast_batch_to_workers(converted_obs) + + batch = Batch(obs=converted_obs) + dist.barrier() + with torch.no_grad(): + result_batch, video_pred = _run_policy_forward( + self._policy, batch + ) + dist.barrier() + + video_chunk = self._prepare_video_chunk(video_pred) + if video_chunk is not None: + self.video_across_time.append(video_chunk.detach().cpu()) + action = self._convert_action(self._extract_action_dict(result_batch.act)) + self._after_infer() + return action + + def _reset_state(self, save_video: bool = True) -> None: + if save_video and len(self.video_across_time) > 0 and self._output_dir: + try: + frame_list = [] + action_head = self._policy.trained_model.action_head + device = getattr(action_head, "_device", None) + if device is None: + device = next(self._policy.trained_model.parameters()).device + video_across_time_cat = torch.cat(self.video_across_time, dim=2).to(device=device, dtype=torch.bfloat16) + frames = action_head.vae.decode( + video_across_time_cat, + tiled=action_head.tiled, + tile_size=(action_head.tile_size_height, action_head.tile_size_width), + tile_stride=(action_head.tile_stride_height, action_head.tile_stride_width), + ) + frames = rearrange(frames, 'B C T H W -> B T H W C') + frames = frames[0] + frames = ((frames.float() + 1) * 127.5).clip(0, 255).cpu().numpy().astype(np.uint8) + for frame in frames: + frame_list.append(frame) + + if frame_list: + sample_frame = frame_list[0] + if len(sample_frame.shape) == 3 and sample_frame.shape[2] in [1, 3, 4]: + save_dir = self._output_dir + os.makedirs(save_dir, exist_ok=True) + all_mp4_files = [f for f in os.listdir(save_dir) if f.endswith('.mp4')] + timestamp = datetime.datetime.now().strftime('%Y%m%d_%H%M%S') + num_frames = len(frame_list) + output_path = os.path.join(save_dir, f'{timestamp}_{len(all_mp4_files):06}_f{num_frames}.mp4') + imageio.mimsave(output_path, frame_list, fps=self._video_save_fps(), codec='libx264') + logger.info('Saved video on reset to: %s', output_path) + except Exception as exc: + logger.warning('Failed to save video on reset: %s', exc) + + for key in self._frame_buffers: + self._frame_buffers[key] = [] + + self.video_across_time = [] + _reset_policy_inference_cache(self._policy, "wrapper reset_state") + self._reset_custom_state() + + def reset(self, reset_info: dict) -> None: + self._broadcast_signal_to_workers(SIGNAL_RESET_CACHE) + self._reset_state(save_video=True) + + +class ARDroidRoboarenaPolicy(DistributedRoboarenaPolicyBase): + """Wrapper policy that implements roboarena.policy.BasePolicy interface for AR_droid.""" + + FRAMES_PER_CHUNK = 4 + + def __init__( + self, + groot_policy: GrootSimPolicy, + signal_group: dist.ProcessGroup, + output_dir: str | None = None, + video_save_mode: str = "first", + ) -> None: + super().__init__( + groot_policy=groot_policy, + signal_group=signal_group, + output_dir=output_dir, + video_save_mode=video_save_mode, + ) + self._reset_custom_state() + + def _init_frame_buffers(self) -> dict[str, list[np.ndarray]]: + return { + 'video.exterior_image_1_left': [], + 'video.exterior_image_2_left': [], + 'video.wrist_image_left': [], + } + + def _reset_custom_state(self) -> None: + self._is_first_call = True + + def _after_infer(self) -> None: + self._is_first_call = False + + def _convert_observation(self, obs: dict) -> dict: + converted = {} + image_key_mapping = { + 'observation/exterior_image_0_left': 'video.exterior_image_1_left', + 'observation/exterior_image_1_left': 'video.exterior_image_2_left', + 'observation/wrist_image_left': 'video.wrist_image_left', + } + + for roboarena_key, droid_key in image_key_mapping.items(): + if roboarena_key in obs: + data = obs[roboarena_key] + if isinstance(data, np.ndarray): + if data.ndim == 4: + self._frame_buffers[droid_key].extend(list(data)) + else: + self._frame_buffers[droid_key].append(data) + + num_frames = 1 if self._is_first_call else self.FRAMES_PER_CHUNK + + for droid_key, buffer in self._frame_buffers.items(): + if len(buffer) > 0: + if len(buffer) >= num_frames: + frames_to_use = buffer[-num_frames:] + else: + frames_to_use = buffer.copy() + while len(frames_to_use) < num_frames: + frames_to_use.insert(0, buffer[0]) + converted[droid_key] = np.stack(frames_to_use, axis=0) + + joint_pos = obs.get('observation/joint_position', np.zeros(7, dtype=np.float32)) + if joint_pos.ndim == 1: + joint_pos = joint_pos.reshape(1, -1) + converted['state.joint_position'] = joint_pos.astype(np.float64) + + gripper_pos = obs.get('observation/gripper_position', np.zeros(1, dtype=np.float32)) + if gripper_pos.ndim == 1: + gripper_pos = gripper_pos.reshape(1, -1) + converted['state.gripper_position'] = gripper_pos.astype(np.float64) + converted['annotation.language.action_text'] = obs.get('prompt', '') + return converted + + def _convert_action(self, action_dict: dict) -> np.ndarray: + joint_action = None + gripper_action = None + for key, value in action_dict.items(): + if 'joint_position' in key: + joint_action = value + elif 'gripper_position' in key or 'gripper' in key: + gripper_action = value + + if joint_action is None: + return np.zeros((1, 8), dtype=np.float32) + + if isinstance(joint_action, torch.Tensor): + joint_action = joint_action.cpu().numpy() + if joint_action.ndim == 1: + joint_action = joint_action.reshape(1, -1) + + num_steps = joint_action.shape[0] + if gripper_action is not None: + if isinstance(gripper_action, torch.Tensor): + gripper_action = gripper_action.cpu().numpy() + if gripper_action.ndim == 1: + gripper_action = gripper_action.reshape(-1, 1) + elif gripper_action.ndim == 0: + gripper_action = gripper_action.reshape(1, 1) + else: + gripper_action = np.zeros((num_steps, 1), dtype=np.float32) + + return np.concatenate([joint_action, gripper_action], axis=-1).astype(np.float32) + + +class AgiBotRoboarenaPolicy(DistributedRoboarenaPolicyBase): + """Adapter that converts websocket observations into AgiBot modality keys.""" + + VIDEO_KEY_MAPPING = { + 'observation/top_head': 'video.top_head', + 'observation/hand_left': 'video.hand_left', + 'observation/hand_right': 'video.hand_right', + } + STATE_KEY_MAPPING = { + 'observation/left_arm_joint_position': 'state.left_arm_joint_position', + 'observation/right_arm_joint_position': 'state.right_arm_joint_position', + 'observation/left_effector_position': 'state.left_effector_position', + 'observation/right_effector_position': 'state.right_effector_position', + 'observation/head_position': 'state.head_position', + 'observation/waist_pitch': 'state.waist_pitch', + 'observation/waist_lift': 'state.waist_lift', + } + + def __init__( + self, + groot_policy: GrootSimPolicy, + signal_group: dist.ProcessGroup, + output_dir: str | None = None, + video_save_mode: str = "first", + ) -> None: + super().__init__( + groot_policy=groot_policy, + signal_group=signal_group, + output_dir=output_dir, + video_save_mode=video_save_mode, + ) + self._action_keys = list(self._policy.modality_configs.action.modality_keys) + + def _lookup_obs_value(self, obs: dict, source_key: str, target_key: str) -> object: + if source_key in obs: + return obs[source_key] + return obs.get(target_key) + + def _normalize_video(self, value: object, target_key: str) -> np.ndarray: + if isinstance(value, dict) and value.get("__dreamzero_image_encoding__") == "jpeg_sequence": + frames = [] + expected_shape = tuple(value.get("shape", ())) + expected_dtype = np.dtype(value.get("dtype", "uint8")) + for index, frame_bytes in enumerate(value.get("frames", [])): + encoded = np.frombuffer(frame_bytes, dtype=np.uint8) + frame = cv2.imdecode(encoded, cv2.IMREAD_COLOR) + if frame is None: + raise ValueError(f"Failed to decode JPEG frame {index} for {target_key}") + frames.append(frame.astype(expected_dtype, copy=False)) + array = np.stack(frames, axis=0) + if expected_shape and tuple(array.shape) != expected_shape: + raise ValueError( + f"Decoded JPEG video for {target_key} has shape {array.shape}, expected {expected_shape}" + ) + return array + + array = np.asarray(value) + if array.ndim == 3: + return np.expand_dims(array, axis=0) + if array.ndim == 4: + return array + raise ValueError(f'AgiBot video input for {target_key} must have shape (H, W, C) or (T, H, W, C), got {array.shape}') + + def _normalize_state(self, value: object, target_key: str) -> np.ndarray: + array = np.asarray(value) + if array.ndim == 0: + return array.reshape(1, 1).astype(np.float64) + if array.ndim == 1: + return array.reshape(1, -1).astype(np.float64) + if array.ndim == 2: + return array.astype(np.float64) + raise ValueError(f'AgiBot state input for {target_key} must be 1D or 2D, got {array.shape}') + + def _prepare_video_chunk(self, video_pred: torch.Tensor) -> torch.Tensor | None: + if self._video_save_mode == "none": + return None + if video_pred.ndim != 5: + raise ValueError(f'AgiBot video prediction must be 5D (B, C, T, H, W), got {tuple(video_pred.shape)}') + if self._video_save_mode == "first": + return video_pred[:, :, :1].contiguous() + if self._video_save_mode == "full": + return video_pred.contiguous() + raise ValueError(f"Unsupported video_save_mode: {self._video_save_mode!r}; expected none, first, or full") + + def _video_save_fps(self) -> int: + return 20 + + def _convert_observation(self, obs: dict) -> dict: + converted = {} + missing_keys: list[str] = [] + + for source_key, target_key in self.VIDEO_KEY_MAPPING.items(): + value = self._lookup_obs_value(obs, source_key, target_key) + if value is None: + missing_keys.append(source_key) + continue + converted[target_key] = self._normalize_video(value, target_key) + + for source_key, target_key in self.STATE_KEY_MAPPING.items(): + value = self._lookup_obs_value(obs, source_key, target_key) + if value is None: + missing_keys.append(source_key) + continue + converted[target_key] = self._normalize_state(value, target_key) + + if missing_keys: + raise ValueError( + 'AgiBot inference requires the following observation keys: ' + + ', '.join(sorted(missing_keys)) + ) + + converted['annotation.language.action_text'] = obs.get('prompt', obs.get('annotation.language.action_text', '')) + return converted + + def _convert_action(self, action_dict: dict) -> np.ndarray: + flattened_chunks: list[np.ndarray] = [] + expected_horizon: int | None = None + missing_keys = [key for key in self._action_keys if key not in action_dict] + if missing_keys: + raise RuntimeError('Missing AgiBot action outputs: ' + ', '.join(missing_keys)) + + for action_key in self._action_keys: + value = action_dict[action_key] + if isinstance(value, torch.Tensor): + value = value.detach().cpu().numpy() + array = np.asarray(value) + if array.ndim == 0: + array = array.reshape(1, 1) + elif array.ndim == 1: + array = array.reshape(-1, 1) + else: + array = array.reshape(array.shape[0], -1) + + if expected_horizon is None: + expected_horizon = array.shape[0] + elif array.shape[0] != expected_horizon: + raise RuntimeError( + f'Inconsistent AgiBot action horizon for {action_key}: expected {expected_horizon}, got {array.shape[0]}' + ) + flattened_chunks.append(array.astype(np.float32)) + + return np.concatenate(flattened_chunks, axis=-1).astype(np.float32) + + +class G2RoboarenaPolicy(DistributedRoboarenaPolicyBase): + """Adapter for the G2 dual-arm joint-space policy.""" + FRAMES_PER_CHUNK = 4 + VIDEO_KEY_MAPPING = { + 'observation/top_head': 'video.top_head', + 'observation/hand_left': 'video.hand_left', + 'observation/hand_right': 'video.hand_right', + } + + STATE_KEY_MAPPING = { + 'observation/left_joint_position': 'state.left_joint_position', + 'observation/left_gripper_position': 'state.left_gripper_position', + 'observation/right_joint_position': 'state.right_joint_position', + 'observation/right_gripper_position': 'state.right_gripper_position', + } + + PACKED_STATE_KEYS = ( + 'observation/state', + 'observation.state', + 'state', + ) + + def __init__( + self, + groot_policy: GrootSimPolicy, + signal_group: dist.ProcessGroup, + output_dir: str | None = None, + video_save_mode: str = "first", + ) -> None: + super().__init__( + groot_policy=groot_policy, + signal_group=signal_group, + output_dir=output_dir, + video_save_mode=video_save_mode, + ) + self._action_keys = list( + self._policy.modality_configs.action.modality_keys + ) + + @staticmethod + def _lookup_obs_value( + obs: dict, + source_key: str, + target_key: str, + ) -> object: + if source_key in obs: + return obs[source_key] + return obs.get(target_key) + + @staticmethod + def _normalize_video( + value: object, + target_key: str, + ) -> np.ndarray: + if ( + isinstance(value, dict) + and value.get("__dreamzero_image_encoding__") + == "jpeg_sequence" + ): + frames = [] + expected_shape = tuple(value.get("shape", ())) + expected_dtype = np.dtype( + value.get("dtype", "uint8") + ) + for index, frame_bytes in enumerate( + value.get("frames", []) + ): + encoded = np.frombuffer( + frame_bytes, + dtype=np.uint8, + ) + frame = cv2.imdecode( + encoded, + cv2.IMREAD_COLOR, + ) + if frame is None: + raise ValueError( + f"Failed to decode JPEG frame {index} " + f"for {target_key}" + ) + # cv2.imdecode always returns BGR, while DreamZero training + # videos are decoded as RGB. Keep the model input contract RGB. + frame = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB) + frames.append( + frame.astype( + expected_dtype, + copy=False, + ) + ) + if not frames: + raise ValueError( + f"No JPEG frames were provided for {target_key}" + ) + array = np.stack(frames, axis=0) + if expected_shape and tuple(array.shape) != expected_shape: + raise ValueError( + f"Decoded JPEG video for {target_key} has " + f"shape {array.shape}, expected {expected_shape}" + ) + return array + + array = np.asarray(value) + if array.ndim == 3: + return np.expand_dims(array, axis=0) + if array.ndim == 4: + return array + raise ValueError( + f"G2 video input for {target_key} must have " + f"shape (H,W,C) or (T,H,W,C), got {array.shape}" + ) + + @staticmethod + def _normalize_state( + value: object, + target_key: str, + ) -> np.ndarray: + array = np.asarray(value, dtype=np.float64) + if array.ndim == 0: + array = array.reshape(1, 1) + elif array.ndim == 1: + array = array.reshape(1, -1) + elif array.ndim != 2: + raise ValueError( + f"G2 state input for {target_key} must be " + f"1D or 2D, got {array.shape}" + ) + return array + + @staticmethod + def _split_packed_state( + value: object, + ) -> dict[str, np.ndarray]: + packed = np.asarray(value, dtype=np.float64) + if packed.ndim == 1: + packed = packed.reshape(1, -1) + elif packed.ndim != 2: + raise ValueError( + "Packed G2 state must have shape (16,) or " + f"(T,16), got {packed.shape}" + ) + if packed.shape[-1] != 16: + raise ValueError( + f"Packed G2 state must contain 16 values, " + f"got {packed.shape}" + ) + return { + 'state.left_joint_position': packed[:, 0:7], + 'state.left_gripper_position': packed[:, 7:8], + 'state.right_joint_position': packed[:, 8:15], + 'state.right_gripper_position': packed[:, 15:16], + } + + def _prepare_video_chunk( + self, + video_pred: torch.Tensor, + ) -> torch.Tensor | None: + if self._video_save_mode == "none": + return None + if video_pred.ndim != 5: + raise ValueError( + "G2 video prediction must be 5D " + f"(B,C,T,H,W), got {tuple(video_pred.shape)}" + ) + if self._video_save_mode == "first": + return video_pred[:, :, :1].contiguous() + if self._video_save_mode == "full": + return video_pred.contiguous() + raise ValueError( + f"Unsupported video_save_mode: " + f"{self._video_save_mode!r}" + ) + + def _video_save_fps(self) -> int: + return 30 + + def _convert_observation(self, obs: dict) -> dict: + converted: dict[str, object] = {} + missing_video: list[str] = [] + + for source_key, target_key in self.VIDEO_KEY_MAPPING.items(): + value = self._lookup_obs_value( + obs, + source_key, + target_key, + ) + if value is None: + missing_video.append(source_key) + continue + frames = self._normalize_video( + value, + target_key, + ) + + self._frame_buffers[target_key].extend( + list(frames) + ) + + history = self._frame_buffers[target_key][-4:] + + while len(history) < 4: + history.insert(0, history[0]) + + converted[target_key] = np.stack(history, axis=0) + + if missing_video: + raise ValueError( + "G2 inference requires video keys: " + + ", ".join(sorted(missing_video)) + ) + + packed_state = None + for key in self.PACKED_STATE_KEYS: + if key in obs: + packed_state = obs[key] + break + + if packed_state is not None: + converted.update( + self._split_packed_state(packed_state) + ) + else: + missing_state: list[str] = [] + for source_key, target_key in self.STATE_KEY_MAPPING.items(): + value = self._lookup_obs_value( + obs, + source_key, + target_key, + ) + if value is None: + missing_state.append(source_key) + continue + converted[target_key] = self._normalize_state( + value, + target_key, + ) + if missing_state: + raise ValueError( + "G2 inference requires a packed 16-D state " + "under observation/state, observation.state, " + "or state; otherwise all split state keys are " + "required: " + + ", ".join(sorted(missing_state)) + ) + + expected_dims = { + 'state.left_joint_position': 7, + 'state.left_gripper_position': 1, + 'state.right_joint_position': 7, + 'state.right_gripper_position': 1, + } + for key, expected_dim in expected_dims.items(): + array = np.asarray(converted[key]) + if array.shape[-1] != expected_dim: + raise ValueError( + f"{key} must have last dimension " + f"{expected_dim}, got {array.shape}" + ) + + converted['annotation.language.action_text'] = obs.get( + 'prompt', + obs.get( + 'annotation.language.action_text', + '', + ), + ) + return converted + + def _convert_action( + self, + action_dict: dict, + ) -> np.ndarray: + missing = [ + key + for key in self._action_keys + if key not in action_dict + ] + if missing: + raise RuntimeError( + "Missing G2 action outputs: " + + ", ".join(missing) + ) + + arrays: list[np.ndarray] = [] + horizon: int | None = None + for key in self._action_keys: + value = action_dict[key] + if isinstance(value, torch.Tensor): + value = value.detach().cpu().numpy() + array = np.asarray(value) + if array.ndim == 0: + array = array.reshape(1, 1) + elif array.ndim == 1: + array = array.reshape(-1, 1) + else: + array = array.reshape(array.shape[0], -1) + + if horizon is None: + horizon = array.shape[0] + elif array.shape[0] != horizon: + raise RuntimeError( + f"Inconsistent G2 action horizon for {key}: " + f"expected {horizon}, got {array.shape[0]}" + ) + arrays.append(array.astype(np.float32)) + + action = np.concatenate(arrays, axis=-1) + if action.shape != (24, 16): + raise RuntimeError( + "G2 action must have shape (24,16), ordered as " + "[left_joint(7), left_gripper(1), " + "right_joint(7), right_gripper(1)], " + f"got {action.shape}" + ) + if not np.isfinite(action).all(): + raise RuntimeError("G2 action contains NaN or infinity") + if float(np.max(np.abs(action))) < 1e-6: + raise RuntimeError( + "G2 action is entirely zero; refusing a likely " + "checkpoint/config loading failure" + ) + return action.astype(np.float32) + + +class WebsocketPolicyServer: + """Serves a policy using the websocket protocol. See websocket_client_policy.py for a client implementation. + Currently only implements the `load` and `infer` methods. + """ + + def __init__( + self, + policy: _base_policy.BasePolicy, + host: str = "0.0.0.0", + port: int | None = None, + metadata: dict | None = None, + output_dir: str | None = None, + signal_group: dist.ProcessGroup | None = None, + ) -> None: + self._policy = policy + self._host = host + self._port = port + self._metadata = metadata or {} + self._output_dir = output_dir + logging.getLogger("websockets.server").setLevel(logging.INFO) + self.video_across_time = [] + self._msg_index = 0 + self._signal_group = signal_group + if self._output_dir: + os.makedirs(self._output_dir, exist_ok=True) + os.makedirs(os.path.join(self._output_dir, "inputs"), exist_ok=True) + + def serve_forever(self, rank: int = 0) -> None: + asyncio.run(self.run(rank)) + + async def run(self, rank: int = 0): + if rank == 0: + async with _server.serve( + self._handler, + self._host, + self._port, + compression=None, + max_size=None, + process_request=_health_check, + ping_interval=None, + ) as server: + await server.serve_forever() + else: + await self._worker_loop() + + async def _worker_loop(self): + logger.info(f"Worker loop started for rank {dist.get_rank()}") + signal_tensor = torch.zeros(1, dtype=torch.int32, device='cpu') + while True: + try: + dist.broadcast(signal_tensor, src=0, group=self._signal_group) + + signal = signal_tensor.item() + if signal == SIGNAL_SHUTDOWN: + logger.info(f"Rank {dist.get_rank()} received shutdown signal") + break + elif signal == SIGNAL_IDLE: + logger.info(f"Rank {dist.get_rank()} received idle signal. Waiting for next client.") + continue + elif signal == SIGNAL_RESET_CACHE: + logger.info(f"Rank {dist.get_rank()} received inference cache reset signal") + _reset_policy_inference_cache(self._policy, "worker signal") + continue + + batch = self._receive_batch_from_rank0() + dist.barrier() + with torch.no_grad(): + result_batch, video_pred = _run_policy_forward( + self._policy, batch + ) + dist.barrier() + + except Exception as e: + logger.error(f"Worker loop error on rank {dist.get_rank()}: {e}") + traceback.print_exc() + break + + def _receive_batch_from_rank0(self): + import pickle + + size_tensor = torch.zeros(1, dtype=torch.int64, device='cuda') + dist.broadcast(size_tensor, src=0) + data_size = size_tensor.item() + + data_tensor = torch.zeros(data_size, dtype=torch.uint8, device='cuda') + dist.broadcast(data_tensor, src=0) + + obs = pickle.loads(data_tensor.cpu().numpy().tobytes()) + return Batch(obs=obs) + + def _broadcast_batch_to_workers(self, obs): + import pickle + + serialized = pickle.dumps(obs) + data_size = len(serialized) + + size_tensor = torch.tensor([data_size], dtype=torch.int64, device='cuda') + dist.broadcast(size_tensor, src=0) + + data_tensor = torch.frombuffer(serialized, dtype=torch.uint8).clone().cuda() + dist.broadcast(data_tensor, src=0) + + async def _handler(self, websocket: _server.ServerConnection): + logger.info(f"Connection from {websocket.remote_address} opened") + packer = msgpack_numpy.Packer() + + await websocket.send(packer.pack(self._metadata)) + + signal_tensor = torch.zeros(1, dtype=torch.int32, device='cpu') + + try: + while True: + try: + data = await websocket.recv() + obs = msgpack_numpy.unpackb(data) + self._msg_index += 1 + + signal_tensor.zero_() + dist.broadcast(signal_tensor, src=0, group=self._signal_group) + + self._broadcast_batch_to_workers(obs) + batch = Batch(obs=obs) + + dist.barrier() + with torch.no_grad(): + result_batch, video_pred = _run_policy_forward( + self._policy, batch + ) + dist.barrier() + + action_chunk_dict = result_batch.act + + def batch_to_dict(batch): + out = {} + for k in dir(batch): + if not k.startswith("action."): + continue + out[k] = getattr(batch, k) + return out + + action_chunk_dict = batch_to_dict(action_chunk_dict) + await websocket.send(packer.pack(action_chunk_dict)) + + except websockets.ConnectionClosed: + logger.info(f"Connection from {websocket.remote_address} closed") + self.video_across_time = [] + break + except Exception: + await websocket.send(traceback.format_exc()) + await websocket.close( + code=websockets.frames.CloseCode.INTERNAL_ERROR, + reason="Internal server error. Traceback included in previous frame.", + ) + raise + finally: + logger.info("Rank 0: Client session ended. Sending idle signal (2) to workers.") + signal_tensor.fill_(2) + dist.broadcast(signal_tensor, src=0, group=self._signal_group) + + +def init_mesh() -> DeviceMesh: + local_rank = int(os.environ.get("LOCAL_RANK", "0")) + + torch.cuda.set_device(local_rank) + _ = torch.cuda.is_available() + _ = torch.cuda.device_count() + + dist.init_process_group("nccl") + rank = dist.get_rank() + world_size = dist.get_world_size() + if world_size not in (1, 2): + raise ValueError( + f"This DreamZero inference path only supports 1 or 2 GPUs, got world_size={world_size}. " + "The action head parallelization code explicitly supports ip_size 1 or 2 only. " + "Please launch with --nproc_per_node=2 (or 1)." + ) + print(f"Rank {rank}/{world_size} (PID: {os.getpid()}) setting device to local_rank={local_rank}") + + torch.cuda.set_device(local_rank) + device = torch.device(f"cuda:{local_rank}") + + mesh = init_device_mesh( + device_type="cuda", + mesh_shape=(world_size,), + mesh_dim_names=("ip",), + ) + print(f"Rank {rank}/{world_size} (PID: {os.getpid()}) using device {device}") + + return mesh + +def _health_check(connection: _server.ServerConnection, request: _server.Request) -> _server.Response | None: + if request.path == "/healthz": + return connection.respond(http.HTTPStatus.OK, "OK\n") + return None + + +def _create_wrapper_policy( + embodiment_tag: str, + groot_policy: GrootSimPolicy, + signal_group: dist.ProcessGroup, + output_dir: str | None, + video_save_mode: str, +) -> DistributedRoboarenaPolicyBase: + if embodiment_tag == 'oxe_droid': + return ARDroidRoboarenaPolicy( + groot_policy=groot_policy, + signal_group=signal_group, + output_dir=output_dir, + video_save_mode=video_save_mode, + ) + if embodiment_tag == 'agibot': + return AgiBotRoboarenaPolicy( + groot_policy=groot_policy, + signal_group=signal_group, + output_dir=output_dir, + video_save_mode=video_save_mode, + ) + if embodiment_tag == 'g2': + return G2RoboarenaPolicy( + groot_policy=groot_policy, + signal_group=signal_group, + output_dir=output_dir, + video_save_mode=video_save_mode, + ) + raise ValueError(f'Unsupported embodiment_tag: {embodiment_tag}') + + +def _create_server_config(embodiment_tag: str) -> PolicyServerConfig: + if embodiment_tag == 'oxe_droid': + return PolicyServerConfig( + image_resolution=(180, 320), + needs_wrist_camera=True, + n_external_cameras=2, + needs_stereo_camera=False, + needs_session_id=True, + action_space='joint_position', + ) + if embodiment_tag == 'agibot': + return PolicyServerConfig( + image_resolution=(640, 480), + needs_wrist_camera=False, + n_external_cameras=3, + needs_stereo_camera=False, + needs_session_id=True, + action_space='agibot_flattened', + ) + if embodiment_tag == 'g2': + return PolicyServerConfig( + image_resolution=(176, 320), + needs_wrist_camera=False, + n_external_cameras=3, + needs_stereo_camera=False, + needs_session_id=True, + action_space='joint_position', + ) + raise ValueError(f'Unsupported embodiment_tag: {embodiment_tag}') + + +def _build_path_overrides(args: Args) -> tuple[list[str], list[str]]: + model_config_overrides: list[str] = [] + train_config_overrides: list[str] = [] + + if args.wan_ckpt_dir: + wan_ckpt_dir = os.path.abspath(args.wan_ckpt_dir) + required_files = [ + os.path.join(wan_ckpt_dir, "models_t5_umt5-xxl-enc-bf16.pth"), + os.path.join(wan_ckpt_dir, "models_clip_open-clip-xlm-roberta-large-vit-huge-14.pth"), + os.path.join(wan_ckpt_dir, "Wan2.1_VAE.pth"), + ] + missing = [path for path in required_files if not os.path.exists(path)] + if missing: + raise FileNotFoundError( + "Missing Wan checkpoint component(s): " + ", ".join(missing) + ) + model_config_overrides.extend( + [ + f"action_head_cfg.config.diffusion_model_cfg.diffusion_model_pretrained_path={wan_ckpt_dir}", + f"action_head_cfg.config.text_encoder_cfg.text_encoder_pretrained_path={wan_ckpt_dir}/models_t5_umt5-xxl-enc-bf16.pth", + f"action_head_cfg.config.image_encoder_cfg.image_encoder_pretrained_path={wan_ckpt_dir}/models_clip_open-clip-xlm-roberta-large-vit-huge-14.pth", + f"action_head_cfg.config.vae_cfg.vae_pretrained_path={wan_ckpt_dir}/Wan2.1_VAE.pth", + ] + ) + + if args.tokenizer_path: + tokenizer_path = os.path.abspath(args.tokenizer_path) + if not os.path.exists(tokenizer_path): + raise FileNotFoundError(f"Tokenizer path does not exist: {tokenizer_path}") + if args.embodiment_tag.lower() == "agibot": + train_config_overrides.append( + f"transforms.agibot.transforms.10.tokenizer_path={tokenizer_path}" + ) + elif args.embodiment_tag.lower() == "oxe_droid": + train_config_overrides.append( + f"transforms.oxe_droid.transforms.10.tokenizer_path={tokenizer_path}" + ) + elif args.embodiment_tag.lower() == "g2": + train_config_overrides.append( + f"transforms.g2.transforms.10.tokenizer_path={tokenizer_path}" + ) + + return model_config_overrides, train_config_overrides + + +def main(args: Args) -> None: + os.environ["DREAMZERO_RESET_AR_EACH_REQUEST"] = ( + "true" if args.reset_cache_each_request else "false" + ) + os.environ["ENABLE_DIT_CACHE"] = "true" if args.enable_dit_cache else "false" + if args.num_dit_steps is not None: + os.environ["NUM_DIT_STEPS"] = str(args.num_dit_steps) + elif args.enable_dit_cache: + os.environ.setdefault("NUM_DIT_STEPS", "8") + os.environ.setdefault("ATTENTION_BACKEND", "FA2") + if args.video_save_mode not in {"none", "first", "full"}: + raise ValueError(f"--video-save-mode must be one of none, first, full; got {args.video_save_mode!r}") + torch._dynamo.config.recompile_limit = 800 + + embodiment_tag = args.embodiment_tag.lower() + if embodiment_tag not in {'oxe_droid', 'agibot', 'g2'}: + raise ValueError(f'Unsupported embodiment_tag: {args.embodiment_tag}') + model_path = args.model_path + model_config_overrides, train_config_overrides = _build_path_overrides(args) + policy_metadata = { + "embodiment": embodiment_tag, + "model_name": "dreamzero", + "model_path": model_path, + "wan_ckpt_dir": args.wan_ckpt_dir, + "tokenizer_path": args.tokenizer_path, + } + + device_mesh = init_mesh() + rank = dist.get_rank() + + timeout_delta = datetime.timedelta(seconds=args.timeout_seconds) + signal_group = dist.new_group(backend="gloo", timeout=timeout_delta) + logger.info(f"Rank {rank} initialized signal_group (gloo)") + + policy = GrootSimPolicy( + embodiment_tag=EmbodimentTag(embodiment_tag), + model_path=model_path, + device="cuda" if torch.cuda.is_available() else "cpu", + device_mesh=device_mesh, + model_config_overrides=model_config_overrides, + train_config_overrides=train_config_overrides, + ) + action_head = policy.trained_model.action_head + if args.num_inference_timesteps is not None: + if args.num_inference_timesteps == 0: + logging.info("Keeping checkpoint diffusion inference steps because --num-inference-timesteps=0") + elif args.num_inference_timesteps < 0: + raise ValueError( + f"--num-inference-timesteps must be non-negative, got {args.num_inference_timesteps}" + ) + else: + action_head.num_inference_steps = int(args.num_inference_timesteps) + action_head.num_inference_timesteps = int(args.num_inference_timesteps) + if hasattr(action_head, "config"): + action_head.config.num_inference_timesteps = int(args.num_inference_timesteps) + logging.info( + "Overrode action_head diffusion inference steps to %s on rank %s", + args.num_inference_timesteps, + rank, + ) + logging.info( + "[CONFIG CHECK] rank=%s action_head.num_inference_steps=%s " + "action_head.num_inference_timesteps=%s action_head.num_frame_per_block=%s " + "action_head.model.num_frame_per_block=%s NUM_DIT_STEPS=%s ENABLE_DIT_CACHE=%s", + rank, + getattr(action_head, "num_inference_steps", None), + getattr(action_head, "num_inference_timesteps", None), + getattr(action_head, "num_frame_per_block", None), + getattr(getattr(action_head, "model", None), "num_frame_per_block", None), + os.getenv("NUM_DIT_STEPS"), + os.getenv("ENABLE_DIT_CACHE"), + ) + + hostname = socket.gethostname() + local_ip = socket.gethostbyname(hostname) + + if rank == 0: + logging.info("Creating server (host: %s, ip: %s)", hostname, local_ip) + output_dir = None if args.video_save_mode == "none" else args.output_dir + if output_dir is not None: + os.makedirs(output_dir, exist_ok=True) + logging.info("Videos will be saved to: %s", output_dir) + else: + logging.info("Video saving disabled; no output directory will be created.") + else: + output_dir = None + logging.info(f"Rank {rank} starting as worker for distributed inference...") + + wrapper_policy = _create_wrapper_policy( + embodiment_tag=embodiment_tag, + groot_policy=policy, + signal_group=signal_group, + output_dir=output_dir, + video_save_mode=args.video_save_mode, + ) + + server_config = _create_server_config(embodiment_tag) + + if rank == 0: + logging.info("Using roboarena policy server interface for %s", embodiment_tag) + logging.info(f"Server config: {server_config}") + roboarena_server = RoboarenaServer( + policy=wrapper_policy, + server_config=server_config, + host="0.0.0.0", + port=args.port, + ) + roboarena_server.serve_forever() + else: + server = WebsocketPolicyServer( + policy=policy, + host="0.0.0.0", + port=args.port, + metadata=policy_metadata, + output_dir=output_dir, + signal_group=signal_group, + ) + asyncio.run(server._worker_loop()) + + +def cli() -> None: + logging.basicConfig(level=logging.INFO, force=True) + main(tyro.cli(Args)) + + +if __name__ == "__main__": + cli() diff --git a/socket_test_optimized_AR_agibot_fruit_video.py b/socket_test_optimized_AR_agibot_fruit_video.py new file mode 100644 index 00000000..38aad35b --- /dev/null +++ b/socket_test_optimized_AR_agibot_fruit_video.py @@ -0,0 +1,978 @@ +import dataclasses +import logging +import socket +import asyncio +import os +import http +import logging +import time +import traceback +import torch +import tyro +from einops import rearrange +import datetime +import cv2 + +from groot.vla.model.n1_5.sim_policy import GrootSimPolicy +from groot.vla.data.schema import EmbodimentTag +import imageio +import numpy as np + +from openpi_client import base_policy as _base_policy +from openpi_client import msgpack_numpy +import websockets.asyncio.server as _server +import websockets.frames +from tianshou.data import Batch +import torch.distributed as dist +from torch.distributed.device_mesh import DeviceMesh, init_device_mesh + +# Use roboarena policy server interface +from eval_utils.policy_server import WebsocketPolicyServer as RoboarenaServer +from eval_utils.policy_server import PolicyServerConfig + +logger = logging.getLogger(__name__) + +SIGNAL_INFER = 0 +SIGNAL_SHUTDOWN = 1 +SIGNAL_IDLE = 2 +SIGNAL_RESET_CACHE = 3 + + +def _reset_policy_inference_cache(policy: object, reason: str) -> None: + trained_model = getattr(policy, "trained_model", None) + action_head = getattr(trained_model, "action_head", None) + reset_fn = getattr(action_head, "reset_inference_cache", None) + if callable(reset_fn): + reset_fn() + logger.info("Reset action-head inference cache on rank %s (%s)", dist.get_rank() if dist.is_initialized() else "?", reason) + else: + logger.warning("Policy action head does not expose reset_inference_cache(); cache reset skipped (%s)", reason) + +@dataclasses.dataclass +class Args: + port: int = 8000 + timeout_seconds: int = 50000 # 10 hours default, configurable + model_path: str = "./checkpoints/dreamzero" + wan_ckpt_dir: str | None = None # Override Wan2.1 component paths stored in config.json. + tokenizer_path: str | None = None # Override tokenizer_path stored in experiment_cfg/conf.yaml. + enable_dit_cache: bool = False # Backward-compatible alias for num_dit_steps=8. + num_dit_steps: int | None = None # Actual DiT compute steps. Supported fast masks: 5, 6, 7, 8. None keeps model default. + num_inference_timesteps: int | None = None # Positive values override diffusion steps; 0 keeps checkpoint default. + index: int = 0 + embodiment_tag: str = "oxe_droid" + max_chunk_size: int | None = None # If None, use config value. Otherwise override max_chunk_size for inference. + video_save_mode: str = "first" # Generated model video: none, first, or full. + input_video_save_mode: str = "top_head" # Real camera input: none, top_head, or all. + input_video_fps: int = 30 + output_dir: str = "/data/wangk/dreamzero/video_rollout" + + +class DistributedRoboarenaPolicyBase: + """Shared distributed inference plumbing for websocket policy wrappers.""" + + def __init__( + self, + groot_policy: GrootSimPolicy, + signal_group: dist.ProcessGroup, + output_dir: str | None = None, + video_save_mode: str = "first", + input_video_save_mode: str = "top_head", + input_video_fps: int = 30, + ) -> None: + self._policy = groot_policy + self._signal_group = signal_group + self._output_dir = output_dir + self._video_save_mode = video_save_mode + self._input_video_save_mode = input_video_save_mode + self._input_video_fps = int(input_video_fps) + self._frame_buffers = self._init_frame_buffers() + self._current_session_id: str | None = None + self.video_across_time = [] + self._input_video_buffers: dict[str, list[np.ndarray]] = {} + self._msg_index = 0 + + if self._output_dir: + os.makedirs(self._output_dir, exist_ok=True) + + def _init_frame_buffers(self) -> dict[str, list[np.ndarray]]: + return {} + + def _reset_custom_state(self) -> None: + pass + + def _after_infer(self) -> None: + pass + + def _prepare_video_chunk(self, video_pred: torch.Tensor) -> torch.Tensor | None: + if self._video_save_mode == "none": + return None + return video_pred + + def _video_save_fps(self) -> int: + return 5 + + @staticmethod + def _sanitize_filename(value: object) -> str: + text = str(value or "session") + safe = "".join( + character if character.isalnum() or character in "-_" else "_" + for character in text + ) + return safe[:80] or "session" + + def _record_input_video(self, modality_key: str, frames: np.ndarray) -> None: + if self._input_video_save_mode == "none": + return + if self._input_video_save_mode == "top_head" and modality_key != "video.top_head": + return + + array = np.asarray(frames) + if array.ndim == 3: + array = array[None, ...] + if array.ndim != 4: + logger.warning( + "Skipping input video %s with invalid shape %s", + modality_key, + array.shape, + ) + return + + buffer = self._input_video_buffers.setdefault(modality_key, []) + for frame in array: + frame = np.asarray(frame) + if frame.dtype != np.uint8: + frame = np.clip(frame, 0, 255).astype(np.uint8) + buffer.append(frame.copy()) + + def _save_input_videos(self) -> None: + if not self._output_dir or not self._input_video_buffers: + return + + input_dir = os.path.join(self._output_dir, "inputs") + os.makedirs(input_dir, exist_ok=True) + + timestamp = datetime.datetime.now().strftime("%Y%m%d_%H%M%S") + session_name = self._sanitize_filename(self._current_session_id) + + for modality_key, frames in self._input_video_buffers.items(): + if not frames: + continue + + camera_name = modality_key.replace("video.", "") + output_path = os.path.join( + input_dir, + f"{timestamp}_{session_name}_{camera_name}_f{len(frames)}.mp4", + ) + + try: + imageio.mimsave( + output_path, + frames, + fps=self._input_video_fps, + codec="libx264", + macro_block_size=None, + ) + logger.info( + "Saved real input video %s (%s frames) to %s", + modality_key, + len(frames), + output_path, + ) + except Exception as exc: + logger.warning( + "Failed to save real input video %s: %s", + modality_key, + exc, + ) + + self._input_video_buffers = {} + + def _convert_observation(self, obs: dict) -> dict: + raise NotImplementedError + + def _convert_action(self, action_dict: dict) -> np.ndarray: + raise NotImplementedError + + def _broadcast_batch_to_workers(self, obs: dict) -> None: + import pickle + + serialized = pickle.dumps(obs) + data_size = len(serialized) + + size_tensor = torch.tensor([data_size], dtype=torch.int64, device='cuda') + dist.broadcast(size_tensor, src=0) + + data_tensor = torch.frombuffer(serialized, dtype=torch.uint8).clone().cuda() + dist.broadcast(data_tensor, src=0) + + def _extract_action_dict(self, action_chunk_dict: object) -> dict[str, object]: + action_dict: dict[str, object] = {} + for key in dir(action_chunk_dict): + if key.startswith('action.'): + action_dict[key] = getattr(action_chunk_dict, key) + return action_dict + + def _broadcast_signal_to_workers(self, signal: int) -> None: + signal_tensor = torch.tensor([signal], dtype=torch.int32, device='cpu') + dist.broadcast(signal_tensor, src=0, group=self._signal_group) + + def infer(self, obs: dict) -> np.ndarray: + session_id = obs.get('session_id') + if session_id is not None and session_id != self._current_session_id: + if self._current_session_id is not None: + logger.info("Session changed from '%s' to '%s', resetting state", self._current_session_id, session_id) + self._broadcast_signal_to_workers(SIGNAL_RESET_CACHE) + self._reset_state() + else: + logger.info("New session started: '%s'", session_id) + self._current_session_id = session_id + + self._msg_index += 1 + converted_obs = self._convert_observation(obs) + + self._broadcast_signal_to_workers(SIGNAL_INFER) + self._broadcast_batch_to_workers(converted_obs) + + batch = Batch(obs=converted_obs) + dist.barrier() + with torch.no_grad(): + result_batch, video_pred = self._policy.lazy_joint_forward_causal(batch) + dist.barrier() + + video_chunk = self._prepare_video_chunk(video_pred) + if video_chunk is not None: + self.video_across_time.append(video_chunk.detach().cpu()) + action = self._convert_action(self._extract_action_dict(result_batch.act)) + self._after_infer() + return action + + def _reset_state(self, save_video: bool = True) -> None: + if save_video: + self._save_input_videos() + + if save_video and len(self.video_across_time) > 0 and self._output_dir: + try: + frame_list = [] + action_head = self._policy.trained_model.action_head + device = getattr(action_head, "_device", None) + if device is None: + device = next(self._policy.trained_model.parameters()).device + video_across_time_cat = torch.cat(self.video_across_time, dim=2).to(device=device, dtype=torch.bfloat16) + frames = action_head.vae.decode( + video_across_time_cat, + tiled=action_head.tiled, + tile_size=(action_head.tile_size_height, action_head.tile_size_width), + tile_stride=(action_head.tile_stride_height, action_head.tile_stride_width), + ) + frames = rearrange(frames, 'B C T H W -> B T H W C') + frames = frames[0] + frames = ((frames.float() + 1) * 127.5).clip(0, 255).cpu().numpy().astype(np.uint8) + for frame in frames: + frame_list.append(frame) + + if frame_list: + sample_frame = frame_list[0] + if len(sample_frame.shape) == 3 and sample_frame.shape[2] in [1, 3, 4]: + save_dir = self._output_dir + os.makedirs(save_dir, exist_ok=True) + all_mp4_files = [f for f in os.listdir(save_dir) if f.endswith('.mp4')] + timestamp = datetime.datetime.now().strftime('%Y%m%d_%H%M%S') + num_frames = len(frame_list) + output_path = os.path.join(save_dir, f'{timestamp}_{len(all_mp4_files):06}_f{num_frames}.mp4') + imageio.mimsave(output_path, frame_list, fps=self._video_save_fps(), codec='libx264') + logger.info('Saved video on reset to: %s', output_path) + except Exception as exc: + logger.warning('Failed to save video on reset: %s', exc) + + for key in self._frame_buffers: + self._frame_buffers[key] = [] + + self.video_across_time = [] + _reset_policy_inference_cache(self._policy, "wrapper reset_state") + self._reset_custom_state() + + def reset(self, reset_info: dict) -> None: + self._broadcast_signal_to_workers(SIGNAL_RESET_CACHE) + self._reset_state(save_video=True) + + +class ARDroidRoboarenaPolicy(DistributedRoboarenaPolicyBase): + """Wrapper policy that implements roboarena.policy.BasePolicy interface for AR_droid.""" + + FRAMES_PER_CHUNK = 4 + + def __init__( + self, + groot_policy: GrootSimPolicy, + signal_group: dist.ProcessGroup, + output_dir: str | None = None, + video_save_mode: str = "first", + ) -> None: + super().__init__( + groot_policy=groot_policy, + signal_group=signal_group, + output_dir=output_dir, + video_save_mode=video_save_mode, + ) + self._reset_custom_state() + + def _init_frame_buffers(self) -> dict[str, list[np.ndarray]]: + return { + 'video.exterior_image_1_left': [], + 'video.exterior_image_2_left': [], + 'video.wrist_image_left': [], + } + + def _reset_custom_state(self) -> None: + self._is_first_call = True + + def _after_infer(self) -> None: + self._is_first_call = False + + def _convert_observation(self, obs: dict) -> dict: + converted = {} + image_key_mapping = { + 'observation/exterior_image_0_left': 'video.exterior_image_1_left', + 'observation/exterior_image_1_left': 'video.exterior_image_2_left', + 'observation/wrist_image_left': 'video.wrist_image_left', + } + + for roboarena_key, droid_key in image_key_mapping.items(): + if roboarena_key in obs: + data = obs[roboarena_key] + if isinstance(data, np.ndarray): + if data.ndim == 4: + self._frame_buffers[droid_key].extend(list(data)) + else: + self._frame_buffers[droid_key].append(data) + + num_frames = 1 if self._is_first_call else self.FRAMES_PER_CHUNK + + for droid_key, buffer in self._frame_buffers.items(): + if len(buffer) > 0: + if len(buffer) >= num_frames: + frames_to_use = buffer[-num_frames:] + else: + frames_to_use = buffer.copy() + while len(frames_to_use) < num_frames: + frames_to_use.insert(0, buffer[0]) + converted[droid_key] = np.stack(frames_to_use, axis=0) + + joint_pos = obs.get('observation/joint_position', np.zeros(7, dtype=np.float32)) + if joint_pos.ndim == 1: + joint_pos = joint_pos.reshape(1, -1) + converted['state.joint_position'] = joint_pos.astype(np.float64) + + gripper_pos = obs.get('observation/gripper_position', np.zeros(1, dtype=np.float32)) + if gripper_pos.ndim == 1: + gripper_pos = gripper_pos.reshape(1, -1) + converted['state.gripper_position'] = gripper_pos.astype(np.float64) + converted['annotation.language.action_text'] = obs.get('prompt', '') + return converted + + def _convert_action(self, action_dict: dict) -> np.ndarray: + joint_action = None + gripper_action = None + for key, value in action_dict.items(): + if 'joint_position' in key: + joint_action = value + elif 'gripper_position' in key or 'gripper' in key: + gripper_action = value + + if joint_action is None: + return np.zeros((1, 8), dtype=np.float32) + + if isinstance(joint_action, torch.Tensor): + joint_action = joint_action.cpu().numpy() + if joint_action.ndim == 1: + joint_action = joint_action.reshape(1, -1) + + num_steps = joint_action.shape[0] + if gripper_action is not None: + if isinstance(gripper_action, torch.Tensor): + gripper_action = gripper_action.cpu().numpy() + if gripper_action.ndim == 1: + gripper_action = gripper_action.reshape(-1, 1) + elif gripper_action.ndim == 0: + gripper_action = gripper_action.reshape(1, 1) + else: + gripper_action = np.zeros((num_steps, 1), dtype=np.float32) + + return np.concatenate([joint_action, gripper_action], axis=-1).astype(np.float32) + + +class AgiBotRoboarenaPolicy(DistributedRoboarenaPolicyBase): + """Adapter that converts websocket observations into AgiBot modality keys.""" + + VIDEO_KEY_MAPPING = { + 'observation/top_head': 'video.top_head', + 'observation/hand_left': 'video.hand_left', + 'observation/hand_right': 'video.hand_right', + } + STATE_KEY_MAPPING = { + 'observation/left_arm_joint_position': 'state.left_arm_joint_position', + 'observation/right_arm_joint_position': 'state.right_arm_joint_position', + 'observation/left_effector_position': 'state.left_effector_position', + 'observation/right_effector_position': 'state.right_effector_position', + 'observation/head_position': 'state.head_position', + 'observation/waist_pitch': 'state.waist_pitch', + 'observation/waist_lift': 'state.waist_lift', + } + + def __init__( + self, + groot_policy: GrootSimPolicy, + signal_group: dist.ProcessGroup, + output_dir: str | None = None, + video_save_mode: str = "first", + input_video_save_mode: str = "top_head", + input_video_fps: int = 30, + ) -> None: + super().__init__( + groot_policy=groot_policy, + signal_group=signal_group, + output_dir=output_dir, + video_save_mode=video_save_mode, + input_video_save_mode=input_video_save_mode, + input_video_fps=input_video_fps, + ) + self._action_keys = list(self._policy.modality_configs.action.modality_keys) + + def _lookup_obs_value(self, obs: dict, source_key: str, target_key: str) -> object: + if source_key in obs: + return obs[source_key] + return obs.get(target_key) + + def _normalize_video(self, value: object, target_key: str) -> np.ndarray: + if isinstance(value, dict) and value.get("__dreamzero_image_encoding__") == "jpeg_sequence": + frames = [] + expected_shape = tuple(value.get("shape", ())) + expected_dtype = np.dtype(value.get("dtype", "uint8")) + for index, frame_bytes in enumerate(value.get("frames", [])): + encoded = np.frombuffer(frame_bytes, dtype=np.uint8) + frame = cv2.imdecode(encoded, cv2.IMREAD_COLOR) + if frame is None: + raise ValueError(f"Failed to decode JPEG frame {index} for {target_key}") + frames.append(frame.astype(expected_dtype, copy=False)) + array = np.stack(frames, axis=0) + if expected_shape and tuple(array.shape) != expected_shape: + raise ValueError( + f"Decoded JPEG video for {target_key} has shape {array.shape}, expected {expected_shape}" + ) + return array + + array = np.asarray(value) + if array.ndim == 3: + return np.expand_dims(array, axis=0) + if array.ndim == 4: + return array + raise ValueError(f'AgiBot video input for {target_key} must have shape (H, W, C) or (T, H, W, C), got {array.shape}') + + def _normalize_state(self, value: object, target_key: str) -> np.ndarray: + array = np.asarray(value) + if array.ndim == 0: + return array.reshape(1, 1).astype(np.float64) + if array.ndim == 1: + return array.reshape(1, -1).astype(np.float64) + if array.ndim == 2: + return array.astype(np.float64) + raise ValueError(f'AgiBot state input for {target_key} must be 1D or 2D, got {array.shape}') + + def _prepare_video_chunk(self, video_pred: torch.Tensor) -> torch.Tensor | None: + if self._video_save_mode == "none": + return None + if video_pred.ndim != 5: + raise ValueError(f'AgiBot video prediction must be 5D (B, C, T, H, W), got {tuple(video_pred.shape)}') + if self._video_save_mode == "first": + return video_pred[:, :, :1].contiguous() + if self._video_save_mode == "full": + return video_pred.contiguous() + raise ValueError(f"Unsupported video_save_mode: {self._video_save_mode!r}; expected none, first, or full") + + def _video_save_fps(self) -> int: + return 20 + + def _convert_observation(self, obs: dict) -> dict: + converted = {} + missing_keys: list[str] = [] + + for source_key, target_key in self.VIDEO_KEY_MAPPING.items(): + value = self._lookup_obs_value(obs, source_key, target_key) + if value is None: + missing_keys.append(source_key) + continue + normalized_video = self._normalize_video(value, target_key) + converted[target_key] = normalized_video + self._record_input_video(target_key, normalized_video) + + for source_key, target_key in self.STATE_KEY_MAPPING.items(): + value = self._lookup_obs_value(obs, source_key, target_key) + if value is None: + missing_keys.append(source_key) + continue + converted[target_key] = self._normalize_state(value, target_key) + + if missing_keys: + raise ValueError( + 'AgiBot inference requires the following observation keys: ' + + ', '.join(sorted(missing_keys)) + ) + + converted['annotation.language.action_text'] = obs.get('prompt', obs.get('annotation.language.action_text', '')) + return converted + + def _convert_action(self, action_dict: dict) -> np.ndarray: + flattened_chunks: list[np.ndarray] = [] + expected_horizon: int | None = None + missing_keys = [key for key in self._action_keys if key not in action_dict] + if missing_keys: + raise RuntimeError('Missing AgiBot action outputs: ' + ', '.join(missing_keys)) + + for action_key in self._action_keys: + value = action_dict[action_key] + if isinstance(value, torch.Tensor): + value = value.detach().cpu().numpy() + array = np.asarray(value) + if array.ndim == 0: + array = array.reshape(1, 1) + elif array.ndim == 1: + array = array.reshape(-1, 1) + else: + array = array.reshape(array.shape[0], -1) + + if expected_horizon is None: + expected_horizon = array.shape[0] + elif array.shape[0] != expected_horizon: + raise RuntimeError( + f'Inconsistent AgiBot action horizon for {action_key}: expected {expected_horizon}, got {array.shape[0]}' + ) + flattened_chunks.append(array.astype(np.float32)) + + return np.concatenate(flattened_chunks, axis=-1).astype(np.float32) + + +class WebsocketPolicyServer: + """Serves a policy using the websocket protocol. See websocket_client_policy.py for a client implementation. + Currently only implements the `load` and `infer` methods. + """ + + def __init__( + self, + policy: _base_policy.BasePolicy, + host: str = "0.0.0.0", + port: int | None = None, + metadata: dict | None = None, + output_dir: str | None = None, + signal_group: dist.ProcessGroup | None = None, + ) -> None: + self._policy = policy + self._host = host + self._port = port + self._metadata = metadata or {} + self._output_dir = output_dir + logging.getLogger("websockets.server").setLevel(logging.INFO) + self.video_across_time = [] + self._msg_index = 0 + self._signal_group = signal_group + if self._output_dir: + os.makedirs(self._output_dir, exist_ok=True) + os.makedirs(os.path.join(self._output_dir, "inputs"), exist_ok=True) + + def serve_forever(self, rank: int = 0) -> None: + asyncio.run(self.run(rank)) + + async def run(self, rank: int = 0): + if rank == 0: + async with _server.serve( + self._handler, + self._host, + self._port, + compression=None, + max_size=None, + process_request=_health_check, + ping_interval=None, + ) as server: + await server.serve_forever() + else: + await self._worker_loop() + + async def _worker_loop(self): + logger.info(f"Worker loop started for rank {dist.get_rank()}") + signal_tensor = torch.zeros(1, dtype=torch.int32, device='cpu') + while True: + try: + dist.broadcast(signal_tensor, src=0, group=self._signal_group) + + signal = signal_tensor.item() + if signal == SIGNAL_SHUTDOWN: + logger.info(f"Rank {dist.get_rank()} received shutdown signal") + break + elif signal == SIGNAL_IDLE: + logger.info(f"Rank {dist.get_rank()} received idle signal. Waiting for next client.") + continue + elif signal == SIGNAL_RESET_CACHE: + logger.info(f"Rank {dist.get_rank()} received inference cache reset signal") + _reset_policy_inference_cache(self._policy, "worker signal") + continue + + batch = self._receive_batch_from_rank0() + dist.barrier() + with torch.no_grad(): + result_batch, video_pred = self._policy.lazy_joint_forward_causal(batch) + dist.barrier() + + except Exception as e: + logger.error(f"Worker loop error on rank {dist.get_rank()}: {e}") + traceback.print_exc() + break + + def _receive_batch_from_rank0(self): + import pickle + + size_tensor = torch.zeros(1, dtype=torch.int64, device='cuda') + dist.broadcast(size_tensor, src=0) + data_size = size_tensor.item() + + data_tensor = torch.zeros(data_size, dtype=torch.uint8, device='cuda') + dist.broadcast(data_tensor, src=0) + + obs = pickle.loads(data_tensor.cpu().numpy().tobytes()) + return Batch(obs=obs) + + def _broadcast_batch_to_workers(self, obs): + import pickle + + serialized = pickle.dumps(obs) + data_size = len(serialized) + + size_tensor = torch.tensor([data_size], dtype=torch.int64, device='cuda') + dist.broadcast(size_tensor, src=0) + + data_tensor = torch.frombuffer(serialized, dtype=torch.uint8).clone().cuda() + dist.broadcast(data_tensor, src=0) + + async def _handler(self, websocket: _server.ServerConnection): + logger.info(f"Connection from {websocket.remote_address} opened") + packer = msgpack_numpy.Packer() + + await websocket.send(packer.pack(self._metadata)) + + signal_tensor = torch.zeros(1, dtype=torch.int32, device='cpu') + + try: + while True: + try: + data = await websocket.recv() + obs = msgpack_numpy.unpackb(data) + self._msg_index += 1 + + signal_tensor.zero_() + dist.broadcast(signal_tensor, src=0, group=self._signal_group) + + self._broadcast_batch_to_workers(obs) + batch = Batch(obs=obs) + + dist.barrier() + with torch.no_grad(): + result_batch, video_pred = self._policy.lazy_joint_forward_causal(batch) + dist.barrier() + + action_chunk_dict = result_batch.act + + def batch_to_dict(batch): + out = {} + for k in dir(batch): + if not k.startswith("action."): + continue + out[k] = getattr(batch, k) + return out + + action_chunk_dict = batch_to_dict(action_chunk_dict) + await websocket.send(packer.pack(action_chunk_dict)) + + except websockets.ConnectionClosed: + logger.info(f"Connection from {websocket.remote_address} closed") + self.video_across_time = [] + break + except Exception: + await websocket.send(traceback.format_exc()) + await websocket.close( + code=websockets.frames.CloseCode.INTERNAL_ERROR, + reason="Internal server error. Traceback included in previous frame.", + ) + raise + finally: + logger.info("Rank 0: Client session ended. Sending idle signal (2) to workers.") + signal_tensor.fill_(2) + dist.broadcast(signal_tensor, src=0, group=self._signal_group) + + +def init_mesh() -> DeviceMesh: + local_rank = int(os.environ.get("LOCAL_RANK", "0")) + + torch.cuda.set_device(local_rank) + _ = torch.cuda.is_available() + _ = torch.cuda.device_count() + + dist.init_process_group("nccl") + rank = dist.get_rank() + world_size = dist.get_world_size() + if world_size not in (1, 2): + raise ValueError( + f"This DreamZero inference path only supports 1 or 2 GPUs, got world_size={world_size}. " + "The action head parallelization code explicitly supports ip_size 1 or 2 only. " + "Please launch with --nproc_per_node=2 (or 1)." + ) + print(f"Rank {rank}/{world_size} (PID: {os.getpid()}) setting device to local_rank={local_rank}") + + torch.cuda.set_device(local_rank) + device = torch.device(f"cuda:{local_rank}") + + mesh = init_device_mesh( + device_type="cuda", + mesh_shape=(world_size,), + mesh_dim_names=("ip",), + ) + print(f"Rank {rank}/{world_size} (PID: {os.getpid()}) using device {device}") + + return mesh + +def _health_check(connection: _server.ServerConnection, request: _server.Request) -> _server.Response | None: + if request.path == "/healthz": + return connection.respond(http.HTTPStatus.OK, "OK\n") + return None + + +def _create_wrapper_policy( + embodiment_tag: str, + groot_policy: GrootSimPolicy, + signal_group: dist.ProcessGroup, + output_dir: str | None, + video_save_mode: str, + input_video_save_mode: str, + input_video_fps: int, +) -> DistributedRoboarenaPolicyBase: + if embodiment_tag == 'oxe_droid': + return ARDroidRoboarenaPolicy( + groot_policy=groot_policy, + signal_group=signal_group, + output_dir=output_dir, + video_save_mode=video_save_mode, + ) + if embodiment_tag == 'agibot': + return AgiBotRoboarenaPolicy( + groot_policy=groot_policy, + signal_group=signal_group, + output_dir=output_dir, + video_save_mode=video_save_mode, + input_video_save_mode=input_video_save_mode, + input_video_fps=input_video_fps, + ) + raise ValueError(f'Unsupported embodiment_tag: {embodiment_tag}') + + +def _create_server_config(embodiment_tag: str) -> PolicyServerConfig: + if embodiment_tag == 'oxe_droid': + return PolicyServerConfig( + image_resolution=(180, 320), + needs_wrist_camera=True, + n_external_cameras=2, + needs_stereo_camera=False, + needs_session_id=True, + action_space='joint_position', + ) + if embodiment_tag == 'agibot': + return PolicyServerConfig( + image_resolution=(640, 480), + needs_wrist_camera=False, + n_external_cameras=3, + needs_stereo_camera=False, + needs_session_id=True, + action_space='agibot_flattened', + ) + raise ValueError(f'Unsupported embodiment_tag: {embodiment_tag}') + + +def _build_path_overrides(args: Args) -> tuple[list[str], list[str]]: + model_config_overrides: list[str] = [] + train_config_overrides: list[str] = [] + + if args.wan_ckpt_dir: + wan_ckpt_dir = os.path.abspath(args.wan_ckpt_dir) + required_files = [ + os.path.join(wan_ckpt_dir, "models_t5_umt5-xxl-enc-bf16.pth"), + os.path.join(wan_ckpt_dir, "models_clip_open-clip-xlm-roberta-large-vit-huge-14.pth"), + os.path.join(wan_ckpt_dir, "Wan2.1_VAE.pth"), + ] + missing = [path for path in required_files if not os.path.exists(path)] + if missing: + raise FileNotFoundError( + "Missing Wan checkpoint component(s): " + ", ".join(missing) + ) + model_config_overrides.extend( + [ + f"action_head_cfg.config.diffusion_model_cfg.diffusion_model_pretrained_path={wan_ckpt_dir}", + f"action_head_cfg.config.text_encoder_cfg.text_encoder_pretrained_path={wan_ckpt_dir}/models_t5_umt5-xxl-enc-bf16.pth", + f"action_head_cfg.config.image_encoder_cfg.image_encoder_pretrained_path={wan_ckpt_dir}/models_clip_open-clip-xlm-roberta-large-vit-huge-14.pth", + f"action_head_cfg.config.vae_cfg.vae_pretrained_path={wan_ckpt_dir}/Wan2.1_VAE.pth", + ] + ) + + if args.tokenizer_path: + tokenizer_path = os.path.abspath(args.tokenizer_path) + if not os.path.exists(tokenizer_path): + raise FileNotFoundError(f"Tokenizer path does not exist: {tokenizer_path}") + if args.embodiment_tag.lower() == "agibot": + train_config_overrides.append( + f"transforms.agibot.transforms.10.tokenizer_path={tokenizer_path}" + ) + elif args.embodiment_tag.lower() == "oxe_droid": + train_config_overrides.append( + f"transforms.oxe_droid.transforms.10.tokenizer_path={tokenizer_path}" + ) + + return model_config_overrides, train_config_overrides + + +def main(args: Args) -> None: + os.environ["ENABLE_DIT_CACHE"] = "true" if args.enable_dit_cache else "false" + if args.num_dit_steps is not None: + os.environ["NUM_DIT_STEPS"] = str(args.num_dit_steps) + elif args.enable_dit_cache: + os.environ.setdefault("NUM_DIT_STEPS", "8") + os.environ.setdefault("ATTENTION_BACKEND", "FA2") + if args.video_save_mode not in {"none", "first", "full"}: + raise ValueError( + f"--video-save-mode must be one of none, first, full; " + f"got {args.video_save_mode!r}" + ) + if args.input_video_save_mode not in {"none", "top_head", "all"}: + raise ValueError( + f"--input-video-save-mode must be one of none, top_head, all; " + f"got {args.input_video_save_mode!r}" + ) + if args.input_video_fps <= 0: + raise ValueError( + f"--input-video-fps must be positive, got {args.input_video_fps}" + ) + torch._dynamo.config.recompile_limit = 800 + + embodiment_tag = args.embodiment_tag.lower() + if embodiment_tag not in {'oxe_droid', 'agibot'}: + raise ValueError(f'Unsupported embodiment_tag: {args.embodiment_tag}') + model_path = args.model_path + model_config_overrides, train_config_overrides = _build_path_overrides(args) + policy_metadata = { + "embodiment": embodiment_tag, + "model_name": "dreamzero", + "model_path": model_path, + "wan_ckpt_dir": args.wan_ckpt_dir, + "tokenizer_path": args.tokenizer_path, + } + + device_mesh = init_mesh() + rank = dist.get_rank() + + timeout_delta = datetime.timedelta(seconds=args.timeout_seconds) + signal_group = dist.new_group(backend="gloo", timeout=timeout_delta) + logger.info(f"Rank {rank} initialized signal_group (gloo)") + + policy = GrootSimPolicy( + embodiment_tag=EmbodimentTag(embodiment_tag), + model_path=model_path, + device="cuda" if torch.cuda.is_available() else "cpu", + device_mesh=device_mesh, + model_config_overrides=model_config_overrides, + train_config_overrides=train_config_overrides, + ) + action_head = policy.trained_model.action_head + if args.num_inference_timesteps is not None: + if args.num_inference_timesteps == 0: + logging.info("Keeping checkpoint diffusion inference steps because --num-inference-timesteps=0") + elif args.num_inference_timesteps < 0: + raise ValueError( + f"--num-inference-timesteps must be non-negative, got {args.num_inference_timesteps}" + ) + else: + action_head.num_inference_steps = int(args.num_inference_timesteps) + action_head.num_inference_timesteps = int(args.num_inference_timesteps) + if hasattr(action_head, "config"): + action_head.config.num_inference_timesteps = int(args.num_inference_timesteps) + logging.info( + "Overrode action_head diffusion inference steps to %s on rank %s", + args.num_inference_timesteps, + rank, + ) + logging.info( + "[CONFIG CHECK] rank=%s action_head.num_inference_steps=%s " + "action_head.num_inference_timesteps=%s action_head.num_frame_per_block=%s " + "action_head.model.num_frame_per_block=%s NUM_DIT_STEPS=%s ENABLE_DIT_CACHE=%s", + rank, + getattr(action_head, "num_inference_steps", None), + getattr(action_head, "num_inference_timesteps", None), + getattr(action_head, "num_frame_per_block", None), + getattr(getattr(action_head, "model", None), "num_frame_per_block", None), + os.getenv("NUM_DIT_STEPS"), + os.getenv("ENABLE_DIT_CACHE"), + ) + + hostname = socket.gethostname() + local_ip = socket.gethostbyname(hostname) + + if rank == 0: + logging.info("Creating server (host: %s, ip: %s)", hostname, local_ip) + output_dir = None if args.video_save_mode == "none" else args.output_dir + if output_dir is not None: + os.makedirs(output_dir, exist_ok=True) + logging.info("Generated videos will be saved to: %s", output_dir) + logging.info( + "Real input videos will be saved to: %s", + os.path.join(output_dir, "inputs"), + ) + else: + logging.info("Video saving disabled; no output directory will be created.") + else: + output_dir = None + logging.info(f"Rank {rank} starting as worker for distributed inference...") + + wrapper_policy = _create_wrapper_policy( + embodiment_tag=embodiment_tag, + groot_policy=policy, + signal_group=signal_group, + output_dir=output_dir, + video_save_mode=args.video_save_mode, + input_video_save_mode=args.input_video_save_mode, + input_video_fps=args.input_video_fps, + ) + + server_config = _create_server_config(embodiment_tag) + + if rank == 0: + logging.info("Using roboarena policy server interface for %s", embodiment_tag) + logging.info(f"Server config: {server_config}") + roboarena_server = RoboarenaServer( + policy=wrapper_policy, + server_config=server_config, + host="0.0.0.0", + port=args.port, + ) + roboarena_server.serve_forever() + else: + server = WebsocketPolicyServer( + policy=policy, + host="0.0.0.0", + port=args.port, + metadata=policy_metadata, + output_dir=output_dir, + signal_group=signal_group, + ) + asyncio.run(server._worker_loop()) + + +def cli() -> None: + logging.basicConfig(level=logging.INFO, force=True) + main(tyro.cli(Args)) + + +if __name__ == "__main__": + cli()