diff --git a/.gitignore b/.gitignore index af484b145..92e4e258f 100644 --- a/.gitignore +++ b/.gitignore @@ -168,3 +168,4 @@ offline_open_loop_eval_plots/ # Generic model/data folders /models /data +/checkpoints \ No newline at end of file diff --git a/examples/RoboJuDo/README.md b/examples/RoboJuDo/README.md new file mode 100644 index 000000000..3b45e5741 --- /dev/null +++ b/examples/RoboJuDo/README.md @@ -0,0 +1,521 @@ +# RoboJuDo X2 and G1 23-DoF + +This example adapts datasets produced by `robojudo_recorder` for GR00T N1.7. The recorder +writes LeRobot v3.0; GR00T currently trains from its LeRobot v2.1 layout plus +`meta/modality.json`. + +## ZMQ ports at a glance + +The robot pipeline and the deployment client use two separate ZMQ data channels: + +| Port | Direction | ZMQ role | Payload | Socket ownership | +| --- | --- | --- | --- | --- | +| `8561` | robot → deploy client | observation PUB/SUB | One or more JPEG images, measured joints, task, session metadata | RoboJuDo pipeline binds; client connects with `--robot-endpoint` | +| `8559` | deploy client → robot | command PUB/SUB | Joint targets plus `[vx, vy, yaw_rate, height]` | deploy client binds with `--command-endpoint`; robot pipeline connects via `--gr00t-command-endpoint` | + +```text +┌─────────────────────────┐ observation: PUB :8561 ┌────────────────────────────┐ +│ RoboJuDo robot pipeline │ ─────────────────────────▶ │ Deployment client │ +│ scripts/run_pipeline.py │ │ run_robojudo_client.py │ +│ connects SUB :8559 │ ◀───────────────────────── │ binds PUB :8559 │ +└─────────────────────────┘ command: :8559 └──────────────┬─────────────┘ + │ policy RPC :5555 + ▼ + ┌────────────────────────────┐ + │ GR00T policy server │ + │ run_gr00t_server.py │ + └────────────────────────────┘ +``` + +These are independent of the GR00T policy-server port (`5555`). In a same-host setup, use +`tcp://127.0.0.1:8561` for the client observation endpoint and `tcp://127.0.0.1:8559` for the +pipeline command endpoint. Do not reverse the two ports: `8561` carries observations and `8559` +carries commands. For a multi-host setup, replace `127.0.0.1` with the host running the relevant +publisher; keep the deployment client's `--command-endpoint tcp://*:8559` bind address unchanged. + +The two profiles are intentionally separate: + +| Profile | State | Action | Config | +| --- | --- | --- | --- | +| G1 23-DoF | 5 left-arm + 5 right-arm + 10 left-hand + 10 right-hand joints | 30 joint targets + `vx`, `vy`, yaw rate, height | `robojudo_g1_23dof_config.py` | +| X2 | 7 left-arm + 7 right-arm joints | 14 joint targets + `vx`, `vy`, yaw rate, height | `robojudo_x2_config.py` | + +Arm targets are trained as actions relative to the measured joint state. G1 dexterous-hand, +navigation, and height targets remain absolute. X2 hand groups are reserved in the deployment +profile but are omitted from its policy modalities until hand telemetry and commands are connected. +Both profiles use a 16-frame action horizon and the episode task text as language input. + +## Prepare G1 23-DoF data + +Convert the recorder output in the conversion helper's isolated environment: + +```bash +uv run --project scripts/lerobot_conversion \ + python scripts/lerobot_conversion/convert_v3_to_v2.py \ + --repo-id g1_23dof_upper_body \ + --root /home/breeze/Desktop/workplace/Humanoid/RoboJuDo-Plus/record_data +``` + +The converter moves the original directory to `g1_23dof_upper_body_v3.0` and places the v2.1 +dataset at the original path. Install the G1 modality file: + +```bash +cp examples/RoboJuDo/g1_23dof_modality.json \ + /home/breeze/Desktop/workplace/Humanoid/RoboJuDo-Plus/record_data/g1_23dof_upper_body/meta/modality.json +``` + +Fine-tune a separate G1 checkpoint: + +```bash +CUDA_VISIBLE_DEVICES=0 NUM_GPUS=1 uv run bash examples/finetune.sh \ + --base-model-path nvidia/GR00T-N1.7-3B \ + --dataset-path /home/breeze/Desktop/workplace/Humanoid/RoboJuDo-Plus/record_data/g1_23dof_upper_body \ + --modality-config-path examples/RoboJuDo/robojudo_g1_23dof_config.py \ + --embodiment-tag NEW_EMBODIMENT \ + --output-dir /tmp/robojudo_g1_23dof_finetune +``` + +## Prepare X2 data + +Convert the X2 recorder output: + +```bash +uv run --project scripts/lerobot_conversion \ + python scripts/lerobot_conversion/convert_v3_to_v2.py \ + --repo-id x2_upper_body \ + --root /home/breeze/Desktop/workplace/Humanoid/RoboJuDo-Plus/record_data +``` + +Install the X2 modality file: + +```bash +cp examples/RoboJuDo/x2_modality.json \ + /home/breeze/Desktop/workplace/Humanoid/RoboJuDo-Plus/record_data/x2_upper_body/meta/modality.json +``` + +Fine-tune a separate X2 checkpoint: + +```bash +CUDA_VISIBLE_DEVICES=0 NUM_GPUS=1 uv run bash examples/finetune.sh \ + --base-model-path nvidia/GR00T-N1.7-3B \ + --dataset-path /home/breeze/Desktop/workplace/Humanoid/RoboJuDo-Plus/record_data/x2_upper_body \ + --modality-config-path examples/RoboJuDo/robojudo_x2_config.py \ + --embodiment-tag NEW_EMBODIMENT \ + --output-dir /tmp/robojudo_x2_finetune +``` + +## Policy interface + +The trained policies expect observations grouped as follows: + +```python +observation = { + "video": { + "ego_view": head_images, + # Present for a G1 checkpoint trained with the mulcam config. + "left_wrist_view": left_wrist_images, + "right_wrist_view": right_wrist_images, + }, + "state": { + "left_arm": left_joint_positions, + "right_arm": right_joint_positions, + # Present for G1; currently omitted for X2. + "left_hand": left_hand_joint_positions, + "right_hand": right_hand_joint_positions, + }, + "language": {"task": [[instruction]]}, +} +``` + +The G1 policy returns six action groups: `left_arm`, `right_arm`, `left_hand`, `right_hand`, +`navigate_command`, and `base_height_command`. X2 continues to return the original four groups +without hands. `navigate_command` is ordered as `[vx, vy, yaw_rate]`. The decoded arm outputs are +absolute joint targets because the policy converts the learned relative actions back using the +current state. + +X2 and G1 have different state/action dimensions. Do not mix them in one `NEW_EMBODIMENT` +training run or use one robot's checkpoint for the other. + +## Deploy + +Start the policy server with the checkpoint produced by the matching training run. The processor +saved inside the checkpoint contains the custom `NEW_EMBODIMENT` modality configuration, so no +extra config path is needed at inference time: + +```bash +uv run python gr00t/eval/run_gr00t_server.py \ + --model-path /tmp/robojudo_g1_23dof_finetune/checkpoint-10000 \ + --embodiment-tag NEW_EMBODIMENT \ + --host 0.0.0.0 \ + --port 5555 +``` + +Start the matching coupled arm and locomotion pipeline in RoboJuDo-Plus. The optional task +override is published with every observation: + +```bash +cd /home/breeze/Desktop/workplace/Humanoid/RoboJuDo-Plus +python scripts/run_pipeline.py \ + -c g1_23_gr00t_locomanipulation_default_real \ + --gr00t-task "pick up the red cup" +``` + +G1 uses a RealSense camera by default. X2 uses the configured ROS2 compressed image topic. In +both cases, `Gr00tZmqCtrl` publishes a msgpack/JPEG multipart observation stream on port 8561; +the deploy client does not open the robot camera itself. + +Single-camera observations use protocol v1 with two multipart frames: + +```text +[msgpack header, ego_view JPEG] +``` + +The G1 multi-camera deployment uses protocol v2. The publisher must send the header and all three +JPEGs atomically in the declared order: + +```text +[msgpack header, ego_view JPEG, left_wrist_view JPEG, right_wrist_view JPEG] +``` + +The v2 header adds the following fields while retaining all v1 session, task, joint-name, and +joint-position fields: + +```python +{ + "protocol_version": 2, + "image_keys": ["ego_view", "left_wrist_view", "right_wrist_view"], + "image_shapes": { + "ego_view": [480, 640, 3], + "left_wrist_view": [480, 640, 3], + "right_wrist_view": [480, 640, 3], + }, +} +``` + +Do not publish each camera as a separate ZMQ message: a policy observation must contain a coherent +set of views. The robot publisher should skip an observation when a required camera has no usable +frame, and should enforce an application-appropriate maximum timestamp skew between the views. + +On the deploy machine, run the subscriber/client after starting the policy server: + +```bash +uv run python examples/RoboJuDo/run_robojudo_client.py \ + --profile g1_23dof \ + --robot-endpoint tcp://:8561 \ + --policy-host \ + --policy-port 5555 \ + --status-interval 5 +``` + +Use `--profile x2` for X2. The client validates the profile and exact joint order before inference. +One thread receives the latest RoboJuDo observation, one performs policy inference, and the command +loop keeps publishing at 30 Hz. Select the action scheduler with `--execution-mode`. + +For a G1 checkpoint trained with `robojudo_g1_23dof_mulcam_config.py`, select the v2 three-camera +layout explicitly: + +```bash +uv run python examples/RoboJuDo/run_robojudo_client.py \ + --profile g1_23dof \ + --camera-layout mulcam \ + --robot-endpoint tcp://:8561 \ + --policy-host \ + --policy-port 5555 \ + --command-endpoint tcp://*:8559 \ + --execution-mode rtc \ + --execution-horizon 8 +``` + +The default `--camera-layout single` remains compatible with protocol v1 and single-camera +checkpoints. A layout mismatch fails before policy inference instead of silently omitting a view. + +### Execution modes + +The default preserves the original **asynchronous double-buffer** behavior: + +```bash +uv run python examples/RoboJuDo/run_robojudo_client.py \ + --profile x2 \ + --robot-endpoint tcp://127.0.0.1:8561 \ + --policy-host 127.0.0.1 \ + --policy-port 5555 \ + --command-endpoint tcp://*:8559 \ + --execution-mode double_buffer \ + --execution-horizon 8 +``` +在双缓冲模式下, execution_horizon 表示每次连续执行多少步: +``` +chunk A 执行 N 步 + → 切换 chunk B +``` + + +**ACT-style Temporal Ensemble** continuously infers new chunks and blends predictions that cover the +same 30 Hz control tick: + +```bash +uv run python examples/RoboJuDo/run_robojudo_client.py \ + --profile x2 \ + --robot-endpoint tcp://127.0.0.1:8561 \ + --policy-host 127.0.0.1 \ + --policy-port 5555 \ + --command-endpoint tcp://*:8559 \ + --execution-mode temporal_ensemble \ + --execution-horizon 16 \ + --temporal-ensemble-coeff 0.01 +``` + +`--temporal-ensemble-coeff 0` gives a direct average. The ACT value `0.01` exponentially gives +slightly more weight to older predictions. + +**Real-Time Chunking (RTC)** continuously infers while executing the current 16-step chunk. It +re-anchors the unexecuted physical arm targets against the latest measured joints, sends that +normalized prefix into the flow sampler, and replaces the queue after skipping the actual inference +delay: + +```bash +uv run python examples/RoboJuDo/run_robojudo_client.py \ + --profile x2 \ + --robot-endpoint tcp://127.0.0.1:8561 \ + --policy-host 127.0.0.1 \ + --policy-port 5555 \ + --command-endpoint tcp://*:8559 \ + --execution-mode rtc \ + --execution-horizon 8 \ + --rtc-prefix-schedule exp \ + --rtc-max-guidance-weight 10 +``` + +In RTC mode, `--execution-horizon` is the end of the prefix guidance window, not the number of +actions returned by the model; all 16 actions remain available to the queue. The estimated frozen +prefix is based on the maximum of the last `--rtc-latency-window` inference times. The queue is +ultimately sliced using the measured control-tick delay, so an inaccurate estimate changes guidance +strength but does not cause already-expired actions to execute. Use `double_buffer` as the rollback +mode while tuning RTC on hardware. + +Temporal Ensemble assigns each prediction to the command tick at which inference started. If a +request made at tick 0 returns at tick 3, actions 0 through 2 have already expired and action 3 is +the first eligible result. With later overlapping chunks, the time-aligned diagonal is averaged: + +```text +30 Hz tick 0 1 2 3 4 5 6 + +infer chunk A [--------- inference -------->] +A prediction time A0 A1 A2 A3 A4 A5 A6 +published A - - - A3 A4 A5 A6 + +infer chunk B [--------- inference -------->] +B prediction time B0 B1 B2 B3 +published ensemble - - - A3 A4 A5 avg(A6,B3) +``` + +**在 Temporal Ensemble 模式下, execution_horizon不再表示“连续执行 N 步”,而每个 chunk 在 ensemble 时间轴上的有效长度.** +``` +当前默认的参数是: + 有效 execution_horizon H = 16 tick + 推理间隔 Q ≈ 2~3 tick + 推理延迟 D ≈ 2~3 tick + + 由于一个 chunk 返回时已经过去约 2~3 步,它实际能参与 ensemble 的剩余长度约为: + + H - D = 16 - 2~3 = 13~14 tick + + 每隔约 2~3 tick 又产生一个新 chunk,因此稳态 contributor 数量大约是: + + N ≈ (H - D) / Q + ≈ 13.5 / 2.3 + ≈ 5.9 + + 日志中就可以看到ensemble_chunks=6, ensemble_contributors=6 + +按当前约 2.5 tick 一次推理估算: + + execution horizon 最大时间跨度 contributor 数量 特点 + ━━━━━━━━━━━━━━━━━━━ ━━━━━━━━━━━━━━ ━━━━━━━━━━━━━━━━━━ ━━━━━━━━━━━━━━━━━━━━━━ + 8 267 ms 约 2–3 响应更快,平滑较弱 + ─────────────────── ────────────── ────────────────── ────────────────────── + 12 400 ms 约 4 中间选择 + ─────────────────── ────────────── ────────────────── ────────────────────── + 16 533 ms 约 5–6 平滑较强,可能更滞后 + + 最大时间跨度是 chunk 从 query 开始计算的。16 步对应: + + (16 - 1) / 30 = 500 ms + + 最老的有效预测可能基于约 0.5 秒前的 observation +``` + + + +当前模型训练default配置为: + `delta_indices=list(range(16))` + 即模型预测 16 步、约 533 ms 的未来动作。 + + Temporal Ensemble 建议使用全部 16 步,而不是当前截断后的 8 步: 推理延迟约 3 步, 返回后仍剩 13 步可以参与 ensemble, 下一次推理约 3 步后返回 +, 通常会有多个 chunk 重叠 + +--temporal-ensemble-coeff 0.01 的权重比较温和。例如最老和最新预测相差 14 tick 时: +```py + oldest weight = 1.0 + newest weight = exp(-0.01 × 14) ≈ 0.87 +``` + 所以目前接近均匀平均,只是稍微偏向旧预测。 + + 接下来主要需要观察真机动作效果: + - 如果抖动明显减少且响应速度正常:保留 0.01。 + - 如果仍有小幅抖动:可以试 0,直接平均通常会更平滑一些。 + - 如果动作明显滞后:可能是 16 步历史预测参与过多,需限制 ensemble 历史长度,而不是盲目增大 coefficient。 + - 如果动作太依赖旧意图:可以降低最大 contributor 数量,例如只保留最近 3–4 个 chunk。 + +This is temporal alignment, not RTC: expired indices are skipped, but the model is not re-run or +corrected for measured inference delay. Unlike LeRobot ACT's online implementation, which assumes +one inference result **every control step**, this client retains a small set of absolute-tick chunks so +the same diagonal weighting remains valid when GR00T returns a chunk every few control steps. + +### Double-buffer execution + +#### Logic + +The deploy client has three concurrent parts: + +1. The observation thread receives the camera image, measured upper-body joints, and task from + RoboJuDo. It drains already queued ZeroMQ messages and keeps only the newest valid observation, + so inference does not work through a backlog of old camera frames. +2. The inference thread sends the newest observation that has not already been inferred to the + GR00T policy server. The returned commands are stored in a single `pending` chunk. While that + slot is occupied, the thread does not request another chunk. +3. The command loop owns the `active` chunk and publishes one command from it every + `1 / --command-fps` seconds. The default command rate is 30 Hz. + +The buffer transition is: + +```text +latest observation + | + v +GR00T inference ---> pending chunk + | + | active chunk is empty + v + active chunk ---> command[0], command[1], ... + | + | activating it frees the pending slot + v + infer the newest observation while active executes +``` + +The same flow on a time axis looks like this (`A0` means the first command in chunk A): + +```text +time / 30 Hz tick ---> 0 1 2 3 4 5 6 + +observation thread O0 O1 O2 O3 O4 O5 O6 +latest observation O0 O1 O2 O3 O4 O5 O6 + +inference thread [------ infer chunk A ------] [------ infer chunk B ------] +pending slot empty empty empty A empty empty B + transfer A wait for A + | to finish + v +active queue empty empty empty A0..A3 A1..A3 A2..A3 A3 +published command none none none A0 A1 A2 A3 +``` + +Inference B can run while A is active because transferring A clears the pending slot. If B becomes +ready before A finishes, B waits in the pending slot. The command loop still consumes all of A and +then switches directly to B; predictions from A and B are not blended. + +At startup, publishing waits for takeover to be enabled, a fresh observation from that control +session, and its first inferred chunk. Once a pending chunk is activated, its commands are copied +into the active queue and the pending slot is cleared. This immediately permits inference of the +newest observation while the active commands continue to execute. If that inference finishes +early, its result waits in the pending slot until the active queue is exhausted. + +A shared `threading.Condition` coordinates this handoff. The inference thread waits on the +condition until the pending slot is empty and a newer observation is available, then stores its +result in the pending slot while holding the condition lock. When the command loop moves that +pending chunk into its local active queue, it clears the pending slot and calls `notify_all()`, +waking the inference thread so it can start the next request. The active queue itself is accessed +only by the command loop, so it does not need separate locking. + +RoboJuDo publishes `takeover_enabled`, `control_session`, and a process-unique `stream_id` with +every observation. Disabling takeover leaves the camera stream running but makes the deploy client +clear active, pending, and held commands. Each disabled-to-enabled transition increments +`control_session`; the client then starts inference from a fresh observation for that session. An +inference result from an older session is discarded even if it returns after the new session has +started. A changed `stream_id` provides the same reset boundary when the whole RoboJuDo pipeline is +restarted. Every command carries the same `stream_id` and `control_session`; RoboJuDo rejects a +command unless it belongs to the currently enabled session, so an old buffered command cannot be +applied during an enable transition. + +In `double_buffer` mode, chunk replacement happens only at the boundary: the client executes every +command in the active chunk, then activates the pending chunk. It does not align or average +overlapping predictions. A pending chunk is also not replaced by a newer prediction while it is +waiting. + +If the active horizon finishes before another chunk is ready, the client keeps publishing the last +command. This includes its locomotion values, so a non-zero velocity command is held until a new +chunk arrives, the observation stream times out, or the robot-side watchdog stops it. If the +observation age exceeds `--observation-timeout`, the client clears both active and pending actions, +but keeps publishing a safe hold command at `--command-fps`: upper-body joint positions and base +height remain at their last commanded values, while `vx`, `vy`, and `yaw_rate` are set to zero. If +no command has been published yet, there is no previous pose to hold and publishing remains idle. + +In `temporal_ensemble` mode, an uncovered tick instead holds the last arm positions and base height +while forcing `vx`, `vy`, and `yaw_rate` to zero. Takeover disable, session changes, stream changes, +and observation timeout clear all ensemble history; a result returning from an old session is +discarded before it can enter the ensemble. + +### Port binding +The two robot-side ports have opposite directions: + +```text +RoboJuDo PUB tcp://*:8561 -> deploy SUB camera + measured upper joints + task +RoboJuDo SUB deploy:8559 <- deploy PUB upper targets + velocity/height command +``` + +Configure RoboJuDo's command `endpoint` with the deploy machine IP when they run on different +hosts. First confirm that observations are fresh and the command subscriber is connected, then +enter `RL_DEFAULT` and enable upper-body takeover. Enabling creates a new control session; inference +and command publishing begin from its first fresh observation. The controller rejects incomplete, +replayed, non-finite, or wrong-session commands atomically and stops locomotion when the command +watchdog expires. + +### Deployment health logs + +The deploy client prints throttled health reports every `--status-interval` seconds. The default is +5 seconds; use a shorter interval while diagnosing a connection: + +```bash +uv run python examples/RoboJuDo/run_robojudo_client.py \ + --profile x2 \ + --robot-endpoint tcp://127.0.0.1:8561 \ + --policy-host 127.0.0.1 \ + --policy-port 5555 \ + --command-endpoint tcp://127.0.0.1:8559 \ + --execution-horizon 16 \ + --status-interval 2 +``` + +The reports cover each stage of the deployment loop: + +```text +[observation] rate=29.8Hz, received=60, last_sequence=412, sequence_gaps=2 +[command] subscriber connected: tcp://127.0.0.1:8559 +[control] takeover enabled for session 0123abcd:1; waiting for a fresh chunk +[inference] chunk ready: observation_sequence=412, session=1, actions=16, latency=0.184s +[command] activated chunk: observation_sequence=412, session=1, actions=16, age=0.190s, inference=0.184s +[command] rate=30.0Hz, published=60, next_sequence=900, subscriber_connected=True, ... +``` + +`sequence_gaps` counts publisher sequence numbers skipped by the latest-frame subscriber. Some gaps +are expected because the subscriber intentionally drains queued observations before decoding the +newest frame; sustained growth together with a low observation rate indicates that transport or +JPEG decoding cannot keep up. + +`subscriber_connected=True` is a ZeroMQ transport connection signal, not an application-level +acknowledgement that RoboJuDo applied a command. Command application and watchdog state remain +visible in the RoboJuDo process logs. If observations stop, the client reports their age, clears +pending actions, and after `--observation-timeout` continuously publishes the last upper-body and +base-height targets with `vx`, `vy`, and `yaw_rate` set to zero. If inference is late, the client +reports that the action horizon is exhausted and holds the last command until a fresh chunk arrives. diff --git a/examples/RoboJuDo/RTC.md b/examples/RoboJuDo/RTC.md new file mode 100644 index 000000000..cda5278dd --- /dev/null +++ b/examples/RoboJuDo/RTC.md @@ -0,0 +1,599 @@ +# RoboJuDo Inference-time RTC + +本文档说明 RoboJuDo 中已经实现的 inference-time Real-Time Chunking(RTC)。实现参考 +[LeRobot RTC](https://github.com/huggingface/lerobot/blob/main/src/lerobot/policies/rtc/modeling_rtc.py), +使用现有 16-step RoboJuDo checkpoint,不需要重新训练。 + +当前实现范围: + +- 支持 `examples/RoboJuDo/run_robojudo_client.py --execution-mode rtc`。 +- 支持 X2 和 G1 23-DoF RoboJuDo profile。 +- 当前 RTC batch size 固定为 1,并保持单个 inference in-flight。 +- `double_buffer` 和 `temporal_ensemble` 的原有行为保持不变。 +- N1.7 action head 中原有的 inpainting RTC 分支仍作为 legacy 路径保留;RoboJuDo 使用新的梯度 guidance 路径。 + +## 启动方式 + +先启动 Policy Server: + +```bash +uv run python gr00t/eval/run_gr00t_server.py \ + --model-path checkpoints/x2_move_box_center_test/ \ + --embodiment-tag NEW_EMBODIMENT \ + --host 127.0.0.1 \ + --port 5555 +``` + +再启动 RTC client: + +```bash +uv run python examples/RoboJuDo/run_robojudo_client.py \ + --profile x2 \ + --robot-endpoint tcp://127.0.0.1:8561 \ + --policy-host 127.0.0.1 \ + --policy-port 5555 \ + --command-endpoint tcp://*:8559 \ + --execution-mode rtc \ + --execution-horizon 8 \ + --rtc-prefix-schedule exp \ + --rtc-max-guidance-weight 10 +``` + +RTC 参数: + +- `--execution-horizon`:RTC guidance 结束位置 `H`,默认 8;模型仍生成并排队 checkpoint + 定义的完整 action chunk。客户端会从第一次 policy 返回动态读取 action horizon。 +- `--rtc-prefix-schedule`:`zeros`、`ones`、`linear` 或 `exp`,默认 `exp`。 +- `--rtc-max-guidance-weight`:denoising correction 最大增益,默认 10。 +- `--rtc-latency-window`:估计延迟时保留的最近推理耗时数量,默认 10。 +- `--command-fps`:控制频率,默认 30 Hz。 + +## 整体数据流 + +```text +RoboJuDo observation(图像、当前关节、任务) + ↓ +run_robojudo_client.py + 截取旧队列中尚未执行的物理动作 + 估计推理延迟 D_est + ↓ +deploy_adapter.py + 构造 GR00T observation + 通过 PolicyClient options 发送 RTC prefix + ↓ +gr00t_policy.py + 物理 absolute prefix + 最新 state + → relative/absolute 动作处理 + → normalization + ↓ +gr00t_n1d7.py + 在每个 flow denoising step 中执行 RTC guidance + ↓ +gr00t_policy.py + 模型动作 denormalize + relative 双臂动作恢复成物理 absolute target + ↓ +deploy_adapter.py + 同时返回完整物理 action chunk 和 RoboJuDo commands + ↓ +run_robojudo_client.py + 根据实际延迟 D_actual 跳过新 chunk 前若干步 + 原子替换执行队列 +``` + +最重要的动作空间关系: + +```text +客户端 RTC 队列 + 保存物理 absolute action + ↓ 结合最新 observation 重新处理 +模型 RTC prefix + 使用 normalized model-space action + +双臂:物理 absolute target → 相对最新关节状态转换 → normalize +导航:保持 absolute → normalize +高度:保持 absolute → normalize +``` + +客户端没有把 raw absolute action 直接送进 flow sampler。保存 absolute target 的目的是在下一轮推理时, +能根据机器人最新实测状态重新计算 relative action。 + +## 实际代码修改 + +### `examples/RoboJuDo/deploy_adapter.py` + +该文件负责 RoboJuDo 与 GR00T Policy 之间的格式转换,不负责队列调度和 RTC 数学。 + +新增 `PolicyActionChunk`: + +```python +@dataclass(frozen=True) +class PolicyActionChunk: + actions: dict[str, np.ndarray] + commands: list[dict[str, Any]] + info: dict[str, Any] +``` + +一次推理结果同时保留两种表示: + +```text +actions + 分组物理动作,shape 为 (1, 16, D) + 用于下一轮 RTC prefix + +commands + RoboJuDo 可直接执行的 positions + locomotion_command + 用于当前控制循环 +``` + +`actions` 包含: + +```python +{ + "left_arm": (1, 16, arm_dim), + "right_arm": (1, 16, arm_dim), + "navigate_command": (1, 16, 3), + "base_height_command": (1, 16, 1), +} +``` + +主要接口: + +- `_validate_action_chunk()` 验证四个动作组、batch、horizon、动作维度和有限值。 +- `decode_action_chunk()` 将每一步转换成 RoboJuDo command。 +- `get_action_chunk()` 将 RTC options 传给 `PolicyClient.get_action()`,并返回完整物理 chunk。 +- 原来的 `get_action()` 保留,内部调用 `get_action_chunk()` 后只返回 commands,供其他 execution mode 使用。 + +RTC 调用 `get_action_chunk(execution_horizon=None)`,因此保存模型返回的完整 chunk,而不是只截取 +`--execution-horizon` 步。chunk 长度来自 checkpoint 对应 embodiment 的 action +`delta_indices`,客户端不再将其固定为 16。 + +### `examples/RoboJuDo/run_robojudo_client.py` + +该文件负责异步推理、RTC 队列、延迟估计、实际延迟切片和控制安全。 + +`ActionChunk` 新增: + +```text +physical_actions + 完整分组物理动作,用于下一次 guidance + +task + 防止跨任务复用 prefix + +estimated_delay_steps + 本轮 sampler 使用的 D_est +``` + +新增 `RTCActionQueue`,内部只保存: + +```python +self._chunk +self._next_index +``` + +`_next_index` 同时表示: + +- 下一条要执行的 command。 +- 下一条尚未执行的 physical action。 + +主要操作: + +```text +pop() + 返回 commands[next_index] + next_index += 1 + +get_left_over() + 返回 physical_actions[:, next_index:] + 同时校验 stream、control session 和 task + +replace(new_chunk, D_actual) + 保存新 chunk + next_index = D_actual +``` + +推理开始时,推理线程在同一把 condition lock 下完成: + +```text +固定最新 observation +记录 query_tick +固定 stream / control session / task +截取尚未执行的 physical prefix +根据最近延迟计算 D_est +``` + +延迟估计公式: + +```text +D_est = ceil(max(recent_inference_latency) / command_period) + +等价于: +D_est = ceil(max(recent_inference_latency) × command_fps) +``` + +然后裁剪到: + +```text +0 <= D_est <= min(prefix_length, execution_horizon) +``` + +发送给服务端的 RTC options: + +```python +{ + "rtc": { + "prefix_actions": { + "left_arm": ..., + "right_arm": ..., + "navigate_command": ..., + "base_height_command": ..., + }, + "prefix_length": L, + "estimated_delay_steps": D_est, + "guidance_horizon": min(H, L), + "prefix_schedule": "exp", + "max_guidance_weight": 10.0, + } +} +``` + +推理返回后重新读取 control tick: + +```text +D_actual = ready_tick - query_tick +``` + +`D_est` 和 `D_actual` 的职责不同: + +```text +D_est + 告诉模型哪些 prefix step 应该强约束 + +D_actual + 告诉客户端新 chunk 的哪些 step 已经过期 +``` + +新队列从 `new_chunk[D_actual]` 开始执行。第一次推理没有旧 prefix,因此 `skipped_steps=0`,不会错误地 +跳过第一批动作。 + +控制安全: + +- observation timeout、takeover disabled、stream 改变、control session 改变或 task 改变时清空 RTC 队列。 +- in-flight 结果若属于旧 session 或旧 task,会在进入队列前被丢弃。 +- RTC 队列耗尽时保持双臂和高度,并将 `vx`、`vy`、`yaw_rate` 置零。 +- 若 `D_actual >= 16`,整段新预测已经过期,丢弃该 chunk 并等待无 prefix 刷新。 + +### `gr00t/policy/gr00t_policy.py` + +该文件负责把客户端传来的物理 absolute prefix 转换成模型使用的 normalized action tensor。 + +`_to_vla_step_data()` 现在可以接收 actions: + +```python +VLAStepData( + states=current_state, + actions=physical_prefix, + ..., +) +``` + +`_prepare_rtc_options()` 负责: + +- 当前只允许 batch size 1。 +- 校验 action groups、`float32`、shape `(1, T, D)` 和有限值。 +- 校验 `0 <= D_est <= H <= action_horizon`。 +- 从 model options 中移除大型 `prefix_actions` 数组。 +- 把物理 prefix 放入 `VLAStepData.actions`,让现有 processor 处理。 + +RoboJuDo 的动作配置是: + +```text +left_arm RELATIVE +right_arm RELATIVE +navigate_command ABSOLUTE +base_height_command ABSOLUTE +``` + +例如旧队列保存的双臂物理目标为 35°: + +```text +最新实测关节为 30° → model prefix = +5° +最新实测关节为 28° → model prefix = +7° +``` + +物理目标仍然是 35°,但 relative 表示根据最新实测状态重新锚定。随后 processor 使用 checkpoint +statistics 归一化动作,并生成 action padding 和 action mask。 + +#### 短 prefix 与 16-step statistics + +RoboJuDo checkpoint 的双臂 relative-action statistics 是按时间步保存的: + +```text +left_arm relative statistics (16, 7) +right_arm relative statistics (16, 7) +``` + +RTC 运行一轮后,leftover prefix 通常短于 16。例如实际推理延迟为 3: + +```text +原始 chunk 16 steps +跳过 3 steps +leftover 13 steps +``` + +如果直接把 `(13, 7)` 交给 `(16, 7)` statistics 归一化,会产生 shape mismatch。当前实现会在进入 +processor 前重复最后一个物理目标,补齐到完整 horizon: + +```text +真实 prefix: A3 A4 ... A15 共 13 步 +processor: A3 A4 ... A15 A15 A15 A15 补齐到 16 步 +``` + +真实 `prefix_length` 仍为 13,sampler 会将 13 步之后的 guidance 权重设为 0;padding 只用于满足 +per-timestep normalization shape,不会扩大有效 RTC prefix。 + +普通推理继续使用: + +```python +with torch.inference_mode(): +``` + +RTC 推理使用: + +```python +with torch.no_grad(): +``` + +原因是 `torch.inference_mode()` 不能被内部 `torch.enable_grad()` 覆盖,而 RTC sampler 需要局部 +autograd correction。backbone 和普通推理部分仍然不构建梯度。 + +模型输出后,`processor.decode_action()` 会: + +```text +denormalize +双臂 relative action → 当前状态下的物理 absolute target +导航和高度保持 absolute +``` + +因此形成闭环: + +```text +normalized model output + ↓ decode +physical absolute queue + ↓ 下一轮结合最新 state +normalized model prefix +``` + +### `gr00t/model/gr00t_n1d7/gr00t_n1d7.py` + +该文件实现每个 flow denoising step 中的 RTC correction。 + +`get_rtc_prefix_weights()` 构造时间权重。设: + +```text +D_est = 3 +H = 8 +T = 16 +``` + +则: + +```text +step 0 1 2 3 4 5 6 7 8 ... 15 +weight 1 1 1 ↓ ↓ ↓ ↓ ↓ 0 ... 0 + └ 强约束 ┘ └ 平滑接管区域 ┘ └ 自由生成 ┘ +``` + +不同 schedule: + +- `zeros`:只有 `[0, D_est)` 为 1。 +- `ones`:`[0, H)` 全部为 1。 +- `linear`:`[D_est, H)` 线性衰减。 +- `exp`:`[D_est, H)` 按 LeRobot exponential 形状衰减。 + +时间权重还会乘 processor 的 action mask,屏蔽无效 action dimension。`guidance_horizon` 会裁剪到真实 +`prefix_length`,因此短 prefix 的 padding step 不参与 guidance。 + +GR00T flow 的时间方向是: + +```text +t = 0 noise +t = 1 clean action +``` + +每个 denoising step 首先预测基础 velocity: + +```text +v_t = model(x_t, observation, t) +``` + +当前 latent 的 clean-action estimate: + +```text +x_clean = x_t + (1 - t) × v_t +``` + +计算加权 prefix error: + +```text +error = (prefix - x_clean) × prefix_weights +``` + +使用局部 autograd 计算 correction: + +```text +correction = grad(x_clean, x_t, grad_outputs=error) +``` + +与 LeRobot 当前实现一致,本次 denoiser 输出 `v_t` 在 correction 中视为固定值,不通过整个 denoiser +构建反向图。然后修正 velocity: + +```text +v_guided = v_t + guidance_gain(t) × correction +``` + +最后继续正常 Euler integration: + +```text +x_next = x_t + dt × v_guided +``` + +每一步结束后 detach latent,避免多个 denoising step 的计算图连接起来。 + +LeRobot reverse-time guidance gain 已转换到 GR00T 的正向 flow convention: + +```text +gain(t) = ((1 - t)² + t²) / ((1 - t) × t) +gain(t) = min(gain(t), max_guidance_weight) +``` + +## 完整 RTC 时序可视化 + +下面假设: + +- 控制频率 30 Hz。 +- 每次推理耗时 3 tick。 +- `D_est = 3`。 +- `guidance_horizon H = 8`。 +- 每个模型 chunk 实际仍为 16 步,为了可读性只画前几步。 + +```text +30 Hz tick 0 1 2 3 4 5 6 7 8 9 + +infer chunk A [--------- inference -------->] +A prediction time A0 A1 A2 A3 A4 A5 A6 +published action - - - A0 A1 A2 + +infer chunk B [--------- inference -------->] +old prefix for B A0 A1 A2 A3 A4 A5 A6 +B prediction time B0 B1 B2 B3 B4 B5 B6 +RTC guidance 强 强 强 渐弱 渐弱 渐弱 渐弱 +published action - - - A0 A1 A2 B3 B4 B5 + +infer chunk C [--------- inference -------->] +old prefix for C B3 B4 B5 B6 +C prediction time C0 C1 C2 C3 +RTC guidance 强 强 强 渐弱 +published action - - - A0 A1 A2 B3 B4 B5 C3 +``` + +第一次推理 A 没有旧 prefix: + +```text +tick 0~2 A 正在推理,没有可执行动作 +tick 3 A 返回;因为没有 prefix,所以从 A0 开始执行 +``` + +推理 B 在 tick 3 发起,快照中的旧 prefix 是: + +```text +prefix_B = [A0, A1, A2, A3, ...] +``` + +B 推理期间控制线程继续执行: + +```text +tick 3 → A0 +tick 4 → A1 +tick 5 → A2 +``` + +B 在 tick 6 返回: + +```text +D_actual = ready_tick - query_tick + = 6 - 3 + = 3 +``` + +因此 `B0`、`B1`、`B2` 对应的时间已经过去,新队列从 `B3` 开始: + +```text +旧执行序列: A0 → A1 → A2 +新执行序列: B3 → B4 → B5 +最终发布: A0 → A1 → A2 → B3 → B4 → B5 → C3 ... +``` + +模型内部的 prefix 对齐: + +```text +旧 prefix A0 A1 A2 A3 A4 A5 A6 A7 + │ │ │ │ │ │ │ │ + ▼ ▼ ▼ ▼ ▼ ▼ ▼ ▼ +新 chunk B0 B1 B2 B3 B4 B5 B6 B7 +weight 1.0 1.0 1.0 ↓ ↓ ↓ ↓ ↓ + └── D_est=3 ──┘ └──── transition,直到 H=8 ────────────┘ +``` + +RTC 与 Temporal Ensemble 的区别: + +```text +Temporal Ensemble + A 和 B 独立生成 + 客户端按同一 control tick 做 avg(A, B) + +RTC + 生成 B 的过程中已经使用 A 作为 prefix guidance + 客户端按 D_actual 跳过过期 step 后直接执行 B3 +``` + +## 日志解读 + +示例: + +```text +[inference] chunk ready: actions=16, latency=0.190s, +query_tick=0, ready_tick=5, skipped_steps=0, +rtc_prefix=0, rtc_estimated_delay=0 + +[inference] chunk ready: actions=16, latency=0.101s, +query_tick=5, ready_tick=8, skipped_steps=3, +rtc_prefix=16, rtc_estimated_delay=6 +``` + +第一行: + +- 首次推理没有 prefix,所以 `rtc_prefix=0`。 +- 即使推理期间 control tick 增长,首次结果仍从 step 0 开始,因此 `skipped_steps=0`。 + +第二行: + +- 推理开始时旧 chunk 尚有完整 16 步,所以 `rtc_prefix=16`。 +- 首轮 warm-up latency 为 0.190 秒,在 30 Hz 下得到 `ceil(0.190 × 30)=6`,所以 `D_est=6`。 +- 第二轮实际只经过 3 tick,所以客户端使用 `D_actual=3`,从新 chunk 的 step 3 开始执行。 +- 下一次推理看到的 leftover 通常为 `16 - 3 = 13` 步;Policy 会为 normalization 补齐到 16, + 但日志和 sampler 中的真实 `prefix_length` 仍是 13。 + +## 当前限制与调参建议 + +- 16-step checkpoint 可以运行 RTC,但比 32-step chunk 留给 transition/free region 的空间更小。 +- warm-up 推理通常比稳态慢,rolling maximum 可能让最初几轮 `D_est` 偏大,这是保守行为。 +- `D_est` 偏大只会让更多 prefix step 被强约束;实际队列切片始终使用 `D_actual`。 +- 如果动作过于依赖旧轨迹,可减小 `--execution-horizon` 或降低 `--rtc-max-guidance-weight`。 +- 如果 chunk 边界仍明显,可增大 `--execution-horizon`,但不能超过客户端从 policy 首个 chunk + 动态读取到的 action horizon。 +- 实机调参期间可随时切回 `--execution-mode double_buffer` 作为回滚模式。 + +建议重点观察: + +- chunk 切换时双臂 position jump。 +- 速度和加速度突变。 +- `D_est` 与 `skipped_steps` 的长期差异。 +- RTC queue 是否频繁耗尽并进入 safe hold。 +- task、takeover、session 或 stream 切换时是否正确清空历史动作。 + +## 验证覆盖 + +当前测试覆盖: + +- prefix weight 的 frozen、transition 和 free 区域。 +- guidance gain 的边界截断。 +- guidance weight 为 0 时与普通 sampling 等价。 +- Policy 物理 prefix 传递、短 prefix 补齐和 model options 清理。 +- RTC queue 按 `D_actual` 替换、leftover 与 commands 锁步。 +- task/session 不匹配时拒绝复用 prefix。 +- 推理线程发送 leftover prefix,并按实际 tick 跳过新 chunk。 +- PolicyClient/PolicyServer 对包含 NumPy prefix 的 RTC options 往返传输。 +- `double_buffer` 和 `temporal_ensemble` 相关回归行为。 diff --git a/examples/RoboJuDo/deploy_adapter.py b/examples/RoboJuDo/deploy_adapter.py new file mode 100644 index 000000000..ebd25f1e3 --- /dev/null +++ b/examples/RoboJuDo/deploy_adapter.py @@ -0,0 +1,284 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""GR00T policy client adapter for RoboJuDo X2 and G1 23-DoF control.""" + +from __future__ import annotations + +from dataclasses import dataclass +from typing import Any, Mapping, Sequence + +from gr00t.policy.server_client import PolicyClient +import numpy as np + + +@dataclass(frozen=True) +class RobotProfile: + left_arm_joint_names: tuple[str, ...] + right_arm_joint_names: tuple[str, ...] + # Reserved as empty tuples for robots whose dexterous hands are not wired yet. + left_hand_joint_names: tuple[str, ...] = () + right_hand_joint_names: tuple[str, ...] = () + + @property + def joint_groups(self) -> tuple[tuple[str, tuple[str, ...]], ...]: + """Policy joint modalities that are active for this deployment profile.""" + groups = ( + ("left_arm", self.left_arm_joint_names), + ("right_arm", self.right_arm_joint_names), + ("left_hand", self.left_hand_joint_names), + ("right_hand", self.right_hand_joint_names), + ) + return tuple((key, names) for key, names in groups if names) + + @property + def joint_names(self) -> tuple[str, ...]: + return tuple(name for _, names in self.joint_groups for name in names) + + +@dataclass(frozen=True) +class PolicyActionChunk: + """Physical policy actions together with their RoboJuDo command encoding.""" + + actions: dict[str, np.ndarray] # For RTC prefix guidance + commands: list[dict[str, Any]] # For loop execution + info: dict[str, Any] + + +PROFILES = { + "g1_23dof": RobotProfile( + left_arm_joint_names=( + "left_shoulder_pitch_joint", + "left_shoulder_roll_joint", + "left_shoulder_yaw_joint", + "left_elbow_joint", + "left_wrist_roll_joint", + ), + right_arm_joint_names=( + "right_shoulder_pitch_joint", + "right_shoulder_roll_joint", + "right_shoulder_yaw_joint", + "right_elbow_joint", + "right_wrist_roll_joint", + ), + left_hand_joint_names=( + "left_thumb_proximal", + "left_thumb_intermediate", + "left_index_proximal", + "left_middle_proximal", + "left_ring_proximal", + "left_pinky_proximal", + "left_index_intermediate", + "left_middle_intermediate", + "left_ring_intermediate", + "left_pinky_intermediate", + ), + right_hand_joint_names=( + "right_thumb_proximal", + "right_thumb_intermediate", + "right_index_proximal", + "right_middle_proximal", + "right_ring_proximal", + "right_pinky_proximal", + "right_index_intermediate", + "right_middle_intermediate", + "right_ring_intermediate", + "right_pinky_intermediate", + ), + ), + "x2": RobotProfile( + left_arm_joint_names=( + "left_shoulder_pitch_joint", + "left_shoulder_roll_joint", + "left_shoulder_yaw_joint", + "left_elbow_joint", + "left_wrist_yaw_joint", + "left_wrist_pitch_joint", + "left_wrist_roll_joint", + ), + right_arm_joint_names=( + "right_shoulder_pitch_joint", + "right_shoulder_roll_joint", + "right_shoulder_yaw_joint", + "right_elbow_joint", + "right_wrist_yaw_joint", + "right_wrist_pitch_joint", + "right_wrist_roll_joint", + ), + # X2 hand names stay empty until its observation/command transport is connected. + left_hand_joint_names=(), + right_hand_joint_names=(), + ), +} + +CAMERA_LAYOUTS = { + "single": ("ego_view",), + "mulcam": ("ego_view", "left_wrist_view", "right_wrist_view"), +} + + +class RoboJuDoPolicyAdapter: + """Translate RoboJuDo observations and GR00T action chunks without changing units.""" + + def __init__(self, policy_client: PolicyClient, profile: str, + video_keys: Sequence[str] = CAMERA_LAYOUTS["single"], + ): + self.policy_client = policy_client + self.profile = PROFILES[profile] + self.video_keys = tuple(video_keys) + if not self.video_keys or len(set(self.video_keys)) != len(self.video_keys): + raise ValueError("video_keys must be a non-empty sequence of unique names") + + def _ordered_joint_positions( + self, joint_positions: Mapping[str, float] | Sequence[float] | np.ndarray + ) -> np.ndarray: + if isinstance(joint_positions, Mapping): + missing = [name for name in self.profile.joint_names if name not in joint_positions] + if missing: + raise ValueError(f"Missing RoboJuDo joint positions: {missing}") + values = [joint_positions[name] for name in self.profile.joint_names] + else: + values = joint_positions + positions = np.asarray(values, dtype=np.float32) + expected_shape = (len(self.profile.joint_names),) + if positions.shape != expected_shape: + raise ValueError( + f"Joint positions have shape {positions.shape}, expected {expected_shape}" + ) + if not np.isfinite(positions).all(): + raise ValueError("Joint positions contain non-finite values") + return positions + + def build_observation( + self, + images: Mapping[str, np.ndarray] | np.ndarray, + joint_positions: Mapping[str, float] | Sequence[float] | np.ndarray, + instruction: str, + ) -> dict[str, Any]: + if isinstance(images, np.ndarray): + images = {"ego_view": images} + missing = [key for key in self.video_keys if key not in images] + unexpected = [key for key in images if key not in self.video_keys] + if missing or unexpected: + raise ValueError( + f"Video keys do not match deployment layout: missing={missing}, " + f"unexpected={unexpected}" + ) + video = {} + for key in self.video_keys: + image = np.asarray(images[key]) + if image.ndim != 3 or image.shape[-1] != 3 or image.dtype != np.uint8: + raise ValueError(f"Video {key!r} must be an HWC uint8 RGB array") + video[key] = image[None, None] + if not instruction: + raise ValueError("instruction must not be empty") + positions = self._ordered_joint_positions(joint_positions) + return { + "video": video, + "state": self._split_joint_groups(positions), + "language": {"task": [[instruction]]}, + } + + def _split_joint_groups(self, positions: np.ndarray) -> dict[str, np.ndarray]: + groups = {} + start = 0 + for key, names in self.profile.joint_groups: + end = start + len(names) + groups[key] = positions[start:end][None, None] + start = end + assert start == len(positions) + return groups + + def _validate_action_chunk( + self, action_chunk: Mapping[str, np.ndarray] + ) -> tuple[dict[str, np.ndarray], int]: + required = { + **{key: len(names) for key, names in self.profile.joint_groups}, + "navigate_command": 3, + "base_height_command": 1, + } + arrays = {} + for key, width in required.items(): + if key not in action_chunk: + raise ValueError(f"Policy response is missing action group {key!r}") + value = np.asarray(action_chunk[key], dtype=np.float32) + if value.ndim != 3 or value.shape[0] != 1 or value.shape[2] != width: + raise ValueError( + f"Action group {key!r} has shape {value.shape}, expected (1, T, {width})" + ) + if not np.isfinite(value).all(): + raise ValueError(f"Action group {key!r} contains non-finite values") + arrays[key] = value + + available_horizon = min(value.shape[1] for value in arrays.values()) + return arrays, available_horizon + + def decode_action_chunk( + self, action_chunk: Mapping[str, np.ndarray], execution_horizon: int + ) -> list[dict[str, Any]]: + arrays, available_horizon = self._validate_action_chunk(action_chunk) + if not 1 <= execution_horizon <= available_horizon: + raise ValueError( + f"execution_horizon must be in [1, {available_horizon}], got {execution_horizon}" + ) + + commands = [] + for step in range(execution_horizon): + joint_positions = np.concatenate( + [arrays[key][0, step] for key, _ in self.profile.joint_groups] + ) + locomotion_command = np.concatenate( + ( + arrays["navigate_command"][0, step], + arrays["base_height_command"][0, step], + ) + ) + commands.append( + { + "positions": dict( + zip(self.profile.joint_names, joint_positions.tolist(), strict=True) + ), + "locomotion_command": locomotion_command, + } + ) + return commands + + def get_action_chunk( + self, + images: Mapping[str, np.ndarray] | np.ndarray, + joint_positions: Mapping[str, float] | Sequence[float] | np.ndarray, + instruction: str, + *, + execution_horizon: int | None = None, + options: dict[str, Any] | None = None, + ) -> PolicyActionChunk: + """Return the physical chunk and commands; RTC uses all available steps.""" + observation = self.build_observation(images, joint_positions, instruction) + + # Inference, pass the options to the Policy client + action_chunk, info = self.policy_client.get_action(observation, options=options) + + arrays, available_horizon = self._validate_action_chunk(action_chunk) + # When using RTC, the execution_horizon is None, and we use all available steps. + horizon = available_horizon if execution_horizon is None else execution_horizon + actions = {key: value[:, :horizon].copy() for key, value in arrays.items()} + return PolicyActionChunk( + actions=actions, + commands=self.decode_action_chunk(actions, horizon), + info=info, + ) + + def get_action( + self, + images: Mapping[str, np.ndarray] | np.ndarray, + joint_positions: Mapping[str, float] | Sequence[float] | np.ndarray, + instruction: str, + *, + execution_horizon: int = 8, + ) -> list[dict[str, Any]]: + return self.get_action_chunk( # Only need the commands except for RTC, which uses get_action_chunk() to get the actions for prefix guidance + images, + joint_positions, + instruction, + execution_horizon=execution_horizon, + ).commands diff --git a/examples/RoboJuDo/g1_23dof_modality.json b/examples/RoboJuDo/g1_23dof_modality.json new file mode 100644 index 000000000..77e617a10 --- /dev/null +++ b/examples/RoboJuDo/g1_23dof_modality.json @@ -0,0 +1,51 @@ +{ + "state": { + "left_arm": { + "start": 0, + "end": 5 + }, + "right_arm": { + "start": 5, + "end": 10 + }, + "left_hand": { + "start": 10, + "end": 20 + }, + "right_hand": { + "start": 20, + "end": 30 + } + }, + "action": { + "left_arm": { + "start": 0, + "end": 5 + }, + "right_arm": { + "start": 5, + "end": 10 + }, + "left_hand": { + "start": 10, + "end": 20 + }, + "right_hand": { + "start": 20, + "end": 30 + }, + "navigate_command": { + "start": 30, + "end": 33 + }, + "base_height_command": { + "start": 33, + "end": 34 + } + }, + "video": { + "ego_view": { + "original_key": "observation.images.head_rgb" + } + } +} diff --git a/examples/RoboJuDo/g1_23dof_mulcam_modality.json b/examples/RoboJuDo/g1_23dof_mulcam_modality.json new file mode 100644 index 000000000..83e0b75a7 --- /dev/null +++ b/examples/RoboJuDo/g1_23dof_mulcam_modality.json @@ -0,0 +1,57 @@ +{ + "state": { + "left_arm": { + "start": 0, + "end": 5 + }, + "right_arm": { + "start": 5, + "end": 10 + }, + "left_hand": { + "start": 10, + "end": 20 + }, + "right_hand": { + "start": 20, + "end": 30 + } + }, + "action": { + "left_arm": { + "start": 0, + "end": 5 + }, + "right_arm": { + "start": 5, + "end": 10 + }, + "left_hand": { + "start": 10, + "end": 20 + }, + "right_hand": { + "start": 20, + "end": 30 + }, + "navigate_command": { + "start": 30, + "end": 33 + }, + "base_height_command": { + "start": 33, + "end": 34 + } + }, + "video": { + "ego_view": { + "original_key": "observation.images.head_rgb" + }, + "left_wrist_view": { + "original_key": "observation.images.left_wrist_rgb" + }, + "right_wrist_view": { + "original_key": "observation.images.right_wrist_rgb" + } + } +} diff --git a/examples/RoboJuDo/robojudo_g1_23dof_config.py b/examples/RoboJuDo/robojudo_g1_23dof_config.py new file mode 100644 index 000000000..7d35de249 --- /dev/null +++ b/examples/RoboJuDo/robojudo_g1_23dof_config.py @@ -0,0 +1,73 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +from gr00t.configs.data.embodiment_configs import register_modality_config +from gr00t.data.embodiment_tags import EmbodimentTag +from gr00t.data.types import ( + ActionConfig, + ActionFormat, + ActionRepresentation, + ActionType, + ModalityConfig, +) + + +robojudo_g1_23dof_config = { + "video": ModalityConfig(delta_indices=[0], modality_keys=["ego_view"]), + "state": ModalityConfig( + delta_indices=[0], + modality_keys=["left_arm", "right_arm", "left_hand", "right_hand"], + sin_cos_embedding_keys=["left_arm", "right_arm"], + ), + "action": ModalityConfig( + delta_indices=list(range(16)), + modality_keys=[ + "left_arm", + "right_arm", + "left_hand", + "right_hand", + "navigate_command", + "base_height_command", + ], + action_configs=[ + ActionConfig( + rep=ActionRepresentation.RELATIVE, + type=ActionType.NON_EEF, + format=ActionFormat.DEFAULT, + ), + ActionConfig( + rep=ActionRepresentation.RELATIVE, + type=ActionType.NON_EEF, + format=ActionFormat.DEFAULT, + ), + # Dexterous-hand actions are absolute actuator targets; unlike arm + # actions, they must not be offset by the measured joint feedback. + ActionConfig( + rep=ActionRepresentation.ABSOLUTE, + type=ActionType.NON_EEF, + format=ActionFormat.DEFAULT, + ), + ActionConfig( + rep=ActionRepresentation.ABSOLUTE, + type=ActionType.NON_EEF, + format=ActionFormat.DEFAULT, + ), + ActionConfig( + rep=ActionRepresentation.ABSOLUTE, + type=ActionType.NON_EEF, + format=ActionFormat.DEFAULT, + ), + ActionConfig( + rep=ActionRepresentation.ABSOLUTE, + type=ActionType.NON_EEF, + format=ActionFormat.DEFAULT, + ), + ], + ), + "language": ModalityConfig(delta_indices=[0], modality_keys=["task"]), +} + +register_modality_config( + robojudo_g1_23dof_config, + embodiment_tag=EmbodimentTag.NEW_EMBODIMENT, +) diff --git a/examples/RoboJuDo/robojudo_g1_23dof_mulcam_config.py b/examples/RoboJuDo/robojudo_g1_23dof_mulcam_config.py new file mode 100644 index 000000000..dba4f1da2 --- /dev/null +++ b/examples/RoboJuDo/robojudo_g1_23dof_mulcam_config.py @@ -0,0 +1,80 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +from gr00t.configs.data.embodiment_configs import register_modality_config +from gr00t.data.embodiment_tags import EmbodimentTag +from gr00t.data.types import ( + ActionConfig, + ActionFormat, + ActionRepresentation, + ActionType, + ModalityConfig, +) + + +robojudo_g1_23dof_config = { + "video": ModalityConfig( + delta_indices=[0], + modality_keys=[ + "ego_view", + "left_wrist_view", + "right_wrist_view", + ], + ), + "state": ModalityConfig( + delta_indices=[0], + modality_keys=["left_arm", "right_arm", "left_hand", "right_hand"], + sin_cos_embedding_keys=["left_arm", "right_arm"], + ), + "action": ModalityConfig( + delta_indices=list(range(16)), + modality_keys=[ + "left_arm", + "right_arm", + "left_hand", + "right_hand", + "navigate_command", + "base_height_command", + ], + action_configs=[ + ActionConfig( + rep=ActionRepresentation.RELATIVE, + type=ActionType.NON_EEF, + format=ActionFormat.DEFAULT, + ), + ActionConfig( + rep=ActionRepresentation.RELATIVE, + type=ActionType.NON_EEF, + format=ActionFormat.DEFAULT, + ), + # Dexterous-hand actions are absolute actuator targets; unlike arm + # actions, they must not be offset by the measured joint feedback. + ActionConfig( + rep=ActionRepresentation.ABSOLUTE, + type=ActionType.NON_EEF, + format=ActionFormat.DEFAULT, + ), + ActionConfig( + rep=ActionRepresentation.ABSOLUTE, + type=ActionType.NON_EEF, + format=ActionFormat.DEFAULT, + ), + ActionConfig( + rep=ActionRepresentation.ABSOLUTE, + type=ActionType.NON_EEF, + format=ActionFormat.DEFAULT, + ), + ActionConfig( + rep=ActionRepresentation.ABSOLUTE, + type=ActionType.NON_EEF, + format=ActionFormat.DEFAULT, + ), + ], + ), + "language": ModalityConfig(delta_indices=[0], modality_keys=["task"]), +} + +register_modality_config( + robojudo_g1_23dof_config, + embodiment_tag=EmbodimentTag.NEW_EMBODIMENT, +) diff --git a/examples/RoboJuDo/robojudo_x2_config.py b/examples/RoboJuDo/robojudo_x2_config.py new file mode 100644 index 000000000..9f4943fc9 --- /dev/null +++ b/examples/RoboJuDo/robojudo_x2_config.py @@ -0,0 +1,61 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +from gr00t.configs.data.embodiment_configs import register_modality_config +from gr00t.data.embodiment_tags import EmbodimentTag +from gr00t.data.types import ( + ActionConfig, + ActionFormat, + ActionRepresentation, + ActionType, + ModalityConfig, +) + + +robojudo_x2_config = { + "video": ModalityConfig(delta_indices=[0], modality_keys=["ego_view"]), + # left_hand/right_hand are intentionally omitted until X2 hand telemetry and + # commands are connected. The deployment profile reserves those group names. + "state": ModalityConfig( + delta_indices=[0], + modality_keys=["left_arm", "right_arm"], + sin_cos_embedding_keys=["left_arm", "right_arm"], + ), + "action": ModalityConfig( + delta_indices=list(range(16)), + modality_keys=[ + "left_arm", + "right_arm", + "navigate_command", + "base_height_command", + ], + action_configs=[ + ActionConfig( + rep=ActionRepresentation.RELATIVE, + type=ActionType.NON_EEF, + format=ActionFormat.DEFAULT, + ), + ActionConfig( + rep=ActionRepresentation.RELATIVE, + type=ActionType.NON_EEF, + format=ActionFormat.DEFAULT, + ), + ActionConfig( + rep=ActionRepresentation.ABSOLUTE, + type=ActionType.NON_EEF, + format=ActionFormat.DEFAULT, + ), + ActionConfig( + rep=ActionRepresentation.ABSOLUTE, + type=ActionType.NON_EEF, + format=ActionFormat.DEFAULT, + ), + ], + ), + "language": ModalityConfig(delta_indices=[0], modality_keys=["task"]), +} + +register_modality_config( + robojudo_x2_config, + embodiment_tag=EmbodimentTag.NEW_EMBODIMENT, +) diff --git a/examples/RoboJuDo/run_robojudo_client.py b/examples/RoboJuDo/run_robojudo_client.py new file mode 100644 index 000000000..1458d0dbf --- /dev/null +++ b/examples/RoboJuDo/run_robojudo_client.py @@ -0,0 +1,1150 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Run a selectable asynchronous RoboJuDo observation-to-command deployment loop.""" + +from __future__ import annotations + +import argparse +from collections import deque +from dataclasses import dataclass +import math +import threading +import time + +import cv2 +from deploy_adapter import CAMERA_LAYOUTS, PROFILES, RoboJuDoPolicyAdapter +from gr00t.policy.server_client import PolicyClient +import msgpack +import numpy as np +import zmq +from zmq.utils.monitor import recv_monitor_message + + +EXECUTION_MODES = ("double_buffer", "temporal_ensemble", "rtc") +RTC_PREFIX_SCHEDULES = ("zeros", "ones", "linear", "exp") + + +@dataclass(frozen=True) +class Observation: + stream_id: str + control_session: int + takeover_enabled: bool + sequence: int + images: dict[str, np.ndarray] + joint_positions: dict[str, float] + task: str + + +@dataclass(frozen=True) +class ActionChunk: + stream_id: str + control_session: int + observation_sequence: int + observation_received_at: float + inference_seconds: float + start_tick: int # The control tick at the start of inference + commands: list[dict] + physical_actions: dict[str, np.ndarray] | None = None # Provide Prefix action for RTC mode + task: str = "" + estimated_delay_steps: int = 0 + + +class RTCActionQueue: + """Lockstep physical-action/command queue; caller supplies synchronization.""" + + def __init__(self): + self._chunk: ActionChunk | None = None + self._next_index = 0 # Index of the next command and next physical action simultaneously + + @property + def chunk(self) -> ActionChunk | None: + return self._chunk + + def clear(self): + self._chunk = None + self._next_index = 0 + + def qsize(self) -> int: + if self._chunk is None: + return 0 + return max(0, len(self._chunk.commands) - self._next_index) + + def pop(self) -> dict | None: + if self._chunk is None or self._next_index >= len(self._chunk.commands): + return None + command = self._chunk.commands[self._next_index] + self._next_index += 1 + return command + + def get_left_over(self, session: tuple[str, int], task: str) -> dict[str, np.ndarray] | None: + # Get remaining physical action as prefix actions for RTC reference + chunk = self._chunk + if ( + chunk is None + or chunk.physical_actions is None + or (chunk.stream_id, chunk.control_session) != session + or chunk.task != task + or self._next_index >= len(chunk.commands) + ): + return None + return { + key: value[:, self._next_index :].copy() + for key, value in chunk.physical_actions.items() + } + + def replace(self, chunk: ActionChunk, skipped_steps: int) -> bool: + # Start at the skipped steps rather than the 0 of the chunk + if chunk.physical_actions is None: + raise ValueError("RTC action chunks must include physical_actions") + horizon = len(chunk.commands) + if any(value.shape[1] != horizon for value in chunk.physical_actions.values()): + raise ValueError("RTC physical actions and commands must have the same horizon") + if skipped_steps >= horizon: # delay is too long, the chunk is expired + self.clear() + return False + self._chunk = chunk + self._next_index = max(0, skipped_steps) + return True + + +@dataclass(frozen=True) +class _EnsembleChunk: + start_tick: int + actions: np.ndarray + + +class ACTTemporalEnsembler: + """ACT-style temporal ensemble generalized to asynchronously returned chunks.""" + + def __init__(self, joint_names: tuple[str, ...], temporal_ensemble_coeff: float): + self.joint_names = joint_names + self.temporal_ensemble_coeff = temporal_ensemble_coeff + self._chunks: deque[_EnsembleChunk] = deque() + + @property + def active_chunk_count(self) -> int: + return len(self._chunks) + + def reset(self): + self._chunks.clear() + + def _pack_command(self, command: dict) -> np.ndarray: + positions = command.get("positions") + if not isinstance(positions, dict): + raise ValueError("Temporal Ensemble command positions must be a dictionary") + missing = [name for name in self.joint_names if name not in positions] + if missing: + raise ValueError(f"Temporal Ensemble command is missing joints: {missing}") + locomotion = np.asarray(command.get("locomotion_command"), dtype=np.float32) + if locomotion.shape != (4,): + raise ValueError( + f"Temporal Ensemble locomotion command has shape {locomotion.shape}, expected (4,)" + ) + action = np.asarray( + [positions[name] for name in self.joint_names] + locomotion.tolist(), + dtype=np.float32, + ) + if not np.isfinite(action).all(): + raise ValueError("Temporal Ensemble command contains non-finite values") + return action + + def _unpack_command(self, action: np.ndarray) -> dict: + joint_count = len(self.joint_names) + return { + "positions": dict(zip(self.joint_names, action[:joint_count].tolist(), strict=True)), + "locomotion_command": action[joint_count:].astype(np.float32, copy=True), + } + + def add_chunk(self, chunk: ActionChunk): + actions = np.stack([self._pack_command(command) for command in chunk.commands]) + self._chunks.append(_EnsembleChunk(start_tick=chunk.start_tick, actions=actions)) + + def get_action(self, current_tick: int) -> tuple[dict | None, int]: + while ( + self._chunks + and self._chunks[0].start_tick + len(self._chunks[0].actions) <= current_tick + ): + self._chunks.popleft() + + predictions = [] + prediction_start_ticks = [] + for chunk in self._chunks: + action_index = current_tick - chunk.start_tick + if 0 <= action_index < len(chunk.actions): + predictions.append(chunk.actions[action_index]) + prediction_start_ticks.append(chunk.start_tick) + if not predictions: + return None, 0 + + stacked = np.stack(predictions) + prediction_offsets = np.asarray(prediction_start_ticks, dtype=np.float32) + prediction_offsets -= prediction_offsets[0] + weights = np.exp(-self.temporal_ensemble_coeff * prediction_offsets) + if not np.isfinite(weights).all(): + raise ValueError("Temporal Ensemble produced non-finite weights") + ensembled = np.average(stacked, axis=0, weights=weights).astype(np.float32) + return self._unpack_command(ensembled), len(predictions) + + +def make_safe_hold_command(command: dict) -> dict: + """Hold arm/height while stopping planar locomotion after prediction exhaustion.""" + locomotion = np.asarray(command["locomotion_command"], dtype=np.float32).copy() + if locomotion.shape != (4,): + raise ValueError(f"locomotion command has shape {locomotion.shape}, expected (4,)") + locomotion[:3] = 0.0 + return { + "positions": dict(command["positions"]), + "locomotion_command": locomotion, + } + + +class ObservationSubscriber: + def __init__(self, endpoint: str, profile: str, image_keys: tuple[str, ...]): + self.endpoint = endpoint + self.profile = profile + self.expected_joint_names = PROFILES[profile].joint_names + self.expected_image_keys = tuple(image_keys) + self._context = zmq.Context() + self._socket = None + + def connect(self): + self._socket = self._context.socket(zmq.SUB) + self._socket.setsockopt(zmq.LINGER, 0) + self._socket.setsockopt(zmq.RCVHWM, 2) + self._socket.setsockopt(zmq.SUBSCRIBE, b"") + self._socket.connect(self.endpoint) + + def receive(self, timeout_ms: int = 100) -> Observation | None: + if self._socket is None: + raise RuntimeError("RoboJuDo observation subscriber is not connected") + if self._socket.poll(timeout_ms, zmq.POLLIN) == 0: + return None + parts = self._socket.recv_multipart() + while self._socket.poll(0, zmq.POLLIN): + parts = self._socket.recv_multipart() + return self._decode_observation(parts) + + def _decode_observation(self, parts: list[bytes]) -> Observation: + if not parts: + raise ValueError("RoboJuDo observation multipart message is empty") + header = msgpack.unpackb(parts[0], raw=False) + protocol_version = header.get("protocol_version") + if protocol_version == 1: + image_keys = ("ego_view",) + image_shapes = {"ego_view": header.get("shape", ())} + elif protocol_version == 2: + raw_image_keys = header.get("image_keys") + if not isinstance(raw_image_keys, list) or not all( + isinstance(key, str) and key for key in raw_image_keys + ): + raise ValueError("RoboJuDo protocol v2 image_keys must be a list of names") + image_keys = tuple(raw_image_keys) + raw_image_shapes = header.get("image_shapes", {}) + if not isinstance(raw_image_shapes, dict): + raise ValueError("RoboJuDo protocol v2 image_shapes must be a dictionary") + if set(raw_image_shapes) != set(image_keys): + raise ValueError( + "RoboJuDo protocol v2 image_shapes keys must exactly match image_keys" + ) + image_shapes = raw_image_shapes + else: + raise ValueError(f"unsupported RoboJuDo protocol version {protocol_version!r}") + if image_keys != self.expected_image_keys: + raise ValueError( + f"RoboJuDo observation image order {image_keys} does not match camera layout " + f"{self.expected_image_keys}" + ) + expected_part_count = 1 + len(image_keys) + if len(parts) != expected_part_count: + raise ValueError( + f"RoboJuDo observation has {len(parts)} parts, expected {expected_part_count} " + f"for images {image_keys}" + ) + if header.get("profile") != self.profile: + raise ValueError( + f"RoboJuDo profile {header.get('profile')!r} does not match {self.profile!r}" + ) + joint_names = tuple(header.get("joint_names", ())) + if joint_names != self.expected_joint_names: + raise ValueError( + f"RoboJuDo observation joint order {joint_names} does not match the deployment profile {self.expected_joint_names}" + ) + positions = np.asarray(header.get("joint_positions"), dtype=np.float32) + if positions.shape != (len(joint_names),) or not np.isfinite(positions).all(): + raise ValueError("RoboJuDo observation contains invalid joint positions") + images = {} + for key, jpeg in zip(image_keys, parts[1:], strict=True): + bgr = cv2.imdecode(np.frombuffer(jpeg, dtype=np.uint8), cv2.IMREAD_COLOR) + if bgr is None: + raise ValueError(f"failed to decode RoboJuDo observation JPEG for {key!r}") + image = cv2.cvtColor(bgr, cv2.COLOR_BGR2RGB) + expected_shape = tuple(image_shapes.get(key, ())) + if expected_shape and image.shape != expected_shape: + raise ValueError( + f"RoboJuDo image {key!r} shape {image.shape} does not match {expected_shape}" + ) + images[key] = image + task = str(header.get("task", "")).strip() + if not task: + raise ValueError("RoboJuDo observation task must not be empty") + stream_id = header.get("stream_id") + if not isinstance(stream_id, str) or not stream_id.strip(): + raise ValueError("RoboJuDo observation stream_id must be a non-empty string") + control_session = header.get("control_session") + if ( + isinstance(control_session, bool) + or not isinstance(control_session, int) + or control_session < 0 + ): + raise ValueError("RoboJuDo observation control_session must be a non-negative integer") + takeover_enabled = header.get("takeover_enabled") + if not isinstance(takeover_enabled, bool): + raise ValueError("RoboJuDo observation takeover_enabled must be a boolean") + return Observation( + stream_id=stream_id, + control_session=control_session, + takeover_enabled=takeover_enabled, + sequence=int(header["sequence"]), + images=images, + joint_positions=dict(zip(joint_names, positions.tolist(), strict=True)), + task=task, + ) + + def close(self): + if self._socket is not None: + self._socket.close(linger=0) + self._socket = None + self._context.term() + + +class DoubleBufferedPolicyRunner: + def __init__( + self, + profile: str, + policy_host: str, + policy_port: int, + subscriber: ObservationSubscriber, + command_endpoint: str, + execution_horizon: int, + command_fps: float, + observation_timeout: float, + status_interval: float, + task_override: str | None, + execution_mode: str = "double_buffer", + temporal_ensemble_coeff: float = 0.01, + rtc_prefix_schedule: str = "exp", + rtc_max_guidance_weight: float = 10.0, + rtc_latency_window: int = 10, + video_keys: tuple[str, ...] = CAMERA_LAYOUTS["single"], + ): + self.profile = profile + self.policy_host = policy_host + self.policy_port = policy_port + self.subscriber = subscriber + self.execution_horizon = execution_horizon + self.command_period = 1.0 / command_fps + self.observation_timeout = observation_timeout + self.status_interval = status_interval + self.task_override = task_override + self.execution_mode = execution_mode + self.temporal_ensemble_coeff = temporal_ensemble_coeff + self.rtc_prefix_schedule = rtc_prefix_schedule + self.rtc_max_guidance_weight = rtc_max_guidance_weight + self.rtc_latency_window = rtc_latency_window + self.video_keys = video_keys + # Learned from the first full chunk returned by the policy. The action + # horizon belongs to the checkpoint/embodiment, not to RoboJuDo. + self._policy_action_horizon: int | None = None + # Coordinates observations, inference results, and the command-loop tick. + self._condition = threading.Condition() + self._stopping = False + self._latest_observation: Observation | None = None + self._latest_observation_at = float("-inf") + self._last_inferred_session: tuple[str, int] | None = None + self._last_inferred_sequence = -1 + self._pending_commands: ActionChunk | None = None + self._ready_chunks: deque[ActionChunk] = deque() + self._rtc_queue = RTCActionQueue() + # Store the last N inference latencies for RTC mode to estimate the delay steps + # e.g. D_est = cel(delay_ms/command_period_ms) = ceil(82 / 33.3) = 3 steps + self._rtc_inference_latencies: deque[float] = deque(maxlen=rtc_latency_window) + self._control_tick = 0 + self._control_tick_session: tuple[str, int] | None = None + self._error: Exception | None = None + self._context = zmq.Context() + self._publisher = self._context.socket(zmq.PUB) + self._publisher.setsockopt(zmq.LINGER, 0) + self._publisher.setsockopt(zmq.SNDHWM, 16) + self._publisher_monitor = self._publisher.get_monitor_socket( + events=zmq.EVENT_ACCEPTED | zmq.EVENT_DISCONNECTED + ) + self._publisher.bind(command_endpoint) + self._command_subscriber_connected = False + self._observation_thread = threading.Thread(target=self._receive_loop, daemon=True) + self._inference_thread = threading.Thread(target=self._inference_loop, daemon=True) + + def _receive_loop(self): + first_observation = True + connected_at = time.monotonic() + report_started_at = connected_at + report_observations = 0 + report_sequence_gaps = 0 + last_stream_id = None + last_sequence = None + + def report_health(now: float): + nonlocal report_started_at, report_observations, report_sequence_gaps + report_elapsed = now - report_started_at + if report_elapsed < self.status_interval: + return + if report_observations: + print( + f"[observation] rate={report_observations / report_elapsed:.1f}Hz, " + f"received={report_observations}, last_sequence={last_sequence}, " + f"sequence_gaps={report_sequence_gaps}", + flush=True, + ) + elif last_sequence is None: + print( + f"[observation] still waiting for first frame " + f"({now - connected_at:.1f}s elapsed)", + flush=True, + ) + else: + with self._condition: + latest_at = self._latest_observation_at + print( + f"[observation] no frame for {now - latest_at:.2f}s; " + f"last_sequence={last_sequence}", + flush=True, + ) + report_started_at = now + report_observations = 0 + report_sequence_gaps = 0 + + try: + self.subscriber.connect() + print(f"Connected to RoboJuDo observations at {self.subscriber.endpoint}", flush=True) + while not self._stopping: + observation = self.subscriber.receive() + if observation is None: + report_health(time.monotonic()) + continue + now = time.monotonic() + if observation.stream_id != last_stream_id: + if last_stream_id is not None: + print( + f"[observation] stream changed: {last_stream_id} -> " + f"{observation.stream_id}", + flush=True, + ) + last_stream_id = observation.stream_id + last_sequence = None + if last_sequence is not None: + if observation.sequence <= last_sequence: + print( + f"[observation] ignored non-increasing sequence " + f"{observation.sequence} after {last_sequence}", + flush=True, + ) + continue + report_sequence_gaps += observation.sequence - last_sequence - 1 + last_sequence = observation.sequence + report_observations += 1 + if first_observation: + image_shapes = {key: image.shape for key, image in observation.images.items()} + print( + f"Received first observation: sequence={observation.sequence}, " + f"session={observation.control_session}, " + f"takeover_enabled={observation.takeover_enabled}, " + f"images={image_shapes}, joints={len(observation.joint_positions)}", + flush=True, + ) + first_observation = False + with self._condition: + self._latest_observation = observation + self._latest_observation_at = time.monotonic() + self._condition.notify_all() + report_health(now) + except Exception as exc: + self._set_error(exc) + finally: + self.subscriber.close() + + def _inference_loop(self): + client = PolicyClient(host=self.policy_host, port=self.policy_port) + adapter = RoboJuDoPolicyAdapter(client, self.profile, self.video_keys) + try: + if not client.ping(): + raise ConnectionError( + f"GR00T policy server is unavailable at {self.policy_host}:{self.policy_port}" + ) + print( + f"Connected to GR00T policy server at {self.policy_host}:{self.policy_port}", + flush=True, + ) + while True: + with self._condition: + self._condition.wait_for( + lambda: ( + self._stopping + or self._error is not None + or ( + ( + self.execution_mode in ("temporal_ensemble", "rtc") + or self._pending_commands is None + ) + and self._latest_observation is not None + and self._latest_observation.takeover_enabled + and ( + ( + self._latest_observation.stream_id, + self._latest_observation.control_session, + ) + != self._last_inferred_session + or self._latest_observation.sequence + > self._last_inferred_sequence + ) + ) + ) + ) + if self._stopping or self._error is not None: + return + observation = self._latest_observation + observation_received_at = self._latest_observation_at + observation_session = ( + observation.stream_id, + observation.control_session, + ) + query_tick = ( + self._control_tick + if self._control_tick_session == observation_session + else 0 + ) + + # RTC mode: get prefix actions related(rtc_prefix, estimated_delay_steps) for RTC reference + instruction = self.task_override or observation.task + rtc_prefix = None + prefix_length = 0 + estimated_delay_steps = 0 + if self.execution_mode == "rtc": + rtc_prefix = self._rtc_queue.get_left_over(observation_session, instruction) + if rtc_prefix is not None: + prefix_length = min( + value.shape[1] for value in rtc_prefix.values() + ) # L + if self._rtc_inference_latencies: + estimated_delay_steps = math.ceil( + max(self._rtc_inference_latencies) / self.command_period + ) + estimated_delay_steps = min( # D_est + estimated_delay_steps, + prefix_length, + self.execution_horizon, + ) + + self._last_inferred_session = observation_session + self._last_inferred_sequence = observation.sequence + inference_started_at = time.monotonic() + + # Construct RTC options and start inference. Always request the + # complete policy chunk so its checkpoint-defined horizon can be + # discovered and validated at runtime. + physical_actions = None + if self.execution_mode == "rtc": + rtc_options = None + if rtc_prefix is not None: + rtc_options = { + "rtc": { + "prefix_actions": rtc_prefix, + "prefix_length": prefix_length, # L, Remaining prefix length of old action + "estimated_delay_steps": estimated_delay_steps, # D_est + "guidance_horizon": min( + self.execution_horizon, prefix_length + ), # H, execution_horizon is the max guidance horizon for RTC + "prefix_schedule": self.rtc_prefix_schedule, + "max_guidance_weight": self.rtc_max_guidance_weight, + } + } + policy_chunk = adapter.get_action_chunk( + images=observation.images, + joint_positions=observation.joint_positions, + instruction=instruction, + execution_horizon=None, + options=rtc_options, + ) + commands = policy_chunk.commands + physical_actions = policy_chunk.actions + else: + policy_chunk = adapter.get_action_chunk( + images=observation.images, + joint_positions=observation.joint_positions, + instruction=instruction, + execution_horizon=None, + ) + commands = policy_chunk.commands[: self.execution_horizon] + policy_action_horizon = len(policy_chunk.commands) + known_action_horizon = getattr(self, "_policy_action_horizon", None) + if known_action_horizon is None: + if self.execution_horizon > policy_action_horizon: + raise ValueError( + f"--execution-horizon={self.execution_horizon} exceeds the policy " + f"action horizon {policy_action_horizon} discovered from its first chunk" + ) + self._policy_action_horizon = policy_action_horizon + print( + f"[inference] detected policy action horizon: {policy_action_horizon}", + flush=True, + ) + elif policy_action_horizon != known_action_horizon: + raise ValueError( + "Policy action horizon changed between chunks: " + f"expected {known_action_horizon}, got {policy_action_horizon}" + ) + if not commands: + raise ValueError( + f"GR00T returned an empty action chunk for observation {observation.sequence}" + ) + inference_seconds = time.monotonic() - inference_started_at + with self._condition: + # After Inference, check if the observation is still valid for the current session and task + if self.execution_mode == "rtc": + self._rtc_inference_latencies.append(inference_seconds) + latest = self._latest_observation + latest_instruction = ( + None if latest is None else self.task_override or latest.task + ) + if ( + latest is None + or not latest.takeover_enabled + or (latest.stream_id, latest.control_session) != observation_session + or latest_instruction != instruction + or ( + self.execution_mode == "rtc" + and time.monotonic() - self._latest_observation_at + > self.observation_timeout + ) + ): + print( + f"[inference] discarded chunk from inactive session " + f"{observation.stream_id}:{observation.control_session}", + flush=True, + ) + self._condition.notify_all() + continue + chunk = ActionChunk( + stream_id=observation.stream_id, + control_session=observation.control_session, + observation_sequence=observation.sequence, + observation_received_at=observation_received_at, + inference_seconds=inference_seconds, + start_tick=query_tick, + commands=commands, + physical_actions=physical_actions, + task=instruction, + estimated_delay_steps=estimated_delay_steps, + ) + ready_tick = ( + self._control_tick + if self._control_tick_session == observation_session + else query_tick + ) + skipped_steps = ( + max(0, ready_tick - query_tick) if rtc_prefix is not None else 0 + ) # D_actual + if self.execution_mode == "temporal_ensemble": + self._ready_chunks.append(chunk) + elif self.execution_mode == "rtc": + activated = self._rtc_queue.replace( + chunk, skipped_steps + ) # Skipped the expired steps due to the inference delay + if not activated: + print( + f"[inference] RTC chunk expired before activation: " + f"actual_delay={skipped_steps}, horizon={len(commands)}; " + "waiting for an unguided refresh", + flush=True, + ) + else: + self._pending_commands = chunk + self._condition.notify_all() + print( + f"[inference] chunk ready: observation_sequence={observation.sequence}, " + f"session={observation.control_session}, actions={len(commands)}, " + f"latency={inference_seconds:.3f}s, query_tick={query_tick}, " + f"ready_tick={ready_tick}, skipped_steps={skipped_steps}, " + f"rtc_prefix={prefix_length}, " + f"rtc_estimated_delay={estimated_delay_steps}", + flush=True, + ) + except Exception as exc: + self._set_error(exc) + finally: + client.close() + + def _set_error(self, exc: Exception): + with self._condition: + self._error = exc + self._condition.notify_all() + + def _poll_command_subscriber(self): + while self._publisher_monitor.poll(0, zmq.POLLIN): + event = recv_monitor_message(self._publisher_monitor) + endpoint = event.get("endpoint", b"") + if isinstance(endpoint, bytes): + endpoint = endpoint.decode(errors="replace") + if event["event"] == zmq.EVENT_ACCEPTED: + self._command_subscriber_connected = True + print(f"[command] subscriber connected: {endpoint}", flush=True) + elif event["event"] == zmq.EVENT_DISCONNECTED: + self._command_subscriber_connected = False + print(f"[command] subscriber disconnected: {endpoint}", flush=True) + + def run(self): + print("Waiting for the first RoboJuDo observation and GR00T action chunk...", flush=True) + self._observation_thread.start() + self._inference_thread.start() + active_commands: deque[dict] = deque() + temporal_ensembler = ( + ACTTemporalEnsembler( + PROFILES[self.profile].joint_names, + self.temporal_ensemble_coeff, + ) + if self.execution_mode == "temporal_ensemble" + else None + ) + last_command = None + command_sequence = 0 + command_stream_started = False + observation_was_fresh = False + control_was_enabled = False + active_session: tuple[str, int] | None = None + active_task: str | None = None + holding_last_command = False + report_started_at = time.monotonic() + report_commands = 0 + next_command_at = time.monotonic() + ensemble_contributors = 0 + while True: + self._poll_command_subscriber() + ready_chunks = [] + rtc_command = None + current_tick = 0 + with self._condition: + if self._error is not None: + raise RuntimeError("RoboJuDo deployment worker failed") from self._error + now = time.monotonic() + observation_fresh = now - self._latest_observation_at <= self.observation_timeout + latest = self._latest_observation + control_enabled = bool( + observation_fresh and latest is not None and latest.takeover_enabled + ) + current_session = ( + (latest.stream_id, latest.control_session) if latest is not None else None + ) + current_task = None if latest is None else self.task_override or latest.task + + # Handle observation timeout, control takeover, and session/task changes + if not observation_fresh: + if observation_was_fresh: + if last_command is None or active_session is None: + print( + "RoboJuDo observation timed out; no previous command to hold", + flush=True, + ) + else: + print( + "RoboJuDo observation timed out; holding arm/height and " + "setting vx/vy/yaw_rate to zero", + flush=True, + ) + active_commands.clear() + self._pending_commands = None + self._ready_chunks.clear() + self._rtc_queue.clear() + self._rtc_inference_latencies.clear() + if temporal_ensembler is not None: + temporal_ensembler.reset() + self._control_tick = 0 + if last_command is not None and active_session is not None: + last_command = make_safe_hold_command(last_command) + self._control_tick_session = active_session + holding_last_command = True + else: + last_command = None + self._control_tick_session = None + holding_last_command = False + active_session = None + active_task = None + elif not observation_was_fresh: + print("RoboJuDo observation stream is fresh", flush=True) + if observation_fresh and not control_enabled: + cleared_pending = self._pending_commands is not None + if control_was_enabled: + print( + "[control] takeover disabled; cleared active and pending commands", + flush=True, + ) + active_commands.clear() + last_command = None + self._pending_commands = None + self._ready_chunks.clear() + self._rtc_queue.clear() + self._rtc_inference_latencies.clear() + if temporal_ensembler is not None: + temporal_ensembler.reset() + self._control_tick = 0 + self._control_tick_session = None + holding_last_command = False + active_session = None + active_task = None + if control_was_enabled or cleared_pending: + self._condition.notify_all() + elif control_enabled and ( + current_session != active_session or current_task != active_task + ): + active_commands.clear() + last_command = None + holding_last_command = False + active_session = current_session + active_task = current_task + self._control_tick = 0 + self._control_tick_session = current_session + if temporal_ensembler is not None: + temporal_ensembler.reset() + rtc_chunk = self._rtc_queue.chunk + if ( + rtc_chunk is None + or (rtc_chunk.stream_id, rtc_chunk.control_session) != current_session + or rtc_chunk.task != current_task + ): + self._rtc_queue.clear() + self._rtc_inference_latencies.clear() + pending_session = ( + ( + self._pending_commands.stream_id, + self._pending_commands.control_session, + ) + if self._pending_commands is not None + else None + ) + if self._pending_commands is not None and pending_session != current_session: + self._pending_commands = None + self._ready_chunks = deque( + chunk + for chunk in self._ready_chunks + if (chunk.stream_id, chunk.control_session) == current_session + ) + self._condition.notify_all() + print( + f"[control] takeover enabled for session " + f"{current_session[0]}:{current_session[1]}; waiting for a fresh chunk", + flush=True, + ) + if ( + self.execution_mode == "double_buffer" + and control_enabled + and not active_commands + and self._pending_commands is not None + ): + chunk = self._pending_commands + self._pending_commands = None + self._condition.notify_all() + chunk_age = now - chunk.observation_received_at + chunk_session = (chunk.stream_id, chunk.control_session) + if chunk_session != current_session: + print( + f"[command] discarded chunk from inactive session " + f"{chunk.stream_id}:{chunk.control_session}", + flush=True, + ) + elif chunk_age <= self.observation_timeout: + active_commands.extend(chunk.commands) + holding_last_command = False + print( + f"[command] activated chunk: observation_sequence={chunk.observation_sequence}, " + f"session={chunk.control_session}, " + f"actions={len(chunk.commands)}, age={chunk_age:.3f}s, " + f"inference={chunk.inference_seconds:.3f}s", + flush=True, + ) + else: + print( + f"[command] discarded stale chunk for observation_sequence=" + f"{chunk.observation_sequence} (age={chunk_age:.3f}s)", + flush=True, + ) + if self.execution_mode == "temporal_ensemble" and control_enabled: + while self._ready_chunks: + chunk = self._ready_chunks.popleft() + chunk_session = (chunk.stream_id, chunk.control_session) + chunk_age = now - chunk.observation_received_at + if chunk_session != current_session: + print( + f"[command] discarded chunk from inactive session " + f"{chunk.stream_id}:{chunk.control_session}", + flush=True, + ) + elif chunk_age <= self.observation_timeout: + ready_chunks.append(chunk) + else: + print( + f"[command] discarded stale chunk for observation_sequence=" + f"{chunk.observation_sequence} (age={chunk_age:.3f}s)", + flush=True, + ) + current_tick = self._control_tick + if self.execution_mode == "rtc" and control_enabled: + rtc_command = self._rtc_queue.pop() + current_tick = self._control_tick + observation_was_fresh = observation_fresh + control_was_enabled = control_enabled + if self.execution_mode == "temporal_ensemble": + for chunk in ready_chunks: + temporal_ensembler.add_chunk(chunk) + print( + f"[command] added ensemble chunk: " + f"observation_sequence={chunk.observation_sequence}, " + f"session={chunk.control_session}, start_tick={chunk.start_tick}, " + f"actions={len(chunk.commands)}", + flush=True, + ) + ensembled_command, ensemble_contributors = temporal_ensembler.get_action( + current_tick + ) + if ensembled_command is not None: + last_command = ensembled_command + holding_last_command = False + elif last_command is not None: + if not holding_last_command: + print( + "[command] Temporal Ensemble horizon exhausted; holding arm/height " + "and setting vx/vy/yaw_rate to zero", + flush=True, + ) + last_command = make_safe_hold_command(last_command) + holding_last_command = True + elif self.execution_mode == "rtc": + if rtc_command is not None: + last_command = rtc_command + holding_last_command = False + elif last_command is not None: + if not holding_last_command: + print( + "[command] RTC queue exhausted; holding arm/height " + "and setting vx/vy/yaw_rate to zero", + flush=True, + ) + last_command = make_safe_hold_command(last_command) + holding_last_command = True + else: + if active_commands: + last_command = active_commands.popleft() + holding_last_command = False + elif last_command is not None and not holding_last_command: + print( + "[command] action horizon exhausted; holding the last command " + "until the next chunk is ready", + flush=True, + ) + holding_last_command = True + if last_command is not None: + if active_session is None: + raise RuntimeError("cannot publish a command without an active control session") + self._publisher.send_json( + { + "sequence": command_sequence, + "stream_id": active_session[0], + "control_session": active_session[1], + "positions": last_command["positions"], + "locomotion_command": last_command["locomotion_command"].tolist(), + } + ) + if not command_stream_started: + print( + f"Publishing GR00T commands on {self._publisher.getsockopt_string(zmq.LAST_ENDPOINT)}", + flush=True, + ) + command_stream_started = True + command_sequence += 1 + report_commands += 1 + if control_enabled and active_session is not None: + with self._condition: + if self._control_tick_session == active_session: + self._control_tick += 1 + report_now = time.monotonic() + report_elapsed = report_now - report_started_at + if report_elapsed >= self.status_interval: + with self._condition: + observation_age = report_now - self._latest_observation_at + ready_chunk_count = len(self._ready_chunks) + inference_pending = ( + self._pending_commands is not None + if self.execution_mode == "double_buffer" + else bool(ready_chunk_count) + ) + latest = self._latest_observation + takeover_enabled = bool(latest and latest.takeover_enabled) + control_session = None if latest is None else latest.control_session + rtc_queue_size = self._rtc_queue.qsize() + if self.execution_mode == "temporal_ensemble": + mode_status = ( + f"control_tick={current_tick}, " + f"ensemble_chunks={temporal_ensembler.active_chunk_count}, " + f"ensemble_contributors={ensemble_contributors}, " + f"ready_chunks={ready_chunk_count}" + ) + elif self.execution_mode == "rtc": + mode_status = ( + f"control_tick={current_tick}, rtc_queue_remaining={rtc_queue_size}, " + f"holding={holding_last_command}, " + f"latency_samples={len(self._rtc_inference_latencies)}" + ) + else: + mode_status = ( + f"chunk_remaining={len(active_commands)}, " + f"holding={holding_last_command}, " + f"inference_chunk_pending={inference_pending}" + ) + print( + f"[command] rate={report_commands / report_elapsed:.1f}Hz, " + f"published={report_commands}, next_sequence={command_sequence}, " + f"subscriber_connected={self._command_subscriber_connected}, " + f"mode={self.execution_mode}, {mode_status}, " + f"observation_age={observation_age:.3f}s, " + f"takeover_enabled={takeover_enabled}, control_session={control_session}", + flush=True, + ) + report_started_at = report_now + report_commands = 0 + next_command_at += self.command_period + time.sleep(max(0.0, next_command_at - time.monotonic())) + + def close(self): + with self._condition: + self._stopping = True + self._condition.notify_all() + self._observation_thread.join(timeout=2) + self._inference_thread.join(timeout=2) + self._publisher.disable_monitor() + self._publisher_monitor.close(linger=0) + self._publisher.close(linger=0) + self._context.term() + + +def parse_args(): + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--profile", choices=sorted(PROFILES), required=True) + parser.add_argument( + "--camera-layout", + choices=sorted(CAMERA_LAYOUTS), + default="single", + help="Expected observation image layout; mulcam requires protocol v2", + ) + parser.add_argument("--robot-endpoint", required=True, help="RoboJuDo observation endpoint") + parser.add_argument("--command-endpoint", default="tcp://*:8559") + parser.add_argument("--policy-host", default="127.0.0.1") + parser.add_argument("--policy-port", type=int, default=5555) + parser.add_argument( + "--execution-mode", + choices=EXECUTION_MODES, + default="double_buffer", + help="Action execution strategy; RTC continuously replaces a guided action queue", + ) + parser.add_argument( + "--execution-horizon", + type=int, + default=8, + help=( + "Execution/guidance horizon; its upper bound is discovered from the first " + "action chunk returned by the policy" + ), + ) + parser.add_argument( + "--temporal-ensemble-coeff", + type=float, + default=0.01, + help="ACT exponential weight coefficient; 0 gives an equal average", + ) + parser.add_argument( + "--rtc-prefix-schedule", + choices=RTC_PREFIX_SCHEDULES, + default="exp", + help="RTC prefix attention schedule between estimated delay and execution horizon", + ) + parser.add_argument( + "--rtc-max-guidance-weight", + type=float, + default=10.0, + help="Maximum per-denoising-step RTC correction gain", + ) + parser.add_argument( + "--rtc-latency-window", + type=int, + default=10, + help="Number of recent inference latencies used by RTC's rolling maximum", + ) + parser.add_argument("--command-fps", type=float, default=30.0) + parser.add_argument("--observation-timeout", type=float, default=1.0) + parser.add_argument("--status-interval", type=float, default=5.0) + parser.add_argument("--task", default=None, help="Override the task sent by RoboJuDo") + args = parser.parse_args() + if args.camera_layout == "mulcam" and args.profile != "g1_23dof": + parser.error("--camera-layout=mulcam is currently supported only for --profile=g1_23dof") + if args.execution_horizon < 1: + parser.error("--execution-horizon must be positive") + if args.command_fps <= 0: + parser.error("--command-fps must be positive") + if args.observation_timeout <= 0: + parser.error("--observation-timeout must be positive") + if args.status_interval <= 0: + parser.error("--status-interval must be positive") + if not np.isfinite(args.temporal_ensemble_coeff): + parser.error("--temporal-ensemble-coeff must be finite") + if not np.isfinite(args.rtc_max_guidance_weight) or args.rtc_max_guidance_weight < 0: + parser.error("--rtc-max-guidance-weight must be finite and non-negative") + if args.rtc_latency_window <= 0: + parser.error("--rtc-latency-window must be positive") + return args + + +def main(): + args = parse_args() + print( + f"Starting RoboJuDo deploy client: profile={args.profile}, " + f"camera_layout={args.camera_layout}, " + f"mode={args.execution_mode}, observations={args.robot_endpoint}, " + f"commands={args.command_endpoint}", + flush=True, + ) + video_keys = CAMERA_LAYOUTS[args.camera_layout] + subscriber = ObservationSubscriber(args.robot_endpoint, args.profile, video_keys) + runner = DoubleBufferedPolicyRunner( + profile=args.profile, + policy_host=args.policy_host, + policy_port=args.policy_port, + subscriber=subscriber, + command_endpoint=args.command_endpoint, + execution_horizon=args.execution_horizon, + command_fps=args.command_fps, + observation_timeout=args.observation_timeout, + status_interval=args.status_interval, + task_override=args.task, + execution_mode=args.execution_mode, + temporal_ensemble_coeff=args.temporal_ensemble_coeff, + rtc_prefix_schedule=args.rtc_prefix_schedule, + rtc_max_guidance_weight=args.rtc_max_guidance_weight, + rtc_latency_window=args.rtc_latency_window, + video_keys=video_keys, + ) + try: + runner.run() + except KeyboardInterrupt: + pass + finally: + runner.close() + + +if __name__ == "__main__": + main() diff --git a/examples/RoboJuDo/x2_modality.json b/examples/RoboJuDo/x2_modality.json new file mode 100644 index 000000000..476d823cc --- /dev/null +++ b/examples/RoboJuDo/x2_modality.json @@ -0,0 +1,35 @@ +{ + "state": { + "left_arm": { + "start": 0, + "end": 7 + }, + "right_arm": { + "start": 7, + "end": 14 + } + }, + "action": { + "left_arm": { + "start": 0, + "end": 7 + }, + "right_arm": { + "start": 7, + "end": 14 + }, + "navigate_command": { + "start": 14, + "end": 17 + }, + "base_height_command": { + "start": 17, + "end": 18 + } + }, + "video": { + "ego_view": { + "original_key": "observation.images.head_rgb" + } + } +} diff --git a/getting_started/real_world_deployment.md b/getting_started/real_world_deployment.md index d85da41ff..7af869889 100644 --- a/getting_started/real_world_deployment.md +++ b/getting_started/real_world_deployment.md @@ -410,18 +410,18 @@ When direct optimization is insufficient, use one or more of the following: **Recommended strategy**: `Asynchronous Inference + RTC` is usually the most effective. -> **RTC status (experimental):** Asynchronous inference is supported today. RTC is currently only a low-level model primitive: `action_head.get_action(..., options={"rtc_overlap_steps": ..., "rtc_frozen_steps": ..., "rtc_ramp_rate": ...})` with the previous action fed back in (`gr00t/model/gr00t_n1d7/gr00t_n1d7.py`). It is **not wired into `Gr00tPolicy` or the server-client path** (there `options` is currently unused), and it has no tests or ready-made example — so the RTC steps below require manual integration. +> **RTC status (experimental):** The RoboJuDo example provides an end-to-end inference-time RTC prototype through `--execution-mode rtc`. It sends the unexecuted physical prefix through `Gr00tPolicy`, re-anchors relative action groups against the latest state, and applies LeRobot-style gradient guidance in the N1.7 flow sampler. Other deployment clients still require their own queue and scheduling integration. #### Real-Time Chunking (RTC) Details **Principle** -RTC treats action prediction as an inpainting problem: overlapping the start of the new prediction with unexecuted steps from the previous one ensures smooth transitions. +RTC guides the clean-action estimate toward unexecuted steps from the previous chunk. A fully weighted prefix covers the estimated inference delay, followed by a decaying transition region where the new policy can smoothly take over. **Applicability** - Validated for **diffusion / flow-based** VLA policies. -- Requires `Action Chunk` length ≥ 32 steps. +- Benefits from longer chunks; the RoboJuDo prototype supports its existing 16-step checkpoint, with less transition room than a 32-step model. - Should be combined with asynchronous inference. **Implementation essentials** @@ -446,7 +446,7 @@ In the RTC (Real-Time Chunking) framework, two key parameters control how adjace - **`overlap`**: The number of action steps retained from the previous prediction to constrain the current one, ensuring temporal consistency between consecutive chunks. - **`frozen`**: The number of steps that remain completely frozen (i.e., not updated by the new prediction), typically set to match the inference latency. -Below is a simplified async inference + RTC loop. Note that official RTC support for GR00T is coming soon; the current implementation may require manual adaptation. +Below is a simplified async inference + RTC loop. Deployment clients other than RoboJuDo still need to adapt this queue behavior. ``` actions = policy.infer(obs) # blocking first call @@ -461,4 +461,3 @@ loop: break # discard frozen tail ``` - diff --git a/gr00t/model/gr00t_n1d7/gr00t_n1d7.py b/gr00t/model/gr00t_n1d7/gr00t_n1d7.py index 346b597a4..fc317a7a5 100644 --- a/gr00t/model/gr00t_n1d7/gr00t_n1d7.py +++ b/gr00t/model/gr00t_n1d7/gr00t_n1d7.py @@ -14,6 +14,7 @@ # limitations under the License. import logging +import math from typing import Any, Tuple import torch @@ -35,6 +36,72 @@ logger = logging.getLogger(__name__) +RTC_PREFIX_SCHEDULES = ("zeros", "ones", "linear", "exp") + + +def get_rtc_prefix_weights( + frozen_steps: int, + guidance_horizon: int, + total_steps: int, + schedule: str, + *, + device: torch.device | None = None, + dtype: torch.dtype = torch.float32, +) -> torch.Tensor: + """Build LeRobot-compatible RTC prefix weights for one action chunk.""" + if schedule not in RTC_PREFIX_SCHEDULES: + raise ValueError( + f"RTC prefix schedule must be one of {RTC_PREFIX_SCHEDULES}, got {schedule!r}" + ) + if total_steps <= 0: + raise ValueError("RTC total_steps must be positive") + if not 0 <= frozen_steps <= guidance_horizon <= total_steps: + raise ValueError( + "RTC horizons must satisfy " + f"0 <= frozen_steps <= guidance_horizon <= total_steps, got " + f"{frozen_steps}, {guidance_horizon}, {total_steps}" + ) + + weights = torch.zeros(total_steps, device=device, dtype=dtype) + if schedule == "ones": + weights[:guidance_horizon] = 1.0 + return weights + weights[:frozen_steps] = 1.0 + if schedule == "zeros" or guidance_horizon == frozen_steps: + return weights + + transition_steps = guidance_horizon - frozen_steps + transition = torch.linspace( + 1.0, + 0.0, + transition_steps + 2, + device=device, + dtype=dtype, + )[1:-1] + if schedule == "exp": + transition = transition * torch.expm1(transition) / math.expm1(1.0) + weights[frozen_steps:guidance_horizon] = transition + return weights + + +def get_rtc_guidance_weight( + flow_time: float, max_guidance_weight: float, *, device: torch.device +) -> torch.Tensor: + """Translate LeRobot's reverse-time RTC gain to GR00T's 0->1 flow time.""" + if not 0.0 <= flow_time <= 1.0: + raise ValueError(f"RTC flow_time must be in [0, 1], got {flow_time}") + if not math.isfinite(max_guidance_weight) or max_guidance_weight < 0.0: + raise ValueError("RTC max_guidance_weight must be finite and non-negative") + progress = torch.as_tensor(flow_time, dtype=torch.float32, device=device) + remaining = 1.0 - progress + gain = (remaining.square() + progress.square()) / (remaining * progress) + return torch.nan_to_num( + gain, + nan=max_guidance_weight, + posinf=max_guidance_weight, + ).clamp(max=max_guidance_weight) + + class Gr00tN1d7ActionHead(nn.Module): """Action head component for flow matching diffusion policy.""" @@ -354,8 +421,41 @@ def get_action_with_features( dt = 1.0 / self.num_inference_timesteps vel_strength = torch.ones_like(actions) + rtc_options = options.get("rtc") if options is not None else None + rtc_prefix = None + rtc_weights = None + + if rtc_options is not None: + if "action" not in action_input: + raise ValueError("RTC requires a normalized action prefix") + prefix_length = int(rtc_options["prefix_length"]) + frozen_steps = min(int(rtc_options["estimated_delay_steps"]), prefix_length) + guidance_horizon = min(int(rtc_options["guidance_horizon"]), prefix_length) + if not 0 < prefix_length <= self.action_horizon: + raise ValueError( + f"RTC prefix_length must be in [1, {self.action_horizon}], got {prefix_length}" + ) + if not 0 <= frozen_steps <= guidance_horizon: + raise ValueError( + "RTC horizons must satisfy 0 <= estimated_delay_steps <= " + "guidance_horizon after prefix clipping" + ) + rtc_prefix = action_input["action"][:, : self.action_horizon].detach() + prefix_weights = get_rtc_prefix_weights( + frozen_steps, + guidance_horizon, + self.action_horizon, + str(rtc_options.get("prefix_schedule", "exp")), + device=device, + dtype=actions.dtype, + ) + rtc_weights = prefix_weights[None, :, None] + if "action_mask" in action_input: + rtc_weights = rtc_weights * action_input["action_mask"][ + :, : self.action_horizon + ].to(dtype=actions.dtype) - if "action" in action_input: + elif "action" in action_input: # If action in input when doing get action, it means we want to use RTC. # action_horizon is the action horizon of the input action. # rtc_overlap_steps is the number of steps to overlap with the previous action chunks. @@ -393,26 +493,13 @@ def get_action_with_features( :, ] = ramp[None, :, None].to(device) - # Run denoising steps. - for t in range(self.num_inference_timesteps): - t_cont = t / float(self.num_inference_timesteps) # e.g. goes 0, 1/N, 2/N, ... - t_discretized = int(t_cont * self.num_timestep_buckets) - - # Embed noised action trajectory. - timesteps_tensor = torch.full( - size=(batch_size,), fill_value=t_discretized, device=device - ) - action_features = self.action_encoder(actions, timesteps_tensor, embodiment_id) - # Add position embedding. + def predict_velocity(current_actions: torch.Tensor, timesteps_tensor: torch.Tensor): + """Run one denoising network evaluation.""" + action_features = self.action_encoder(current_actions, timesteps_tensor, embodiment_id) if self.config.add_pos_embed: pos_ids = torch.arange(action_features.shape[1], dtype=torch.long, device=device) - pos_embs = self.position_embedding(pos_ids).unsqueeze(0) - action_features = action_features + pos_embs - - # Join vision, language, state and action embedding along sequence dimension. + action_features = action_features + self.position_embedding(pos_ids).unsqueeze(0) sa_embs = torch.cat((state_features, action_features), dim=1) - - # Run model forward. if self.config.use_alternate_vl_dit: model_output = self.model( hidden_states=sa_embs, @@ -428,11 +515,40 @@ def get_action_with_features( timestep=timesteps_tensor, ) pred = self.action_decoder(model_output, embodiment_id) + return pred[:, -self.action_horizon :] + + # Run denoising steps. + for t in range(self.num_inference_timesteps): + t_cont = t / float(self.num_inference_timesteps) # e.g. goes 0, 1/N, 2/N, ... + t_discretized = int(t_cont * self.num_timestep_buckets) - pred_velocity = pred[:, -self.action_horizon :] + timesteps_tensor = torch.full( + size=(batch_size,), fill_value=t_discretized, device=device + ) + pred_velocity = predict_velocity(actions, timesteps_tensor) + + if rtc_prefix is not None: + max_guidance_weight = float(rtc_options.get("max_guidance_weight", 10.0)) + # LeRobot intentionally treats the denoiser output as fixed while taking + # this gradient. This computes the same vector-Jacobian correction without + # retaining a graph through the denoising network. + with torch.enable_grad(): + differentiable_actions = actions.detach().requires_grad_(True) + clean_actions = differentiable_actions + (1.0 - t_cont) * pred_velocity + prefix_error = (rtc_prefix - clean_actions) * rtc_weights + correction = torch.autograd.grad( + clean_actions, + differentiable_actions, + grad_outputs=prefix_error.detach(), + retain_graph=False, + )[0] + guidance_weight = get_rtc_guidance_weight( + t_cont, max_guidance_weight, device=device + ).to(dtype=pred_velocity.dtype) + pred_velocity = pred_velocity + guidance_weight * correction # Update actions using euler integration. - actions = actions + dt * pred_velocity * vel_strength + actions = (actions + dt * pred_velocity * vel_strength).detach() return BatchFeature( data={ diff --git a/gr00t/policy/gr00t_policy.py b/gr00t/policy/gr00t_policy.py index 6f5a46b10..b97dc525a 100644 --- a/gr00t/policy/gr00t_policy.py +++ b/gr00t/policy/gr00t_policy.py @@ -197,7 +197,11 @@ def _unbatch_observation(self, value: dict[str, Any]) -> list[dict[str, Any]]: unbatched_obs.append(unbatched_value) return unbatched_obs - def _to_vla_step_data(self, observation: dict[str, Any]) -> VLAStepData: + def _to_vla_step_data( + self, + observation: dict[str, Any], + actions: dict[str, np.ndarray] | None = None, + ) -> VLAStepData: """Convert a single observation into a VLAStepData object for processing. Args: @@ -209,11 +213,109 @@ def _to_vla_step_data(self, observation: dict[str, Any]) -> VLAStepData: return VLAStepData( images=observation["video"], states=observation["state"], - actions={}, # No ground truth actions during inference + actions={} if actions is None else actions, text=observation["language"][self.language_key][0], embodiment=self.embodiment_tag, ) + def _prepare_rtc_options( + self, + options: dict[str, Any] | None, + batch_size: int, + ) -> tuple[list[dict[str, np.ndarray]] | None, dict[str, Any] | None]: + """Validate physical RTC prefixes and strip them from model options.""" + if options is None or options.get("rtc") is None: + return None, options + if batch_size != 1: + raise ValueError("RTC currently supports batch size 1 only") + rtc = options["rtc"] + if not isinstance(rtc, dict): + raise ValueError("options['rtc'] must be a dictionary") + + prefix_actions = rtc.get("prefix_actions") + if not isinstance(prefix_actions, dict): + raise ValueError("RTC prefix_actions must be a dictionary") + action_keys = self.modality_configs["action"].modality_keys + missing = [key for key in action_keys if key not in prefix_actions] + if missing: + raise ValueError(f"RTC prefix_actions is missing action groups: {missing}") + + prefix_length = rtc.get("prefix_length") + if isinstance(prefix_length, bool) or not isinstance(prefix_length, int): + raise ValueError("RTC prefix_length must be an integer") + action_horizon = len(self.modality_configs["action"].delta_indices) + if not 1 <= prefix_length <= action_horizon: + raise ValueError( + f"RTC prefix_length must be in [1, {action_horizon}], got {prefix_length}" + ) + + per_sample_actions: dict[str, np.ndarray] = {} + for key in action_keys: + value = np.asarray(prefix_actions[key]) + if value.dtype != np.float32: + raise ValueError(f"RTC prefix action {key!r} must have dtype float32") + if value.ndim != 3 or value.shape[0] != batch_size: + raise ValueError( + f"RTC prefix action {key!r} must have shape (1, T, D), got {value.shape}" + ) + if value.shape[1] < prefix_length: + raise ValueError( + f"RTC prefix action {key!r} has only {value.shape[1]} steps, " + f"expected at least {prefix_length}" + ) + if not np.isfinite(value).all(): + raise ValueError(f"RTC prefix action {key!r} contains non-finite values") + action = value[0, :prefix_length].copy() + # Some checkpoints store per-timestep action statistics with shape + # (action_horizon, action_dim). StateActionProcessor therefore expects a + # full-horizon array even when RTC only has a shorter leftover prefix. + # Repeat the final physical target for normalization; the sampler still + # uses prefix_length/guidance_horizon to mask every padded timestep. + if prefix_length < action_horizon: + padding = np.repeat( + action[-1:], + action_horizon - prefix_length, + axis=0, + ) + action = np.concatenate((action, padding), axis=0) + per_sample_actions[key] = action + + estimated_delay = rtc.get("estimated_delay_steps") + guidance_horizon = rtc.get("guidance_horizon") + for name, value in ( + ("estimated_delay_steps", estimated_delay), + ("guidance_horizon", guidance_horizon), + ): + if isinstance(value, bool) or not isinstance(value, int): + raise ValueError(f"RTC {name} must be an integer") + if not 0 <= estimated_delay <= guidance_horizon <= action_horizon: + raise ValueError( + "RTC horizons must satisfy 0 <= estimated_delay_steps <= " + f"guidance_horizon <= {action_horizon}" + ) + prefix_schedule = rtc.get("prefix_schedule", "exp") + if prefix_schedule not in ("zeros", "ones", "linear", "exp"): + raise ValueError(f"Unsupported RTC prefix_schedule {prefix_schedule!r}") + max_guidance_weight = rtc.get("max_guidance_weight", 10.0) + if ( + isinstance(max_guidance_weight, bool) + or not isinstance(max_guidance_weight, (int, float)) + or not np.isfinite(max_guidance_weight) + or max_guidance_weight < 0 + ): + raise ValueError("RTC max_guidance_weight must be finite and non-negative") + + model_rtc = { + "prefix_length": prefix_length, + "estimated_delay_steps": min(estimated_delay, prefix_length), + "guidance_horizon": min(guidance_horizon, prefix_length), + "prefix_schedule": prefix_schedule, + "max_guidance_weight": float(max_guidance_weight), + } + model_options = {key: value for key, value in options.items() if key != "rtc"} + model_options["rtc"] = model_rtc + return [per_sample_actions], model_options + def check_observation(self, observation: dict[str, Any]) -> None: """Validate that the observation has the correct structure and types. @@ -391,19 +493,21 @@ def _get_action( Args: observation: Batched observation dictionary - options: Optional parameters (currently unused) + options: Optional inference parameters, including an RTC physical action prefix Returns: Tuple of (actions_dict, info_dict) """ # Step 1: Split batched observation into individual observations unbatched_observations = self._unbatch_observation(observation) + rtc_actions, model_options = self._prepare_rtc_options(options, len(unbatched_observations)) processed_inputs = [] # Step 2: Process each observation through the VLA processor states = [] - for obs in unbatched_observations: - vla_step_data = self._to_vla_step_data(obs) + for index, obs in enumerate(unbatched_observations): + actions = None if rtc_actions is None else rtc_actions[index] + vla_step_data = self._to_vla_step_data(obs, actions) states.append(vla_step_data.states) # dict[str, np.ndarray[np.float32, (T, D)]] messages = [{"type": MessageType.EPISODE_STEP.value, "content": vla_step_data}] processed_inputs.append(self.processor(messages)) @@ -413,8 +517,14 @@ def _get_action( collated_inputs = _rec_to_dtype(collated_inputs, dtype=torch.bfloat16) # Step 4: Run model inference to predict actions - with torch.inference_mode(): - model_pred = self.model.get_action(**collated_inputs) + if rtc_actions is None: + with torch.inference_mode(): + model_pred = self.model.get_action(**collated_inputs, options=model_options) + else: + # torch.inference_mode() cannot be locally overridden by the RTC autograd + # correction. no_grad() can, while still keeping the rest of inference cheap. + with torch.no_grad(): + model_pred = self.model.get_action(**collated_inputs, options=model_options) normalized_action = model_pred["action_pred"].float() # Step 5: Decode actions from normalized space back to physical units @@ -429,7 +539,10 @@ def _get_action( casted_action = { key: value.astype(np.float32) for key, value in unnormalized_action.items() } - return casted_action, {} + info = {} + if rtc_actions is not None: + info["rtc"] = model_options["rtc"] + return casted_action, info def check_action(self, action: dict[str, Any]) -> None: """Validate that the action has the correct structure and types. diff --git a/tests/gr00t/model/test_action_head.py b/tests/gr00t/model/test_action_head.py index a3cf9aa8c..e739b75d3 100644 --- a/tests/gr00t/model/test_action_head.py +++ b/tests/gr00t/model/test_action_head.py @@ -23,7 +23,11 @@ import math from gr00t.configs.model.gr00t_n1d7 import Gr00tN1d7Config -from gr00t.model.gr00t_n1d7.gr00t_n1d7 import Gr00tN1d7ActionHead +from gr00t.model.gr00t_n1d7.gr00t_n1d7 import ( + Gr00tN1d7ActionHead, + get_rtc_guidance_weight, + get_rtc_prefix_weights, +) import pytest import torch from transformers.feature_extraction_utils import BatchFeature @@ -152,6 +156,59 @@ def test_get_action_single_sample(self, action_head): ) assert out["action_pred"].shape[0] == 1 + def test_rtc_zero_guidance_matches_ordinary_sampling(self, action_head): + head, config = action_head + backbone_output = _make_backbone_output(config, batch_size=1) + ordinary_input = _make_action_input(config, batch_size=1) + del ordinary_input["action"] + del ordinary_input["action_mask"] + rtc_input = BatchFeature( + data={ + "state": ordinary_input["state"].clone(), + "embodiment_id": ordinary_input["embodiment_id"].clone(), + "action": torch.randn(1, config.action_horizon, config.max_action_dim), + "action_mask": torch.ones(1, config.action_horizon, config.max_action_dim), + } + ) + + torch.manual_seed(123) + ordinary = head.get_action(backbone_output, ordinary_input)["action_pred"] + + torch.manual_seed(123) + guided = head.get_action( + backbone_output, + rtc_input, + options={ + "rtc": { + "prefix_length": config.action_horizon, + "estimated_delay_steps": 1, + "guidance_horizon": 3, + "prefix_schedule": "exp", + "max_guidance_weight": 0.0, + } + }, + )["action_pred"] + + torch.testing.assert_close(guided, ordinary) + assert not guided.requires_grad + + +class TestRTCUtilities: + def test_prefix_weights_have_frozen_transition_and_free_regions(self): + weights = get_rtc_prefix_weights(2, 5, 8, "exp") + + torch.testing.assert_close(weights[:2], torch.ones(2)) + assert torch.all(weights[2:5] < 1.0) + assert torch.all(weights[2:5] > 0.0) + assert torch.all(weights[2:4] > weights[3:5]) + torch.testing.assert_close(weights[5:], torch.zeros(3)) + + def test_guidance_weight_is_capped_at_flow_boundaries(self): + device = torch.device("cpu") + assert get_rtc_guidance_weight(0.0, 10.0, device=device).item() == 10.0 + assert get_rtc_guidance_weight(1.0, 10.0, device=device).item() == 10.0 + assert get_rtc_guidance_weight(0.5, 10.0, device=device).item() == 2.0 + class TestActionHeadEncodeFeatures: """Test feature encoding helper.""" diff --git a/tests/gr00t/policy/test_gr00t_policy.py b/tests/gr00t/policy/test_gr00t_policy.py index 54de946f8..577972158 100644 --- a/tests/gr00t/policy/test_gr00t_policy.py +++ b/tests/gr00t/policy/test_gr00t_policy.py @@ -169,6 +169,40 @@ def test_get_action_returns_dict(self, policy): assert isinstance(action, dict) assert isinstance(info, dict) + def test_rtc_prefix_is_forwarded_as_processed_actions(self, policy): + obs = _make_observation() + prefix_actions = { + key: np.arange(4, dtype=np.float32).reshape(1, 4, 1) for key in ACTION_KEYS + } + + _, info = policy.get_action( + obs, + options={ + "rtc": { + "prefix_actions": prefix_actions, + "prefix_length": 4, + "estimated_delay_steps": 2, + "guidance_horizon": 4, + "prefix_schedule": "exp", + "max_guidance_weight": 10.0, + } + }, + ) + + message = policy.processor.call_args.args[0][0] + for key in ACTION_KEYS: + processed_prefix = message["content"].actions[key] + assert processed_prefix.shape == (16, 1) + np.testing.assert_array_equal(processed_prefix[:4], prefix_actions[key][0]) + np.testing.assert_array_equal( + processed_prefix[4:], + np.repeat(prefix_actions[key][0, -1:], 12, axis=0), + ) + model_options = policy.model.get_action.call_args.kwargs["options"] + assert "prefix_actions" not in model_options["rtc"] + assert model_options["rtc"]["estimated_delay_steps"] == 2 + assert info["rtc"] == model_options["rtc"] + class _NumpyLanguageSimPolicy: def __init__(self): diff --git a/tests/gr00t/policy/test_policy_service.py b/tests/gr00t/policy/test_policy_service.py index a176808d1..e115108ac 100644 --- a/tests/gr00t/policy/test_policy_service.py +++ b/tests/gr00t/policy/test_policy_service.py @@ -37,8 +37,10 @@ class MockPolicy: def __init__(self): self.strict = False self._reset_count = 0 + self.last_options = None def get_action(self, observation, options=None): + self.last_options = options # Echo back a dummy action dict derived from observation keys action = {"joint_pos": np.zeros(7, dtype=np.float32)} info = {"mock": True} @@ -106,12 +108,22 @@ def test_ping(self, server_client): assert client.ping() is True def test_get_action_roundtrip(self, server_client): - client, _, _ = server_client + client, _, policy = server_client obs = {"state": {"joint_pos": np.zeros(7, dtype=np.float32)}} - result = client.call_endpoint("get_action", {"observation": obs}) + options = { + "rtc": { + "prefix_actions": {"joint_pos": np.ones((1, 4, 7), dtype=np.float32)}, + "prefix_length": 4, + } + } + result = client.get_action(obs, options=options) action, info = result assert "joint_pos" in action np.testing.assert_array_equal(action["joint_pos"], np.zeros(7, dtype=np.float32)) + np.testing.assert_array_equal( + policy.last_options["rtc"]["prefix_actions"]["joint_pos"], + options["rtc"]["prefix_actions"]["joint_pos"], + ) def test_reset(self, server_client): client, _, policy = server_client diff --git a/tests/test_robojudo_client.py b/tests/test_robojudo_client.py new file mode 100644 index 000000000..2d5c882dd --- /dev/null +++ b/tests/test_robojudo_client.py @@ -0,0 +1,572 @@ +from collections import deque +from pathlib import Path +import sys +import threading +import time +from types import SimpleNamespace +from unittest.mock import patch + +import numpy as np +import pytest + + +ROBOJUDO_EXAMPLE = Path(__file__).parents[1] / "examples" / "RoboJuDo" +sys.path.insert(0, str(ROBOJUDO_EXAMPLE)) + +import deploy_adapter as adapter_module # noqa: E402 +import run_robojudo_client as client_module # noqa: E402 + + +def _observation(*, session: int, enabled: bool, sequence: int): + return client_module.Observation( + stream_id="test-stream", + control_session=session, + takeover_enabled=enabled, + sequence=sequence, + images={"ego_view": np.zeros((2, 2, 3), dtype=np.uint8)}, + joint_positions={}, + task="test task", + ) + + +def _command(joint_name: str, value: float): + return { + "positions": {joint_name: value}, + "locomotion_command": np.full(4, value, dtype=np.float32), + } + + +def _chunk(*, start_tick: int, values: list[float], session: int = 1): + return client_module.ActionChunk( + stream_id="test-stream", + control_session=session, + observation_sequence=start_tick, + observation_received_at=0.0, + inference_seconds=0.1, + start_tick=start_tick, + commands=[_command("joint", value) for value in values], + ) + + +def _rtc_chunk(*, values: list[float], task: str = "test task"): + actions = np.asarray(values, dtype=np.float32)[None, :, None] + return client_module.ActionChunk( + stream_id="test-stream", + control_session=1, + observation_sequence=1, + observation_received_at=0.0, + inference_seconds=0.1, + start_tick=0, + commands=[_command("joint", value) for value in values], + physical_actions={"joint": actions}, + task=task, + ) + + +def test_g1_adapter_splits_and_decodes_dexterous_hand_joints(): + profile = adapter_module.PROFILES["g1_23dof"] + adapter = adapter_module.RoboJuDoPolicyAdapter(object(), "g1_23dof") + positions = np.arange(30, dtype=np.float32) + + observation = adapter.build_observation( + np.zeros((4, 5, 3), dtype=np.uint8), + positions, + "pick up the bag", + ) + + assert tuple(observation["state"]) == ( + "left_arm", + "right_arm", + "left_hand", + "right_hand", + ) + np.testing.assert_array_equal(observation["state"]["left_arm"], positions[0:5][None, None]) + np.testing.assert_array_equal(observation["state"]["right_arm"], positions[5:10][None, None]) + np.testing.assert_array_equal(observation["state"]["left_hand"], positions[10:20][None, None]) + np.testing.assert_array_equal(observation["state"]["right_hand"], positions[20:30][None, None]) + + action_chunk = { + "left_arm": positions[0:5][None, None], + "right_arm": positions[5:10][None, None], + "left_hand": positions[10:20][None, None], + "right_hand": positions[20:30][None, None], + "navigate_command": np.asarray([[[0.1, 0.2, 0.3]]], dtype=np.float32), + "base_height_command": np.asarray([[[0.75]]], dtype=np.float32), + } + command = adapter.decode_action_chunk(action_chunk, execution_horizon=1)[0] + + assert tuple(command["positions"]) == profile.joint_names + np.testing.assert_array_equal(list(command["positions"].values()), positions) + np.testing.assert_allclose(command["locomotion_command"], [0.1, 0.2, 0.3, 0.75]) + + +def test_g1_mulcam_adapter_builds_all_video_modalities(): + video_keys = adapter_module.CAMERA_LAYOUTS["mulcam"] + adapter = adapter_module.RoboJuDoPolicyAdapter(object(), "g1_23dof", video_keys) + images = { + key: np.full((4, 5, 3), index, dtype=np.uint8) for index, key in enumerate(video_keys) + } + + observation = adapter.build_observation( + images, + np.arange(30, dtype=np.float32), + "pick up the bag", + ) + + assert tuple(observation["video"]) == video_keys + for key in video_keys: + assert observation["video"][key].shape == (1, 1, 4, 5, 3) + np.testing.assert_array_equal(observation["video"][key][0, 0], images[key]) + + +def _encoded_observation_parts(protocol_version: int, image_keys: tuple[str, ...]): + profile = "x2" + joint_names = adapter_module.PROFILES[profile].joint_names + images = { + key: np.full((6, 8, 3), index * 20, dtype=np.uint8) for index, key in enumerate(image_keys) + } + header = { + "protocol_version": protocol_version, + "profile": profile, + "joint_names": list(joint_names), + "joint_positions": [0.0] * len(joint_names), + "task": "test task", + "stream_id": "test-stream", + "control_session": 1, + "takeover_enabled": True, + "sequence": 7, + } + if protocol_version == 1: + header["shape"] = list(images["ego_view"].shape) + else: + header["image_keys"] = list(image_keys) + header["image_shapes"] = {key: list(image.shape) for key, image in images.items()} + parts = [client_module.msgpack.packb(header, use_bin_type=True)] + for key in image_keys: + ok, jpeg = client_module.cv2.imencode(".jpg", images[key]) + assert ok + parts.append(jpeg.tobytes()) + return profile, parts + + +def test_observation_subscriber_decodes_protocol_v1_single_camera(): + profile, parts = _encoded_observation_parts(1, ("ego_view",)) + subscriber = client_module.ObservationSubscriber.__new__(client_module.ObservationSubscriber) + subscriber.profile = profile + subscriber.expected_joint_names = adapter_module.PROFILES[profile].joint_names + subscriber.expected_image_keys = adapter_module.CAMERA_LAYOUTS["single"] + + observation = subscriber._decode_observation(parts) + + assert tuple(observation.images) == ("ego_view",) + assert observation.images["ego_view"].shape == (6, 8, 3) + + +def test_observation_subscriber_decodes_protocol_v2_mulcam(): + image_keys = adapter_module.CAMERA_LAYOUTS["mulcam"] + profile, parts = _encoded_observation_parts(2, image_keys) + subscriber = client_module.ObservationSubscriber.__new__(client_module.ObservationSubscriber) + subscriber.profile = profile + subscriber.expected_joint_names = adapter_module.PROFILES[profile].joint_names + subscriber.expected_image_keys = image_keys + + observation = subscriber._decode_observation(parts) + + assert tuple(observation.images) == image_keys + assert all(image.shape == (6, 8, 3) for image in observation.images.values()) + + +def test_observation_subscriber_rejects_missing_mulcam_part(): + image_keys = adapter_module.CAMERA_LAYOUTS["mulcam"] + profile, parts = _encoded_observation_parts(2, image_keys) + subscriber = client_module.ObservationSubscriber.__new__(client_module.ObservationSubscriber) + subscriber.profile = profile + subscriber.expected_joint_names = adapter_module.PROFILES[profile].joint_names + subscriber.expected_image_keys = image_keys + + with pytest.raises(ValueError, match="has 3 parts, expected 4"): + subscriber._decode_observation(parts[:-1]) + + +def test_x2_adapter_reserves_hands_without_requiring_them(): + profile = adapter_module.PROFILES["x2"] + assert profile.left_hand_joint_names == () + assert profile.right_hand_joint_names == () + assert tuple(key for key, _ in profile.joint_groups) == ("left_arm", "right_arm") + + adapter = adapter_module.RoboJuDoPolicyAdapter(object(), "x2") + observation = adapter.build_observation( + np.zeros((4, 5, 3), dtype=np.uint8), + np.arange(14, dtype=np.float32), + "test x2", + ) + assert tuple(observation["state"]) == ("left_arm", "right_arm") + + +def test_rtc_queue_replaces_using_actual_delay_and_keeps_prefix_in_lockstep(): + queue = client_module.RTCActionQueue() + chunk = _rtc_chunk(values=[10.0, 11.0, 12.0, 13.0]) + + assert queue.replace(chunk, skipped_steps=2) + assert queue.qsize() == 2 + left_over = queue.get_left_over(("test-stream", 1), "test task") + np.testing.assert_allclose(left_over["joint"], [[[12.0], [13.0]]]) + assert queue.pop()["positions"]["joint"] == 12.0 + np.testing.assert_allclose( + queue.get_left_over(("test-stream", 1), "test task")["joint"], + [[[13.0]]], + ) + + +def test_rtc_queue_rejects_cross_task_prefix_and_expires_old_chunk(): + queue = client_module.RTCActionQueue() + chunk = _rtc_chunk(values=[1.0, 2.0]) + assert queue.replace(chunk, skipped_steps=0) + assert queue.get_left_over(("test-stream", 1), "different task") is None + + assert not queue.replace(chunk, skipped_steps=2) + assert queue.qsize() == 0 + assert queue.pop() is None + + +def test_parse_args_accepts_horizon_above_previous_fixed_limit(): + argv = [ + "run_robojudo_client.py", + "--profile", + "x2", + "--robot-endpoint", + "tcp://127.0.0.1:8561", + "--execution-horizon", + "32", + ] + with patch.object(sys, "argv", argv): + assert client_module.parse_args().execution_horizon == 32 + + +def test_parse_args_rejects_non_positive_horizon(): + argv = [ + "run_robojudo_client.py", + "--profile", + "x2", + "--robot-endpoint", + "tcp://127.0.0.1:8561", + "--execution-horizon", + "0", + ] + with patch.object(sys, "argv", argv), pytest.raises(SystemExit): + client_module.parse_args() + + +def test_parse_args_accepts_g1_mulcam_layout(): + argv = [ + "run_robojudo_client.py", + "--profile", + "g1_23dof", + "--camera-layout", + "mulcam", + "--robot-endpoint", + "tcp://127.0.0.1:8561", + ] + with patch.object(sys, "argv", argv): + assert client_module.parse_args().camera_layout == "mulcam" + + +def test_temporal_ensemble_equal_average_uses_aligned_chunk_diagonal(): + ensembler = client_module.ACTTemporalEnsembler(("joint",), temporal_ensemble_coeff=0.0) + ensembler.add_chunk(_chunk(start_tick=0, values=[0.0, 1.0, 2.0, 3.0])) + ensembler.add_chunk(_chunk(start_tick=1, values=[10.0, 11.0, 12.0, 13.0])) + + action, contributors = ensembler.get_action(current_tick=2) + + assert contributors == 2 + assert action["positions"]["joint"] == 6.5 + np.testing.assert_allclose(action["locomotion_command"], np.full(4, 6.5)) + + +def test_temporal_ensemble_matches_act_exponential_weighting_for_sparse_queries(): + coeff = 0.01 + ensembler = client_module.ACTTemporalEnsembler(("joint",), coeff) + ensembler.add_chunk(_chunk(start_tick=0, values=[0.0, 1.0, 2.0, 3.0])) + ensembler.add_chunk(_chunk(start_tick=3, values=[9.0, 10.0, 11.0, 12.0])) + + action, contributors = ensembler.get_action(current_tick=3) + + newer_weight = np.exp(-coeff * 3) + expected = (3.0 + 9.0 * newer_weight) / (1.0 + newer_weight) + assert contributors == 2 + np.testing.assert_allclose(action["positions"]["joint"], expected) + + +def test_temporal_ensemble_skips_elapsed_prefix_and_prunes_expired_chunks(): + ensembler = client_module.ACTTemporalEnsembler(("joint",), temporal_ensemble_coeff=0.0) + ensembler.add_chunk(_chunk(start_tick=4, values=[4.0, 5.0, 6.0, 7.0])) + + action, contributors = ensembler.get_action(current_tick=7) + assert contributors == 1 + assert action["positions"]["joint"] == 7.0 + + action, contributors = ensembler.get_action(current_tick=8) + assert action is None + assert contributors == 0 + assert ensembler.active_chunk_count == 0 + + +def test_temporal_ensemble_reset_and_safe_hold(): + ensembler = client_module.ACTTemporalEnsembler(("joint",), temporal_ensemble_coeff=0.0) + ensembler.add_chunk(_chunk(start_tick=0, values=[1.0])) + ensembler.reset() + assert ensembler.get_action(0) == (None, 0) + + command = { + "positions": {"joint": 1.5}, + "locomotion_command": np.asarray([0.2, -0.3, 0.4, 0.65], dtype=np.float32), + } + held = client_module.make_safe_hold_command(command) + assert held["positions"] == command["positions"] + np.testing.assert_allclose(held["locomotion_command"], [0.0, 0.0, 0.0, 0.65]) + + +def test_inference_discards_disabled_session_and_uses_reenabled_session(): + first_inference_started = threading.Event() + release_first_inference = threading.Event() + inference_calls = 0 + + class FakePolicyClient: + def __init__(self, host, port): + del host, port + + def ping(self): + return True + + def close(self): + return None + + class FakeAdapter: + def __init__(self, policy_client, profile, video_keys): + del policy_client, profile, video_keys + + def get_action_chunk(self, **kwargs): + nonlocal inference_calls + del kwargs + inference_calls += 1 + if inference_calls == 1: + first_inference_started.set() + assert release_first_inference.wait(timeout=1) + return SimpleNamespace( + commands=[ + { + "positions": {}, + "locomotion_command": np.zeros(4, dtype=np.float32), + } + ] + ) + + runner = client_module.DoubleBufferedPolicyRunner.__new__( + client_module.DoubleBufferedPolicyRunner + ) + runner.policy_host = "test" + runner.policy_port = 0 + runner.profile = "x2" + runner.task_override = None + runner.execution_horizon = 1 + runner.execution_mode = "double_buffer" + runner.video_keys = adapter_module.CAMERA_LAYOUTS["single"] + runner._condition = threading.Condition() + runner._stopping = False + runner._error = None + runner._latest_observation = _observation(session=1, enabled=True, sequence=1) + runner._latest_observation_at = 1.0 + runner._last_inferred_session = None + runner._last_inferred_sequence = -1 + runner._pending_commands = None + runner._ready_chunks = deque() + runner._control_tick = 0 + runner._control_tick_session = None + + with ( + patch.object(client_module, "PolicyClient", FakePolicyClient), + patch.object(client_module, "RoboJuDoPolicyAdapter", FakeAdapter), + ): + inference_thread = threading.Thread(target=runner._inference_loop) + inference_thread.start() + assert first_inference_started.wait(timeout=1) + + with runner._condition: + runner._latest_observation = _observation(session=1, enabled=False, sequence=2) + runner._latest_observation_at = 2.0 + runner._condition.notify_all() + release_first_inference.set() + + with runner._condition: + runner._latest_observation = _observation(session=2, enabled=True, sequence=3) + runner._latest_observation_at = 3.0 + runner._condition.notify_all() + assert runner._condition.wait_for( + lambda: runner._pending_commands is not None, + timeout=1, + ) + assert runner._pending_commands.control_session == 2 + assert runner._pending_commands.observation_sequence == 3 + runner._stopping = True + runner._condition.notify_all() + + inference_thread.join(timeout=1) + + assert not inference_thread.is_alive() + assert inference_calls == 2 + + +def test_temporal_ensemble_inference_does_not_wait_for_ready_queue_to_drain(): + inference_calls = 0 + + class FakePolicyClient: + def __init__(self, host, port): + del host, port + + def ping(self): + return True + + def close(self): + return None + + class FakeAdapter: + def __init__(self, policy_client, profile, video_keys): + del policy_client, profile, video_keys + + def get_action_chunk(self, **kwargs): + nonlocal inference_calls + del kwargs + inference_calls += 1 + return SimpleNamespace(commands=[_command("joint", float(inference_calls))]) + + runner = client_module.DoubleBufferedPolicyRunner.__new__( + client_module.DoubleBufferedPolicyRunner + ) + runner.policy_host = "test" + runner.policy_port = 0 + runner.profile = "x2" + runner.task_override = None + runner.execution_horizon = 1 + runner.execution_mode = "temporal_ensemble" + runner.video_keys = adapter_module.CAMERA_LAYOUTS["single"] + runner._condition = threading.Condition() + runner._stopping = False + runner._error = None + runner._latest_observation = _observation(session=1, enabled=True, sequence=1) + runner._latest_observation_at = 1.0 + runner._last_inferred_session = None + runner._last_inferred_sequence = -1 + runner._pending_commands = None + runner._ready_chunks = deque() + runner._control_tick = 0 + runner._control_tick_session = ("test-stream", 1) + + with ( + patch.object(client_module, "PolicyClient", FakePolicyClient), + patch.object(client_module, "RoboJuDoPolicyAdapter", FakeAdapter), + ): + inference_thread = threading.Thread(target=runner._inference_loop) + inference_thread.start() + with runner._condition: + assert runner._condition.wait_for(lambda: len(runner._ready_chunks) == 1, timeout=1) + runner._latest_observation = _observation(session=1, enabled=True, sequence=2) + runner._latest_observation_at = 2.0 + runner._condition.notify_all() + assert runner._condition.wait_for(lambda: len(runner._ready_chunks) == 2, timeout=1) + runner._stopping = True + runner._condition.notify_all() + inference_thread.join(timeout=1) + + assert not inference_thread.is_alive() + assert inference_calls == 2 + + +def test_rtc_inference_sends_leftover_prefix_and_uses_actual_delay_on_replace(): + inference_started = threading.Event() + release_inference = threading.Event() + received_options = None + + class FakePolicyClient: + def __init__(self, host, port): + del host, port + + def ping(self): + return True + + def close(self): + return None + + class FakeAdapter: + def __init__(self, policy_client, profile, video_keys): + del policy_client, profile, video_keys + + def get_action_chunk(self, **kwargs): + nonlocal received_options + received_options = kwargs["options"] + inference_started.set() + assert release_inference.wait(timeout=1) + values = [100.0, 101.0, 102.0, 103.0] + return SimpleNamespace( + commands=[_command("joint", value) for value in values], + actions={"joint": np.asarray(values, dtype=np.float32)[None, :, None]}, + ) + + runner = client_module.DoubleBufferedPolicyRunner.__new__( + client_module.DoubleBufferedPolicyRunner + ) + runner.policy_host = "test" + runner.policy_port = 0 + runner.profile = "x2" + runner.task_override = None + runner.execution_horizon = 3 + runner.execution_mode = "rtc" + runner.video_keys = adapter_module.CAMERA_LAYOUTS["single"] + runner.rtc_prefix_schedule = "exp" + runner.rtc_max_guidance_weight = 10.0 + runner.command_period = 0.1 + runner.observation_timeout = 10.0 + runner._condition = threading.Condition() + runner._stopping = False + runner._error = None + runner._latest_observation = _observation(session=1, enabled=True, sequence=1) + runner._latest_observation_at = time.monotonic() + runner._last_inferred_session = None + runner._last_inferred_sequence = -1 + runner._pending_commands = None + runner._ready_chunks = deque() + runner._rtc_queue = client_module.RTCActionQueue() + runner._rtc_queue.replace(_rtc_chunk(values=[10.0, 11.0, 12.0, 13.0]), 1) + runner._rtc_inference_latencies = deque([0.05], maxlen=10) + runner._control_tick = 3 + runner._control_tick_session = ("test-stream", 1) + + with ( + patch.object(client_module, "PolicyClient", FakePolicyClient), + patch.object(client_module, "RoboJuDoPolicyAdapter", FakeAdapter), + ): + inference_thread = threading.Thread(target=runner._inference_loop) + inference_thread.start() + assert inference_started.wait(timeout=1) + with runner._condition: + runner._control_tick = 5 + release_inference.set() + with runner._condition: + assert runner._condition.wait_for( + lambda: ( + runner._rtc_queue.chunk is not None + and runner._rtc_queue.chunk.observation_sequence == 1 + and runner._rtc_queue.chunk.commands[0]["positions"]["joint"] == 100.0 + ), + timeout=1, + ) + runner._stopping = True + runner._condition.notify_all() + inference_thread.join(timeout=1) + + rtc = received_options["rtc"] + np.testing.assert_allclose(rtc["prefix_actions"]["joint"], [[[11.0], [12.0], [13.0]]]) + assert rtc["estimated_delay_steps"] == 1 + assert runner._policy_action_horizon == 4 + assert runner._rtc_queue.qsize() == 2 + assert runner._rtc_queue.pop()["positions"]["joint"] == 102.0 + assert not inference_thread.is_alive()