mirror of
https://github.com/huggingface/lerobot.git
synced 2026-07-26 11:16:00 +00:00
fix(robomme): align docs and tests with nested pixels obs layout
Addresses PR #3311 review feedback: - Docs: correct observation keys to `pixels/image` / `pixels/wrist_image` (mapped to `observation.images.image` / `observation.images.wrist_image`) and drop the now-obsolete column-rename snippet. - Tests: assert `result["pixels"]["image"]` instead of flat `pixels/image`, matching the nested layout required by `preprocess_observation()`. Co-Authored-By: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
This commit is contained in:
@@ -109,15 +109,9 @@ dataset = LeRobotDataset("lerobot/robomme")
|
|||||||
|
|
||||||
### Feature key alignment (training)
|
### Feature key alignment (training)
|
||||||
|
|
||||||
The env wrapper exposes `front_rgb` and `wrist_rgb` as observation keys. The `features_map` in `RoboMMEEnv` maps these to `obs_images.front` and `obs_images.wrist` for the policy.
|
The env wrapper exposes `pixels/image` and `pixels/wrist_image` as observation keys. The `features_map` in `RoboMMEEnv` maps these to `observation.images.image` and `observation.images.wrist_image` for the policy. State is exposed as `agent_pos` and maps to `observation.state`.
|
||||||
|
|
||||||
The dataset uses `image` and `wrist_image` as column names. When fine-tuning from the dataset, you may need to rename columns to match your policy's expected input keys:
|
The dataset's `image` and `wrist_image` columns already align with the policy input keys, so no renaming is needed when fine-tuning.
|
||||||
|
|
||||||
```python
|
|
||||||
dataset = LeRobotDataset("lerobot/robomme")
|
|
||||||
# rename dataset columns if needed:
|
|
||||||
# dataset.hf_dataset = dataset.hf_dataset.rename_column("image", "front_rgb")
|
|
||||||
```
|
|
||||||
|
|
||||||
## Action Spaces
|
## Action Spaces
|
||||||
|
|
||||||
|
|||||||
@@ -118,7 +118,11 @@ def test_robomme_features_action_dim_ee_pose():
|
|||||||
|
|
||||||
def test_convert_obs_list_format():
|
def test_convert_obs_list_format():
|
||||||
"""_convert_obs takes the last element from list-format obs fields and
|
"""_convert_obs takes the last element from list-format obs fields and
|
||||||
emits the lerobot-canonical keys (pixels/image, pixels/wrist_image, agent_pos)."""
|
emits a nested ``pixels`` dict (image, wrist_image) plus ``agent_pos``.
|
||||||
|
|
||||||
|
The nested layout is required so ``preprocess_observation()`` in
|
||||||
|
``envs/utils.py`` maps each camera to ``observation.images.<cam>``.
|
||||||
|
"""
|
||||||
_install_robomme_stub()
|
_install_robomme_stub()
|
||||||
try:
|
try:
|
||||||
from lerobot.envs.robomme import RoboMMEGymEnv
|
from lerobot.envs.robomme import RoboMMEGymEnv
|
||||||
@@ -138,8 +142,8 @@ def test_convert_obs_list_format():
|
|||||||
}
|
}
|
||||||
|
|
||||||
result = env._convert_obs(obs_raw)
|
result = env._convert_obs(obs_raw)
|
||||||
np.testing.assert_array_equal(result["pixels/image"], front)
|
np.testing.assert_array_equal(result["pixels"]["image"], front)
|
||||||
np.testing.assert_array_equal(result["pixels/wrist_image"], wrist)
|
np.testing.assert_array_equal(result["pixels"]["wrist_image"], wrist)
|
||||||
assert result["agent_pos"].shape == (8,)
|
assert result["agent_pos"].shape == (8,)
|
||||||
np.testing.assert_array_almost_equal(result["agent_pos"][:7], joints)
|
np.testing.assert_array_almost_equal(result["agent_pos"][:7], joints)
|
||||||
assert result["agent_pos"][7] == gripper[0]
|
assert result["agent_pos"][7] == gripper[0]
|
||||||
@@ -163,8 +167,8 @@ def test_convert_obs_array_format():
|
|||||||
"gripper_state_list": np.zeros(2, dtype=np.float32),
|
"gripper_state_list": np.zeros(2, dtype=np.float32),
|
||||||
}
|
}
|
||||||
result = env._convert_obs(obs_raw)
|
result = env._convert_obs(obs_raw)
|
||||||
assert result["pixels/image"].shape == (256, 256, 3)
|
assert result["pixels"]["image"].shape == (256, 256, 3)
|
||||||
assert result["pixels/wrist_image"].shape == (256, 256, 3)
|
assert result["pixels"]["wrist_image"].shape == (256, 256, 3)
|
||||||
assert result["agent_pos"].shape == (8,)
|
assert result["agent_pos"].shape == (8,)
|
||||||
finally:
|
finally:
|
||||||
_uninstall_robomme_stub()
|
_uninstall_robomme_stub()
|
||||||
|
|||||||
Reference in New Issue
Block a user