feat(depth maps): adding support for depth in LeRobot (#3644)

* feat(depth): add depth quantization helpers and tests

* feat(video): add ffv1 to supported codecs

* feat(depth): persist depth metadata

* feat(depth): extend quantization tools to better fit the encoding/decoding pipeline

* feat(depth): plumb DepthEncoderConfig through LeRobotDataset and DatasetWriter

* feat(depth): wire StreamingVideoEncoder + writer to depth encoder

* feat(depth): wire DatasetReader to decode_depth_frames

* feat(cameras/realsense): expose async depth in metric meters

* feat(features): route 2D camera shapes to observation.depth.<key>

* feat(robots/so_follower): emit + populate depth keys when use_depth

* feat(record): plumb DepthEncoderConfig through lerobot-record

* feat(viz): render depth observations as rr.DepthImage in Viridis

* feat(depth maps writer): adding support for raw depth maps recording with image writer

* chore(format): format code

* feat(depth shape): ensuring depth maps shape is always including the channel

* feat(is_depth): simplifying is_depth nested name + legacy support

* fix(stop_event): fixing stop_event race condition in camera classes

* fix(plumbing): fixing missing parts in the depth maps pipeline

* chore(typos): fixing typos

* test(fix): fixing exisiting tests to still work with latest features

* tests(depth): adding new tests for depth integration validation

* feat(pix_fmt channels): use PyAv to check get pixel formats number of channels

* feat(refactor): refactor DepthEncoderConfig quantization pipeline, so that the methods do not live in the config class. Add pixel format - channels validation.Move the default pixel format for depth in the config file.

* fix(pre-commit): fixing mutable defautl value

* fix(info): fixing info metadata update when is_depth_map was set

* tests(typos): fixing typos in tests

* fix(realsense): fixing typo in realsense serial number

* fix(normalization): restricting 255 normalization to non depth/uint8 images only

* fix(typo): fixing typo

* fix(TIFF): add missing quantization and cleanup for TIFF files

* feat(batched dequantization): optimizing dequantize_depth for torch based batched dequantization

* feat(tools): adding depth support in LeRobotDataset edition tools

* test(aggregate): extending aggregation tests to depth frames

* test(cleaning): cleaning up tests

* fix(from_video_info): fixing early validation issue in from_video_info

* fix(typo): fixing typo

* fix(is_depth): adding missing doctrings and is_depth arguments in video decoding functions

Co-authored-by: Wensi (Vince) Ai <59036629+wensi-ai@users.noreply.github.com>

* fix(depth units): fixing depth units output for the realsense cameras

* feat(output unit): adding support for output unit specification at dataset reading/training time

Co-authored-by: Wensi (Vince) Ai <59036629+wensi-ai@users.noreply.github.com>

* test(depth): cleaning up depth tests

* test(depth encoding): updating and cleaning video/depth encoding tests

* chore(format): formatting code

* docs(depth): improving depth maps docs

* test(fix): fixing depth tests

* test(dataset tools): adding missing tests for new dataset edition tools features

* chore(format): formatting code

* fix(pyav check): fixing PyAV option validation for integer codec options by normalizing
numeric values before calling `is_integer()`

Co-authored-by: Wensi (Vince) Ai <59036629+wensi-ai@users.noreply.github.com>

* docs(mermaid): fixing mermaid diagram

* fix(rebase): rebase follow up corrections

* feat(dataset tools): adding missing docstrings and features for depth fill support in dataset edition tools

* docs(docstring): updating docstrings

* docs(dataset tools): updating docs

* fix(save images): fixing image saving in dataset tools

* fix(update video info): fixing update video info logic to match the recording and editing use cases

* test(reencode): fixing reencoding monkeypatch

* fix(review): add Claude review

* chore(format): format code

* fix(update video info): ditching the differentiated approahces for video info update - video info are always updated unless for preserved keys.

* chore(rebase): fixing rebase merge conflicts

* test(visualization): fixing visualization tests

* feat(docstrings): adding explicit docstring for encoding parameters. Docstrigns will now show up as description in the CLI --help.

* feat(mm as default): adding a global DEFAULT_DEPTH_UNIT variable setting mm as default depth unit

* fix(RGB <-> camera): renaming camera_encoder to rgb_encoder for clarity

* chore(TODO): removing deprecated TODO

* doc(write_u16_plane): improving docstrings for write_u16_plane

* feat(units): adding constants for depth frames units (m and mm)

* fix(spam): replacing spamming warning but a debug log

* feat(leagcy metadata): adding automatic metadata update for legacy 'video.is_depth_map' feature

* fix(copy&reindex): fixing metadat reshaping for single channel frames

* fix(ImageNet): excluding dpeth frames from ImageNet stats

* fix(PyAV container seek): fixing initial  PyAV container seek to be robust againsy codec choice

* feat(lerobot-dataset-viz): adding support for depth in lerobot-dataset-viz

* fix(compress): removing rerun compression for DepthImages

* fix(signle channel squeeze): fixing single channel squeezing

* chore(format): format code

* fix(streaming): adding support for dequantization in streaming_dataset.py

* refactor(read depth): factorizing depth reading methods for realsense camera and adding support for depth-only usage

* chore(renaming): fixing missed RGBEncoderConfig renamings

* docs(renaming): reflecting renamings in a clearer way in the docs

* chore(annotation): excluding depth from the annotation pipeline

* feat(robots): adding depth support in compatible follower robots

* feat(LeSadKiwi): excluding LeKiwi from depth support (for now)

* chore(fail): removing misplaced file

* chore(fail): removing misplaced file

* fix(remove ffv1): removing ffv1 as it does not support MP4

* docs(cheat sheet): adding depth and video encoding to the cheat sheet

* fix(lossless): tuning depth encoding parameters for lossless depth storage

* test(fix): fixing failing tests

* depth(ZMQ): excluding ZMQ from depth support

* Revert "depth(ZMQ): excluding ZMQ from depth support"

This reverts commit b95cf4e4c2.

* fix(image transforms): excluding depth frames from images transforms

* fix(typo): typo

* fix(stats): fixing stats computation for depth frames

* fix(TIFF vs. pytorch): adding an extra uint16 to float32 conversion for depth maps stored as raw TIFF images

* fix(typos): fixing typos

* test(dtype): fixing stats computation typing tests

---------

Signed-off-by: Steven Palma <imstevenpmwork@ieee.org>
Co-authored-by: Wensi (Vince) Ai <59036629+wensi-ai@users.noreply.github.com>
Co-authored-by: Steven Palma <imstevenpmwork@ieee.org>
Co-authored-by: Wensi Ai <wsai@stanford.edu>
This commit is contained in:
Caroline Pascal
2026-06-27 14:21:21 +02:00
committed by GitHub
parent 6a788fbdb0
commit 3dd19d043e
69 changed files with 2740 additions and 679 deletions
+8
View File
@@ -157,6 +157,14 @@ finally:
</hfoption> </hfoption>
</hfoptions> </hfoptions>
### Working with depth
The Intel RealSense and Reachy 2 cameras can capture both color and depth in lockstep. Calling `read()` returns the **color** frame as `(H, W, 3)` `uint8`. Calling `read_depth()` returns the **depth map** as `(H, W, 1)` `uint16`, where each pixel value is the distance from the sensor expressed in **millimetres**. A pixel value of `0` typically means "no measurement available" (out-of-range, occluded, or low-confidence).
During recording, the control loop peeks the freshest buffered frames non-blockingly via `read_latest()` (color) and `read_latest_depth()` (depth), adding the depth map as a sibling feature (e.g. `front_depth` next to `front`).
For how depth streams are stored and encoded when recording a dataset, see the [Depth streams](./video_encoding_parameters#depth-streams) section of the video encoding guide.
## Use your phone's camera ## Use your phone's camera
<hfoptions id="use phone"> <hfoptions id="use phone">
+30
View File
@@ -89,6 +89,36 @@ Control the data recording flow using keyboard shortcuts:
- Press **Left Arrow (`←`)**: Delete current episode and retry. - Press **Left Arrow (`←`)**: Delete current episode and retry.
- Press **Escape (`ESC`)**: Stop, encode videos, and upload. - Press **Escape (`ESC`)**: Stop, encode videos, and upload.
### Recording depth
Intel RealSense cameras (`type: intelrealsense`) record a depth stream when you set `use_depth: true`. Depth is quantized to 12-bit codes and stored as its own video.
```bash
lerobot-record \
... \
--robot.cameras="{ head: {type: intelrealsense, serial_number_or_name: \"0123456789\", width: 640, height: 480, fps: 30, use_depth: true} }" \
--dataset.repo_id=${HF_USER}/so101_depth_test \
--dataset.single_task="put the red brick in a bowl" \
--dataset.depth_encoder.depth_min=0.01 \
--dataset.depth_encoder.depth_max=10.0 \
--dataset.depth_encoder.shift=0.0 \
--dataset.depth_encoder.use_log=true
```
### Video encoding parameters
RGB and depth streams are encoded independently via the `--dataset.rgb_encoder.*` and `--dataset.depth_encoder.*` keys.
```bash
lerobot-record \
... \
--dataset.rgb_encoder.vcodec=h264 \
--dataset.rgb_encoder.pix_fmt=yuv420p \
--dataset.rgb_encoder.crf=23 \
--dataset.depth_encoder.vcodec=hevc \
--dataset.depth_encoder.extra_options='{"x265-params": "lossless=1"}'
```
### Training ### Training
Depending on your hardware training the policy might take a few hours. That's how you train simple `ACT` policy: Depending on your hardware training the policy might take a few hours. That's how you train simple `ACT` policy:
+1 -1
View File
@@ -194,7 +194,7 @@ lerobot-record \
--dataset.single_task="Navigate around obstacles" \ --dataset.single_task="Navigate around obstacles" \
--dataset.streaming_encoding=true \ --dataset.streaming_encoding=true \
--dataset.encoder_threads=2 \ --dataset.encoder_threads=2 \
# --dataset.camera_encoder.vcodec=auto \ # --dataset.rgb_encoder.vcodec=auto \
--display_data=true --display_data=true
``` ```
+1 -1
View File
@@ -124,7 +124,7 @@ lerobot-rollout\
--dataset.single_task="Grab and handover the red cube to the other arm" \ --dataset.single_task="Grab and handover the red cube to the other arm" \
--dataset.streaming_encoding=true \ --dataset.streaming_encoding=true \
--dataset.encoder_threads=2 \ --dataset.encoder_threads=2 \
# --dataset.camera_encoder.vcodec=auto \ # --dataset.rgb_encoder.vcodec=auto \
--policy.path=<user>/groot-bimanual \ # your trained model --policy.path=<user>/groot-bimanual \ # your trained model
--duration=600 --duration=600
``` ```
+2 -2
View File
@@ -232,7 +232,7 @@ lerobot-record \
--dataset.private=true \ --dataset.private=true \
--dataset.streaming_encoding=true \ --dataset.streaming_encoding=true \
--dataset.encoder_threads=2 \ --dataset.encoder_threads=2 \
# --dataset.camera_encoder.vcodec=auto \ # --dataset.rgb_encoder.vcodec=auto \
--display_data=true --display_data=true
``` ```
@@ -278,6 +278,6 @@ lerobot-record \
--dataset.num_episodes=10 \ --dataset.num_episodes=10 \
--dataset.streaming_encoding=true \ --dataset.streaming_encoding=true \
--dataset.encoder_threads=2 \ --dataset.encoder_threads=2 \
# --dataset.camera_encoder.vcodec=auto \ # --dataset.rgb_encoder.vcodec=auto \
--policy.path=outputs/train/hopejr_hand/checkpoints/last/pretrained_model --policy.path=outputs/train/hopejr_hand/checkpoints/last/pretrained_model
``` ```
+1 -1
View File
@@ -207,7 +207,7 @@ lerobot-record \
--dataset.num_episodes=5 \ --dataset.num_episodes=5 \
--dataset.single_task="Grab the black cube" \ --dataset.single_task="Grab the black cube" \
--dataset.streaming_encoding=true \ --dataset.streaming_encoding=true \
# --dataset.camera_encoder.vcodec=auto \ # --dataset.rgb_encoder.vcodec=auto \
--dataset.encoder_threads=2 --dataset.encoder_threads=2
``` ```
</hfoption> </hfoption>
+1 -1
View File
@@ -44,7 +44,7 @@ lerobot-record \
--dataset.num_episodes=5 \ --dataset.num_episodes=5 \
--dataset.single_task="Grab the black cube" \ --dataset.single_task="Grab the black cube" \
--dataset.streaming_encoding=true \ --dataset.streaming_encoding=true \
# --dataset.camera_encoder.vcodec=auto \ # --dataset.rgb_encoder.vcodec=auto \
--dataset.encoder_threads=2 --dataset.encoder_threads=2
``` ```
+2 -2
View File
@@ -161,7 +161,7 @@ lerobot-record \
--dataset.private=true \ --dataset.private=true \
--dataset.streaming_encoding=true \ --dataset.streaming_encoding=true \
--dataset.encoder_threads=2 \ --dataset.encoder_threads=2 \
# --dataset.camera_encoder.vcodec=auto \ # --dataset.rgb_encoder.vcodec=auto \
--display_data=true --display_data=true
``` ```
@@ -203,7 +203,7 @@ lerobot-record \
--dataset.private=true \ --dataset.private=true \
--dataset.streaming_encoding=true \ --dataset.streaming_encoding=true \
--dataset.encoder_threads=2 \ --dataset.encoder_threads=2 \
# --dataset.camera_encoder.vcodec=auto \ # --dataset.rgb_encoder.vcodec=auto \
--display_data=true --display_data=true
``` ```
+20 -20
View File
@@ -17,7 +17,7 @@ This makes `save_episode()` near-instant (the video is already encoded by the ti
| Parameter | CLI Flag | Type | Default | Description | | Parameter | CLI Flag | Type | Default | Description |
| ----------------------- | --------------------------------- | ------------- | ------------- | ----------------------------------------------------------------- | | ----------------------- | --------------------------------- | ------------- | ------------- | ----------------------------------------------------------------- |
| `streaming_encoding` | `--dataset.streaming_encoding` | `bool` | `True` | Enable real-time encoding during capture | | `streaming_encoding` | `--dataset.streaming_encoding` | `bool` | `True` | Enable real-time encoding during capture |
| `vcodec` | `--dataset.camera_encoder.vcodec` | `str` | `"libsvtav1"` | Video codec. `"auto"` detects best HW encoder | | `vcodec` | `--dataset.rgb_encoder.vcodec` | `str` | `"libsvtav1"` | Video codec. `"auto"` detects best HW encoder |
| `encoder_threads` | `--dataset.encoder_threads` | `int \| None` | `None` (auto) | Threads per encoder instance. `None` will leave the vcoded decide | | `encoder_threads` | `--dataset.encoder_threads` | `int \| None` | `None` (auto) | Threads per encoder instance. `None` will leave the vcoded decide |
| `encoder_queue_maxsize` | `--dataset.encoder_queue_maxsize` | `int` | `30` | Max buffered frames per camera (~1s at 30fps). Consumes RAM | | `encoder_queue_maxsize` | `--dataset.encoder_queue_maxsize` | `int` | `30` | Max buffered frames per camera (~1s at 30fps). Consumes RAM |
@@ -82,15 +82,15 @@ Use HW encoding when:
### Available HW Encoders ### Available HW Encoders
| Encoder | Platform | Hardware | CLI Value | | Encoder | Platform | Hardware | CLI Value |
| ------------------- | ------------- | ------------------------------------------------------------------------------------------------ | --------------------------------------------------- | | ------------------- | ------------- | ------------------------------------------------------------------------------------------------ | ------------------------------------------------ |
| `h264_videotoolbox` | macOS | Apple Silicon / Intel | `--dataset.camera_encoder.vcodec=h264_videotoolbox` | | `h264_videotoolbox` | macOS | Apple Silicon / Intel | `--dataset.rgb_encoder.vcodec=h264_videotoolbox` |
| `hevc_videotoolbox` | macOS | Apple Silicon / Intel | `--dataset.camera_encoder.vcodec=hevc_videotoolbox` | | `hevc_videotoolbox` | macOS | Apple Silicon / Intel | `--dataset.rgb_encoder.vcodec=hevc_videotoolbox` |
| `h264_nvenc` | Linux/Windows | NVIDIA GPU | `--dataset.camera_encoder.vcodec=h264_nvenc` | | `h264_nvenc` | Linux/Windows | NVIDIA GPU | `--dataset.rgb_encoder.vcodec=h264_nvenc` |
| `hevc_nvenc` | Linux/Windows | NVIDIA GPU | `--dataset.camera_encoder.vcodec=hevc_nvenc` | | `hevc_nvenc` | Linux/Windows | NVIDIA GPU | `--dataset.rgb_encoder.vcodec=hevc_nvenc` |
| `h264_vaapi` | Linux | Intel/AMD GPU | `--dataset.camera_encoder.vcodec=h264_vaapi` | | `h264_vaapi` | Linux | Intel/AMD GPU | `--dataset.rgb_encoder.vcodec=h264_vaapi` |
| `h264_qsv` | Linux/Windows | Intel Quick Sync | `--dataset.camera_encoder.vcodec=h264_qsv` | | `h264_qsv` | Linux/Windows | Intel Quick Sync | `--dataset.rgb_encoder.vcodec=h264_qsv` |
| `auto` | Any | Probes the system for available HW encoders. Falls back to `libsvtav1` if no HW encoder is found | `--dataset.camera_encoder.vcodec=auto` | | `auto` | Any | Probes the system for available HW encoders. Falls back to `libsvtav1` if no HW encoder is found | `--dataset.rgb_encoder.vcodec=auto` |
> [!NOTE] > [!NOTE]
> In order to use the HW accelerated encoders you might need to upgrade your GPU drivers. > In order to use the HW accelerated encoders you might need to upgrade your GPU drivers.
@@ -100,15 +100,15 @@ Use HW encoding when:
## 5. Troubleshooting ## 5. Troubleshooting
| Symptom | Likely Cause | Fix | | Symptom | Likely Cause | Fix |
| ------------------------------------------------------------------ | -------------------------------------------- | --------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------- | | ------------------------------------------------------------------ | -------------------------------------------- | ------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------ |
| System freezes or choppy robot movement or Rerun visualization lag | CPU starved (100% load usage) | Close other apps, reduce encoding throughput, lower `encoder_threads`, use `h264`, use `display_data=False`. If the CPU continues to be at 100% then it might be insufficient for your setup, consider `--dataset.streaming_encoding=false` or HW encoding (`--dataset.camera_encoder.vcodec=auto`) | | System freezes or choppy robot movement or Rerun visualization lag | CPU starved (100% load usage) | Close other apps, reduce encoding throughput, lower `encoder_threads`, use `h264`, use `display_data=False`. If the CPU continues to be at 100% then it might be insufficient for your setup, consider `--dataset.streaming_encoding=false` or HW encoding (`--dataset.rgb_encoder.vcodec=auto`) |
| "Encoder queue full" warnings or dropped frames in dataset | Encoder can't keep up (Queue overflow) | If CPU is not at 100%: Increase `encoder_threads`, increase `encoder_queue_maxsize` or use HW encoding (`--dataset.camera_encoder.vcodec=auto`). | | "Encoder queue full" warnings or dropped frames in dataset | Encoder can't keep up (Queue overflow) | If CPU is not at 100%: Increase `encoder_threads`, increase `encoder_queue_maxsize` or use HW encoding (`--dataset.rgb_encoder.vcodec=auto`). |
| High RAM usage | Queue filling faster than encoding | `encoder_threads` too low or CPU insufficient. Reduce `encoder_queue_maxsize` or use HW encoding | | High RAM usage | Queue filling faster than encoding | `encoder_threads` too low or CPU insufficient. Reduce `encoder_queue_maxsize` or use HW encoding |
| Large video files | Using HW encoder or H.264 | Expected trade-off. Switch to `libsvtav1` if CPU allows | | Large video files | Using HW encoder or H.264 | Expected trade-off. Switch to `libsvtav1` if CPU allows |
| `save_episode()` still slow | `streaming_encoding` is `False` | Set `--dataset.streaming_encoding=true` | | `save_episode()` still slow | `streaming_encoding` is `False` | Set `--dataset.streaming_encoding=true` |
| Encoder thread crash | Codec not available or invalid settings | Check `vcodec` is installed, try `--dataset.camera_encoder.vcodec=auto` | | Encoder thread crash | Codec not available or invalid settings | Check `vcodec` is installed, try `--dataset.rgb_encoder.vcodec=auto` |
| Recorded dataset is missing frames | CPU/GPU starvation or occasional load spikes | If ~5% of frames are missing, your system is likely overloaded — follow the recommendations above. If fewer frames are missing (~2%), they are probably due to occasional transient load spikes (often at startup) and can be considered expected. | | Recorded dataset is missing frames | CPU/GPU starvation or occasional load spikes | If ~5% of frames are missing, your system is likely overloaded — follow the recommendations above. If fewer frames are missing (~2%), they are probably due to occasional transient load spikes (often at startup) and can be considered expected. |
## 6. Recommended Configurations ## 6. Recommended Configurations
@@ -146,7 +146,7 @@ On very constrained systems, streaming encoding may compete too heavily with the
# 2camsx 640x480x3 @30fps: Requires some tuning. # 2camsx 640x480x3 @30fps: Requires some tuning.
# Use H.264, disable streaming, consider batching encoding # Use H.264, disable streaming, consider batching encoding
lerobot-record --dataset.camera_encoder.vcodec=h264 --dataset.streaming_encoding=false ... lerobot-record --dataset.rgb_encoder.vcodec=h264 --dataset.streaming_encoding=false ...
``` ```
## 7. Closing note ## 7. Closing note
+49 -8
View File
@@ -11,8 +11,9 @@ LeRobot provides several utilities for manipulating datasets:
3. **Merge Datasets** - Combine multiple datasets into one. The datasets must have identical features, and episodes are concatenated in the order specified in `repo_ids` 3. **Merge Datasets** - Combine multiple datasets into one. The datasets must have identical features, and episodes are concatenated in the order specified in `repo_ids`
4. **Add Features** - Add new features to a dataset 4. **Add Features** - Add new features to a dataset
5. **Remove Features** - Remove features from a dataset 5. **Remove Features** - Remove features from a dataset
6. **Convert to Video** - Convert image-based datasets to video format for efficient storage 6. **Convert to Video** - Convert image-based datasets to video format for efficient storage (RGB and depth cameras are encoded with separate encoders)
7. **Show the Info of Datasets** - Show the summary of datasets information such as number of episode etc. 7. **Re-encode Videos** - Re-encode an existing video dataset's RGB and/or depth streams with new encoder settings
8. **Show the Info of Datasets** - Show the summary of datasets information such as number of episode etc.
The core implementation is in `lerobot.datasets.dataset_tools`. The core implementation is in `lerobot.datasets.dataset_tools`.
An example script detailing how to use the tools API is available in `examples/dataset/use_dataset_tools.py`. An example script detailing how to use the tools API is available in `examples/dataset/use_dataset_tools.py`.
@@ -117,10 +118,19 @@ lerobot-edit-dataset \
--repo_id lerobot/pusht_image \ --repo_id lerobot/pusht_image \
--operation.type convert_image_to_video \ --operation.type convert_image_to_video \
--operation.output_dir outputs/pusht_video \ --operation.output_dir outputs/pusht_video \
--operation.camera_encoder.vcodec libsvtav1 \ --operation.rgb_encoder.vcodec libsvtav1 \
--operation.camera_encoder.pix_fmt yuv420p \ --operation.rgb_encoder.pix_fmt yuv420p \
--operation.camera_encoder.g 2 \ --operation.rgb_encoder.g 2 \
--operation.camera_encoder.crf 30 --operation.rgb_encoder.crf 30
# Convert a dataset that includes depth maps, customizing the depth encoder
lerobot-edit-dataset \
--repo_id lerobot/pusht_image \
--operation.type convert_image_to_video \
--operation.output_dir outputs/pusht_video \
--operation.depth_encoder.depth_min 0.01 \
--operation.depth_encoder.depth_max 10.0 \
--operation.depth_encoder.use_log true
# Convert only specific episodes # Convert only specific episodes
lerobot-edit-dataset \ lerobot-edit-dataset \
@@ -147,11 +157,42 @@ lerobot-edit-dataset \
**Parameters:** **Parameters:**
- `output_dir`: Custom output directory (optional - by default uses `new_repo_id` or `{repo_id}_video`) - `output_dir`: Custom output directory (optional - by default uses `new_repo_id` or `{repo_id}_video`)
- `camera_encoder`: Video encoder settings — all sub-fields accessible via `--operation.camera_encoder.<field>. See [Video Encoding Parameters](./video_encoding_parameters) for more details. - `rgb_encoder`: Video encoder settings applied to RGB cameras — all sub-fields accessible via `--operation.rgb_encoder.<field>`. See [Video Encoding Parameters](./video_encoding_parameters) for more details.
- `depth_encoder`: Video encoder settings applied to depth-map cameras (e.g. from an Intel RealSense). In addition to the standard encoder fields it exposes the depth quantization knobs (`depth_min`, `depth_max`, `shift`, `use_log`), accessible via `--operation.depth_encoder.<field>`. These quantization settings are persisted to the dataset metadata so depth can be dequantized back to physical units on load. See the [Depth streams](./video_encoding_parameters#depth-streams) section for details.
- `episode_indices`: List of specific episodes to convert (default: all episodes) - `episode_indices`: List of specific episodes to convert (default: all episodes)
- `num_workers`: Number of parallel workers for processing (default: 4) - `num_workers`: Number of parallel workers for processing (default: 4)
**Note:** The resulting dataset will be a proper LeRobotDataset with all cameras encoded as videos in the `videos/` directory, with parquet files containing only metadata (no raw image data). All episodes, stats, and tasks are preserved. **Note:** The resulting dataset will be a proper LeRobotDataset with all cameras encoded as videos in the `videos/` directory, with parquet files containing only metadata (no raw image data). Depth-map cameras are detected automatically and routed to the `depth_encoder`, while RGB cameras use the `rgb_encoder`. All episodes, stats, and tasks are preserved.
#### Re-encode Videos
Re-encode the videos of an existing video dataset with different encoder settings, without going back to raw frames. RGB videos use the `rgb_encoder` and depth videos use the `depth_encoder`. Provide only the encoder(s) you want to re-encode; the other stream type is left untouched.
```bash
# Re-encode all RGB videos with new settings (saves to lerobot/pusht_reencoded by default)
lerobot-edit-dataset \
--repo_id lerobot/pusht \
--operation.type reencode_videos \
--operation.rgb_encoder.vcodec h264 \
--operation.rgb_encoder.pix_fmt yuv420p \
--operation.rgb_encoder.crf 23
# Re-encode both RGB and depth videos in a dataset with depth maps
lerobot-edit-dataset \
--repo_id lerobot/pusht_depth \
--operation.type reencode_videos \
--operation.rgb_encoder.vcodec h264 \
--operation.depth_encoder.crf 50
```
**Parameters:**
- `rgb_encoder`: Encoder settings applied to every RGB video. Omit to skip re-encoding RGB videos.
- `depth_encoder`: Encoder settings applied to every depth video. Omit to skip re-encoding depth videos.
- `num_workers`: Number of parallel workers for processing.
> [!NOTE]
> When re-encoding depth videos, the existing depth quantization parameters (`depth_min`, `depth_max`, `shift`, `use_log`) and the `is_depth_map` flag are **preserved** — re-encoding only changes the codec/quality of the stored stream, not how depth is dequantized on load.
### Show the information of datasets ### Show the information of datasets
+84 -13
View File
@@ -2,15 +2,15 @@
When video storage is enabled, LeRobot stores each camera stream as an **MP4** file instead of saving one image file per timestep. Video encoding compresses across time, which usually cuts dataset size and I/O compared to a pile of PNG, while keeping MP4 — a format every player and loader understands. When video storage is enabled, LeRobot stores each camera stream as an **MP4** file instead of saving one image file per timestep. Video encoding compresses across time, which usually cuts dataset size and I/O compared to a pile of PNG, while keeping MP4 — a format every player and loader understands.
Encoding frames into an MP4 is a full FFmpeg pipeline: choice of encoder, pixel format, GOP/keyframes, quality vs. speed, and optional extra encoder flags. Most of these knobs are user-tunable through `camera_encoder`, a nested `VideoEncoderConfig` (`lerobot.configs.video.VideoEncoderConfig`) passed through PyAV. Encoding frames into an MP4 is a full FFmpeg pipeline: choice of encoder, pixel format, GOP/keyframes, quality vs. speed, and optional extra encoder flags. Most of these knobs are user-tunable through `rgb_encoder`, a nested `RGBEncoderConfig` (`lerobot.configs.video.RGBEncoderConfig`) passed through PyAV.
You can set these parameters from the CLI with `--dataset.camera_encoder.<field>` (e.g. with `lerobot-record` or `lerobot-rollout`). The same block applies to every camera video stream in that run. You can set these parameters from the CLI with `--dataset.rgb_encoder.<field>` (e.g. with `lerobot-record` or `lerobot-rollout`). The same block applies to every camera video stream in that run.
<Tip> <Tip>
Video storage must be on for `camera_encoder` to have any effect — Video storage must be on for `rgb_encoder` to have any effect —
`use_videos=True` in Python APIs, or `--dataset.video=true` on the CLI (the `use_videos=True` in Python APIs, or `--dataset.video=true` on the CLI (the
recording default). With video off, inputs stay as images and `camera_encoder` recording default). With video off, inputs stay as images and `rgb_encoder` is
is ignored. ignored.
</Tip> </Tip>
For details on **when** frames are written vs. encoded (streaming vs. post-episode), queues, and other top-level `--dataset.*` switches, see [Streaming Video Encoding](./streaming_video_encoding). For an encoding-parameter comparison and experiments, see the [video-benchmark Space](https://huggingface.co/spaces/lerobot/video-benchmark). For details on **when** frames are written vs. encoded (streaming vs. post-episode), queues, and other top-level `--dataset.*` switches, see [Streaming Video Encoding](./streaming_video_encoding). For an encoding-parameter comparison and experiments, see the [video-benchmark Space](https://huggingface.co/spaces/lerobot/video-benchmark).
@@ -33,9 +33,9 @@ lerobot-record \
--dataset.single_task="Grab the cube" \ --dataset.single_task="Grab the cube" \
--dataset.streaming_encoding=true \ --dataset.streaming_encoding=true \
--dataset.encoder_threads=2 \ --dataset.encoder_threads=2 \
--dataset.camera_encoder.vcodec=h264 \ --dataset.rgb_encoder.vcodec=h264 \
--dataset.camera_encoder.preset=fast \ --dataset.rgb_encoder.preset=fast \
--dataset.camera_encoder.extra_options={"tune": "film", "profile:v": "high", "bf": 2} \ --dataset.rgb_encoder.extra_options={"tune": "film", "profile:v": "high", "bf": 2} \
--display_data=true --display_data=true
``` ```
@@ -50,7 +50,7 @@ Only override these parameters if you have a specific reason to, and measure the
</Tip> </Tip>
All flags below are prefixed with `--dataset.camera_encoder.` on the CLI. All flags below are prefixed with `--dataset.rgb_encoder.` on the CLI.
| Parameter | Type | Default | Description | | Parameter | Type | Default | Description |
| --------------- | ---------------- | ------------- | -------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------- | | --------------- | ---------------- | ------------- | -------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------- |
@@ -65,6 +65,77 @@ All flags below are prefixed with `--dataset.camera_encoder.` on the CLI.
--- ---
## Depth streams
Depth maps (Intel RealSense, Reachy 2) are stored as their **own video streams** alongside the RGB streams. Raw depth (`uint16` millimetres or `float32` metres) can't survive an 8-bit codec, so LeRobot **quantizes** each map to a 12-bit code (`[0, 4095]`) — logarithmically by default, to match the `1/depth` error profile of depth sensors — then packs it into a high-bit-depth pixel format (`gray12le`) and encodes it with a 12-bit codec.
```mermaid
flowchart LR
A["Raw depth (uint16 mm / float32 m)"] --> B["Clip to depth_min, depth_max"]
B --> C["Quantize to 12-bit code 04095 (log or linear)"]
C --> D["Pack into gray12le"]
D --> E["Encode video (hevc Main 12)"]
E --> F[("MP4 + metadata: depth_min/max, shift, use_log")]
F -. "load time (depth_output_unit)" .-> G["Dequantize to mm or m"]
classDef input fill:#e3f2fd,stroke:#1565c0,color:#0d47a1;
classDef encode fill:#ede7f6,stroke:#5e35b1,color:#311b92;
classDef store fill:#fff8e1,stroke:#f9a825,color:#e65100;
classDef load fill:#e8f5e9,stroke:#2e7d32,color:#1b5e20;
class A input;
class B,C,D,E encode;
class F store;
class G load;
```
Configure the depth pipeline through a parallel **`depth_encoder`** block (`DepthEncoderConfig`). It shares every `RGBEncoderConfig` field (`vcodec`, `pix_fmt`, `crf`, …) and adds four quantizer knobs, set via `--dataset.depth_encoder.<field>`:
```bash
lerobot-record \
... \
--dataset.depth_encoder.vcodec=hevc \
--dataset.depth_encoder.depth_min=0.05 \
--dataset.depth_encoder.depth_max=5.0 \
--dataset.depth_encoder.use_log=true
```
| Parameter | Type | Default | Description |
| --------------- | ------- | ------------------------------- | -------------------------------------------------------------------------------------------------------------------------------------- |
| `vcodec` | `str` | `"hevc"` | HEVC Main 12 (a 12-bit-capable codec, MP4-compatible). |
| `extra_options` | `dict` | `{"x265-params": "lossless=1"}` | **Depth defaults to lossless** (exact round-trip); `crf` is ignored. Pass `extra_options={}` and set `crf` for a smaller lossy stream. |
| `pix_fmt` | `str` | `"gray12le"` | Single-channel 12-bit pixel format used to carry the quantized codes. |
| `depth_min` | `float` | `0.01` | Depth in metres mapped to quantum `0`. Values below are clipped on decode. |
| `depth_max` | `float` | `10.0` | Depth in metres mapped to quantum `4095`. Values above are clipped on decode. |
| `shift` | `float` | `3.5` | Pre-log offset (metres) used in logarithmic quantization for numerical stability near zero. Must satisfy `depth_min + shift > 0`. |
| `use_log` | `bool` | `True` | If `true`, quantize in log-space (recommended for typical depth sensors). Set to `false` for uniform/linear quantization. |
> [!TIP]
> `depth_min`, `depth_max`, and `shift` are always interpreted in **metres**, regardless of the input depth's unit. Inputs are auto-detected: integer arrays (e.g. `uint16` millimetres straight from a RealSense) are treated as millimetres, floating arrays as metres.
> Pick `depth_min` / `depth_max` to bracket the actual working range of your sensor — quanta outside that range saturate, which can crush detail at the boundaries.
Depth features are flagged with `"is_depth_map": true` in `meta/info.json`, and their quantizer settings (`video.depth_min`, `video.depth_max`, `video.shift`, `video.use_log`) are persisted — which is what lets depth be **dequantized back to physical units** on load.
### Output unit at load time
`depth_encoder` is a **record-time** concern. The unit that depth maps are dequantized to on _load_ (e.g. during training) is set separately by the read-time flag `--dataset.depth_output_unit`:
```bash
lerobot-train \
--dataset.repo_id=<my_username>/<my_dataset_name> \
--dataset.depth_output_unit=m \
--policy.type=act
```
| Parameter | Type | Default | Description |
| ------------------- | ----- | ------- | -------------------------------------------------------------------------------------------- |
| `depth_output_unit` | `str` | `"mm"` | Physical unit depth maps are dequantized to on load: `"mm"` (millimetres) or `"m"` (metres). |
> [!TIP]
> This is purely a decode-time presentation choice — it does **not** alter the stored video or its metadata, so the same dataset can be read as `mm` or `m` without re-encoding. It has no effect on datasets without depth cameras.
---
## Persistence in dataset metadata ## Persistence in dataset metadata
After the first episode of a video stream is encoded, the encoder configuration is **persisted into the dataset metadata** (`meta/info.json`) under each video feature, alongside the values probed from the file itself. For a video feature `observation.images.<camera>`, the layout in `info.json` is: After the first episode of a video stream is encoded, the encoder configuration is **persisted into the dataset metadata** (`meta/info.json`) under each video feature, alongside the values probed from the file itself. For a video feature `observation.images.<camera>`, the layout in `info.json` is:
@@ -82,7 +153,7 @@ After the first episode of a video stream is encoded, the encoder configuration
"video.pix_fmt": "yuv420p", "video.pix_fmt": "yuv420p",
"video.fps": 30, "video.fps": 30,
"video.channels": 3, "video.channels": 3,
"video.is_depth_map": false, "is_depth_map": false,
"video.g": 2, "video.g": 2,
"video.crf": 30, "video.crf": 30,
"video.preset": "fast", "video.preset": "fast",
@@ -97,12 +168,12 @@ After the first episode of a video stream is encoded, the encoder configuration
Two sources contribute to the `info` block: Two sources contribute to the `info` block:
- **Stream-derived** (read back from the encoded MP4 with PyAV): `video.height`, `video.width`, `video.codec`, `video.pix_fmt`, `video.fps`, `video.channels`, `video.is_depth_map`, plus `audio.*` if an audio stream is present. - **Stream-derived** (read back from the encoded MP4 with PyAV): `video.height`, `video.width`, `video.codec`, `video.pix_fmt`, `video.fps`, `video.channels`, `is_depth_map`, plus `audio.*` if an audio stream is present.
- **Encoder-derived** (taken from `VideoEncoderConfig`): `video.g`, `video.crf`, `video.preset`, `video.fast_decode`, `video.video_backend`, `video.extra_options`. - **Encoder-derived** (taken from `RGBEncoderConfig` or `DepthEncoderConfig`): `video.g`, `video.crf`, `video.preset`, `video.fast_decode`, `video.video_backend`, `video.extra_options`.
<Tip> <Tip>
This block is populated **once**, from the **first** episode. It assumes every This block is populated **once**, from the **first** episode. It assumes every
episode in the dataset was encoded with the same `camera_encoder`. Changing episode in the dataset was encoded with the same `rgb_encoder`. Changing
encoder settings partway through a recording is not supported — the encoder settings partway through a recording is not supported — the
`info.json` will only reflect the parameters used for the first episode. `info.json` will only reflect the parameters used for the first episode.
</Tip> </Tip>
@@ -36,7 +36,7 @@ from typing import Any, Protocol
import PIL.Image import PIL.Image
import torch import torch
from lerobot.configs.video import VideoEncoderConfig from lerobot.configs import RGBEncoderConfig
from lerobot.datasets.video_utils import decode_video_frames, reencode_video from lerobot.datasets.video_utils import decode_video_frames, reencode_video
from .reader import EpisodeRecord, snap_to_frame from .reader import EpisodeRecord, snap_to_frame
@@ -164,7 +164,9 @@ class VideoFrameProvider:
# only for video-stored cameras. Image-stored cameras (also in # only for video-stored cameras. Image-stored cameras (also in
# ``camera_keys``) would KeyError, so restrict the list — and the # ``camera_keys``) would KeyError, so restrict the list — and the
# default — to video keys. # default — to video keys.
keys = list(self._meta.video_keys) # Depth cameras are excluded from the annotation pipeline for now.
depth_keys = set(self._meta.depth_keys)
keys = [key for key in self._meta.video_keys if key not in depth_keys]
# Last-resort fallback: if metadata didn't surface any video keys but # Last-resort fallback: if metadata didn't surface any video keys but
# the caller explicitly named a camera (``--vlm.camera_key=...``), # the caller explicitly named a camera (``--vlm.camera_key=...``),
# trust them — the key is by definition known to exist on the dataset. # trust them — the key is by definition known to exist on the dataset.
@@ -276,12 +278,12 @@ class VideoFrameProvider:
from_timestamp = float(ep[f"videos/{self.camera_key}/from_timestamp"]) from_timestamp = float(ep[f"videos/{self.camera_key}/from_timestamp"])
to_timestamp = float(ep[f"videos/{self.camera_key}/to_timestamp"]) to_timestamp = float(ep[f"videos/{self.camera_key}/to_timestamp"])
src = self.root / self._meta.get_video_file_path(record.episode_index, self.camera_key) src = self.root / self._meta.get_video_file_path(record.episode_index, self.camera_key)
encoder = VideoEncoderConfig(vcodec="h264", pix_fmt="yuv420p", g=None, crf=23, preset="ultrafast") encoder = RGBEncoderConfig(vcodec="h264", pix_fmt="yuv420p", g=None, crf=23, preset="ultrafast")
try: try:
reencode_video( reencode_video(
src, src,
out_path, out_path,
camera_encoder=encoder, video_encoder=encoder,
overwrite=True, overwrite=True,
start_time_s=from_timestamp, start_time_s=from_timestamp,
end_time_s=to_timestamp, end_time_s=to_timestamp,
+3 -2
View File
@@ -105,8 +105,9 @@ def raw_observation_to_observation(
def prepare_image(image: torch.Tensor) -> torch.Tensor: def prepare_image(image: torch.Tensor) -> torch.Tensor:
"""Minimal preprocessing to turn int8 images to float32 in [0, 1], and create a memory-contiguous tensor""" """Minimal preprocessing to turn RGB uint8 images to float32 in [0, 1], and create a memory-contiguous tensor"""
image = image.type(torch.float32) / 255 if image.dtype == torch.uint8:
image = image.type(torch.float32) / 255
image = image.contiguous() image = image.contiguous()
return image return image
+3 -1
View File
@@ -436,7 +436,7 @@ class OpenCVCamera(Camera):
Internal loop run by the background thread for asynchronous reading. Internal loop run by the background thread for asynchronous reading.
On each iteration: On each iteration:
1. Reads a color frame 1. Reads a color frame (blocking call)
2. Stores result in latest_frame and updates timestamp (thread-safe) 2. Stores result in latest_frame and updates timestamp (thread-safe)
3. Sets new_frame_event to notify listeners 3. Sets new_frame_event to notify listeners
@@ -485,6 +485,8 @@ class OpenCVCamera(Camera):
if self.thread is not None and self.thread.is_alive(): if self.thread is not None and self.thread.is_alive():
self.thread.join(timeout=2.0) self.thread.join(timeout=2.0)
if self.thread.is_alive():
logger.warning(f"{self} read thread did not terminate within timeout.")
self.thread = None self.thread = None
self.stop_event = None self.stop_event = None
+120 -63
View File
@@ -128,6 +128,7 @@ class RealSenseCamera(Camera):
self.fps = config.fps self.fps = config.fps
self.color_mode = config.color_mode self.color_mode = config.color_mode
self.use_rgb = config.use_rgb
self.use_depth = config.use_depth self.use_depth = config.use_depth
self.warmup_s = config.warmup_s self.warmup_s = config.warmup_s
@@ -195,12 +196,15 @@ class RealSenseCamera(Camera):
# NOTE(Steven/Caroline): Enforcing at least one second of warmup as RS cameras need a bit of time before the first read. If we don't wait, the first read from the warmup will raise. # NOTE(Steven/Caroline): Enforcing at least one second of warmup as RS cameras need a bit of time before the first read. If we don't wait, the first read from the warmup will raise.
self.warmup_s = max(self.warmup_s, 1) self.warmup_s = max(self.warmup_s, 1)
warmup_read = self.async_read if self.use_rgb else self.async_read_depth
start_time = time.time() start_time = time.time()
while time.time() - start_time < self.warmup_s: while time.time() - start_time < self.warmup_s:
self.async_read(timeout_ms=self.warmup_s * 1000) warmup_read(timeout_ms=self.warmup_s * 1000)
time.sleep(0.1) time.sleep(0.1)
with self.frame_lock: with self.frame_lock:
if self.latest_color_frame is None or self.use_depth and self.latest_depth_frame is None: if (self.use_rgb and self.latest_color_frame is None) or (
self.use_depth and self.latest_depth_frame is None
):
raise ConnectionError(f"{self} failed to capture frames during warmup.") raise ConnectionError(f"{self} failed to capture frames during warmup.")
logger.info(f"{self} connected.") logger.info(f"{self} connected.")
@@ -268,13 +272,13 @@ class RealSenseCamera(Camera):
) )
if len(found_devices) > 1: if len(found_devices) > 1:
serial_numbers = [dev["serial_number"] for dev in found_devices] serial_numbers = [dev["id"] for dev in found_devices]
raise ValueError( raise ValueError(
f"Multiple RealSense cameras found with name '{name}'. " f"Multiple RealSense cameras found with name '{name}'. "
f"Please use a unique serial number instead. Found SNs: {serial_numbers}" f"Please use a unique serial number instead. Found SNs: {serial_numbers}"
) )
serial_number = str(found_devices[0]["serial_number"]) serial_number = str(found_devices[0]["id"])
return serial_number return serial_number
def _configure_rs_pipeline_config(self, rs_config: Any) -> None: def _configure_rs_pipeline_config(self, rs_config: Any) -> None:
@@ -282,15 +286,17 @@ class RealSenseCamera(Camera):
rs.config.enable_device(rs_config, self.serial_number) rs.config.enable_device(rs_config, self.serial_number)
if self.width and self.height and self.fps: if self.width and self.height and self.fps:
rs_config.enable_stream( if self.use_rgb:
rs.stream.color, self.capture_width, self.capture_height, rs.format.rgb8, self.fps rs_config.enable_stream(
) rs.stream.color, self.capture_width, self.capture_height, rs.format.rgb8, self.fps
)
if self.use_depth: if self.use_depth:
rs_config.enable_stream( rs_config.enable_stream(
rs.stream.depth, self.capture_width, self.capture_height, rs.format.z16, self.fps rs.stream.depth, self.capture_width, self.capture_height, rs.format.z16, self.fps
) )
else: else:
rs_config.enable_stream(rs.stream.color) if self.use_rgb:
rs_config.enable_stream(rs.stream.color)
if self.use_depth: if self.use_depth:
rs_config.enable_stream(rs.stream.depth) rs_config.enable_stream(rs.stream.depth)
@@ -298,8 +304,9 @@ class RealSenseCamera(Camera):
def _configure_capture_settings(self) -> None: def _configure_capture_settings(self) -> None:
"""Sets fps, width, and height from device stream if not already configured. """Sets fps, width, and height from device stream if not already configured.
Uses the color stream profile to update unset attributes. Handles rotation by Uses the color stream profile (or the depth stream profile when the color
swapping width/height when needed. Original capture dimensions are always stored. stream is disabled) to update unset attributes. Handles rotation by swapping
width/height when needed. Original capture dimensions are always stored.
Raises: Raises:
DeviceNotConnectedError: If device is not connected. DeviceNotConnectedError: If device is not connected.
@@ -308,7 +315,8 @@ class RealSenseCamera(Camera):
if self.rs_profile is None: if self.rs_profile is None:
raise RuntimeError(f"{self}: rs_profile must be initialized before use.") raise RuntimeError(f"{self}: rs_profile must be initialized before use.")
stream = self.rs_profile.get_stream(rs.stream.color).as_video_stream_profile() rs_stream = rs.stream.color if self.use_rgb else rs.stream.depth
stream = self.rs_profile.get_stream(rs_stream).as_video_stream_profile()
if self.fps is None: if self.fps is None:
self.fps = stream.fps() self.fps = stream.fps()
@@ -323,6 +331,14 @@ class RealSenseCamera(Camera):
self.width, self.height = actual_width, actual_height self.width, self.height = actual_width, actual_height
self.capture_width, self.capture_height = actual_width, actual_height self.capture_width, self.capture_height = actual_width, actual_height
def _read(self, read_depth: bool = False) -> NDArray[Any]:
"""Shared helper for :meth:`read`/:meth:`read_depth`: wait for a fresh color or depth frame."""
if self.thread is None or not self.thread.is_alive():
raise RuntimeError(f"{self} read thread is not running.")
self.new_frame_event.clear()
return self._async_read(timeout_ms=10000, read_depth=read_depth)
@check_if_not_connected @check_if_not_connected
def read_depth(self, timeout_ms: int = 200) -> NDArray[Any]: def read_depth(self, timeout_ms: int = 200) -> NDArray[Any]:
""" """
@@ -332,8 +348,8 @@ class RealSenseCamera(Camera):
from the camera hardware via the RealSense pipeline. from the camera hardware via the RealSense pipeline.
Returns: Returns:
np.ndarray: The depth map as a NumPy array (height, width) np.ndarray: The depth map as a NumPy array (height, width, 1)
of type `np.uint16` (raw depth values in millimeters) and rotation. of type `np.uint16` (raw depth values in millimeters).
Raises: Raises:
DeviceNotConnectedError: If the camera is not connected. DeviceNotConnectedError: If the camera is not connected.
@@ -349,20 +365,7 @@ class RealSenseCamera(Camera):
f"Failed to capture depth frame '.read_depth()'. Depth stream is not enabled for {self}." f"Failed to capture depth frame '.read_depth()'. Depth stream is not enabled for {self}."
) )
if self.thread is None or not self.thread.is_alive(): return self._read(read_depth=True)
raise RuntimeError(f"{self} read thread is not running.")
self.new_frame_event.clear()
_ = self.async_read(timeout_ms=10000)
with self.frame_lock:
depth_map = self.latest_depth_frame
if depth_map is None:
raise RuntimeError("No depth frame available. Ensure camera is streaming.")
return depth_map
def _read_from_hardware(self): def _read_from_hardware(self):
if self.rs_pipeline is None: if self.rs_pipeline is None:
@@ -405,12 +408,10 @@ class RealSenseCamera(Camera):
f"{self} read() timeout_ms parameter is deprecated and will be removed in future versions." f"{self} read() timeout_ms parameter is deprecated and will be removed in future versions."
) )
if self.thread is None or not self.thread.is_alive(): if not self.use_rgb:
raise RuntimeError(f"{self} read thread is not running.") raise RuntimeError(f"{self}: cannot read color — camera was configured with use_rgb=False.")
self.new_frame_event.clear() frame = self._read()
frame = self.async_read(timeout_ms=10000)
read_duration_ms = (time.perf_counter() - start_time) * 1e3 read_duration_ms = (time.perf_counter() - start_time) * 1e3
logger.debug(f"{self} read took: {read_duration_ms:.1f}ms") logger.debug(f"{self} read took: {read_duration_ms:.1f}ms")
@@ -465,8 +466,8 @@ class RealSenseCamera(Camera):
Internal loop run by the background thread for asynchronous reading. Internal loop run by the background thread for asynchronous reading.
On each iteration: On each iteration:
1. Reads a color frame with 500ms timeout 1. Reads a color/depth frame (blocking call with 10s timeout)
2. Stores result in latest_frame and updates timestamp (thread-safe) 2. Stores result in latest_color_frame/latest_depth_frame and updates timestamp (thread-safe)
3. Sets new_frame_event to notify listeners 3. Sets new_frame_event to notify listeners
Stops on DeviceNotConnectedError, logs other errors and continues. Stops on DeviceNotConnectedError, logs other errors and continues.
@@ -479,19 +480,24 @@ class RealSenseCamera(Camera):
while not stop_event.is_set(): while not stop_event.is_set():
try: try:
frame = self._read_from_hardware() frame = self._read_from_hardware()
color_frame_raw = frame.get_color_frame()
color_frame = np.asanyarray(color_frame_raw.get_data()) if self.use_rgb:
processed_color_frame = self._postprocess_image(color_frame) color_frame_raw = frame.get_color_frame()
color_frame = np.asanyarray(color_frame_raw.get_data())
processed_color_frame = self._postprocess_image(color_frame)
if self.use_depth: if self.use_depth:
depth_frame_raw = frame.get_depth_frame() depth_frame_raw = frame.get_depth_frame()
depth_frame = np.asanyarray(depth_frame_raw.get_data()) depth_frame = np.asanyarray(depth_frame_raw.get_data())
processed_depth_frame = self._postprocess_image(depth_frame, depth_frame=True) processed_depth_frame = self._postprocess_image(depth_frame, depth_frame=True)
if processed_depth_frame.ndim == 2: # (H, W) -> (H, W, 1)
processed_depth_frame = processed_depth_frame[..., np.newaxis]
capture_time = time.perf_counter() capture_time = time.perf_counter()
with self.frame_lock: with self.frame_lock:
self.latest_color_frame = processed_color_frame if self.use_rgb:
self.latest_color_frame = processed_color_frame
if self.use_depth: if self.use_depth:
self.latest_depth_frame = processed_depth_frame self.latest_depth_frame = processed_depth_frame
self.latest_timestamp = capture_time self.latest_timestamp = capture_time
@@ -523,6 +529,8 @@ class RealSenseCamera(Camera):
if self.thread is not None and self.thread.is_alive(): if self.thread is not None and self.thread.is_alive():
self.thread.join(timeout=2.0) self.thread.join(timeout=2.0)
if self.thread.is_alive(): # pragma: no cover
logger.warning(f"{self} read thread did not terminate within timeout.")
self.thread = None self.thread = None
self.stop_event = None self.stop_event = None
@@ -533,7 +541,26 @@ class RealSenseCamera(Camera):
self.latest_timestamp = None self.latest_timestamp = None
self.new_frame_event.clear() self.new_frame_event.clear()
# NOTE(Steven): Missing implementation for depth for now def _async_read(self, timeout_ms: float, read_depth: bool = False) -> NDArray[Any]:
"""Shared helper for :meth:`async_read`/:meth:`async_read_depth`: return the latest buffered frame."""
if self.thread is None or not self.thread.is_alive():
raise RuntimeError(f"{self} read thread is not running.")
if not self.new_frame_event.wait(timeout=timeout_ms / 1000.0):
raise TimeoutError(
f"Timed out waiting for frame from camera {self} after {timeout_ms} ms. "
f"Read thread alive: {self.thread.is_alive()}."
)
with self.frame_lock:
frame = self.latest_depth_frame if read_depth else self.latest_color_frame
self.new_frame_event.clear()
if frame is None:
raise RuntimeError(f"Internal error: Event set but no frame available for {self}.")
return frame
@check_if_not_connected @check_if_not_connected
def async_read(self, timeout_ms: float = 200) -> NDArray[Any]: def async_read(self, timeout_ms: float = 200) -> NDArray[Any]:
""" """
@@ -558,25 +585,31 @@ class RealSenseCamera(Camera):
RuntimeError: If the background thread died unexpectedly or another error occurs. RuntimeError: If the background thread died unexpectedly or another error occurs.
""" """
if not self.use_rgb:
raise RuntimeError(f"{self}: cannot read color — camera was configured with use_rgb=False.")
return self._async_read(timeout_ms=timeout_ms)
def _read_latest(self, max_age_ms: int, read_depth: bool = False) -> NDArray[Any]:
"""Shared helper for :meth:`read_latest`/:meth:`read_latest_depth`: peek the latest buffered frame."""
if self.thread is None or not self.thread.is_alive(): if self.thread is None or not self.thread.is_alive():
raise RuntimeError(f"{self} read thread is not running.") raise RuntimeError(f"{self} read thread is not running.")
if not self.new_frame_event.wait(timeout=timeout_ms / 1000.0):
raise TimeoutError(
f"Timed out waiting for frame from camera {self} after {timeout_ms} ms. "
f"Read thread alive: {self.thread.is_alive()}."
)
with self.frame_lock: with self.frame_lock:
frame = self.latest_color_frame frame = self.latest_depth_frame if read_depth else self.latest_color_frame
self.new_frame_event.clear() timestamp = self.latest_timestamp
if frame is None: if frame is None or timestamp is None:
raise RuntimeError(f"Internal error: Event set but no frame available for {self}.") raise RuntimeError(f"{self} has not captured any frames yet.")
age_ms = (time.perf_counter() - timestamp) * 1e3
if age_ms > max_age_ms:
raise TimeoutError(
f"{self} latest frame is too old: {age_ms:.1f} ms (max allowed: {max_age_ms} ms)."
)
return frame return frame
# NOTE(Steven): Missing implementation for depth for now
@check_if_not_connected @check_if_not_connected
def read_latest(self, max_age_ms: int = 500) -> NDArray[Any]: def read_latest(self, max_age_ms: int = 500) -> NDArray[Any]:
"""Return the most recent (color) frame captured immediately (Peeking). """Return the most recent (color) frame captured immediately (Peeking).
@@ -593,24 +626,48 @@ class RealSenseCamera(Camera):
DeviceNotConnectedError: If the camera is not connected. DeviceNotConnectedError: If the camera is not connected.
RuntimeError: If the camera is connected but has not captured any frames yet. RuntimeError: If the camera is connected but has not captured any frames yet.
""" """
if not self.use_rgb:
raise RuntimeError(f"{self}: cannot read color — camera was configured with use_rgb=False.")
if self.thread is None or not self.thread.is_alive(): return self._read_latest(max_age_ms=max_age_ms)
raise RuntimeError(f"{self} read thread is not running.")
with self.frame_lock: @check_if_not_connected
frame = self.latest_color_frame def async_read_depth(self, timeout_ms: float = 200) -> NDArray[np.uint16]:
timestamp = self.latest_timestamp """Read the latest depth frame asynchronously, in millimeters.
if frame is None or timestamp is None: Mirrors :meth:`async_read` but returns the depth stream rather than the
raise RuntimeError(f"{self} has not captured any frames yet.") color stream. Output is ``np.uint16`` of shape ``(H, W, 1)``, where each
pixel is the distance from the sensor in millimeters.
age_ms = (time.perf_counter() - timestamp) * 1e3 Raises:
if age_ms > max_age_ms: DeviceNotConnectedError: If the camera is not connected.
raise TimeoutError( RuntimeError: If ``use_depth`` is ``False`` for this camera, or if
f"{self} latest frame is too old: {age_ms:.1f} ms (max allowed: {max_age_ms} ms)." the background read thread is not running.
) TimeoutError: If no frame becomes available within ``timeout_ms``.
"""
if not self.use_depth:
raise RuntimeError(f"{self}: cannot read depth — camera was configured with use_depth=False.")
return frame return self._async_read(timeout_ms=timeout_ms, read_depth=True)
@check_if_not_connected
def read_latest_depth(self, max_age_ms: int = 500) -> NDArray[Any]:
"""Return the most recent depth frame in millimeters (peeking).
Non-blocking counterpart of :meth:`read_latest` for the depth stream.
Output is ``np.uint16`` of shape ``(H, W, 1)``, where each pixel is the
distance from the sensor in millimeters.
Raises:
DeviceNotConnectedError: If the camera is not connected.
RuntimeError: If ``use_depth`` is ``False`` for this camera, or if
no depth frame has been captured yet.
TimeoutError: If the latest depth frame is older than ``max_age_ms``.
"""
if not self.use_depth:
raise RuntimeError(f"{self}: cannot read depth — camera was configured with use_depth=False.")
return self._read_latest(max_age_ms=max_age_ms, read_depth=True)
def disconnect(self) -> None: def disconnect(self) -> None:
""" """
@@ -42,12 +42,14 @@ class RealSenseCameraConfig(CameraConfig):
height: Requested frame height in pixels for the color stream. height: Requested frame height in pixels for the color stream.
serial_number_or_name: Unique serial number or human-readable name to identify the camera. serial_number_or_name: Unique serial number or human-readable name to identify the camera.
color_mode: Color mode for image output (RGB or BGR). Defaults to RGB. color_mode: Color mode for image output (RGB or BGR). Defaults to RGB.
use_rgb: Whether to enable the color stream. Defaults to True.
use_depth: Whether to enable depth stream. Defaults to False. use_depth: Whether to enable depth stream. Defaults to False.
rotation: Image rotation setting (0°, 90°, 180°, or 270°). Defaults to no rotation. rotation: Image rotation setting (0°, 90°, 180°, or 270°). Defaults to no rotation.
warmup_s: Time reading frames before returning from connect (in seconds) warmup_s: Time reading frames before returning from connect (in seconds)
Note: Note:
- Either name or serial_number must be specified. - Either name or serial_number must be specified.
- At least one of `use_rgb` or `use_depth` must be enabled.
- Depth stream configuration (if enabled) will use the same FPS as the color stream. - Depth stream configuration (if enabled) will use the same FPS as the color stream.
- The actual resolution and FPS may be adjusted by the camera to the nearest supported mode. - The actual resolution and FPS may be adjusted by the camera to the nearest supported mode.
- For `fps`, `width` and `height`, either all of them need to be set, or none of them. - For `fps`, `width` and `height`, either all of them need to be set, or none of them.
@@ -55,6 +57,7 @@ class RealSenseCameraConfig(CameraConfig):
serial_number_or_name: str serial_number_or_name: str
color_mode: ColorMode = ColorMode.RGB color_mode: ColorMode = ColorMode.RGB
use_rgb: bool = True
use_depth: bool = False use_depth: bool = False
rotation: Cv2Rotation = Cv2Rotation.NO_ROTATION rotation: Cv2Rotation = Cv2Rotation.NO_ROTATION
warmup_s: int = 1 warmup_s: int = 1
@@ -63,6 +66,9 @@ class RealSenseCameraConfig(CameraConfig):
self.color_mode = ColorMode(self.color_mode) self.color_mode = ColorMode(self.color_mode)
self.rotation = Cv2Rotation(self.rotation) self.rotation = Cv2Rotation(self.rotation)
if not self.use_rgb and not self.use_depth:
raise ValueError("At least one of `use_rgb` or `use_depth` must be enabled.")
values = (self.fps, self.width, self.height) values = (self.fps, self.width, self.height)
if any(v is not None for v in values) and any(v is None for v in values): if any(v is not None for v in values) and any(v is None for v in values):
raise ValueError( raise ValueError(
+2
View File
@@ -293,6 +293,8 @@ class ZMQCamera(Camera):
if self.thread is not None and self.thread.is_alive(): if self.thread is not None and self.thread.is_alive():
self.thread.join(timeout=2.0) self.thread.join(timeout=2.0)
if self.thread.is_alive():
logger.warning(f"{self} read thread did not terminate within timeout.")
self.thread = None self.thread = None
self.stop_event = None self.stop_event = None
+13 -2
View File
@@ -33,10 +33,15 @@ from .types import (
RTCAttentionSchedule, RTCAttentionSchedule,
) )
from .video import ( from .video import (
DEFAULT_DEPTH_UNIT,
VALID_VIDEO_CODECS, VALID_VIDEO_CODECS,
VIDEO_ENCODER_INFO_KEYS, VIDEO_ENCODER_INFO_KEYS,
DepthEncoderConfig,
RGBEncoderConfig,
VideoEncoderConfig, VideoEncoderConfig,
camera_encoder_defaults, depth_encoder_defaults,
encoder_config_from_video_info,
rgb_encoder_defaults,
) )
__all__ = [ __all__ = [
@@ -57,9 +62,15 @@ __all__ = [
"WandBConfig", "WandBConfig",
"load_recipe", "load_recipe",
"VideoEncoderConfig", "VideoEncoderConfig",
"RGBEncoderConfig",
"DepthEncoderConfig",
# Defaults # Defaults
"camera_encoder_defaults", "rgb_encoder_defaults",
"depth_encoder_defaults",
# Factories
"encoder_config_from_video_info",
# Constants # Constants
"DEFAULT_DEPTH_UNIT",
"VALID_VIDEO_CODECS", "VALID_VIDEO_CODECS",
"VIDEO_ENCODER_INFO_KEYS", "VIDEO_ENCODER_INFO_KEYS",
] ]
+5 -3
View File
@@ -18,7 +18,7 @@ from dataclasses import dataclass, field
from datetime import datetime from datetime import datetime
from pathlib import Path from pathlib import Path
from .video import VideoEncoderConfig, camera_encoder_defaults from .video import DepthEncoderConfig, RGBEncoderConfig, depth_encoder_defaults, rgb_encoder_defaults
@dataclass @dataclass
@@ -58,8 +58,10 @@ class DatasetRecordConfig:
# Set to 1 for immediate encoding (default behavior), or higher for batched encoding # Set to 1 for immediate encoding (default behavior), or higher for batched encoding
video_encoding_batch_size: int = 1 video_encoding_batch_size: int = 1
# Video encoder settings for camera MP4s (codec, quality, GOP, etc.). Tuned via CLI nested keys, # Video encoder settings for camera MP4s (codec, quality, GOP, etc.). Tuned via CLI nested keys,
# e.g. ``--dataset.camera_encoder.vcodec=h264`` (see ``VideoEncoderConfig``). # e.g. ``--dataset.rgb_encoder.vcodec=h264`` (see ``RGBEncoderConfig``).
camera_encoder: VideoEncoderConfig = field(default_factory=camera_encoder_defaults) rgb_encoder: RGBEncoderConfig = field(default_factory=rgb_encoder_defaults)
# Video encoder settings for depth-map MP4s (codec, quality, GOP, etc.). Tuned via CLI nested keys.
depth_encoder: DepthEncoderConfig = field(default_factory=depth_encoder_defaults)
# Enable streaming video encoding: encode frames in real-time during capture instead # Enable streaming video encoding: encode frames in real-time during capture instead
# of writing PNG images first. Makes save_episode() near-instant. More info in the documentation: https://huggingface.co/docs/lerobot/streaming_video_encoding # of writing PNG images first. Makes save_episode() near-instant. More info in the documentation: https://huggingface.co/docs/lerobot/streaming_video_encoding
streaming_encoding: bool = False streaming_encoding: bool = False
+10 -1
View File
@@ -19,6 +19,8 @@ from dataclasses import dataclass, field
from lerobot.transforms import ImageTransformsConfig from lerobot.transforms import ImageTransformsConfig
from lerobot.utils.import_utils import get_safe_default_video_backend from lerobot.utils.import_utils import get_safe_default_video_backend
from .video import DEFAULT_DEPTH_UNIT, DEPTH_METER_UNIT, DEPTH_MILLIMETER_UNIT
@dataclass @dataclass
class DatasetConfig: class DatasetConfig:
@@ -35,14 +37,21 @@ class DatasetConfig:
revision: str | None = None revision: str | None = None
use_imagenet_stats: bool = True use_imagenet_stats: bool = True
video_backend: str = field(default_factory=get_safe_default_video_backend) video_backend: str = field(default_factory=get_safe_default_video_backend)
# When True, video frames are returned as uint8 tensors (0-255) instead of float32 (0.0-1.0). # When True, RGB video frames are returned as uint8 tensors (0-255) instead of float32 (0.0-1.0).
# This reduces memory and speeds up DataLoader IPC. The training pipeline handles the conversion. # This reduces memory and speeds up DataLoader IPC. The training pipeline handles the conversion.
return_uint8: bool = False return_uint8: bool = False
# Physical unit depth maps are dequantized to at load time: "mm" (millimeters) or "m" (metres).
# Has no effect on datasets without depth cameras.
depth_output_unit: str = DEFAULT_DEPTH_UNIT
streaming: bool = False streaming: bool = False
# Fraction of episodes held out per task for offline evaluation (0.0 = disabled). # Fraction of episodes held out per task for offline evaluation (0.0 = disabled).
eval_split: float = 0.0 eval_split: float = 0.0
def __post_init__(self) -> None: def __post_init__(self) -> None:
if self.depth_output_unit not in (DEPTH_METER_UNIT, DEPTH_MILLIMETER_UNIT):
raise ValueError(
f"depth_output_unit must be '{DEPTH_METER_UNIT}' or '{DEPTH_MILLIMETER_UNIT}', got {self.depth_output_unit!r}"
)
if not (0.0 <= self.eval_split < 1.0): if not (0.0 <= self.eval_split < 1.0):
raise ValueError(f"eval_split must be in [0.0, 1.0), got {self.eval_split}") raise ValueError(f"eval_split must be in [0.0, 1.0), got {self.eval_split}")
if self.episodes is not None: if self.episodes is not None:
+123 -36
View File
@@ -20,7 +20,7 @@ from __future__ import annotations
import logging import logging
from dataclasses import dataclass, field from dataclasses import dataclass, field
from typing import Any from typing import Any, ClassVar, Self
from lerobot.utils.import_utils import require_package from lerobot.utils.import_utils import require_package
@@ -40,7 +40,6 @@ VALID_VIDEO_CODECS: frozenset[str] = frozenset({"h264", "hevc", "libsvtav1", "au
# Aliases for legacy video codec names. # Aliases for legacy video codec names.
VIDEO_CODECS_ALIASES: dict[str, str] = {"av1": "libsvtav1"} VIDEO_CODECS_ALIASES: dict[str, str] = {"av1": "libsvtav1"}
LIBSVTAV1_DEFAULT_PRESET: int = 12 LIBSVTAV1_DEFAULT_PRESET: int = 12
# Keys persisted under ``features[*]["info"]`` as ``video.<name>`` (from :class:`VideoEncoderConfig`). # Keys persisted under ``features[*]["info"]`` as ``video.<name>`` (from :class:`VideoEncoderConfig`).
@@ -52,40 +51,45 @@ VIDEO_ENCODER_INFO_KEYS: frozenset[str] = frozenset(
f"video.{name}" for name in VIDEO_ENCODER_INFO_FIELD_NAMES f"video.{name}" for name in VIDEO_ENCODER_INFO_FIELD_NAMES
) )
# Default depth quantization and encoding parameters.
DEPTH_QUANT_BITS: int = 12
DEPTH_QMAX: int = (1 << DEPTH_QUANT_BITS) - 1 # 4095
DEFAULT_DEPTH_MIN: float = 0.01
DEFAULT_DEPTH_MAX: float = 10.0
DEFAULT_DEPTH_SHIFT: float = 3.5
DEFAULT_DEPTH_USE_LOG: bool = True
DEFAULT_DEPTH_PIX_FMT: str = "gray12le"
DEPTH_METER_UNIT: str = "m"
DEPTH_MILLIMETER_UNIT: str = "mm"
DEFAULT_DEPTH_UNIT: str = DEPTH_MILLIMETER_UNIT
# Depth-specific tuning fields persisted under ``features[*]["info"]`` as ``video.<name>``.
DEPTH_ENCODER_INFO_FIELD_NAMES: frozenset[str] = frozenset({"depth_min", "depth_max", "shift", "use_log"})
@dataclass @dataclass
class VideoEncoderConfig: class VideoEncoderConfig:
"""Video encoder configuration. """Video encoder configuration."""
Attributes: vcodec: str = "libsvtav1" # Video codec name. "auto" picks a hardware codec if available, else libsvtav1.
vcodec: Video encoder name. ``"auto"`` is resolved during pix_fmt: str = "yuv420p" # Pixel format (e.g. yuv420p).
construction (HW encoder if available, else ``libsvtav1``). g: int | None = 2 # GOP size (keyframe interval).
pix_fmt: Pixel format (e.g. ``"yuv420p"``). crf: int | float | None = 30 # Quality level. Lower means better quality and larger files.
g: GOP size (keyframe interval). preset: int | str | None = None # Speed/quality preset. Accepted values are codec-specific.
crf: Quality level mapped to the native quality parameter of the fast_decode: int = 0 # Fast-decode tuning. Accepted values are codec-specific, 0 disables it.
codec (``crf`` for software, ``qp`` for NVENC/VAAPI,
``q:v`` for VideoToolbox, ``global_quality`` for QSV).
preset: Speed/quality preset. Accepted type is per-codec.
fast_decode: Fast-decode tuning. For ``libsvtav1`` this is a level (0-2)
embedded in ``svtav1-params``. For ``h264`` and ``hevc`` non-zero values
set ``tune=fastdecode``. Ignored for other codecs.
video_backend: Python to be used for encoding. Only ``"pyav"``
is currently supported.
extra_options: Free-form dictionary of additional video encoder options
(e.g. ``{"tune": "film", "profile:v": "high", "bf": 2}``).
"""
vcodec: str = "libsvtav1" # TODO(CarolinePascal): rename to codec ?
pix_fmt: str = "yuv420p"
g: int | None = 2
crf: int | float | None = 30
preset: int | str | None = None
fast_decode: int = 0
# TODO(CarolinePascal): add torchcodec support + find a way to unify the # TODO(CarolinePascal): add torchcodec support + find a way to unify the
# two backends (encoding and decoding). # two backends (encoding and decoding).
video_backend: str = "pyav" video_backend: str = "pyav" # Encoding backend. Only "pyav" is currently supported.
# Extra codec options merged last, e.g. {"tune": "film"}.
extra_options: dict[str, Any] = field(default_factory=dict) extra_options: dict[str, Any] = field(default_factory=dict)
# Source-data channel count this encoder is expected to handle. ``None``
# disables the pix_fmt channel-count check; concrete subclasses set it
# (3 for RGB, 1 for depth, etc.).
_DEFAULT_CHANNELS: ClassVar[int | None] = None
def __post_init__(self) -> None: def __post_init__(self) -> None:
self.resolve_vcodec() self.resolve_vcodec()
# Empty-constructor ergonomics: ``VideoEncoderConfig()`` must "just work". # Empty-constructor ergonomics: ``VideoEncoderConfig()`` must "just work".
@@ -94,9 +98,9 @@ class VideoEncoderConfig:
self.validate() self.validate()
@classmethod @classmethod
def from_video_info(cls, video_info: dict | None) -> VideoEncoderConfig: def _kwargs_from_video_info(cls, video_info: dict | None) -> dict[str, Any]:
"""Reconstruct a :class:`VideoEncoderConfig` from a video feature's ``info`` block. """Parse the ``video.*`` keys of a feature ``info`` block into
Missing or ``None`` values fall back to the class defaults. constructor kwargs.
""" """
video_info = video_info or {} video_info = video_info or {}
kwargs: dict[str, Any] = {} kwargs: dict[str, Any] = {}
@@ -115,7 +119,15 @@ class VideoEncoderConfig:
continue continue
kwargs[field_name] = value kwargs[field_name] = value
return cls(**kwargs) return kwargs
@classmethod
def from_video_info(cls, video_info: dict | None) -> Self:
"""Reconstruct an encoder config from a video feature's ``info`` block.
Missing or ``None`` values fall back to the class defaults.
"""
return cls(**cls._kwargs_from_video_info(video_info))
def detect_available_encoders(self, encoders: list[str] | str) -> list[str]: def detect_available_encoders(self, encoders: list[str] | str) -> list[str]:
"""Return the subset of available encoders based on the specified video backend. """Return the subset of available encoders based on the specified video backend.
@@ -138,7 +150,9 @@ class VideoEncoderConfig:
require_package("av", extra="dataset") require_package("av", extra="dataset")
from lerobot.datasets import check_video_encoder_parameters_pyav from lerobot.datasets import check_video_encoder_parameters_pyav
check_video_encoder_parameters_pyav(self.vcodec, self.pix_fmt, self.get_codec_options()) check_video_encoder_parameters_pyav(
self.vcodec, self.pix_fmt, self.get_codec_options(), channels=self._DEFAULT_CHANNELS
)
def resolve_vcodec(self) -> None: def resolve_vcodec(self) -> None:
"""Check ``vcodec`` and, when it is ``"auto"``, pick a concrete encoder. """Check ``vcodec`` and, when it is ``"auto"``, pick a concrete encoder.
@@ -230,6 +244,79 @@ class VideoEncoderConfig:
return opts return opts
def camera_encoder_defaults() -> VideoEncoderConfig: @dataclass
"""Return a :class:`VideoEncoderConfig` with RGB-camera defaults.""" class RGBEncoderConfig(VideoEncoderConfig):
return VideoEncoderConfig() """Encoder configuration for RGB camera streams.
Identical to :class:`VideoEncoderConfig` but declares the 3-channel
source-data layout so ``pix_fmt`` is validated against RGB inputs.
"""
_DEFAULT_CHANNELS: ClassVar[int] = 3
def rgb_encoder_defaults() -> RGBEncoderConfig:
"""Return a :class:`RGBEncoderConfig` with RGB-camera defaults."""
return RGBEncoderConfig()
@dataclass
class DepthEncoderConfig(VideoEncoderConfig):
"""Encoder configuration for depth-map streams.
Inherits the full :class:`VideoEncoderConfig` surface (codec, GOP, CRF,
preset, ``extra_options``) and adds the parameters of the depth quantizer.
Defaults flip ``vcodec`` to ``"hevc"`` (Main 12 profile) and ``pix_fmt`` to
``"gray12le"``.
"""
vcodec: str = "hevc" # Video codec name. Defaults to HEVC Main 12 (a 12-bit-capable codec).
pix_fmt: str = "gray12le" # Pixel format. Defaults to 12-bit grayscale.
extra_options: dict[str, Any] = field(default_factory=lambda: {"x265-params": "lossless=1"})
depth_min: float = DEFAULT_DEPTH_MIN # Minimum depth in meters, mapped to the lowest quantum.
depth_max: float = DEFAULT_DEPTH_MAX # Maximum depth in meters, mapped to the highest quantum.
shift: float = DEFAULT_DEPTH_SHIFT # Pre-log offset in meters for numerical stability near zero.
use_log: bool = DEFAULT_DEPTH_USE_LOG # Use logarithmic quantization (True) or linear (False).
_DEFAULT_CHANNELS: ClassVar[int] = 1
@classmethod
def _kwargs_from_video_info(cls, video_info: dict | None) -> dict[str, Any]:
"""Layer the depth-specific tuning (``depth_min`` / ``depth_max`` /
``shift`` / ``use_log``) on top of the base parser. Missing keys
fall back to the class defaults.
"""
kwargs = super()._kwargs_from_video_info(video_info)
video_info = video_info or {}
for name in DEPTH_ENCODER_INFO_FIELD_NAMES:
value = video_info.get(f"video.{name}")
if value is not None:
kwargs[name] = value
return kwargs
def depth_encoder_defaults() -> DepthEncoderConfig:
"""Return a :class:`DepthEncoderConfig` with depth-camera defaults."""
return DepthEncoderConfig()
def encoder_config_from_video_info(video_info: dict | None) -> VideoEncoderConfig:
"""Build the appropriate encoder config from a feature's ``info`` block.
Dispatches to :class:`DepthEncoderConfig` when the dict marks the feature
as a depth map and to :class:`RGBEncoderConfig`
otherwise.
Args:
video_info: A feature's ``info`` dict as persisted in ``info.json``,
or ``None`` (treated as an empty dict).
Returns:
A :class:`DepthEncoderConfig` for depth features, otherwise a
:class:`RGBEncoderConfig`.
"""
video_info = video_info or {}
is_depth = bool(video_info.get("is_depth_map") or video_info.get("video.is_depth_map"))
cls: type[VideoEncoderConfig] = DepthEncoderConfig if is_depth else RGBEncoderConfig
return cls.from_video_info(video_info)
+15 -7
View File
@@ -242,12 +242,12 @@ def sample_images(image_paths: list[str]) -> np.ndarray:
images = None images = None
for i, idx in enumerate(sampled_indices): for i, idx in enumerate(sampled_indices):
path = image_paths[idx] path = image_paths[idx]
# we load as uint8 to reduce memory usage # we load RGB images as uint8 to reduce memory usage; depth keeps its native dtype
img = load_image_as_numpy(path, dtype=np.uint8, channel_first=True) img = load_image_as_numpy(path, dtype=np.uint8, channel_first=True)
img = auto_downsample_height_width(img) img = auto_downsample_height_width(img)
if images is None: if images is None:
images = np.empty((len(sampled_indices), *img.shape), dtype=np.uint8) images = np.empty((len(sampled_indices), *img.shape), dtype=img.dtype)
images[i] = img images[i] = img
@@ -506,8 +506,10 @@ def compute_episode_stats(
Each statistics dictionary contains min, max, mean, std, count, and quantiles. Each statistics dictionary contains min, max, mean, std, count, and quantiles.
Note: Note:
Image statistics are normalized to [0,1] range and have shape (3,1,1) for For 'image'/'video' features, stats are computed per channel and kept with a
per-channel values when dtype is 'image' or 'video'. leading channel axis (e.g. shape (3, 1, 1) for RGB). RGB stats are divided by
255 to land in [0, 1]; depth maps (features flagged with ``is_depth_map``) skip
this rescaling and remain in their stored units.
""" """
if quantile_list is None: if quantile_list is None:
quantile_list = DEFAULT_QUANTILES quantile_list = DEFAULT_QUANTILES
@@ -531,8 +533,12 @@ def compute_episode_stats(
) )
if features[key]["dtype"] in ["image", "video"]: if features[key]["dtype"] in ["image", "video"]:
normalization_factor = (
255.0 if not (features[key].get("info") or {}).get("is_depth_map", False) else 1.0
)
ep_stats[key] = { ep_stats[key] = {
k: v if k == "count" else np.squeeze(v / 255.0, axis=0) for k, v in ep_stats[key].items() k: v if k == "count" else np.squeeze(v / normalization_factor, axis=0)
for k, v in ep_stats[key].items()
} }
return ep_stats return ep_stats
@@ -552,8 +558,10 @@ def _validate_stat_value(value: np.ndarray, key: str, feature_key: str) -> None:
if key == "count" and value.shape != (1,): if key == "count" and value.shape != (1,):
raise ValueError(f"Shape of 'count' must be (1), but is {value.shape} instead.") raise ValueError(f"Shape of 'count' must be (1), but is {value.shape} instead.")
if "image" in feature_key and key != "count" and value.shape != (3, 1, 1): if "image" in feature_key and key != "count" and value.shape not in ((3, 1, 1), (1, 1, 1)):
raise ValueError(f"Shape of quantile '{key}' must be (3,1,1), but is {value.shape} instead.") raise ValueError(
f"Shape of quantile '{key}' must be (3,1,1) or (1,1,1) but is {value.shape} instead."
)
def _assert_type_and_shape(stats_list: list[dict[str, dict]]): def _assert_type_and_shape(stats_list: list[dict[str, dict]]):
+46 -7
View File
@@ -14,7 +14,8 @@
# See the License for the specific language governing permissions and # See the License for the specific language governing permissions and
# limitations under the License. # limitations under the License.
import contextlib import contextlib
from collections.abc import Callable import logging
from collections.abc import Callable, Iterable
from copy import deepcopy from copy import deepcopy
from pathlib import Path from pathlib import Path
@@ -338,6 +339,25 @@ class LeRobotDatasetMetadata:
"""Keys to access visual modalities stored as videos.""" """Keys to access visual modalities stored as videos."""
return [key for key, ft in self.features.items() if ft["dtype"] == "video"] return [key for key, ft in self.features.items() if ft["dtype"] == "video"]
@property
def depth_keys(self) -> list[str]:
"""Keys to access depth-map modalities stored as videos or images.
A depth key is a feature whose ``info`` dict carries ``"is_depth_map": True``
(or the legacy ``"video.is_depth_map"`` inside ``info`` or ``video_info``).
"""
def _is_depth(ft: dict) -> bool:
info = ft.get("info") or {}
video_info = ft.get("video_info") or {}
return (
info.get("is_depth_map", False)
or info.get("video.is_depth_map", False)
or video_info.get("video.is_depth_map", False)
)
return [key for key, ft in self.features.items() if _is_depth(ft)]
@property @property
def camera_keys(self) -> list[str]: def camera_keys(self) -> list[str]:
"""Keys to access visual modalities (regardless of their storage method).""" """Keys to access visual modalities (regardless of their storage method)."""
@@ -581,29 +601,48 @@ class LeRobotDatasetMetadata:
def update_video_info( def update_video_info(
self, self,
video_key: str | None = None, video_key: str | None = None,
camera_encoder: VideoEncoderConfig | None = None, video_encoder: VideoEncoderConfig | None = None,
preserve_keys: Iterable[str] | None = None,
) -> None: ) -> None:
"""Populate per-feature video info in ``info.json``. """Populate or refresh per-feature video info in ``info.json``.
Warning: this function writes info from first episode videos, implicitly assuming that all videos have Warning: this function writes info from first episode videos, implicitly assuming that all videos have
been encoded the same way. Also, this means it assumes the first episode exists. been encoded the same way. Also, this means it assumes the first episode exists.
Always re-probes the videos and overwrites existing info for every recomputed
key. ``preserve_keys`` lists keys whose existing values must be kept (e.g.
data-intrinsic entries like ``is_depth_map`` and depth quantization params)
instead of being recomputed.
Args: Args:
video_key: If provided, only update this video key. Otherwise update video_key: If provided, only update this video key. Otherwise update
all video keys in the dataset. all video keys in the dataset.
camera_encoder: Encoder configuration used to produce the video_encoder: Encoder configuration used to produce the
videos. When provided, its fields are recorded as videos. When provided, its fields are recorded as
``video.<field>`` entries alongside the stream-derived ``video.<field>`` entries alongside the stream-derived
``video.*`` entries (see :func:`get_video_info`). ``video.*`` entries (see :func:`get_video_info`).
preserve_keys: Keys whose existing values are kept instead of being
recomputed. ``None`` (default) recomputes every key.
""" """
if video_key is not None and video_key not in self.video_keys: if video_key is not None and video_key not in self.video_keys:
raise ValueError(f"Video key {video_key} not found in dataset") raise ValueError(f"Video key {video_key} not found in dataset")
video_keys = [video_key] if video_key is not None else self.video_keys video_keys = [video_key] if video_key is not None else self.video_keys
preserve_set = set(preserve_keys or ())
for key in video_keys: for key in video_keys:
if not self.features[key].get("info", None): existing = self.features[key].get("info") or {}
video_path = self.root / self.video_path.format(video_key=key, chunk_index=0, file_index=0) video_path = self.root / self.video_path.format(video_key=key, chunk_index=0, file_index=0)
self.info.features[key]["info"] = get_video_info(video_path, camera_encoder=camera_encoder) new_info = get_video_info(video_path, video_encoder=video_encoder)
# Drop preserved keys so the existing values win on merge.
new_info = {k: v for k, v in new_info.items() if k not in preserve_set}
merged = {**existing, **new_info}
# Migrate the legacy depth marker to the canonical key.
if "video.is_depth_map" in merged:
logging.warning(
f"Migrating legacy 'video.is_depth_map' to 'is_depth_map' for feature {key!r}."
)
merged.setdefault("is_depth_map", merged.pop("video.is_depth_map"))
self.info.features[key]["info"] = merged
def update_chunk_settings( def update_chunk_settings(
self, self,
+28 -2
View File
@@ -22,7 +22,10 @@ from pathlib import Path
import datasets import datasets
import torch import torch
from lerobot.configs import DEFAULT_DEPTH_UNIT, DepthEncoderConfig
from .dataset_metadata import LeRobotDatasetMetadata from .dataset_metadata import LeRobotDatasetMetadata
from .depth_utils import dequantize_depth
from .feature_utils import ( from .feature_utils import (
check_delta_timestamps, check_delta_timestamps,
get_delta_indices, get_delta_indices,
@@ -51,6 +54,7 @@ class DatasetReader:
delta_timestamps: dict[str, list[float]] | None, delta_timestamps: dict[str, list[float]] | None,
image_transforms: Callable | None, image_transforms: Callable | None,
return_uint8: bool = False, return_uint8: bool = False,
depth_output_unit: str = DEFAULT_DEPTH_UNIT,
): ):
"""Initialize the reader with metadata, filtering, and transform config. """Initialize the reader with metadata, filtering, and transform config.
@@ -68,6 +72,10 @@ class DatasetReader:
relative timestamp offsets for temporal context windows. relative timestamp offsets for temporal context windows.
image_transforms: Optional torchvision v2 transform applied to image_transforms: Optional torchvision v2 transform applied to
visual features. visual features.
return_uint8: If True, return RGB video frames as raw uint8 tensors
instead of normalized float32.
depth_output_unit: Physical unit depth maps are dequantized to
(``"m"`` or ``"mm"``). Defaults to ``"mm"``.
""" """
self._meta = meta self._meta = meta
self.root = root self.root = root
@@ -78,6 +86,7 @@ class DatasetReader:
raise TypeError("image_transforms must be callable or None.") raise TypeError("image_transforms must be callable or None.")
self._image_transforms = image_transforms self._image_transforms = image_transforms
self._return_uint8 = return_uint8 self._return_uint8 = return_uint8
self._depth_output_unit = depth_output_unit
self.hf_dataset: datasets.Dataset | None = None self.hf_dataset: datasets.Dataset | None = None
self._absolute_to_relative_idx: dict[int, int] | None = None self._absolute_to_relative_idx: dict[int, int] | None = None
@@ -88,6 +97,11 @@ class DatasetReader:
check_delta_timestamps(delta_timestamps, meta.fps, tolerance_s) check_delta_timestamps(delta_timestamps, meta.fps, tolerance_s)
self.delta_indices = get_delta_indices(delta_timestamps, meta.fps) self.delta_indices = get_delta_indices(delta_timestamps, meta.fps)
self._depth_encoder_configs: dict[str, DepthEncoderConfig] = {
vid_key: DepthEncoderConfig.from_video_info(self._meta.features[vid_key].get("info"))
for vid_key in self._meta.depth_keys
}
def set_image_transforms(self, image_transforms: Callable | None) -> None: def set_image_transforms(self, image_transforms: Callable | None) -> None:
"""Replace the transform applied to visual observations.""" """Replace the transform applied to visual observations."""
if image_transforms is not None and not callable(image_transforms): if image_transforms is not None and not callable(image_transforms):
@@ -259,7 +273,18 @@ class DatasetReader:
self._tolerance_s, self._tolerance_s,
self._video_backend, self._video_backend,
return_uint8=self._return_uint8, return_uint8=self._return_uint8,
is_depth=vid_key in self._meta.depth_keys,
) )
if vid_key in self._meta.depth_keys:
depth_encoder = self._depth_encoder_configs[vid_key]
frames = dequantize_depth(
frames,
depth_min=depth_encoder.depth_min,
depth_max=depth_encoder.depth_max,
shift=depth_encoder.shift,
use_log=depth_encoder.use_log,
output_unit=self._depth_output_unit,
)
return vid_key, frames.squeeze(0) return vid_key, frames.squeeze(0)
items = list(query_timestamps.items()) items = list(query_timestamps.items())
@@ -299,8 +324,9 @@ class DatasetReader:
item = {**video_frames, **item} item = {**video_frames, **item}
if self._image_transforms is not None: if self._image_transforms is not None:
image_keys = self._meta.camera_keys for cam in self._meta.camera_keys:
for cam in image_keys: if cam in self._meta.depth_keys:
continue
item[cam] = self._image_transforms(item[cam]) item[cam] = self._image_transforms(item[cam])
# Add task as a string # Add task as a string
+113 -72
View File
@@ -37,7 +37,15 @@ import pyarrow.parquet as pq
import torch import torch
from tqdm import tqdm from tqdm import tqdm
from lerobot.configs import VideoEncoderConfig, camera_encoder_defaults from lerobot.configs import (
DepthEncoderConfig,
RGBEncoderConfig,
VideoEncoderConfig,
depth_encoder_defaults,
encoder_config_from_video_info,
rgb_encoder_defaults,
)
from lerobot.configs.video import DEPTH_ENCODER_INFO_FIELD_NAMES
from lerobot.utils.constants import ACTION, HF_LEROBOT_HOME, OBS_IMAGE, OBS_STATE from lerobot.utils.constants import ACTION, HF_LEROBOT_HOME, OBS_IMAGE, OBS_STATE
from lerobot.utils.utils import flatten_dict from lerobot.utils.utils import flatten_dict
@@ -48,6 +56,7 @@ from .compute_stats import (
compute_relative_action_stats, compute_relative_action_stats,
) )
from .dataset_metadata import LeRobotDatasetMetadata from .dataset_metadata import LeRobotDatasetMetadata
from .image_writer import write_image
from .io_utils import ( from .io_utils import (
get_parquet_file_size_in_mb, get_parquet_file_size_in_mb,
load_episodes, load_episodes,
@@ -62,12 +71,13 @@ from .utils import (
DEFAULT_DATA_FILE_SIZE_IN_MB, DEFAULT_DATA_FILE_SIZE_IN_MB,
DEFAULT_DATA_PATH, DEFAULT_DATA_PATH,
DEFAULT_EPISODES_PATH, DEFAULT_EPISODES_PATH,
DEPTH_FILE_PATTERN,
IMAGE_FILE_PATTERN,
VIDEO_DIR, VIDEO_DIR,
update_chunk_file_indices, update_chunk_file_indices,
) )
from .video_utils import ( from .video_utils import (
encode_video_frames, encode_video_frames,
get_video_info,
reencode_video, reencode_video,
) )
@@ -601,7 +611,7 @@ def _keep_episodes_from_video_with_av(
output_path: Path, output_path: Path,
episodes_to_keep: list[tuple[int, int]], episodes_to_keep: list[tuple[int, int]],
fps: float, fps: float,
camera_encoder: VideoEncoderConfig, video_encoder: VideoEncoderConfig,
) -> None: ) -> None:
"""Keep only specified episodes from a video file using PyAV. """Keep only specified episodes from a video file using PyAV.
@@ -615,7 +625,7 @@ def _keep_episodes_from_video_with_av(
Ranges are half-open intervals: [start_frame, end_frame), where start_frame Ranges are half-open intervals: [start_frame, end_frame), where start_frame
is inclusive and end_frame is exclusive. is inclusive and end_frame is exclusive.
fps: Frame rate of the video. fps: Frame rate of the video.
camera_encoder: Video encoder settings used to re-encode the kept frames. video_encoder: Video encoder settings used to re-encode the kept frames.
""" """
from fractions import Fraction from fractions import Fraction
@@ -640,13 +650,13 @@ def _keep_episodes_from_video_with_av(
# Convert fps to Fraction for PyAV compatibility. # Convert fps to Fraction for PyAV compatibility.
fps_fraction = Fraction(fps).limit_denominator(1000) fps_fraction = Fraction(fps).limit_denominator(1000)
codec_options = camera_encoder.get_codec_options(as_strings=True) codec_options = video_encoder.get_codec_options(as_strings=True)
v_out = out.add_stream(camera_encoder.vcodec, rate=fps_fraction, options=codec_options) v_out = out.add_stream(video_encoder.vcodec, rate=fps_fraction, options=codec_options)
# PyAV type stubs don't distinguish video streams from audio/subtitle streams. # PyAV type stubs don't distinguish video streams from audio/subtitle streams.
v_out.width = v_in.codec_context.width v_out.width = v_in.codec_context.width
v_out.height = v_in.codec_context.height v_out.height = v_in.codec_context.height
v_out.pix_fmt = camera_encoder.pix_fmt v_out.pix_fmt = video_encoder.pix_fmt
# Set time_base to match the frame rate for proper timestamp handling. # Set time_base to match the frame rate for proper timestamp handling.
v_out.time_base = Fraction(1, int(fps)) v_out.time_base = Fraction(1, int(fps))
@@ -733,7 +743,7 @@ def _copy_and_reindex_videos(
for video_key in src_dataset.meta.video_keys: for video_key in src_dataset.meta.video_keys:
logging.info(f"Processing videos for {video_key}") logging.info(f"Processing videos for {video_key}")
camera_encoder = VideoEncoderConfig.from_video_info( video_encoder = encoder_config_from_video_info(
src_dataset.meta.info.features.get(video_key, {}).get("info") src_dataset.meta.info.features.get(video_key, {}).get("info")
) )
@@ -817,7 +827,7 @@ def _copy_and_reindex_videos(
dst_video_path, dst_video_path,
episodes_to_keep_ranges, episodes_to_keep_ranges,
src_dataset.meta.fps, src_dataset.meta.fps,
camera_encoder, video_encoder,
) )
cumulative_ts = 0.0 cumulative_ts = 0.0
@@ -874,11 +884,11 @@ def _copy_and_reindex_episodes_metadata(
episode_meta.update(video_metadata[new_idx]) episode_meta.update(video_metadata[new_idx])
# Extract episode statistics from parquet metadata. # Extract episode statistics from parquet metadata.
# Note (maractingi): When pandas/pyarrow serializes numpy arrays with shape (3, 1, 1) to parquet, # When pandas/pyarrow serializes numpy arrays with shape (C, 1, 1) to parquet,
# they are being deserialized as nested object arrays like: # they are being deserialized as nested object arrays like:
# array([array([array([0.])]), array([array([0.])]), array([array([0.])])]) # array([array([array([0.])]), array([array([0.])]), array([array([0.])])])
# This happens particularly with image/video statistics. We need to detect and flatten # This happens particularly with image/video statistics. We need to detect and flatten
# these nested structures back to proper (3, 1, 1) arrays so aggregate_stats can process them. # these nested structures back to proper (C, 1, 1) arrays so aggregate_stats can process them.
episode_stats = {} episode_stats = {}
for key in src_episode_full: for key in src_episode_full:
if key.startswith("stats/"): if key.startswith("stats/"):
@@ -894,15 +904,16 @@ def _copy_and_reindex_episodes_metadata(
if feature_name in src_dataset.meta.features: if feature_name in src_dataset.meta.features:
feature_dtype = src_dataset.meta.features[feature_name]["dtype"] feature_dtype = src_dataset.meta.features[feature_name]["dtype"]
if feature_dtype in ["image", "video"] and stat_name != "count": if feature_dtype in ["image", "video"] and stat_name != "count":
# Stats are channel-first (C, 1, 1)
if isinstance(value, np.ndarray) and value.dtype == object: if isinstance(value, np.ndarray) and value.dtype == object:
flat_values = [] flat_values = []
for item in value: for item in value:
while isinstance(item, np.ndarray): while isinstance(item, np.ndarray):
item = item.flatten()[0] item = item.flatten()[0]
flat_values.append(item) flat_values.append(item)
value = np.array(flat_values, dtype=np.float64).reshape(3, 1, 1) value = np.array(flat_values, dtype=np.float64).reshape(-1, 1, 1)
elif isinstance(value, np.ndarray) and value.shape == (3,): elif isinstance(value, np.ndarray) and value.ndim == 1:
value = value.reshape(3, 1, 1) value = value.reshape(-1, 1, 1)
episode_stats[feature_name][stat_name] = value episode_stats[feature_name][stat_name] = value
@@ -1153,15 +1164,15 @@ def _save_episode_images_for_video(
# Get all items for this episode # Get all items for this episode
episode_dataset = imgs_dataset.select(range(from_idx, to_idx)) episode_dataset = imgs_dataset.select(range(from_idx, to_idx))
is_depth = img_key in dataset.meta.depth_keys
frame_pattern = DEPTH_FILE_PATTERN if is_depth else IMAGE_FILE_PATTERN
# Define function to save a single image # Define function to save a single image
def save_single_image(i_item_tuple): def save_single_image(i_item_tuple):
i, item = i_item_tuple i, item = i_item_tuple
img = item[img_key] write_image(item[img_key], imgs_dir / frame_pattern.format(frame_index=i))
# Use frame-XXXXXX.png format to match encode_video_frames expectations
img.save(str(imgs_dir / f"frame-{i:06d}.png"), quality=100)
return i return i
# Save images with proper naming convention for encode_video_frames (frame-XXXXXX.png)
items = list(enumerate(episode_dataset)) items = list(enumerate(episode_dataset))
with ThreadPoolExecutor(max_workers=num_workers) as executor: with ThreadPoolExecutor(max_workers=num_workers) as executor:
@@ -1193,13 +1204,14 @@ def _save_batch_episodes_images(
hf_dataset = dataset.hf_dataset.with_format(None) hf_dataset = dataset.hf_dataset.with_format(None)
imgs_dataset = hf_dataset.select_columns(img_key) imgs_dataset = hf_dataset.select_columns(img_key)
is_depth = img_key in dataset.meta.depth_keys
frame_pattern = DEPTH_FILE_PATTERN if is_depth else IMAGE_FILE_PATTERN
# Define function to save a single image with global frame index # Define function to save a single image with global frame index
# Defined once outside the loop to avoid repeated closure creation # Defined once outside the loop to avoid repeated closure creation
def save_single_image(i_item_tuple, base_frame_idx, img_key_param): def save_single_image(i_item_tuple, base_frame_idx, img_key_param):
i, item = i_item_tuple i, item = i_item_tuple
img = item[img_key_param] write_image(item[img_key_param], imgs_dir / frame_pattern.format(frame_index=base_frame_idx + i))
# Use global frame index for naming
img.save(str(imgs_dir / f"frame-{base_frame_idx + i:06d}.png"), quality=100)
return i return i
episode_durations = [] episode_durations = []
@@ -1290,7 +1302,7 @@ def _estimate_frame_size_via_calibration(
episode_indices: list[int], episode_indices: list[int],
temp_dir: Path, temp_dir: Path,
fps: int, fps: int,
camera_encoder: VideoEncoderConfig, video_encoder: VideoEncoderConfig,
num_calibration_frames: int = 30, num_calibration_frames: int = 30,
) -> float: ) -> float:
"""Estimate MB per frame by encoding a small calibration sample. """Estimate MB per frame by encoding a small calibration sample.
@@ -1304,7 +1316,7 @@ def _estimate_frame_size_via_calibration(
episode_indices: List of episode indices being processed. episode_indices: List of episode indices being processed.
temp_dir: Temporary directory for calibration files. temp_dir: Temporary directory for calibration files.
fps: Frames per second for video encoding. fps: Frames per second for video encoding.
camera_encoder: Video encoder settings used for calibration encoding. video_encoder: Video encoder settings used for calibration encoding.
num_calibration_frames: Number of frames to use for calibration (default: 30). num_calibration_frames: Number of frames to use for calibration (default: 30).
Returns: Returns:
@@ -1329,10 +1341,11 @@ def _estimate_frame_size_via_calibration(
hf_dataset = dataset.hf_dataset.with_format(None) hf_dataset = dataset.hf_dataset.with_format(None)
sample_indices = range(from_idx, from_idx + num_frames) sample_indices = range(from_idx, from_idx + num_frames)
# Save calibration frames # Save calibration frames using the suffix/format the encoder expects.
is_depth = img_key in dataset.meta.depth_keys
frame_pattern = DEPTH_FILE_PATTERN if is_depth else IMAGE_FILE_PATTERN
for i, idx in enumerate(sample_indices): for i, idx in enumerate(sample_indices):
img = hf_dataset[idx][img_key] write_image(hf_dataset[idx][img_key], calibration_dir / frame_pattern.format(frame_index=i))
img.save(str(calibration_dir / f"frame-{i:06d}.png"), quality=100)
# Encode calibration video # Encode calibration video
calibration_video_path = calibration_dir / "calibration.mp4" calibration_video_path = calibration_dir / "calibration.mp4"
@@ -1340,7 +1353,7 @@ def _estimate_frame_size_via_calibration(
imgs_dir=calibration_dir, imgs_dir=calibration_dir,
video_path=calibration_video_path, video_path=calibration_video_path,
fps=fps, fps=fps,
camera_encoder=camera_encoder, video_encoder=video_encoder,
overwrite=True, overwrite=True,
) )
@@ -1613,6 +1626,7 @@ def recompute_stats(
raise ValueError(f"No parquet files found in {data_dir}") raise ValueError(f"No parquet files found in {data_dir}")
all_episode_stats = [] all_episode_stats = []
# TODO: enable image and video stats re-computation
numeric_keys = [k for k, v in features_to_compute.items() if v["dtype"] not in ["image", "video"]] numeric_keys = [k for k, v in features_to_compute.items() if v["dtype"] not in ["image", "video"]]
for parquet_path in tqdm(parquet_files, desc="Computing stats from data files"): for parquet_path in tqdm(parquet_files, desc="Computing stats from data files"):
@@ -1658,7 +1672,8 @@ def convert_image_to_video_dataset(
dataset: LeRobotDataset, dataset: LeRobotDataset,
output_dir: Path | None = None, output_dir: Path | None = None,
repo_id: str | None = None, repo_id: str | None = None,
camera_encoder: VideoEncoderConfig | None = None, rgb_encoder: RGBEncoderConfig | None = None,
depth_encoder: DepthEncoderConfig | None = None,
episode_indices: list[int] | None = None, episode_indices: list[int] | None = None,
num_workers: int = 4, num_workers: int = 4,
max_episodes_per_batch: int | None = None, max_episodes_per_batch: int | None = None,
@@ -1670,21 +1685,32 @@ def convert_image_to_video_dataset(
LeRobot dataset structure with videos stored in chunked MP4 files. LeRobot dataset structure with videos stored in chunked MP4 files.
Args: Args:
dataset: The source LeRobot dataset with images dataset: The source LeRobot dataset with images.
output_dir: Root directory where the edited dataset will be stored. If not specified, defaults to $HF_LEROBOT_HOME/repo_id. Equivalent to new_root in EditDatasetConfig. output_dir: Root directory where the converted dataset will be stored. When
repo_id: Edited dataset identifier. Equivalent to new_repo_id in EditDatasetConfig. ``None``, defaults to ``$HF_LEROBOT_HOME/repo_id``. Equivalent to
camera_encoder: Video encoder settings ``new_root`` in ``EditDatasetConfig``.
(``None`` uses :func:`~lerobot.configs.camera_encoder_defaults`). repo_id: Converted dataset identifier. Equivalent to ``new_repo_id`` in
episode_indices: List of episode indices to convert (None = all episodes) ``EditDatasetConfig``.
num_workers: Number of threads for parallel processing (default: 4) rgb_encoder: Video encoder settings applied to RGB cameras. When ``None``,
max_episodes_per_batch: Maximum episodes per video batch to avoid memory issues (None = no limit) :func:`~lerobot.configs.video.rgb_encoder_defaults` is used.
max_frames_per_batch: Maximum frames per video batch to avoid memory issues (None = no limit) depth_encoder: Video encoder settings applied to depth-map cameras, including
the quantization parameters persisted to the dataset metadata. When
``None``, :func:`~lerobot.configs.video.depth_encoder_defaults` is used.
episode_indices: Episode indices to convert. When ``None``, all episodes are
converted.
num_workers: Number of threads for parallel processing.
max_episodes_per_batch: Maximum episodes per video batch, to bound memory use.
``None`` means no limit.
max_frames_per_batch: Maximum frames per video batch, to bound memory use.
``None`` means no limit.
Returns: Returns:
New LeRobotDataset with images encoded as videos A new :class:`LeRobotDataset` with images encoded as videos.
""" """
if camera_encoder is None: if rgb_encoder is None:
camera_encoder = camera_encoder_defaults() rgb_encoder = rgb_encoder_defaults()
if depth_encoder is None:
depth_encoder = depth_encoder_defaults()
# Check that it's an image dataset # Check that it's an image dataset
if len(dataset.meta.video_keys) > 0: if len(dataset.meta.video_keys) > 0:
@@ -1709,10 +1735,7 @@ def convert_image_to_video_dataset(
logging.info( logging.info(
f"Converting {len(episode_indices)} episodes with {len(img_keys)} cameras from {dataset.repo_id}" f"Converting {len(episode_indices)} episodes with {len(img_keys)} cameras from {dataset.repo_id}"
) )
logging.info( logging.info(f"RGB video encoder: {rgb_encoder}, depth video encoder: {depth_encoder}")
f"Video codec: {camera_encoder.vcodec}, pixel format: {camera_encoder.pix_fmt}, "
f"GOP: {camera_encoder.g}, CRF: {camera_encoder.crf}"
)
# Create new features dict, converting image features to video features # Create new features dict, converting image features to video features
new_features = {} new_features = {}
@@ -1774,6 +1797,8 @@ def convert_image_to_video_dataset(
episode_lengths = {ep_idx: dataset.meta.episodes["length"][ep_idx] for ep_idx in episode_indices} episode_lengths = {ep_idx: dataset.meta.episodes["length"][ep_idx] for ep_idx in episode_indices}
for img_key in tqdm(img_keys, desc="Processing cameras"): for img_key in tqdm(img_keys, desc="Processing cameras"):
target_encoder = depth_encoder if img_key in dataset.meta.depth_keys else rgb_encoder
# Estimate size per frame by encoding a small calibration sample # Estimate size per frame by encoding a small calibration sample
# This provides accurate compression ratio for the specific codec parameters # This provides accurate compression ratio for the specific codec parameters
size_per_frame_mb = _estimate_frame_size_via_calibration( size_per_frame_mb = _estimate_frame_size_via_calibration(
@@ -1782,7 +1807,7 @@ def convert_image_to_video_dataset(
episode_indices=episode_indices, episode_indices=episode_indices,
temp_dir=temp_dir, temp_dir=temp_dir,
fps=fps, fps=fps,
camera_encoder=camera_encoder, video_encoder=target_encoder,
) )
logging.info(f"Processing camera: {img_key}") logging.info(f"Processing camera: {img_key}")
@@ -1824,7 +1849,7 @@ def convert_image_to_video_dataset(
imgs_dir=imgs_dir, imgs_dir=imgs_dir,
video_path=video_path, video_path=video_path,
fps=fps, fps=fps,
camera_encoder=camera_encoder, video_encoder=target_encoder,
overwrite=True, overwrite=True,
) )
@@ -1863,16 +1888,11 @@ def convert_image_to_video_dataset(
new_meta.info.total_tasks = dataset.meta.total_tasks new_meta.info.total_tasks = dataset.meta.total_tasks
new_meta.info.splits = {"train": f"0:{len(episode_indices)}"} new_meta.info.splits = {"train": f"0:{len(episode_indices)}"}
# Update video info for all image keys (now videos) # Update video info for all image keys (now videos). They are registered as
# We need to manually set video info since update_video_info() checks video_keys first # video features above, so update_video_info populates their (still-empty) info.
for img_key in img_keys: for img_key in img_keys:
if not new_meta.features[img_key].get("info", None): target_encoder = depth_encoder if img_key in dataset.meta.depth_keys else rgb_encoder
video_path = new_meta.root / new_meta.video_path.format( new_meta.update_video_info(video_key=img_key, video_encoder=target_encoder)
video_key=img_key, chunk_index=0, file_index=0
)
new_meta.info.features[img_key]["info"] = get_video_info(
video_path, camera_encoder=camera_encoder
)
write_info(new_meta.info, new_meta.root) write_info(new_meta.info, new_meta.root)
@@ -1899,11 +1919,11 @@ def convert_image_to_video_dataset(
def _reencode_video_worker(args: tuple) -> Path: def _reencode_video_worker(args: tuple) -> Path:
"""Picklable worker for :func:`reencode_dataset`'s process pool.""" """Picklable worker for :func:`reencode_dataset`'s process pool."""
video_path, camera_encoder, encoder_threads = args video_path, video_encoder, encoder_threads = args
reencode_video( reencode_video(
input_video_path=video_path, input_video_path=video_path,
output_video_path=video_path, output_video_path=video_path,
camera_encoder=camera_encoder, video_encoder=video_encoder,
encoder_threads=encoder_threads, encoder_threads=encoder_threads,
overwrite=True, overwrite=True,
) )
@@ -1912,7 +1932,8 @@ def _reencode_video_worker(args: tuple) -> Path:
def reencode_dataset( def reencode_dataset(
dataset: LeRobotDataset, dataset: LeRobotDataset,
camera_encoder: VideoEncoderConfig, rgb_encoder: RGBEncoderConfig | None = None,
depth_encoder: DepthEncoderConfig | None = None,
encoder_threads: int | None = None, encoder_threads: int | None = None,
num_workers: int | None = None, num_workers: int | None = None,
) -> LeRobotDataset: ) -> LeRobotDataset:
@@ -1923,8 +1944,11 @@ def reencode_dataset(
Args: Args:
dataset: An existing :class:`LeRobotDataset` whose videos will be dataset: An existing :class:`LeRobotDataset` whose videos will be
re-encoded. re-encoded.
camera_encoder: Target encoder configuration applied to every video rgb_encoder: Target encoder configuration applied to every RGB video
file. file. If ``None``, re-encoding is skipped for RGB videos.
depth_encoder: Target encoder configuration applied to every depth video
file. If ``None``, re-encoding is skipped for depth videos.
Quantization parameters will not override the ones in the current dataset.
encoder_threads: Per-encoder thread count forwarded to encoder_threads: Per-encoder thread count forwarded to
:func:`reencode_video`. ``None`` lets the codec decide. :func:`reencode_video`. ``None`` lets the codec decide.
num_workers: Number of parallel processes. ``None`` or ``0`` means num_workers: Number of parallel processes. ``None`` or ``0`` means
@@ -1936,23 +1960,35 @@ def reencode_dataset(
on disk. on disk.
""" """
meta = dataset.meta meta = dataset.meta
video_paths_list = [] video_keys_encoders_dict = {}
video_keys_paths_dict = {}
if rgb_encoder is None and depth_encoder is None:
raise ValueError("Either rgb_encoder or depth_encoder must be provided")
# Only re-encode if the videos are not already encoded with the given video encoding parameters # Only re-encode if the videos are not already encoded with the given video encoding parameters
for video_key in meta.video_keys: for video_key in meta.video_keys:
current_info = meta.info.features[video_key].get("info", {}) current_info = meta.info.features[video_key].get("info", {})
current_encoder = VideoEncoderConfig.from_video_info(current_info) current_encoder = encoder_config_from_video_info(current_info)
if current_encoder != camera_encoder: target_encoder = depth_encoder if video_key in meta.depth_keys else rgb_encoder
video_paths_list.extend((meta.root / VIDEO_DIR / video_key).rglob("*.mp4")) if target_encoder is None:
logging.info(f"No encoder provided for {video_key} video. Skipping re-encoding.")
elif current_encoder != target_encoder:
video_keys_paths_dict[video_key] = list((meta.root / VIDEO_DIR / video_key).rglob("*.mp4"))
video_keys_encoders_dict[video_key] = target_encoder
else: else:
logging.info(f"{video_key} videos are already encoded with {camera_encoder}. Nothing to do.") logging.info(f"{video_key} videos are already encoded with {target_encoder}. Nothing to do.")
if len(video_paths_list) == 0: if len(video_keys_paths_dict) == 0:
logging.warning("Dataset has no videos to re-encode.") logging.warning("Dataset has no videos to re-encode.")
return dataset return dataset
logging.info(f"Re-encoding {len(video_paths_list)} video file(s) with {camera_encoder}") logging.info(f"Re-encoding {sum(len(paths) for paths in video_keys_paths_dict.values())} video file(s).")
worker_args = [(vp, camera_encoder, encoder_threads) for vp in video_paths_list] worker_args = [
(path, encoder, encoder_threads)
for video_key, encoder in video_keys_encoders_dict.items()
for path in video_keys_paths_dict[video_key]
]
if num_workers and num_workers > 1: if num_workers and num_workers > 1:
with ProcessPoolExecutor(max_workers=num_workers) as pool: with ProcessPoolExecutor(max_workers=num_workers) as pool:
futures = [pool.submit(_reencode_video_worker, args) for args in worker_args] futures = [pool.submit(_reencode_video_worker, args) for args in worker_args]
@@ -1966,10 +2002,15 @@ def reencode_dataset(
for args in tqdm(worker_args, desc="Re-encoding videos"): for args in tqdm(worker_args, desc="Re-encoding videos"):
_reencode_video_worker(args) _reencode_video_worker(args)
# Refresh video info in metadata for every video key. # Refresh video info in metadata for every re-encoded key. Re-encoding only
for vid_key in meta.video_keys: # changes codec/container params, so for depth videos we preserve ``is_depth_map``
video_path = meta.root / meta.get_video_file_path(0, vid_key) # and the depth quantization params (``video.depth_min`` / ``video.depth_max`` /
meta.info.features[vid_key]["info"] = get_video_info(video_path, camera_encoder=camera_encoder) # ...), which describe the data rather than the codec and must survive a transcode.
# RGB videos pass an empty set: still a refresh, but nothing to preserve.
depth_preserve_keys = {"is_depth_map", *(f"video.{n}" for n in DEPTH_ENCODER_INFO_FIELD_NAMES)}
for video_key, encoder in video_keys_encoders_dict.items():
preserve_keys = depth_preserve_keys if video_key in meta.depth_keys else set()
meta.update_video_info(video_key=video_key, video_encoder=encoder, preserve_keys=preserve_keys)
write_info(meta.info, meta.root) write_info(meta.info, meta.root)
logging.info("Dataset metadata updated.") logging.info("Dataset metadata updated.")
+42 -14
View File
@@ -31,7 +31,13 @@ import PIL.Image
import pyarrow.parquet as pq import pyarrow.parquet as pq
import torch import torch
from lerobot.configs import VideoEncoderConfig, camera_encoder_defaults from lerobot.configs import (
DepthEncoderConfig,
RGBEncoderConfig,
VideoEncoderConfig,
depth_encoder_defaults,
rgb_encoder_defaults,
)
from .compute_stats import compute_episode_stats from .compute_stats import compute_episode_stats
from .dataset_metadata import LeRobotDatasetMetadata from .dataset_metadata import LeRobotDatasetMetadata
@@ -48,6 +54,7 @@ from .io_utils import (
write_info, write_info,
) )
from .utils import ( from .utils import (
DEFAULT_DEPTH_PATH,
DEFAULT_EPISODES_PATH, DEFAULT_EPISODES_PATH,
DEFAULT_IMAGE_PATH, DEFAULT_IMAGE_PATH,
update_chunk_file_indices, update_chunk_file_indices,
@@ -67,17 +74,22 @@ def _encode_video_worker(
episode_index: int, episode_index: int,
root: Path, root: Path,
fps: int, fps: int,
camera_encoder: VideoEncoderConfig | None = None, video_encoder: VideoEncoderConfig | None = None,
encoder_threads: int | None = None, encoder_threads: int | None = None,
) -> Path: ) -> Path:
temp_path = Path(tempfile.mkdtemp(dir=root)) / f"{video_key}_{episode_index:03d}.mp4" temp_path = Path(tempfile.mkdtemp(dir=root)) / f"{video_key}_{episode_index:03d}.mp4"
fpath = DEFAULT_IMAGE_PATH.format(image_key=video_key, episode_index=episode_index, frame_index=0) path_template = (
DEFAULT_DEPTH_PATH
if video_encoder is not None and isinstance(video_encoder, DepthEncoderConfig)
else DEFAULT_IMAGE_PATH
)
fpath = path_template.format(image_key=video_key, episode_index=episode_index, frame_index=0)
img_dir = (root / fpath).parent img_dir = (root / fpath).parent
encode_video_frames( encode_video_frames(
img_dir, img_dir,
temp_path, temp_path,
fps, fps,
camera_encoder=camera_encoder, video_encoder=video_encoder,
encoder_threads=encoder_threads, encoder_threads=encoder_threads,
overwrite=True, overwrite=True,
) )
@@ -96,7 +108,8 @@ class DatasetWriter:
self, self,
meta: LeRobotDatasetMetadata, meta: LeRobotDatasetMetadata,
root: Path, root: Path,
camera_encoder: VideoEncoderConfig | None, rgb_encoder: RGBEncoderConfig | None,
depth_encoder: DepthEncoderConfig | None,
encoder_threads: int | None, encoder_threads: int | None,
batch_encoding_size: int, batch_encoding_size: int,
streaming_encoder: StreamingVideoEncoder | None = None, streaming_encoder: StreamingVideoEncoder | None = None,
@@ -108,8 +121,11 @@ class DatasetWriter:
meta: Dataset metadata instance (used for feature schema, chunk meta: Dataset metadata instance (used for feature schema, chunk
settings, and episode persistence). settings, and episode persistence).
root: Local dataset root directory. root: Local dataset root directory.
camera_encoder: Video encoder settings applied to all cameras. rgb_encoder: Video encoder settings applied to RGB cameras. When
``None`` uses :func:`~lerobot.configs.camera_encoder_defaults`. ``None``, :func:`~lerobot.configs.video.rgb_encoder_defaults` is used.
depth_encoder: Video encoder settings applied to depth cameras, including
the quantization parameters. When ``None``,
:func:`~lerobot.configs.video.depth_encoder_defaults` is used.
encoder_threads: Number of encoder threads (global). ``None`` encoder_threads: Number of encoder threads (global). ``None``
lets the codec decide. lets the codec decide.
batch_encoding_size: Number of episodes to accumulate before batch_encoding_size: Number of episodes to accumulate before
@@ -120,7 +136,8 @@ class DatasetWriter:
""" """
self._meta = meta self._meta = meta
self._root = root self._root = root
self._camera_encoder = camera_encoder or camera_encoder_defaults() self._rgb_encoder = rgb_encoder or rgb_encoder_defaults()
self._depth_encoder = depth_encoder or depth_encoder_defaults()
self._encoder_threads = encoder_threads self._encoder_threads = encoder_threads
self._batch_encoding_size = batch_encoding_size self._batch_encoding_size = batch_encoding_size
self._streaming_encoder = streaming_encoder self._streaming_encoder = streaming_encoder
@@ -145,7 +162,8 @@ class DatasetWriter:
return ep_buffer return ep_buffer
def _get_image_file_path(self, episode_index: int, image_key: str, frame_index: int) -> Path: def _get_image_file_path(self, episode_index: int, image_key: str, frame_index: int) -> Path:
fpath = DEFAULT_IMAGE_PATH.format( path_template = DEFAULT_DEPTH_PATH if image_key in self._meta.depth_keys else DEFAULT_IMAGE_PATH
fpath = path_template.format(
image_key=image_key, episode_index=episode_index, frame_index=frame_index image_key=image_key, episode_index=episode_index, frame_index=frame_index
) )
return self._root / fpath return self._root / fpath
@@ -195,6 +213,7 @@ class DatasetWriter:
if frame_index == 0 and self._streaming_encoder is not None: if frame_index == 0 and self._streaming_encoder is not None:
self._streaming_encoder.start_episode( self._streaming_encoder.start_episode(
video_keys=list(self._meta.video_keys), video_keys=list(self._meta.video_keys),
depth_video_keys=list(self._meta.depth_keys),
temp_dir=self._root, temp_dir=self._root,
) )
@@ -282,10 +301,13 @@ class DatasetWriter:
if use_streaming: if use_streaming:
streaming_results = self._streaming_encoder.finish_episode() streaming_results = self._streaming_encoder.finish_episode()
for video_key in self._meta.video_keys: for video_key in self._meta.video_keys:
normalization_factor = 255.0 if video_key not in self._meta.depth_keys else 1.0
temp_path, video_stats = streaming_results[video_key] temp_path, video_stats = streaming_results[video_key]
if video_stats is not None: if video_stats is not None:
ep_stats[video_key] = { ep_stats[video_key] = {
k: v if k == "count" else np.squeeze(v.reshape(1, -1, 1, 1) / 255.0, axis=0) k: v
if k == "count"
else np.squeeze(v.reshape(1, -1, 1, 1) / normalization_factor, axis=0)
for k, v in video_stats.items() for k, v in video_stats.items()
} }
ep_metadata.update(self._save_episode_video(video_key, episode_index, temp_path=temp_path)) ep_metadata.update(self._save_episode_video(video_key, episode_index, temp_path=temp_path))
@@ -300,7 +322,7 @@ class DatasetWriter:
episode_index, episode_index,
self._root, self._root,
self._meta.fps, self._meta.fps,
self._camera_encoder, self._depth_encoder if video_key in self._meta.depth_keys else self._rgb_encoder,
self._encoder_threads, self._encoder_threads,
): video_key ): video_key
for video_key in self._meta.video_keys for video_key in self._meta.video_keys
@@ -511,7 +533,12 @@ class DatasetWriter:
# Update video info (only needed when first episode is encoded) # Update video info (only needed when first episode is encoded)
if episode_index == 0: if episode_index == 0:
self._meta.update_video_info(video_key, camera_encoder=self._camera_encoder) self._meta.update_video_info(
video_key,
video_encoder=self._depth_encoder
if video_key in self._meta.depth_keys
else self._rgb_encoder,
)
write_info(self._meta.info, self._meta.root) write_info(self._meta.info, self._meta.root)
metadata = { metadata = {
@@ -578,13 +605,14 @@ class DatasetWriter:
self.image_writer.wait_until_done() self.image_writer.wait_until_done()
def _encode_temporary_episode_video(self, video_key: str, episode_index: int) -> Path: def _encode_temporary_episode_video(self, video_key: str, episode_index: int) -> Path:
"""Use ffmpeg to convert frames stored as png into mp4 videos.""" """Use ffmpeg to convert frames stored as png/tiff into mp4 videos."""
is_depth = video_key in self._meta.depth_keys
return _encode_video_worker( return _encode_video_worker(
video_key, video_key,
episode_index, episode_index,
self._root, self._root,
self._meta.fps, self._meta.fps,
self._camera_encoder, self._depth_encoder if is_depth else self._rgb_encoder,
self._encoder_threads, self._encoder_threads,
) )
+268
View File
@@ -0,0 +1,268 @@
#!/usr/bin/env python
# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""
Depth encoding/decoding helpers for :class:`DepthEncoderConfig`.
"""
import math
from typing import Literal
import av
import numpy as np
import torch
from numpy.typing import NDArray
from lerobot.configs.video import (
DEFAULT_DEPTH_MAX,
DEFAULT_DEPTH_MIN,
DEFAULT_DEPTH_PIX_FMT,
DEFAULT_DEPTH_SHIFT,
DEFAULT_DEPTH_USE_LOG,
DEPTH_METER_UNIT,
DEPTH_MILLIMETER_UNIT,
DEPTH_QMAX,
)
from .image_writer import squeeze_single_channel
from .pyav_utils import write_u16_plane
_MM_PER_METRE = 1000.0
_UINT16_MAX = 65535
def _validate_log_quant_params(depth_min: float, shift: float) -> None:
"""Ensure ``log(depth_min + shift)`` is finite."""
if depth_min + shift <= 0:
raise ValueError(
f"depth_min + shift must be positive for logarithmic quantization, "
f"got depth_min={depth_min} + shift={shift} = {depth_min + shift}"
)
def _depth_input_to_float32_and_unit(
depth: NDArray[np.integer] | NDArray[np.floating],
input_unit: Literal["auto", DEPTH_METER_UNIT, DEPTH_MILLIMETER_UNIT],
) -> tuple[NDArray[np.float32], Literal[DEPTH_METER_UNIT, DEPTH_MILLIMETER_UNIT]]:
"""Convert depth to float32 in the chosen unit, and return the resolved unit."""
resolved_unit = (
(DEPTH_METER_UNIT if np.issubdtype(depth.dtype, np.floating) else DEPTH_MILLIMETER_UNIT)
if input_unit == "auto"
else input_unit
)
return depth.astype(np.float32, order="K"), resolved_unit
def quantize_depth(
depth: NDArray[np.uint16] | NDArray[np.float32] | torch.Tensor,
depth_min: float = DEFAULT_DEPTH_MIN,
depth_max: float = DEFAULT_DEPTH_MAX,
shift: float = DEFAULT_DEPTH_SHIFT,
use_log: bool = DEFAULT_DEPTH_USE_LOG,
pix_fmt: str = DEFAULT_DEPTH_PIX_FMT,
video_backend: str | None = "pyav",
input_unit: Literal["auto", DEPTH_METER_UNIT, DEPTH_MILLIMETER_UNIT] = "auto",
) -> NDArray[np.uint16] | av.VideoFrame:
"""Quantize depth to 12-bit codes (``uint16``, values ``0…DEPTH_QMAX``).
Depth maps are packed into 12-bit integer frames so they fit in standard
high-bit-depth pixel formats (e.g. ``yuv420p12le`` / ``gray12le``)
and can be encoded by widely supported video codecs (e.g. HEVC Main 12).
Logarithmic quantization is the default because it allocates more quanta
to near-range depth, which matches the (1/depth) error profile of typical
depth sensors. Math is ported from BEHAVIOR-1K's ``obs_utils.py``.
**Input units**:
- ``input_unit="auto"`` (default): infer from dtype (floating = m, non-floating = mm).
- ``input_unit="mm"``: interpret input values as millimetres.
- ``input_unit="m"``: interpret input values as metres.
Quantization math runs in the **resolved input unit**.
``depth_min``, ``depth_max``, and ``shift`` are always in **metres**.
Args:
depth: Depth map; ``torch.Tensor`` is moved to CPU for conversion.
depth_min: Depth (metres) at quantum ``0``.
depth_max: Depth (metres) at quantum :data:`DEPTH_QMAX`.
shift: Depth shift (metres); used in log mode. Must satisfy ``depth_min + shift > 0``.
use_log: If ``True`` (default), quantize in log space.
video_backend: Video backend to use for encoding. Defaults to "pyav".
input_unit: Input unit policy (``"auto"``, ``"mm"``, ``"m"``).
Returns:
``numpy.ndarray``, ``dtype=uint16``, same shape as ``depth``, values in
``[0, DEPTH_QMAX]``.
Raises:
ValueError: If ``input_unit`` is not ``"auto"``, ``"mm"``, or ``"m"``.
ValueError: If ``use_log=True`` and ``depth_min + shift <= 0``.
"""
if input_unit not in ("auto", DEPTH_METER_UNIT, DEPTH_MILLIMETER_UNIT):
raise ValueError(
f"input_unit must be 'auto', '{DEPTH_METER_UNIT}', or '{DEPTH_MILLIMETER_UNIT}', got {input_unit!r}"
)
if isinstance(depth, torch.Tensor):
depth = depth.detach().cpu().numpy()
# Squeeze single-channel dim: (H, W, 1) or (1, H, W) → (H, W)
depth = squeeze_single_channel(depth)
depth_f, resolved_unit = _depth_input_to_float32_and_unit(depth, input_unit=input_unit)
# Convert depth_min, depth_max, and shift to the resolved input unit.
depth_min_u = (
np.float32(depth_min) if resolved_unit == DEPTH_METER_UNIT else np.float32(depth_min * _MM_PER_METRE)
)
depth_max_u = (
np.float32(depth_max) if resolved_unit == DEPTH_METER_UNIT else np.float32(depth_max * _MM_PER_METRE)
)
shift_u = np.float32(shift) if resolved_unit == DEPTH_METER_UNIT else np.float32(shift * _MM_PER_METRE)
# Normalization and quantization is performed in the resolved input unit.
if use_log:
_validate_log_quant_params(depth_min, shift)
log_min = math.log(float(depth_min_u + shift_u))
log_max = math.log(float(depth_max_u + shift_u))
norm = (np.log(depth_f + shift_u) - log_min) / (log_max - log_min)
else:
norm = (depth_f - depth_min_u) / (depth_max_u - depth_min_u)
quantized = np.rint(norm * DEPTH_QMAX).clip(0, DEPTH_QMAX).astype(np.uint16, copy=False)
if video_backend == "pyav":
frame = av.VideoFrame.from_ndarray(quantized, format=pix_fmt)
write_u16_plane(frame.planes[0], quantized)
return frame
else:
return quantized
def dequantize_depth(
quantized: NDArray[np.uint16] | av.VideoFrame | torch.Tensor,
depth_min: float = DEFAULT_DEPTH_MIN,
depth_max: float = DEFAULT_DEPTH_MAX,
shift: float = DEFAULT_DEPTH_SHIFT,
use_log: bool = DEFAULT_DEPTH_USE_LOG,
pix_fmt: str = DEFAULT_DEPTH_PIX_FMT,
output_unit: Literal[DEPTH_METER_UNIT, DEPTH_MILLIMETER_UNIT] = DEPTH_MILLIMETER_UNIT,
output_tensor: bool = True,
output_channel_last: bool = False,
) -> NDArray[np.uint16] | NDArray[np.float32] | torch.Tensor:
"""Inverse of :func:`quantize_depth`.
Decoding inverts the same normalized code mapping as :func:`quantize_depth`
using ``depth_min`` / ``depth_max`` / ``shift`` (in metres), then returns
the requested output unit. Tuning arguments **must match** :func:`quantize_depth`.
Accepted input layouts :
- ``(H, W, 1)`` or ``(H, W)`` single frame with channel-last.
- ``(..., 1, H, W)`` batched frames with channel-first.
- ``(..., H, W, 1)`` batched frames with channel-last.
Output layout is determined by ``output_channel_last``.
Args:
quantized: 12-bit codes in ``[0, DEPTH_QMAX]``. ``np.ndarray``,
``av.VideoFrame``, or ``torch.Tensor`` (any integer or float dtype).
depth_min, depth_max, shift, use_log: Same as :func:`quantize_depth` (metres).
pix_fmt: Pixel format used to extract the plane from an ``av.VideoFrame``.
output_unit: ``"mm"`` returns ``uint16`` millimetres (rint, clip
``[0, 65535]``) when returning a numpy array, or ``float32`` mm when
``output_tensor=True``. ``"m"`` returns ``float32`` metres in
``[depth_min, depth_max]``.
output_tensor: If True, return a ``torch.Tensor`` instead of a numpy array.
Returns:
Depth map in the requested unit and dtype.
Raises:
ValueError: If ``output_unit`` is not ``"m"`` or ``"mm"``.
ValueError: If ``use_log=True`` and ``depth_min + shift <= 0``.
"""
if output_unit not in (DEPTH_METER_UNIT, DEPTH_MILLIMETER_UNIT):
raise ValueError(
f"output_unit must be '{DEPTH_METER_UNIT}' or '{DEPTH_MILLIMETER_UNIT}', got {output_unit!r}"
)
if use_log:
_validate_log_quant_params(depth_min, shift)
if isinstance(quantized, av.VideoFrame):
quantized = quantized.to_ndarray(format=pix_fmt)
# Compute the scale and offset first.
depth_min_m = float(depth_min)
depth_max_m = float(depth_max)
shift_m = float(shift)
if use_log:
log_min = math.log(depth_min_m + shift_m)
log_max = math.log(depth_max_m + shift_m)
scale = (log_max - log_min) / DEPTH_QMAX
offset = log_min
else:
scale = (depth_max_m - depth_min_m) / DEPTH_QMAX
offset = depth_min_m
# ── Torch path: stay on the input device, single fp32 allocation. ────────
if isinstance(quantized, torch.Tensor):
if quantized.ndim >= 3:
# Drop the single-channel dimension so the math runs on (..., H, W).
quantized = quantized.squeeze(-3) if quantized.shape[-3] == 1 else quantized.squeeze(-1)
# Single allocation we own; everything else is in-place.
buf = quantized.to(dtype=torch.float32, copy=True)
buf.mul_(scale).add_(offset)
if use_log:
buf.exp_().sub_(shift_m)
buf.clamp_(depth_min_m, depth_max_m)
buf.unsqueeze_(-1) if output_channel_last else buf.unsqueeze_(-3)
if output_unit == DEPTH_METER_UNIT:
return buf if output_tensor else buf.cpu().numpy()
# mm path: round + clamp in float32, skipping the uint16 round-trip
# when returning a tensor (torch.uint16 is poorly supported).
buf.mul_(_MM_PER_METRE).round_().clamp_(0.0, _UINT16_MAX)
if output_tensor:
return buf
return buf.cpu().numpy().astype(np.uint16, copy=False)
# ── NumPy path: single fp32 allocation, ``out=`` for in-place math. ─────
arr = np.asarray(quantized)
if arr.ndim >= 3:
# Drop the single-channel dimension so the math runs on (..., H, W).
arr = np.squeeze(arr, axis=-3) if arr.shape[-3] == 1 else np.squeeze(arr, axis=-1)
buf = np.empty(arr.shape, dtype=np.float32)
np.multiply(arr, scale, out=buf)
np.add(buf, offset, out=buf)
if use_log:
np.exp(buf, out=buf)
np.subtract(buf, shift_m, out=buf)
np.clip(buf, depth_min_m, depth_max_m, out=buf)
buf = np.expand_dims(buf, axis=-1) if output_channel_last else np.expand_dims(buf, axis=-3)
if output_unit == DEPTH_METER_UNIT:
return torch.from_numpy(buf) if output_tensor else buf
np.multiply(buf, _MM_PER_METRE, out=buf)
np.rint(buf, out=buf)
np.clip(buf, 0.0, _UINT16_MAX, out=buf)
if output_tensor:
# torch.uint16 support is very limited; return float32 millimetres.
return torch.from_numpy(buf)
return buf.astype(np.uint16, copy=False)
+3
View File
@@ -97,6 +97,7 @@ def make_dataset(cfg: TrainPipelineConfig) -> LeRobotDataset | MultiLeRobotDatas
revision=cfg.dataset.revision, revision=cfg.dataset.revision,
video_backend=cfg.dataset.video_backend, video_backend=cfg.dataset.video_backend,
return_uint8=True, return_uint8=True,
depth_output_unit=cfg.dataset.depth_output_unit,
tolerance_s=cfg.tolerance_s, tolerance_s=cfg.tolerance_s,
) )
else: else:
@@ -127,6 +128,8 @@ def make_dataset(cfg: TrainPipelineConfig) -> LeRobotDataset | MultiLeRobotDatas
if cfg.dataset.use_imagenet_stats: if cfg.dataset.use_imagenet_stats:
for key in dataset.meta.camera_keys: for key in dataset.meta.camera_keys:
if key in dataset.meta.depth_keys:
continue # Exclude depth keys from ImageNet stats
for stats_type, stats in IMAGENET_STATS.items(): for stats_type, stats in IMAGENET_STATS.items():
dataset.meta.stats[key][stats_type] = torch.tensor(stats, dtype=torch.float32) dataset.meta.stats[key][stats_type] = torch.tensor(stats, dtype=torch.float32)
+1 -1
View File
@@ -336,7 +336,7 @@ def validate_feature_image_or_video(
Args: Args:
name (str): The name of the feature. name (str): The name of the feature.
expected_shape (list[str]): The expected shape (C, H, W). expected_shape (list[str]): The expected shape, e.g. (C, H, W) or (H, W, C).
value: The image data to validate. value: The image data to validate.
Returns: Returns:
+62 -6
View File
@@ -41,11 +41,51 @@ def safe_stop_image_writer(func):
return wrapper return wrapper
def image_array_to_pil_image(image_array: np.ndarray, range_check: bool = True) -> PIL.Image.Image: def squeeze_single_channel(array: np.ndarray) -> np.ndarray:
# TODO(aliberts): handle 1 channel and 4 for depth images """Drop a leading or trailing singleton channel dim: ``(1, H, W)`` / ``(H, W, 1)`` -> ``(H, W)``.
if image_array.ndim != 3:
raise ValueError(f"The array has {image_array.ndim} dimensions, but 3 is expected for an image.")
Unlike ``array.squeeze()``, this only removes the channel axis, never an ``H`` or ``W`` of size 1.
"""
if array.ndim == 3:
if array.shape[0] == 1:
return array[0]
if array.shape[-1] == 1:
return array[..., 0]
return array
def image_array_to_pil_image(image_array: np.ndarray, range_check: bool = True) -> PIL.Image.Image:
"""Convert a NumPy array to a PIL Image, preserving precision for grayscale.
Behaviour by shape:
- ``(H, W)`` or ``(1, H, W)`` / ``(H, W, 1)``: single-channel grayscale.
The native dtype is preserved using the matching PIL mode
(``I;16`` / ``F``). This is the path used for raw depth maps (no rescaling, clamping, or downcasting)
- ``(3, H, W)`` / ``(H, W, 3)``: RGB. Channels-first inputs are transposed
to channels-last. Float inputs in ``[0, 1]`` are scaled to ``uint8``
(existing behaviour, gated by ``range_check``).
Other shapes / channel counts raise ``NotImplementedError`` or
``ValueError``.
"""
# TODO(CarolinePascal): 4 dimensions RGB-D images
if image_array.ndim not in (2, 3):
raise ValueError(f"The array has {image_array.ndim} dimensions, but 2 or 3 is expected for an image.")
# Squeeze 3D single-channel inputs to 2D so depth maps work whether the
# caller emits (H, W), (1, H, W), or (H, W, 1).
image_array = squeeze_single_channel(image_array)
if image_array.ndim == 2:
if image_array.dtype not in [np.uint16, np.float32]:
raise ValueError(
f"Unsupported single-channel image dtype: {image_array.dtype}. "
f"Supported dtypes: {sorted(str(d) for d in [np.uint16, np.float32])}."
)
return PIL.Image.fromarray(np.ascontiguousarray(image_array))
# 3D path: must be RGB (3 channels), channels-first or channels-last.
if image_array.shape[0] == 3: if image_array.shape[0] == 3:
# Transpose from pytorch convention (C, H, W) to (H, W, C) # Transpose from pytorch convention (C, H, W) to (H, W, C)
image_array = image_array.transpose(1, 2, 0) image_array = image_array.transpose(1, 2, 0)
@@ -71,13 +111,29 @@ def image_array_to_pil_image(image_array: np.ndarray, range_check: bool = True)
return PIL.Image.fromarray(image_array) return PIL.Image.fromarray(image_array)
def save_kwargs_for_path(fpath: Path, compress_level: int) -> dict:
"""Pick the right format-specific kwargs for :meth:`PIL.Image.Image.save`.
PNG uses ``compress_level`` (0-9, zlib). TIFF uses ``compression`` (raw) for lossless raw depth maps.
"""
suffix = Path(fpath).suffix.lower()
if suffix == ".png":
return {"compress_level": compress_level}
if suffix in (".tif", ".tiff"):
return {"compression": "raw"}
else:
raise ValueError(f"Unsupported image file extension: {suffix}")
def write_image(image: np.ndarray | PIL.Image.Image, fpath: Path, compress_level: int = 1): def write_image(image: np.ndarray | PIL.Image.Image, fpath: Path, compress_level: int = 1):
""" """
Saves a NumPy array or PIL Image to a file. Saves a NumPy array or PIL Image to a file.
This function handles both NumPy arrays and PIL Image objects, converting This function handles both NumPy arrays and PIL Image objects, converting
the former to a PIL Image before saving. It includes error handling for the former to a PIL Image before saving. It includes error handling for
the save operation. the save operation. The output format is inferred from the *fpath*
extension: ``.png`` PNG with ``compress_level``, ``.tiff`` / ``.tif``
lossless raw depth maps (TIFF).
Args: Args:
image (np.ndarray | PIL.Image.Image): The image data to save. image (np.ndarray | PIL.Image.Image): The image data to save.
@@ -101,7 +157,7 @@ def write_image(image: np.ndarray | PIL.Image.Image, fpath: Path, compress_level
img = image img = image
else: else:
raise TypeError(f"Unsupported image type: {type(image)}") raise TypeError(f"Unsupported image type: {type(image)}")
img.save(fpath, compress_level=compress_level) img.save(fpath, **save_kwargs_for_path(fpath, compress_level))
except Exception as e: except Exception as e:
logger.error("Error writing image %s: %s", fpath, e) logger.error("Error writing image %s: %s", fpath, e)
+36 -11
View File
@@ -226,28 +226,50 @@ def load_image_as_numpy(
Args: Args:
fpath (str | Path): Path to the image file. fpath (str | Path): Path to the image file.
dtype (np.dtype): The desired data type of the output array. If floating, dtype (np.dtype): The desired data type of the output array. If floating,
pixels are scaled to [0, 1]. pixels are scaled to [0, 1]. Only used for RGB images.
channel_first (bool): If True, converts the image to (C, H, W) format. channel_first (bool): If True, converts the image to (C, H, W) format.
Otherwise, it remains in (H, W, C) format. Otherwise, it remains in (H, W, C) format.
Returns: Returns:
np.ndarray: The image as a numpy array. np.ndarray: The image as a numpy array.
""" """
img = PILImage.open(fpath).convert("RGB") is_depth = fpath.endswith(".tiff") or fpath.endswith(".tif")
img_array = np.array(img, dtype=dtype) if is_depth:
# Preserve the native depth dtype (uint16 -> "I;16", float32 -> "F").
img = PILImage.open(fpath)
img_array = np.array(img)
else:
img = PILImage.open(fpath).convert("RGB")
img_array = np.array(img, dtype=dtype)
if np.issubdtype(dtype, np.floating):
img_array /= 255.0
if channel_first: # (H, W, C) -> (C, H, W) if channel_first: # (H, W, C) -> (C, H, W)
img_array = np.transpose(img_array, (2, 0, 1)) img_array = img_array[np.newaxis, ...] if img_array.ndim == 2 else np.transpose(img_array, (2, 0, 1))
if np.issubdtype(dtype, np.floating):
img_array /= 255.0
return img_array return img_array
# PIL modes for 16-bit unsigned depth maps.
UINT16_PIL_MODES = {"I;16", "I;16B", "I;16L"}
def pil_to_chw_tensor(img: PILImage.Image) -> torch.Tensor:
"""Convert a PIL image to a channel-first tensor.
``uint16`` depth maps become ``float32 (1, H, W)`` in native units (``ToTensor``
would overflow them to ``int16``); all other modes use the standard ``ToTensor`` path.
"""
if img.mode in UINT16_PIL_MODES:
return torch.from_numpy(np.array(img, dtype=np.float32))[None, ...]
return transforms.ToTensor()(img)
def hf_transform_to_torch(items_dict: dict[str, list[Any]]) -> dict[str, list[torch.Tensor | str]]: def hf_transform_to_torch(items_dict: dict[str, list[Any]]) -> dict[str, list[torch.Tensor | str]]:
"""Convert a batch from a Hugging Face dataset to torch tensors. """Convert a batch from a Hugging Face dataset to torch tensors.
This transform function converts items from Hugging Face dataset format (pyarrow) This transform function converts items from Hugging Face dataset format (pyarrow)
to torch tensors. Importantly, images are converted from PIL objects (H, W, C, uint8) to torch tensors. RGB images are converted from PIL objects (H, W, C, uint8)
to a torch image representation (C, H, W, float32) in the range [0, 1]. Other to a torch image representation (C, H, W, float32) in the range [0, 1]. Depth
maps are returned as float32 (1, H, W) in their native units. Other
types are converted to torch.tensor. types are converted to torch.tensor.
Args: Args:
@@ -262,8 +284,7 @@ def hf_transform_to_torch(items_dict: dict[str, list[Any]]) -> dict[str, list[to
continue continue
first_item = items_dict[key][0] first_item = items_dict[key][0]
if isinstance(first_item, PILImage.Image): if isinstance(first_item, PILImage.Image):
to_tensor = transforms.ToTensor() items_dict[key] = [pil_to_chw_tensor(img) for img in items_dict[key]]
items_dict[key] = [to_tensor(img) for img in items_dict[key]]
elif first_item is None or isinstance(first_item, dict): elif first_item is None or isinstance(first_item, dict):
pass pass
else: else:
@@ -329,7 +350,11 @@ def item_to_torch(item: dict) -> dict:
""" """
skip_keys = {"task", *LANGUAGE_COLUMNS} skip_keys = {"task", *LANGUAGE_COLUMNS}
for key, val in item.items(): for key, val in item.items():
if isinstance(val, (np.ndarray | list)) and key not in skip_keys: if key in skip_keys:
continue
if isinstance(val, PILImage.Image):
item[key] = pil_to_chw_tensor(val)
elif isinstance(val, (np.ndarray | list)):
# Convert numpy arrays and lists to torch tensors # Convert numpy arrays and lists to torch tensors
item[key] = torch.tensor(val) item[key] = torch.tensor(val)
return item return item
+40 -18
View File
@@ -24,7 +24,7 @@ import torch.utils
from huggingface_hub import HfApi, snapshot_download from huggingface_hub import HfApi, snapshot_download
from huggingface_hub.errors import RevisionNotFoundError from huggingface_hub.errors import RevisionNotFoundError
from lerobot.configs import VideoEncoderConfig from lerobot.configs import DEFAULT_DEPTH_UNIT, DepthEncoderConfig, RGBEncoderConfig
from lerobot.utils.constants import HF_LEROBOT_HUB_CACHE from lerobot.utils.constants import HF_LEROBOT_HUB_CACHE
from .dataset_metadata import CODEBASE_VERSION, LeRobotDatasetMetadata from .dataset_metadata import CODEBASE_VERSION, LeRobotDatasetMetadata
@@ -58,8 +58,10 @@ class LeRobotDataset(torch.utils.data.Dataset):
download_videos: bool = True, download_videos: bool = True,
video_backend: str | None = None, video_backend: str | None = None,
return_uint8: bool = False, return_uint8: bool = False,
depth_output_unit: str = DEFAULT_DEPTH_UNIT,
batch_encoding_size: int = 1, batch_encoding_size: int = 1,
camera_encoder: VideoEncoderConfig | None = None, rgb_encoder: RGBEncoderConfig | None = None,
depth_encoder: DepthEncoderConfig | None = None,
encoder_threads: int | None = None, encoder_threads: int | None = None,
streaming_encoding: bool = False, streaming_encoding: bool = False,
encoder_queue_maxsize: int = 30, encoder_queue_maxsize: int = 30,
@@ -183,8 +185,11 @@ class LeRobotDataset(torch.utils.data.Dataset):
You can also use the 'pyav' decoder used by Torchvision, which used to be the default option, or 'video_reader' which is another decoder of Torchvision. You can also use the 'pyav' decoder used by Torchvision, which used to be the default option, or 'video_reader' which is another decoder of Torchvision.
batch_encoding_size (int, optional): Number of episodes to accumulate before batch encoding videos. batch_encoding_size (int, optional): Number of episodes to accumulate before batch encoding videos.
Set to 1 for immediate encoding (default), or higher for batched encoding. Defaults to 1. Set to 1 for immediate encoding (default), or higher for batched encoding. Defaults to 1.
camera_encoder (VideoEncoderConfig | None, optional): Video encoder settings for cameras rgb_encoder (RGBEncoderConfig | None, optional): Video encoder settings for cameras
(codec, quality, etc.). When ``None``, :func:`~lerobot.configs.video.camera_encoder_defaults` (codec, quality, etc.). When ``None``, :func:`~lerobot.configs.video.rgb_encoder_defaults`
is used by the writer.
depth_encoder (DepthEncoderConfig | None, optional): Video encoder settings for depth cameras
(codec, quality, etc.). When ``None``, :func:`~lerobot.configs.video.depth_encoder_defaults`
is used by the writer. is used by the writer.
encoder_threads (int | None, optional): Number of encoder threads (global). ``None`` lets the encoder_threads (int | None, optional): Number of encoder threads (global). ``None`` lets the
codec decide. codec decide.
@@ -206,6 +211,7 @@ class LeRobotDataset(torch.utils.data.Dataset):
self.revision = revision if revision else CODEBASE_VERSION self.revision = revision if revision else CODEBASE_VERSION
self._video_backend = video_backend if video_backend else get_safe_default_video_backend() self._video_backend = video_backend if video_backend else get_safe_default_video_backend()
self._return_uint8 = return_uint8 self._return_uint8 = return_uint8
self._depth_output_unit = depth_output_unit
self._batch_encoding_size = batch_encoding_size self._batch_encoding_size = batch_encoding_size
self._encoder_threads = encoder_threads self._encoder_threads = encoder_threads
@@ -246,6 +252,7 @@ class LeRobotDataset(torch.utils.data.Dataset):
delta_timestamps=delta_timestamps, delta_timestamps=delta_timestamps,
image_transforms=image_transforms, image_transforms=image_transforms,
return_uint8=self._return_uint8, return_uint8=self._return_uint8,
depth_output_unit=self._depth_output_unit,
) )
self.image_transforms = image_transforms self.image_transforms = image_transforms
@@ -271,14 +278,16 @@ class LeRobotDataset(torch.utils.data.Dataset):
if streaming_encoding and len(self.meta.video_keys) > 0: if streaming_encoding and len(self.meta.video_keys) > 0:
streaming_enc = self._build_streaming_encoder( streaming_enc = self._build_streaming_encoder(
self.meta.fps, self.meta.fps,
camera_encoder, rgb_encoder,
depth_encoder,
encoder_queue_maxsize, encoder_queue_maxsize,
encoder_threads, encoder_threads,
) )
self.writer = DatasetWriter( self.writer = DatasetWriter(
meta=self.meta, meta=self.meta,
root=self.root, root=self.root,
camera_encoder=camera_encoder, rgb_encoder=rgb_encoder,
depth_encoder=depth_encoder,
encoder_threads=encoder_threads, encoder_threads=encoder_threads,
batch_encoding_size=batch_encoding_size, batch_encoding_size=batch_encoding_size,
streaming_encoder=streaming_enc, streaming_encoder=streaming_enc,
@@ -314,19 +323,22 @@ class LeRobotDataset(torch.utils.data.Dataset):
delta_timestamps=self.delta_timestamps, delta_timestamps=self.delta_timestamps,
image_transforms=self.image_transforms, image_transforms=self.image_transforms,
return_uint8=self._return_uint8, return_uint8=self._return_uint8,
depth_output_unit=self._depth_output_unit,
) )
return self.reader return self.reader
@staticmethod @staticmethod
def _build_streaming_encoder( def _build_streaming_encoder(
fps: int, fps: int,
camera_encoder: VideoEncoderConfig | None, rgb_encoder: RGBEncoderConfig | None,
depth_encoder: DepthEncoderConfig | None,
encoder_queue_maxsize: int, encoder_queue_maxsize: int,
encoder_threads: int | None, encoder_threads: int | None,
) -> StreamingVideoEncoder: ) -> StreamingVideoEncoder:
return StreamingVideoEncoder( return StreamingVideoEncoder(
fps=fps, fps=fps,
camera_encoder=camera_encoder, rgb_encoder=rgb_encoder,
depth_encoder=depth_encoder,
queue_maxsize=encoder_queue_maxsize, queue_maxsize=encoder_queue_maxsize,
encoder_threads=encoder_threads, encoder_threads=encoder_threads,
) )
@@ -655,7 +667,8 @@ class LeRobotDataset(torch.utils.data.Dataset):
image_writer_threads: int = 0, image_writer_threads: int = 0,
video_backend: str | None = None, video_backend: str | None = None,
batch_encoding_size: int = 1, batch_encoding_size: int = 1,
camera_encoder: VideoEncoderConfig | None = None, rgb_encoder: RGBEncoderConfig | None = None,
depth_encoder: DepthEncoderConfig | None = None,
metadata_buffer_size: int = 10, metadata_buffer_size: int = 10,
streaming_encoding: bool = False, streaming_encoding: bool = False,
encoder_queue_maxsize: int = 30, encoder_queue_maxsize: int = 30,
@@ -686,8 +699,10 @@ class LeRobotDataset(torch.utils.data.Dataset):
video_backend: Video decoding backend (used when reading back). video_backend: Video decoding backend (used when reading back).
batch_encoding_size: Number of episodes to accumulate before batch_encoding_size: Number of episodes to accumulate before
batch-encoding videos. ``1`` means encode immediately. batch-encoding videos. ``1`` means encode immediately.
camera_encoder: Video encoder settings for cameras (codec, quality, etc.). rgb_encoder: Video encoder settings for cameras (codec, quality, etc.).
When ``None``, :func:`~lerobot.configs.video.camera_encoder_defaults` is used. When ``None``, :func:`~lerobot.configs.video.rgb_encoder_defaults` is used.
depth_encoder: Video encoder settings for depth cameras (codec, quality, etc.).
When ``None``, :func:`~lerobot.configs.video.depth_encoder_defaults` is used.
encoder_threads: Number of encoder threads (global). ``None`` encoder_threads: Number of encoder threads (global). ``None``
lets the codec decide. lets the codec decide.
metadata_buffer_size: Number of episode metadata records to buffer metadata_buffer_size: Number of episode metadata records to buffer
@@ -722,6 +737,7 @@ class LeRobotDataset(torch.utils.data.Dataset):
obj.episodes = None obj.episodes = None
obj._video_backend = video_backend if video_backend is not None else get_safe_default_video_backend() obj._video_backend = video_backend if video_backend is not None else get_safe_default_video_backend()
obj._return_uint8 = False obj._return_uint8 = False
obj._depth_output_unit = DEFAULT_DEPTH_UNIT
obj._batch_encoding_size = batch_encoding_size obj._batch_encoding_size = batch_encoding_size
obj._encoder_threads = encoder_threads obj._encoder_threads = encoder_threads
@@ -731,12 +747,13 @@ class LeRobotDataset(torch.utils.data.Dataset):
streaming_enc = None streaming_enc = None
if streaming_encoding and len(obj.meta.video_keys) > 0: if streaming_encoding and len(obj.meta.video_keys) > 0:
streaming_enc = cls._build_streaming_encoder( streaming_enc = cls._build_streaming_encoder(
fps, camera_encoder, encoder_queue_maxsize, encoder_threads fps, rgb_encoder, depth_encoder, encoder_queue_maxsize, encoder_threads
) )
obj.writer = DatasetWriter( obj.writer = DatasetWriter(
meta=obj.meta, meta=obj.meta,
root=obj.root, root=obj.root,
camera_encoder=camera_encoder, rgb_encoder=rgb_encoder,
depth_encoder=depth_encoder,
encoder_threads=encoder_threads, encoder_threads=encoder_threads,
batch_encoding_size=batch_encoding_size, batch_encoding_size=batch_encoding_size,
streaming_encoder=streaming_enc, streaming_encoder=streaming_enc,
@@ -759,7 +776,8 @@ class LeRobotDataset(torch.utils.data.Dataset):
force_cache_sync: bool = False, force_cache_sync: bool = False,
video_backend: str | None = None, video_backend: str | None = None,
batch_encoding_size: int = 1, batch_encoding_size: int = 1,
camera_encoder: VideoEncoderConfig | None = None, rgb_encoder: RGBEncoderConfig | None = None,
depth_encoder: DepthEncoderConfig | None = None,
encoder_threads: int | None = None, encoder_threads: int | None = None,
image_writer_processes: int = 0, image_writer_processes: int = 0,
image_writer_threads: int = 0, image_writer_threads: int = 0,
@@ -787,8 +805,10 @@ class LeRobotDataset(torch.utils.data.Dataset):
video_backend: Video decoding backend for reading back data. video_backend: Video decoding backend for reading back data.
batch_encoding_size: Number of episodes to accumulate before batch_encoding_size: Number of episodes to accumulate before
batch-encoding videos. batch-encoding videos.
camera_encoder: Video encoder settings for cameras (codec, quality, etc.). rgb_encoder: Video encoder settings for cameras (codec, quality, etc.).
When ``None``, :func:`~lerobot.configs.video.camera_encoder_defaults` is used. When ``None``, :func:`~lerobot.configs.video.rgb_encoder_defaults` is used.
depth_encoder: Video encoder settings for depth cameras (codec, quality, etc.).
When ``None``, :func:`~lerobot.configs.video.depth_encoder_defaults` is used.
encoder_threads: Number of encoder threads (global). ``None`` encoder_threads: Number of encoder threads (global). ``None``
lets the codec decide. lets the codec decide.
image_writer_processes: Subprocesses for async image writing. image_writer_processes: Subprocesses for async image writing.
@@ -816,6 +836,7 @@ class LeRobotDataset(torch.utils.data.Dataset):
obj.episodes = None obj.episodes = None
obj._video_backend = video_backend if video_backend else get_safe_default_video_backend() obj._video_backend = video_backend if video_backend else get_safe_default_video_backend()
obj._return_uint8 = False obj._return_uint8 = False
obj._depth_output_unit = DEFAULT_DEPTH_UNIT
obj._batch_encoding_size = batch_encoding_size obj._batch_encoding_size = batch_encoding_size
if obj._requested_root is not None: if obj._requested_root is not None:
@@ -835,12 +856,13 @@ class LeRobotDataset(torch.utils.data.Dataset):
streaming_enc = None streaming_enc = None
if streaming_encoding and len(obj.meta.video_keys) > 0: if streaming_encoding and len(obj.meta.video_keys) > 0:
streaming_enc = cls._build_streaming_encoder( streaming_enc = cls._build_streaming_encoder(
obj.meta.fps, camera_encoder, encoder_queue_maxsize, encoder_threads obj.meta.fps, rgb_encoder, depth_encoder, encoder_queue_maxsize, encoder_threads
) )
obj.writer = DatasetWriter( obj.writer = DatasetWriter(
meta=obj.meta, meta=obj.meta,
root=obj.root, root=obj.root,
camera_encoder=camera_encoder, rgb_encoder=rgb_encoder,
depth_encoder=depth_encoder,
encoder_threads=encoder_threads, encoder_threads=encoder_threads,
batch_encoding_size=batch_encoding_size, batch_encoding_size=batch_encoding_size,
streaming_encoder=streaming_enc, streaming_encoder=streaming_enc,
+49 -2
View File
@@ -24,6 +24,7 @@ import logging
from typing import Any from typing import Any
import av import av
import numpy as np
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -31,6 +32,34 @@ FFMPEG_NUMERIC_OPTION_TYPES = ("INT", "INT64", "UINT64", "FLOAT", "DOUBLE")
FFMPEG_INTEGER_OPTION_TYPES = ("INT", "INT64", "UINT64") FFMPEG_INTEGER_OPTION_TYPES = ("INT", "INT64", "UINT64")
def write_u16_plane(plane: av.video.plane.VideoPlane, src: np.ndarray, fill_value: int | None = None) -> None:
"""Copy a 2D ``uint16`` image into the plane's memory buffer, row by row.
For speed, each row is padded to a wider size than ``width``, so the true row width in
memory is ``plane.line_size`` (bytes), not ``width``. Copying as one straight stream
would skew the image, so we write only the first ``width`` columns of each row and
leave the padding untouched.
Args:
plane: Destination 16-bit plane.
src: Source image, shape ``(height, width)``, dtype ``uint16``.
fill_value: If given, every pixel (padding included) is set to this first, so the
padding holds clean data instead of garbage.
"""
height, width = src.shape
stride_u16 = plane.line_size // np.dtype(np.uint16).itemsize
dst = np.frombuffer(plane, dtype=np.uint16).reshape(height, stride_u16)
if fill_value is not None:
dst.fill(fill_value)
dst[:, :width] = src
@functools.cache
def get_pix_fmt_channels(pix_fmt: str) -> int:
"""Return the number of components (channels) for *pix_fmt*."""
return len(av.VideoFormat(pix_fmt).components)
@functools.cache @functools.cache
def get_codec(vcodec: str) -> av.codec.Codec | None: def get_codec(vcodec: str) -> av.codec.Codec | None:
"""PyAV write-mode ``Codec`` for *vcodec*, or ``None`` if unavailable.""" """PyAV write-mode ``Codec`` for *vcodec*, or ``None`` if unavailable."""
@@ -92,7 +121,7 @@ def _check_option_value(vcodec: str, label: str, value: Any, opt: av.option.Opti
f"{label}={value!r} is not numeric; codec {vcodec!r} expects a number for this option." f"{label}={value!r} is not numeric; codec {vcodec!r} expects a number for this option."
) from e ) from e
elif isinstance(value, (float, int)): elif isinstance(value, (float, int)):
num_val = value num_val = float(value)
else: else:
raise ValueError( raise ValueError(
f"{label}={value!r} is not numeric; codec {vcodec!r} expects a number for this option." f"{label}={value!r} is not numeric; codec {vcodec!r} expects a number for this option."
@@ -142,6 +171,16 @@ def _check_pixel_format(vcodec: str, pix_fmt: str) -> None:
) )
def _check_pix_fmt_channels(pix_fmt: str, channels: int) -> None:
"""Ensure *pix_fmt* can carry at least *channels* components."""
pix_fmt_channels = get_pix_fmt_channels(pix_fmt)
if pix_fmt_channels < channels:
raise ValueError(
f"pix_fmt={pix_fmt!r} carries only {pix_fmt_channels} component(s) "
f"but the source data has {channels} channel(s)."
)
def _check_codec_options(vcodec: str, codec_options: dict[str, Any]) -> None: def _check_codec_options(vcodec: str, codec_options: dict[str, Any]) -> None:
"""Validate merged encoder options (typed) against the codec's published AVOptions.""" """Validate merged encoder options (typed) against the codec's published AVOptions."""
supported_options = _get_codec_options_by_name(vcodec) supported_options = _get_codec_options_by_name(vcodec)
@@ -156,12 +195,18 @@ def _check_codec_options(vcodec: str, codec_options: dict[str, Any]) -> None:
_check_option_value(vcodec, key, value, supported_options[key]) _check_option_value(vcodec, key, value, supported_options[key])
def check_video_encoder_parameters_pyav(vcodec: str, pix_fmt: str, codec_options: dict[str, Any]) -> None: def check_video_encoder_parameters_pyav(
vcodec: str,
pix_fmt: str,
codec_options: dict[str, Any],
channels: int | None = None,
) -> None:
"""Verify *config* is compatible with the bundled FFmpeg build. """Verify *config* is compatible with the bundled FFmpeg build.
Checks pixel format, abstract tuning-field compatibility, and each merged Checks pixel format, abstract tuning-field compatibility, and each merged
encoder option from :meth:`~lerobot.configs.video.VideoEncoderConfig.get_codec_options` encoder option from :meth:`~lerobot.configs.video.VideoEncoderConfig.get_codec_options`
against PyAV (including numeric ``extra_options`` present in that dict). against PyAV (including numeric ``extra_options`` present in that dict).
When given, additionally verify that *pix_fmt* carries as many components as the source data channels.
No-op when ``config.vcodec`` isn't in the local FFmpeg build. No-op when ``config.vcodec`` isn't in the local FFmpeg build.
Raises: Raises:
@@ -171,4 +216,6 @@ def check_video_encoder_parameters_pyav(vcodec: str, pix_fmt: str, codec_options
if not options: if not options:
raise ValueError(f"Codec {vcodec!r} is not available in the bundled FFmpeg build") raise ValueError(f"Codec {vcodec!r} is not available in the bundled FFmpeg build")
_check_pixel_format(vcodec, pix_fmt) _check_pixel_format(vcodec, pix_fmt)
if channels is not None:
_check_pix_fmt_channels(pix_fmt, channels)
_check_codec_options(vcodec, codec_options) _check_codec_options(vcodec, codec_options)
+40 -7
View File
@@ -22,9 +22,11 @@ import numpy as np
import torch import torch
from datasets import load_dataset from datasets import load_dataset
from lerobot.configs import DEFAULT_DEPTH_UNIT, DepthEncoderConfig
from lerobot.utils.constants import HF_LEROBOT_HOME, LOOKAHEAD_BACKTRACKTABLE, LOOKBACK_BACKTRACKTABLE from lerobot.utils.constants import HF_LEROBOT_HOME, LOOKAHEAD_BACKTRACKTABLE, LOOKBACK_BACKTRACKTABLE
from .dataset_metadata import CODEBASE_VERSION, LeRobotDatasetMetadata from .dataset_metadata import CODEBASE_VERSION, LeRobotDatasetMetadata
from .depth_utils import dequantize_depth
from .feature_utils import get_delta_indices from .feature_utils import get_delta_indices
from .io_utils import item_to_torch from .io_utils import item_to_torch
from .utils import ( from .utils import (
@@ -35,6 +37,7 @@ from .utils import (
) )
from .video_utils import ( from .video_utils import (
VideoDecoderCache, VideoDecoderCache,
decode_video_frames,
decode_video_frames_torchcodec, decode_video_frames_torchcodec,
) )
@@ -252,6 +255,7 @@ class StreamingLeRobotDataset(torch.utils.data.IterableDataset):
rng: np.random.Generator | None = None, rng: np.random.Generator | None = None,
shuffle: bool = True, shuffle: bool = True,
return_uint8: bool = False, return_uint8: bool = False,
depth_output_unit: str = DEFAULT_DEPTH_UNIT,
): ):
"""Initialize a StreamingLeRobotDataset. """Initialize a StreamingLeRobotDataset.
@@ -272,6 +276,8 @@ class StreamingLeRobotDataset(torch.utils.data.IterableDataset):
seed (int, optional): Reproducibility random seed. seed (int, optional): Reproducibility random seed.
rng (np.random.Generator | None, optional): Random number generator. rng (np.random.Generator | None, optional): Random number generator.
shuffle (bool, optional): Whether to shuffle the dataset across exhaustions. Defaults to True. shuffle (bool, optional): Whether to shuffle the dataset across exhaustions. Defaults to True.
depth_output_unit (str, optional): Physical unit depth maps are dequantized to ("m" or "mm").
Defaults to "mm".
""" """
super().__init__() super().__init__()
self.repo_id = repo_id self.repo_id = repo_id
@@ -290,6 +296,7 @@ class StreamingLeRobotDataset(torch.utils.data.IterableDataset):
self.streaming = streaming self.streaming = streaming
self.buffer_size = buffer_size self.buffer_size = buffer_size
self._return_uint8 = return_uint8 self._return_uint8 = return_uint8
self._depth_output_unit = depth_output_unit
# We cache the video decoders to avoid re-initializing them at each frame (avoiding a ~10x slowdown) # We cache the video decoders to avoid re-initializing them at each frame (avoiding a ~10x slowdown)
self.video_decoder_cache = None self.video_decoder_cache = None
@@ -306,6 +313,11 @@ class StreamingLeRobotDataset(torch.utils.data.IterableDataset):
# Check version # Check version
check_version_compatibility(self.repo_id, self.meta._version, CODEBASE_VERSION) check_version_compatibility(self.repo_id, self.meta._version, CODEBASE_VERSION)
self._depth_encoder_configs: dict[str, DepthEncoderConfig] = {
vid_key: DepthEncoderConfig.from_video_info(self.meta.features[vid_key].get("info"))
for vid_key in self.meta.depth_keys
}
self.delta_timestamps = None self.delta_timestamps = None
self.delta_indices = None self.delta_indices = None
@@ -554,13 +566,34 @@ class StreamingLeRobotDataset(torch.utils.data.IterableDataset):
for video_key, query_ts in query_timestamps.items(): for video_key, query_ts in query_timestamps.items():
root = self.meta.url_root if self.streaming and not self.streaming_from_local else self.root root = self.meta.url_root if self.streaming and not self.streaming_from_local else self.root
video_path = f"{root}/{self.meta.get_video_file_path(ep_idx, video_key)}" video_path = f"{root}/{self.meta.get_video_file_path(ep_idx, video_key)}"
frames = decode_video_frames_torchcodec( if video_key in self.meta.depth_keys:
video_path, # Depth maps are 12-bit quantized and only decodable via pyav; dequantize back
query_ts, # to physical units to match the non-streaming reader.
self.tolerance_s, frames = decode_video_frames(
decoder_cache=self.video_decoder_cache, video_path,
return_uint8=self._return_uint8, query_ts,
) self.tolerance_s,
backend="pyav",
return_uint8=False,
is_depth=True,
)
depth_encoder = self._depth_encoder_configs[video_key]
frames = dequantize_depth(
frames,
depth_min=depth_encoder.depth_min,
depth_max=depth_encoder.depth_max,
shift=depth_encoder.shift,
use_log=depth_encoder.use_log,
output_unit=self._depth_output_unit,
)
else:
frames = decode_video_frames_torchcodec(
video_path,
query_ts,
self.tolerance_s,
decoder_cache=self.video_decoder_cache,
return_uint8=self._return_uint8,
)
item[video_key] = frames.squeeze(0) if len(query_ts) == 1 else frames item[video_key] = frames.squeeze(0) if len(query_ts) == 1 else frames
+4 -1
View File
@@ -87,11 +87,14 @@ DATA_DIR = "data"
VIDEO_DIR = "videos" VIDEO_DIR = "videos"
CHUNK_FILE_PATTERN = "chunk-{chunk_index:03d}/file-{file_index:03d}" CHUNK_FILE_PATTERN = "chunk-{chunk_index:03d}/file-{file_index:03d}"
IMAGE_FILE_PATTERN = "frame-{frame_index:06d}.png"
DEPTH_FILE_PATTERN = "frame-{frame_index:06d}.tiff"
DEFAULT_TASKS_PATH = "meta/tasks.parquet" DEFAULT_TASKS_PATH = "meta/tasks.parquet"
DEFAULT_EPISODES_PATH = EPISODES_DIR + "/" + CHUNK_FILE_PATTERN + ".parquet" DEFAULT_EPISODES_PATH = EPISODES_DIR + "/" + CHUNK_FILE_PATTERN + ".parquet"
DEFAULT_DATA_PATH = DATA_DIR + "/" + CHUNK_FILE_PATTERN + ".parquet" DEFAULT_DATA_PATH = DATA_DIR + "/" + CHUNK_FILE_PATTERN + ".parquet"
DEFAULT_VIDEO_PATH = VIDEO_DIR + "/{video_key}/" + CHUNK_FILE_PATTERN + ".mp4" DEFAULT_VIDEO_PATH = VIDEO_DIR + "/{video_key}/" + CHUNK_FILE_PATTERN + ".mp4"
DEFAULT_IMAGE_PATH = "images/{image_key}/episode-{episode_index:06d}/frame-{frame_index:06d}.png" DEFAULT_IMAGE_PATH = "images/{image_key}/episode-{episode_index:06d}/" + IMAGE_FILE_PATTERN
DEFAULT_DEPTH_PATH = "images/{image_key}/episode-{episode_index:06d}/" + DEPTH_FILE_PATTERN
LEGACY_EPISODES_PATH = "meta/episodes.jsonl" LEGACY_EPISODES_PATH = "meta/episodes.jsonl"
LEGACY_EPISODES_STATS_PATH = "meta/episodes_stats.jsonl" LEGACY_EPISODES_STATS_PATH = "meta/episodes_stats.jsonl"
+163 -76
View File
@@ -39,11 +39,17 @@ from datasets.features.features import register_feature
from PIL import Image from PIL import Image
from lerobot.configs import ( from lerobot.configs import (
DepthEncoderConfig,
RGBEncoderConfig,
VideoEncoderConfig, VideoEncoderConfig,
camera_encoder_defaults, depth_encoder_defaults,
rgb_encoder_defaults,
) )
from lerobot.utils.import_utils import get_safe_default_video_backend from lerobot.utils.import_utils import get_safe_default_video_backend
from .depth_utils import quantize_depth
from .pyav_utils import get_pix_fmt_channels
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -53,6 +59,7 @@ def decode_video_frames(
tolerance_s: float, tolerance_s: float,
backend: str | None = None, backend: str | None = None,
return_uint8: bool = False, return_uint8: bool = False,
is_depth: bool = False,
) -> torch.Tensor: ) -> torch.Tensor:
""" """
Decodes video frames using the specified backend. Decodes video frames using the specified backend.
@@ -64,23 +71,35 @@ def decode_video_frames(
backend (str, optional): Backend to use for decoding. Defaults to "torchcodec" when available backend (str, optional): Backend to use for decoding. Defaults to "torchcodec" when available
in the platform; otherwise, defaults to "pyav". The legacy value "video_reader" is in the platform; otherwise, defaults to "pyav". The legacy value "video_reader" is
accepted for one release as an alias for "pyav" and will be removed in a future version. accepted for one release as an alias for "pyav" and will be removed in a future version.
return_uint8 (bool): If True, return raw uint8 frames without float32 normalization. return_uint8 (bool): For RGB videos, if True return raw uint8 frames without float32 normalization.
This reduces memory for DataLoader IPC; normalization can be done on GPU afterward. This reduces memory for DataLoader IPC; normalization can be done on GPU afterward.
is_depth (bool): Set to True if the video is a depth map (1 channel, uint12).
Returns: Returns:
torch.Tensor: Decoded frames (float32 in [0,1] by default, or uint8 if return_uint8=True). torch.Tensor: Decoded frames (RGB: float32 in [0,1] by default, or uint8 if return_uint8=True, Depth: uint12).
Currently supports torchcodec on cpu and pyav. Currently supports torchcodec on cpu and pyav.
""" """
if backend != "pyav" and is_depth:
logger.debug("Decoding depth maps is only supported with the 'pyav' backend, falling back to pyav.")
# We do not actually return uint8 here, but we avoid the 255 normalization step.
return decode_video_frames_pyav(
video_path, timestamps, tolerance_s, return_uint8=False, is_depth=True
)
if backend is None: if backend is None:
backend = get_safe_default_video_backend() backend = get_safe_default_video_backend()
if backend == "torchcodec": if backend == "torchcodec":
return decode_video_frames_torchcodec(video_path, timestamps, tolerance_s, return_uint8=return_uint8) return decode_video_frames_torchcodec(video_path, timestamps, tolerance_s, return_uint8=return_uint8)
elif backend == "pyav": elif backend == "pyav":
return decode_video_frames_pyav(video_path, timestamps, tolerance_s, return_uint8=return_uint8) return decode_video_frames_pyav(
video_path, timestamps, tolerance_s, return_uint8=return_uint8, is_depth=is_depth
)
elif backend == "video_reader": elif backend == "video_reader":
logger.warning("backend='video_reader' is deprecated and now aliases to 'pyav'.") logger.warning("backend='video_reader' is deprecated and now aliases to 'pyav'.")
return decode_video_frames_pyav(video_path, timestamps, tolerance_s, return_uint8=return_uint8) return decode_video_frames_pyav(
video_path, timestamps, tolerance_s, return_uint8=return_uint8, is_depth=is_depth
)
else: else:
raise ValueError(f"Unsupported video backend: {backend}") raise ValueError(f"Unsupported video backend: {backend}")
@@ -91,6 +110,7 @@ def decode_video_frames_pyav(
tolerance_s: float, tolerance_s: float,
log_loaded_timestamps: bool = False, log_loaded_timestamps: bool = False,
return_uint8: bool = False, return_uint8: bool = False,
is_depth: bool = False,
) -> torch.Tensor: ) -> torch.Tensor:
"""Loads frames associated to the requested timestamps of a video using PyAV. """Loads frames associated to the requested timestamps of a video using PyAV.
@@ -109,8 +129,9 @@ def decode_video_frames_pyav(
tolerance_s: Allowed deviation in seconds between a queried timestamp and the closest tolerance_s: Allowed deviation in seconds between a queried timestamp and the closest
decoded frame. decoded frame.
log_loaded_timestamps: When True, log every decoded frame's timestamp at INFO level. log_loaded_timestamps: When True, log every decoded frame's timestamp at INFO level.
return_uint8: When True, return raw uint8 frames (C, H, W). Otherwise, return float32 in return_uint8: For RGB videos, if True return raw uint8 frames (C, H, W).
[0, 1] range. Otherwise, return float32 in [0, 1] range.
is_depth: Set to True if the video is a depth map (1 channel, uint12).
Returns: Returns:
torch.Tensor of shape (len(timestamps), C, H, W). torch.Tensor of shape (len(timestamps), C, H, W).
@@ -132,7 +153,13 @@ def decode_video_frames_pyav(
# https://pyav.basswood-io.com/docs/stable/api/container.html#av.container.InputContainer.seek # https://pyav.basswood-io.com/docs/stable/api/container.html#av.container.InputContainer.seek
with av.open(video_path) as container: with av.open(video_path) as container:
stream = container.streams.video[0] stream = container.streams.video[0]
container.seek(int(first_ts * av.time_base), backward=True) # Seek to the nearest keyframe at or before `first_ts` with a 1 frame margin
container.seek(
round(first_ts / stream.time_base) - 1,
backward=True,
any_frame=False,
stream=stream,
)
for frame in container.decode(stream): for frame in container.decode(stream):
if frame.pts is None: if frame.pts is None:
@@ -140,9 +167,13 @@ def decode_video_frames_pyav(
current_ts = float(frame.pts * stream.time_base) current_ts = float(frame.pts * stream.time_base)
if log_loaded_timestamps: if log_loaded_timestamps:
logger.info(f"frame loaded at timestamp={current_ts:.4f}") logger.info(f"frame loaded at timestamp={current_ts:.4f}")
# Convert to CHW uint8 to match torchcodec's output layout. if is_depth:
arr = frame.to_ndarray(format="rgb24") # H, W, 3 arr = frame.to_ndarray(format="gray12le") # (H, W) uint12
loaded_frames.append(torch.from_numpy(arr).permute(2, 0, 1).contiguous()) loaded_frames.append(torch.from_numpy(arr).unsqueeze(0).contiguous())
else:
arr = frame.to_ndarray(format="rgb24") # (H, W, 3)
# Convert to CHW uint8 to match torchcodec's output layout.
loaded_frames.append(torch.from_numpy(arr).permute(2, 0, 1).contiguous())
loaded_ts.append(current_ts) loaded_ts.append(current_ts)
if current_ts >= last_ts: if current_ts >= last_ts:
break break
@@ -185,7 +216,7 @@ def decode_video_frames_pyav(
f"number of queried timestamps ({len(timestamps)})" f"number of queried timestamps ({len(timestamps)})"
) )
if return_uint8: if return_uint8 or is_depth:
return closest_frames return closest_frames
# convert to the pytorch format which is float32 in [0,1] range (and channel first) # convert to the pytorch format which is float32 in [0,1] range (and channel first)
@@ -406,17 +437,38 @@ def encode_video_frames(
imgs_dir: Path | str, imgs_dir: Path | str,
video_path: Path | str, video_path: Path | str,
fps: int, fps: int,
camera_encoder: VideoEncoderConfig | None = None, video_encoder: VideoEncoderConfig | None = None,
encoder_threads: int | None = None, encoder_threads: int | None = None,
*, *,
log_level: int | None = av.logging.WARNING, log_level: int | None = av.logging.WARNING,
overwrite: bool = False, overwrite: bool = False,
) -> None: ) -> None:
"""More info on ffmpeg arguments tuning on `benchmark/video/README.md`""" """Encode a directory of image frames into an MP4 video.
if camera_encoder is None:
camera_encoder = camera_encoder_defaults() When ``video_encoder`` is a :class:`~lerobot.configs.video.DepthEncoderConfig`,
vcodec = camera_encoder.vcodec frames are read from ``.tiff`` files and quantized to 12-bit depth codes using the
pix_fmt = camera_encoder.pix_fmt encoder's ``depth_min`` / ``depth_max`` / ``shift`` / ``use_log``; otherwise ``.png``
RGB frames are encoded directly.
Args:
imgs_dir: Directory containing the frames to encode, named ``frame-000000``
onwards (``.png`` for RGB, ``.tiff`` for depth).
video_path: Output path for the encoded ``.mp4`` file.
fps: Frame rate of the output video.
video_encoder: Encoder settings (codec, pixel format, quality, ...). When
``None``, :func:`rgb_encoder_defaults` is used. Pass a
:class:`~lerobot.configs.video.DepthEncoderConfig` to encode depth frames.
encoder_threads: Per-encoder thread count forwarded to the codec. ``None``
lets the codec decide.
log_level: libav log level to set while encoding, or ``None`` to leave the
current logging configuration unchanged.
overwrite: When ``False`` and ``video_path`` already exists, skip encoding and
log a warning. When ``True``, re-encode and replace the existing file.
"""
if video_encoder is None:
video_encoder = rgb_encoder_defaults()
vcodec = video_encoder.vcodec
pix_fmt = video_encoder.pix_fmt
video_path = Path(video_path) video_path = Path(video_path)
imgs_dir = Path(imgs_dir) imgs_dir = Path(imgs_dir)
@@ -428,17 +480,19 @@ def encode_video_frames(
video_path.parent.mkdir(parents=True, exist_ok=True) video_path.parent.mkdir(parents=True, exist_ok=True)
# Get input frames # Get input frames
template = "frame-" + ("[0-9]" * 6) + ".png" is_depth = isinstance(video_encoder, DepthEncoderConfig)
suffix = ".png" if not is_depth else ".tiff"
template = "frame-" + ("[0-9]" * 6) + suffix
input_list = sorted( input_list = sorted(
glob.glob(str(imgs_dir / template)), key=lambda x: int(x.split("-")[-1].split(".")[0]) glob.glob(str(imgs_dir / template)), key=lambda x: int(x.split("-")[-1].split(".")[0])
) )
if len(input_list) == 0: if len(input_list) == 0:
raise FileNotFoundError(f"No images found in {imgs_dir}.") raise FileNotFoundError(f"No images with suffix {suffix} found in {imgs_dir}.")
with Image.open(input_list[0]) as dummy_image: with Image.open(input_list[0]) as dummy_image:
width, height = dummy_image.size width, height = dummy_image.size
video_options = camera_encoder.get_codec_options(encoder_threads, as_strings=True) video_options = video_encoder.get_codec_options(encoder_threads, as_strings=True)
# Set logging level # Set logging level
if log_level is not None: if log_level is not None:
@@ -455,8 +509,19 @@ def encode_video_frames(
# Loop through input frames and encode them # Loop through input frames and encode them
for input_data in input_list: for input_data in input_list:
with Image.open(input_data) as input_image: with Image.open(input_data) as input_image:
input_image = input_image.convert("RGB") if is_depth:
input_frame = av.VideoFrame.from_image(input_image) input_frame = quantize_depth(
np.array(input_image),
depth_min=video_encoder.depth_min,
depth_max=video_encoder.depth_max,
shift=video_encoder.shift,
use_log=video_encoder.use_log,
pix_fmt=video_encoder.pix_fmt,
video_backend="pyav",
)
else:
input_image = input_image.convert("RGB")
input_frame = av.VideoFrame.from_image(input_image)
packet = output_stream.encode(input_frame) packet = output_stream.encode(input_frame)
if packet: if packet:
output.mux(packet) output.mux(packet)
@@ -477,7 +542,7 @@ def encode_video_frames(
def reencode_video( def reencode_video(
input_video_path: Path | str, input_video_path: Path | str,
output_video_path: Path | str, output_video_path: Path | str,
camera_encoder: VideoEncoderConfig | None = None, video_encoder: VideoEncoderConfig | None = None,
encoder_threads: int | None = None, encoder_threads: int | None = None,
log_level: int | None = av.logging.WARNING, log_level: int | None = av.logging.WARNING,
overwrite: bool = False, overwrite: bool = False,
@@ -489,7 +554,7 @@ def reencode_video(
Args: Args:
input_video_path: Existing video file to read. input_video_path: Existing video file to read.
output_video_path: Path for the re-encoded file. output_video_path: Path for the re-encoded file.
camera_encoder: Encoder configuration. Defaults to :func:`camera_encoder_defaults`. video_encoder: Encoder configuration. Defaults to :func:`rgb_encoder_defaults`.
encoder_threads: Optional thread count forwarded to :meth:`VideoEncoderConfig.get_codec_options`. encoder_threads: Optional thread count forwarded to :meth:`VideoEncoderConfig.get_codec_options`.
log_level: libav log level while encoding, or ``None`` to leave logging unchanged. Defaults to WARNING. log_level: libav log level while encoding, or ``None`` to leave logging unchanged. Defaults to WARNING.
overwrite: When ``False`` and ``output_video_path`` already exists, skip and log a warning. overwrite: When ``False`` and ``output_video_path`` already exists, skip and log a warning.
@@ -497,7 +562,7 @@ def reencode_video(
end_time_s: When set, trim the output to end at this timestamp (seconds, exclusive). end_time_s: When set, trim the output to end at this timestamp (seconds, exclusive).
""" """
camera_encoder = camera_encoder or camera_encoder_defaults() video_encoder = video_encoder or rgb_encoder_defaults()
if (start_time_s is not None and start_time_s < 0) or (end_time_s is not None and end_time_s < 0): if (start_time_s is not None and start_time_s < 0) or (end_time_s is not None and end_time_s < 0):
raise ValueError(f"Trim times must be non-negative, got start={start_time_s}, end={end_time_s}.") raise ValueError(f"Trim times must be non-negative, got start={start_time_s}, end={end_time_s}.")
@@ -512,9 +577,9 @@ def reencode_video(
output_video_path.parent.mkdir(parents=True, exist_ok=True) output_video_path.parent.mkdir(parents=True, exist_ok=True)
video_options = camera_encoder.get_codec_options(encoder_threads, as_strings=True) video_options = video_encoder.get_codec_options(encoder_threads, as_strings=True)
vcodec = camera_encoder.vcodec vcodec = video_encoder.vcodec
pix_fmt = camera_encoder.pix_fmt pix_fmt = video_encoder.pix_fmt
with tempfile.NamedTemporaryFile(suffix=".mp4", delete=False) as tmp_named_file: with tempfile.NamedTemporaryFile(suffix=".mp4", delete=False) as tmp_named_file:
tmp_output_video_path = tmp_named_file.name tmp_output_video_path = tmp_named_file.name
@@ -696,22 +761,21 @@ class _CameraEncoderThread(threading.Thread):
self, self,
video_path: Path, video_path: Path,
fps: int, fps: int,
vcodec: str, video_encoder: VideoEncoderConfig,
pix_fmt: str,
codec_options: dict[str, str],
frame_queue: queue.Queue, frame_queue: queue.Queue,
result_queue: queue.Queue, result_queue: queue.Queue,
stop_event: threading.Event, stop_event: threading.Event,
encoder_threads: int | None = None,
): ):
super().__init__(daemon=True) super().__init__(daemon=True)
self.video_path = video_path self.video_path = video_path
self.fps = fps self.fps = fps
self.vcodec = vcodec self.video_encoder = video_encoder
self.pix_fmt = pix_fmt self.is_depth = isinstance(video_encoder, DepthEncoderConfig)
self.codec_options = codec_options
self.frame_queue = frame_queue self.frame_queue = frame_queue
self.result_queue = result_queue self.result_queue = result_queue
self.stop_event = stop_event self.stop_event = stop_event
self.encoder_threads = encoder_threads
def run(self) -> None: def run(self) -> None:
from .compute_stats import RunningQuantileStats, auto_downsample_height_width from .compute_stats import RunningQuantileStats, auto_downsample_height_width
@@ -736,12 +800,12 @@ class _CameraEncoderThread(threading.Thread):
# Sentinel: flush and close # Sentinel: flush and close
break break
# Ensure HWC uint8 numpy array # Ensure HWC (RGB or depth) uint8 (RGB only) numpy array
if isinstance(frame_data, np.ndarray): if isinstance(frame_data, np.ndarray):
if frame_data.ndim == 3 and frame_data.shape[0] == 3: if frame_data.ndim == 3 and frame_data.shape[0] in (1, 3):
# CHW -> HWC # CHW -> HWC
frame_data = frame_data.transpose(1, 2, 0) frame_data = frame_data.transpose(1, 2, 0)
if frame_data.dtype != np.uint8: if not self.is_depth and frame_data.dtype != np.uint8:
frame_data = (frame_data * 255).astype(np.uint8) frame_data = (frame_data * 255).astype(np.uint8)
# Open container on first frame (to get width/height) # Open container on first frame (to get width/height)
@@ -749,15 +813,29 @@ class _CameraEncoderThread(threading.Thread):
height, width = frame_data.shape[:2] height, width = frame_data.shape[:2]
Path(self.video_path).parent.mkdir(parents=True, exist_ok=True) Path(self.video_path).parent.mkdir(parents=True, exist_ok=True)
container = av.open(str(self.video_path), "w") container = av.open(str(self.video_path), "w")
output_stream = container.add_stream(self.vcodec, self.fps, options=self.codec_options) output_stream = container.add_stream(
output_stream.pix_fmt = self.pix_fmt self.video_encoder.vcodec,
self.fps,
options=self.video_encoder.get_codec_options(self.encoder_threads, as_strings=True),
)
output_stream.pix_fmt = self.video_encoder.pix_fmt
output_stream.width = width output_stream.width = width
output_stream.height = height output_stream.height = height
output_stream.time_base = Fraction(1, self.fps) output_stream.time_base = Fraction(1, self.fps)
# Encode frame with explicit timestamps # Encode frame with explicit timestamps
pil_img = Image.fromarray(frame_data) if not self.is_depth:
video_frame = av.VideoFrame.from_image(pil_img) pil_img = Image.fromarray(frame_data)
video_frame = av.VideoFrame.from_image(pil_img)
else:
video_frame = quantize_depth(
frame_data,
depth_min=self.video_encoder.depth_min,
depth_max=self.video_encoder.depth_max,
shift=self.video_encoder.shift,
use_log=self.video_encoder.use_log,
video_backend=self.video_encoder.video_backend,
)
video_frame.pts = frame_count video_frame.pts = frame_count
video_frame.time_base = Fraction(1, self.fps) video_frame.time_base = Fraction(1, self.fps)
packet = output_stream.encode(video_frame) packet = output_stream.encode(video_frame)
@@ -815,22 +893,27 @@ class StreamingVideoEncoder:
def __init__( def __init__(
self, self,
fps: int, fps: int,
camera_encoder: VideoEncoderConfig | None = None, rgb_encoder: RGBEncoderConfig | None = None,
depth_encoder: DepthEncoderConfig | None = None,
queue_maxsize: int = 30, queue_maxsize: int = 30,
encoder_threads: int | None = None, encoder_threads: int | None = None,
): ):
""" """
Args: Args:
fps: Frames per second for the output videos. fps: Frames per second for the output videos.
camera_encoder: Video encoder settings applied to all cameras. rgb_encoder: Video encoder settings applied to all RGB cameras.
When ``None``, :func:`camera_encoder_defaults` is used. When ``None``, :func:`rgb_encoder_defaults` is used.
encoder_threads: Number of encoder threads (global setting). depth_encoder: Video encoder settings applied to all depth cameras,
``None`` lets the codec decide. including the depth quantization parameters. When ``None``,
:func:`depth_encoder_defaults` is used.
queue_maxsize: Max frames to buffer per camera before queue_maxsize: Max frames to buffer per camera before
back-pressure drops frames. back-pressure drops frames.
encoder_threads: Number of encoder threads (global setting).
``None`` lets the codec decide.
""" """
self.fps = fps self.fps = fps
self._camera_encoder = camera_encoder or camera_encoder_defaults() self._rgb_encoder = rgb_encoder or rgb_encoder_defaults()
self._depth_encoder = depth_encoder or depth_encoder_defaults()
self._encoder_threads = encoder_threads self._encoder_threads = encoder_threads
self.queue_maxsize = queue_maxsize self.queue_maxsize = queue_maxsize
@@ -843,18 +926,25 @@ class StreamingVideoEncoder:
self._episode_active = False self._episode_active = False
self._closed = False self._closed = False
def start_episode(self, video_keys: list[str], temp_dir: Path) -> None: def start_episode(
self, video_keys: list[str], temp_dir: Path, depth_video_keys: list[str] | None = None
) -> None:
"""Start encoder threads for a new episode. """Start encoder threads for a new episode.
Args: Args:
video_keys: List of video feature keys (e.g. ["observation.images.laptop"]) video_keys: List of video feature keys (e.g. ["observation.images.laptop"])
temp_dir: Base directory for temporary MP4 files temp_dir: Base directory for temporary MP4 files
depth_video_keys: List of video or image feature keys that carry depth maps (e.g.
["observation.images.laptop_depth"]). Defaults to ``[]`` (no depth keys).
""" """
if self._episode_active: if self._episode_active:
self.cancel_episode() self.cancel_episode()
self._dropped_frames.clear() self._dropped_frames.clear()
if depth_video_keys is None:
depth_video_keys = []
for video_key in video_keys: for video_key in video_keys:
frame_queue: queue.Queue = queue.Queue(maxsize=self.queue_maxsize) frame_queue: queue.Queue = queue.Queue(maxsize=self.queue_maxsize)
result_queue: queue.Queue = queue.Queue(maxsize=1) result_queue: queue.Queue = queue.Queue(maxsize=1)
@@ -863,17 +953,15 @@ class StreamingVideoEncoder:
temp_video_dir = Path(tempfile.mkdtemp(dir=temp_dir)) temp_video_dir = Path(tempfile.mkdtemp(dir=temp_dir))
video_path = temp_video_dir / f"{video_key.replace('/', '_')}_streaming.mp4" video_path = temp_video_dir / f"{video_key.replace('/', '_')}_streaming.mp4"
vcodec = self._camera_encoder.vcodec encoder = self._depth_encoder if video_key in depth_video_keys else self._rgb_encoder
codec_options = self._camera_encoder.get_codec_options(self._encoder_threads, as_strings=True)
encoder_thread = _CameraEncoderThread( encoder_thread = _CameraEncoderThread(
video_path=video_path, video_path=video_path,
fps=self.fps, fps=self.fps,
vcodec=vcodec, video_encoder=encoder,
pix_fmt=self._camera_encoder.pix_fmt,
codec_options=codec_options,
frame_queue=frame_queue, frame_queue=frame_queue,
result_queue=result_queue, result_queue=result_queue,
stop_event=stop_event, stop_event=stop_event,
encoder_threads=self._encoder_threads,
) )
encoder_thread.start() encoder_thread.start()
@@ -1080,15 +1168,23 @@ def get_audio_info(video_path: Path | str) -> dict:
def get_video_info( def get_video_info(
video_path: Path | str, video_path: Path | str,
camera_encoder: VideoEncoderConfig | None = None, video_encoder: VideoEncoderConfig | None = None,
) -> dict: ) -> dict:
"""Build the ``video.*`` / ``audio.*`` info dict persisted in ``info.json``. """Build the ``video.*`` / ``audio.*`` info dict persisted in ``info.json``.
Args: Args:
video_path: Path to the encoded video file to probe. video_path: Path to the encoded video file to probe.
camera_encoder: If provided, record the exact encoder settings used to encode this video_encoder: If provided, record the exact encoder settings used to encode this
video. Stream-derived values take precedence encoder fields are only written for keys video. Stream-derived values take precedence encoder fields are only written for keys
not already populated from the video file itself. not already populated from the video file itself. When a
:class:`~lerobot.configs.video.DepthEncoderConfig` is passed, the depth
quantization parameters (``depth_min`` / ``depth_max`` / ``shift`` /
``use_log``) are recorded so frames can be dequantized on read.
Returns:
The ``video.*`` / ``audio.*`` info dict, including ``is_depth_map`` which is
``True`` only when ``video_encoder`` is a
:class:`~lerobot.configs.video.DepthEncoderConfig`.
""" """
logging.getLogger("libav").setLevel(av.logging.WARNING) logging.getLogger("libav").setLevel(av.logging.WARNING)
@@ -1106,13 +1202,10 @@ def get_video_info(
video_info["video.width"] = video_stream.width video_info["video.width"] = video_stream.width
video_info["video.codec"] = video_stream.codec.canonical_name video_info["video.codec"] = video_stream.codec.canonical_name
video_info["video.pix_fmt"] = video_stream.pix_fmt video_info["video.pix_fmt"] = video_stream.pix_fmt
video_info["video.is_depth_map"] = False
# Calculate fps from r_frame_rate # Calculate fps from r_frame_rate
video_info["video.fps"] = int(video_stream.base_rate) video_info["video.fps"] = int(video_stream.base_rate)
video_info["video.channels"] = get_pix_fmt_channels(video_stream.pix_fmt)
pixel_channels = get_video_pixel_channels(video_stream.pix_fmt)
video_info["video.channels"] = pixel_channels
# Reset logging level # Reset logging level
av.logging.restore_default_callback() av.logging.restore_default_callback()
@@ -1121,27 +1214,18 @@ def get_video_info(
video_info.update(**get_audio_info(video_path)) video_info.update(**get_audio_info(video_path))
# Add additional encoder configuration if provided # Add additional encoder configuration if provided
if camera_encoder is not None: if video_encoder is not None:
for field_name, field_value in asdict(camera_encoder).items(): for field_name, field_value in asdict(video_encoder).items():
# vcodec is already populated from the video stream # vcodec is already populated from the video stream
if field_name == "vcodec": if field_name == "vcodec":
continue continue
video_info.setdefault(f"video.{field_name}", field_value) video_info.setdefault(f"video.{field_name}", field_value)
video_info["is_depth_map"] = isinstance(video_encoder, DepthEncoderConfig)
return video_info return video_info
def get_video_pixel_channels(pix_fmt: str) -> int:
if "gray" in pix_fmt or "depth" in pix_fmt or "monochrome" in pix_fmt:
return 1
elif "rgba" in pix_fmt or "yuva" in pix_fmt:
return 4
elif "rgb" in pix_fmt or "yuv" in pix_fmt:
return 3
else:
raise ValueError("Unknown format")
def get_video_duration_in_s(video_path: Path | str) -> float: def get_video_duration_in_s(video_path: Path | str) -> float:
""" """
Get the duration of a video file in seconds using PyAV. Get the duration of a video file in seconds using PyAV.
@@ -1202,10 +1286,13 @@ class VideoEncodingManager:
img_dir = self.dataset.root / "images" img_dir = self.dataset.root / "images"
if img_dir.exists(): if img_dir.exists():
png_files = list(img_dir.rglob("*.png")) png_files = list(img_dir.rglob("*.png"))
if len(png_files) == 0: tiff_files = list(img_dir.rglob("*.tiff"))
if len(png_files) == 0 and len(tiff_files) == 0:
shutil.rmtree(img_dir) shutil.rmtree(img_dir)
logger.debug("Cleaned up empty images directory") logger.debug("Cleaned up empty images directory")
else: else:
logger.debug(f"Images directory is not empty, containing {len(png_files)} PNG files") logger.debug(
f"Images directory is not empty, containing {len(png_files)} PNG and {len(tiff_files)} TIFF files"
)
return False # Don't suppress the original exception return False # Don't suppress the original exception
+2 -1
View File
@@ -126,7 +126,8 @@ def prepare_observation_for_inference(
for name in observation: for name in observation:
observation[name] = torch.from_numpy(observation[name]) observation[name] = torch.from_numpy(observation[name])
if "image" in name: if "image" in name:
observation[name] = observation[name].type(torch.float32) / 255 if observation[name].dtype == torch.uint8:
observation[name] = observation[name].type(torch.float32) / 255
observation[name] = observation[name].permute(2, 0, 1).contiguous() observation[name] = observation[name].permute(2, 0, 1).contiguous()
observation[name] = observation[name].unsqueeze(0) observation[name] = observation[name].unsqueeze(0)
observation[name] = observation[name].to(device) observation[name] = observation[name].to(device)
+19 -7
View File
@@ -66,9 +66,14 @@ class HopeJrArm(Robot):
@property @property
def _cameras_ft(self) -> dict[str, tuple]: def _cameras_ft(self) -> dict[str, tuple]:
return { features: dict[str, tuple] = {}
cam: (self.config.cameras[cam].height, self.config.cameras[cam].width, 3) for cam in self.cameras for cam in self.cameras:
} cfg = self.config.cameras[cam]
if getattr(cfg, "use_rgb", True):
features[cam] = (cfg.height, cfg.width, 3)
if getattr(cfg, "use_depth", False):
features[f"{cam}_depth"] = (cfg.height, cfg.width, 1)
return features
@cached_property @cached_property
def observation_features(self) -> dict[str, type | tuple]: def observation_features(self) -> dict[str, type | tuple]:
@@ -139,10 +144,17 @@ class HopeJrArm(Robot):
# Capture images from cameras # Capture images from cameras
for cam_key, cam in self.cameras.items(): for cam_key, cam in self.cameras.items():
start = time.perf_counter() if getattr(cam, "use_rgb", True):
obs_dict[cam_key] = cam.read_latest() start = time.perf_counter()
dt_ms = (time.perf_counter() - start) * 1e3 obs_dict[cam_key] = cam.read_latest()
logger.debug(f"{self} read {cam_key}: {dt_ms:.1f}ms") dt_ms = (time.perf_counter() - start) * 1e3
logger.debug(f"{self} read {cam_key}: {dt_ms:.1f}ms")
if getattr(cam, "use_depth", False):
start = time.perf_counter()
obs_dict[f"{cam_key}_depth"] = cam.read_latest_depth()
dt_ms = (time.perf_counter() - start) * 1e3
logger.debug(f"{self} read {cam_key} depth: {dt_ms:.1f}ms")
return obs_dict return obs_dict
+19 -7
View File
@@ -102,9 +102,14 @@ class HopeJrHand(Robot):
@property @property
def _cameras_ft(self) -> dict[str, tuple]: def _cameras_ft(self) -> dict[str, tuple]:
return { features: dict[str, tuple] = {}
cam: (self.config.cameras[cam].height, self.config.cameras[cam].width, 3) for cam in self.cameras for cam in self.cameras:
} cfg = self.config.cameras[cam]
if getattr(cfg, "use_rgb", True):
features[cam] = (cfg.height, cfg.width, 3)
if getattr(cfg, "use_depth", False):
features[f"{cam}_depth"] = (cfg.height, cfg.width, 1)
return features
@cached_property @cached_property
def observation_features(self) -> dict[str, type | tuple]: def observation_features(self) -> dict[str, type | tuple]:
@@ -170,10 +175,17 @@ class HopeJrHand(Robot):
# Capture images from cameras # Capture images from cameras
for cam_key, cam in self.cameras.items(): for cam_key, cam in self.cameras.items():
start = time.perf_counter() if getattr(cam, "use_rgb", True):
obs_dict[cam_key] = cam.read_latest() start = time.perf_counter()
dt_ms = (time.perf_counter() - start) * 1e3 obs_dict[cam_key] = cam.read_latest()
logger.debug(f"{self} read {cam_key}: {dt_ms:.1f}ms") dt_ms = (time.perf_counter() - start) * 1e3
logger.debug(f"{self} read {cam_key}: {dt_ms:.1f}ms")
if getattr(cam, "use_depth", False):
start = time.perf_counter()
obs_dict[f"{cam_key}_depth"] = cam.read_latest_depth()
dt_ms = (time.perf_counter() - start) * 1e3
logger.debug(f"{self} read {cam_key} depth: {dt_ms:.1f}ms")
return obs_dict return obs_dict
@@ -68,9 +68,14 @@ class KochFollower(Robot):
@property @property
def _cameras_ft(self) -> dict[str, tuple]: def _cameras_ft(self) -> dict[str, tuple]:
return { features: dict[str, tuple] = {}
cam: (self.config.cameras[cam].height, self.config.cameras[cam].width, 3) for cam in self.cameras for cam in self.cameras:
} cfg = self.config.cameras[cam]
if getattr(cfg, "use_rgb", True):
features[cam] = (cfg.height, cfg.width, 3)
if getattr(cfg, "use_depth", False):
features[f"{cam}_depth"] = (cfg.height, cfg.width, 1)
return features
@cached_property @cached_property
def observation_features(self) -> dict[str, type | tuple]: def observation_features(self) -> dict[str, type | tuple]:
@@ -192,10 +197,17 @@ class KochFollower(Robot):
# Capture images from cameras # Capture images from cameras
for cam_key, cam in self.cameras.items(): for cam_key, cam in self.cameras.items():
start = time.perf_counter() if getattr(cam, "use_rgb", True):
obs_dict[cam_key] = cam.read_latest() start = time.perf_counter()
dt_ms = (time.perf_counter() - start) * 1e3 obs_dict[cam_key] = cam.read_latest()
logger.debug(f"{self} read {cam_key}: {dt_ms:.1f}ms") dt_ms = (time.perf_counter() - start) * 1e3
logger.debug(f"{self} read {cam_key}: {dt_ms:.1f}ms")
if getattr(cam, "use_depth", False):
start = time.perf_counter()
obs_dict[f"{cam_key}_depth"] = cam.read_latest_depth()
dt_ms = (time.perf_counter() - start) * 1e3
logger.debug(f"{self} read {cam_key} depth: {dt_ms:.1f}ms")
return obs_dict return obs_dict
+6
View File
@@ -72,6 +72,12 @@ class LeKiwi(Robot):
) )
self.arm_motors = [motor for motor in self.bus.motors if motor.startswith("arm")] self.arm_motors = [motor for motor in self.bus.motors if motor.startswith("arm")]
self.base_motors = [motor for motor in self.bus.motors if motor.startswith("base")] self.base_motors = [motor for motor in self.bus.motors if motor.startswith("base")]
depth_cameras = [name for name, cfg in config.cameras.items() if getattr(cfg, "use_depth", False)]
if depth_cameras:
raise NotImplementedError(
f"Depth cameras are not supported on LeKiwi (got depth-enabled cameras: {depth_cameras}). "
"The host/client transport only carries color frames."
)
self.cameras = make_cameras_from_configs(config.cameras) self.cameras = make_cameras_from_configs(config.cameras)
@property @property
@@ -44,6 +44,13 @@ class LeKiwiClient(Robot):
self.id = config.id self.id = config.id
self.robot_type = config.type self.robot_type = config.type
depth_cameras = [name for name, cfg in config.cameras.items() if getattr(cfg, "use_depth", False)]
if depth_cameras:
raise NotImplementedError(
f"Depth cameras are not supported on LeKiwi (got depth-enabled cameras: {depth_cameras}). "
"The host/client transport only carries color frames."
)
self.remote_ip = config.remote_ip self.remote_ip = config.remote_ip
self.port_zmq_cmd = config.port_zmq_cmd self.port_zmq_cmd = config.port_zmq_cmd
self.port_zmq_observations = config.port_zmq_observations self.port_zmq_observations = config.port_zmq_observations
@@ -68,9 +68,14 @@ class OmxFollower(Robot):
@property @property
def _cameras_ft(self) -> dict[str, tuple]: def _cameras_ft(self) -> dict[str, tuple]:
return { features: dict[str, tuple] = {}
cam: (self.config.cameras[cam].height, self.config.cameras[cam].width, 3) for cam in self.cameras for cam in self.cameras:
} cfg = self.config.cameras[cam]
if getattr(cfg, "use_rgb", True):
features[cam] = (cfg.height, cfg.width, 3)
if getattr(cfg, "use_depth", False):
features[f"{cam}_depth"] = (cfg.height, cfg.width, 1)
return features
@cached_property @cached_property
def observation_features(self) -> dict[str, type | tuple]: def observation_features(self) -> dict[str, type | tuple]:
@@ -175,10 +180,17 @@ class OmxFollower(Robot):
# Capture images from cameras # Capture images from cameras
for cam_key, cam in self.cameras.items(): for cam_key, cam in self.cameras.items():
start = time.perf_counter() if getattr(cam, "use_rgb", True):
obs_dict[cam_key] = cam.read_latest() start = time.perf_counter()
dt_ms = (time.perf_counter() - start) * 1e3 obs_dict[cam_key] = cam.read_latest()
logger.debug(f"{self} read {cam_key}: {dt_ms:.1f}ms") dt_ms = (time.perf_counter() - start) * 1e3
logger.debug(f"{self} read {cam_key}: {dt_ms:.1f}ms")
if getattr(cam, "use_depth", False):
start = time.perf_counter()
obs_dict[f"{cam_key}_depth"] = cam.read_latest_depth()
dt_ms = (time.perf_counter() - start) * 1e3
logger.debug(f"{self} read {cam_key} depth: {dt_ms:.1f}ms")
return obs_dict return obs_dict
@@ -101,9 +101,14 @@ class OpenArmFollower(Robot):
@property @property
def _cameras_ft(self) -> dict[str, tuple]: def _cameras_ft(self) -> dict[str, tuple]:
"""Camera features for observation space.""" """Camera features for observation space."""
return { features: dict[str, tuple] = {}
cam: (self.config.cameras[cam].height, self.config.cameras[cam].width, 3) for cam in self.cameras for cam in self.cameras:
} cfg = self.config.cameras[cam]
if getattr(cfg, "use_rgb", True):
features[cam] = (cfg.height, cfg.width, 3)
if getattr(cfg, "use_depth", False):
features[f"{cam}_depth"] = (cfg.height, cfg.width, 1)
return features
@cached_property @cached_property
def observation_features(self) -> dict[str, type | tuple]: def observation_features(self) -> dict[str, type | tuple]:
@@ -242,10 +247,17 @@ class OpenArmFollower(Robot):
# Capture images from cameras # Capture images from cameras
for cam_key, cam in self.cameras.items(): for cam_key, cam in self.cameras.items():
start = time.perf_counter() if getattr(cam, "use_rgb", True):
obs_dict[cam_key] = cam.read_latest() start = time.perf_counter()
dt_ms = (time.perf_counter() - start) * 1e3 obs_dict[cam_key] = cam.read_latest()
logger.debug(f"{self} read {cam_key}: {dt_ms:.1f}ms") dt_ms = (time.perf_counter() - start) * 1e3
logger.debug(f"{self} read {cam_key}: {dt_ms:.1f}ms")
if getattr(cam, "use_depth", False):
start = time.perf_counter()
obs_dict[f"{cam_key}_depth"] = cam.read_latest_depth()
dt_ms = (time.perf_counter() - start) * 1e3
logger.debug(f"{self} read {cam_key} depth: {dt_ms:.1f}ms")
dt_ms = (time.perf_counter() - start) * 1e3 dt_ms = (time.perf_counter() - start) * 1e3
logger.debug(f"{self} get_observation took: {dt_ms:.1f}ms") logger.debug(f"{self} get_observation took: {dt_ms:.1f}ms")
@@ -80,9 +80,14 @@ class RebotB601Follower(Robot):
@property @property
def _cameras_ft(self) -> dict[str, tuple]: def _cameras_ft(self) -> dict[str, tuple]:
return { features: dict[str, tuple] = {}
cam: (self.config.cameras[cam].height, self.config.cameras[cam].width, 3) for cam in self.cameras for cam in self.cameras:
} cfg = self.config.cameras[cam]
if getattr(cfg, "use_rgb", True):
features[cam] = (cfg.height, cfg.width, 3)
if getattr(cfg, "use_depth", False):
features[f"{cam}_depth"] = (cfg.height, cfg.width, 1)
return features
@cached_property @cached_property
def observation_features(self) -> dict[str, type | tuple]: def observation_features(self) -> dict[str, type | tuple]:
@@ -213,10 +218,17 @@ class RebotB601Follower(Robot):
logger.debug(f"{self} read state: {dt_ms:.1f}ms") logger.debug(f"{self} read state: {dt_ms:.1f}ms")
for cam_key, cam in self.cameras.items(): for cam_key, cam in self.cameras.items():
start = time.perf_counter() if getattr(cam, "use_rgb", True):
obs_dict[cam_key] = cam.read_latest() start = time.perf_counter()
dt_ms = (time.perf_counter() - start) * 1e3 obs_dict[cam_key] = cam.read_latest()
logger.debug(f"{self} read {cam_key}: {dt_ms:.1f}ms") dt_ms = (time.perf_counter() - start) * 1e3
logger.debug(f"{self} read {cam_key}: {dt_ms:.1f}ms")
if getattr(cam, "use_depth", False):
start = time.perf_counter()
obs_dict[f"{cam_key}_depth"] = cam.read_latest_depth()
dt_ms = (time.perf_counter() - start) * 1e3
logger.debug(f"{self} read {cam_key} depth: {dt_ms:.1f}ms")
return obs_dict return obs_dict
+18 -7
View File
@@ -68,9 +68,13 @@ class SOFollower(Robot):
@property @property
def _cameras_ft(self) -> dict[str, tuple]: def _cameras_ft(self) -> dict[str, tuple]:
return { features: dict[str, tuple] = {}
cam: (self.config.cameras[cam].height, self.config.cameras[cam].width, 3) for cam in self.cameras for cam in self.cameras:
} if getattr(self.cameras[cam], "use_rgb", True):
features[cam] = (self.cameras[cam].height, self.cameras[cam].width, 3)
if getattr(self.cameras[cam], "use_depth", False):
features[f"{cam}_depth"] = (self.cameras[cam].height, self.cameras[cam].width, 1)
return features
@cached_property @cached_property
def observation_features(self) -> dict[str, type | tuple]: def observation_features(self) -> dict[str, type | tuple]:
@@ -185,10 +189,17 @@ class SOFollower(Robot):
# Capture images from cameras # Capture images from cameras
for cam_key, cam in self.cameras.items(): for cam_key, cam in self.cameras.items():
start = time.perf_counter() if getattr(cam, "use_rgb", True):
obs_dict[cam_key] = cam.read_latest() start = time.perf_counter()
dt_ms = (time.perf_counter() - start) * 1e3 obs_dict[cam_key] = cam.read_latest()
logger.debug(f"{self} read {cam_key}: {dt_ms:.1f}ms") dt_ms = (time.perf_counter() - start) * 1e3
logger.debug(f"{self} read {cam_key}: {dt_ms:.1f}ms")
if getattr(cam, "use_depth", False):
start = time.perf_counter()
obs_dict[f"{cam_key}_depth"] = cam.read_latest_depth()
dt_ms = (time.perf_counter() - start) * 1e3
logger.debug(f"{self} read {cam_key} depth: {dt_ms:.1f}ms")
return obs_dict return obs_dict
+12 -4
View File
@@ -222,9 +222,14 @@ class UnitreeG1(Robot):
@property @property
def _cameras_ft(self) -> dict[str, tuple]: def _cameras_ft(self) -> dict[str, tuple]:
return { features: dict[str, tuple] = {}
cam: (self.config.cameras[cam].height, self.config.cameras[cam].width, 3) for cam in self.cameras for cam in self.cameras:
} cfg = self.config.cameras[cam]
if getattr(cfg, "use_rgb", True):
features[cam] = (cfg.height, cfg.width, 3)
if getattr(cfg, "use_depth", False):
features[f"{cam}_depth"] = (cfg.height, cfg.width, 1)
return features
@cached_property @cached_property
def observation_features(self) -> dict[str, type | tuple]: def observation_features(self) -> dict[str, type | tuple]:
@@ -458,7 +463,10 @@ class UnitreeG1(Robot):
# Cameras - read images from ZMQ cameras # Cameras - read images from ZMQ cameras
for cam_name, cam in self._cameras.items(): for cam_name, cam in self._cameras.items():
obs[cam_name] = cam.read_latest() if getattr(cam, "use_rgb", True):
obs[cam_name] = cam.read_latest()
if getattr(cam, "use_depth", False):
obs[f"{cam_name}_depth"] = cam.read_latest_depth()
return obs return obs
+4 -2
View File
@@ -332,7 +332,8 @@ def build_rollout_context(
cfg.dataset.repo_id, cfg.dataset.repo_id,
root=cfg.dataset.root, root=cfg.dataset.root,
batch_encoding_size=cfg.dataset.video_encoding_batch_size, batch_encoding_size=cfg.dataset.video_encoding_batch_size,
camera_encoder=cfg.dataset.camera_encoder, rgb_encoder=cfg.dataset.rgb_encoder,
depth_encoder=cfg.dataset.depth_encoder,
streaming_encoding=cfg.dataset.streaming_encoding, streaming_encoding=cfg.dataset.streaming_encoding,
encoder_queue_maxsize=cfg.dataset.encoder_queue_maxsize, encoder_queue_maxsize=cfg.dataset.encoder_queue_maxsize,
encoder_threads=cfg.dataset.encoder_threads, encoder_threads=cfg.dataset.encoder_threads,
@@ -367,7 +368,8 @@ def build_rollout_context(
image_writer_threads=cfg.dataset.num_image_writer_threads_per_camera image_writer_threads=cfg.dataset.num_image_writer_threads_per_camera
* len(robot.cameras if hasattr(robot, "cameras") else []), * len(robot.cameras if hasattr(robot, "cameras") else []),
batch_encoding_size=cfg.dataset.video_encoding_batch_size, batch_encoding_size=cfg.dataset.video_encoding_batch_size,
camera_encoder=cfg.dataset.camera_encoder, rgb_encoder=cfg.dataset.rgb_encoder,
depth_encoder=cfg.dataset.depth_encoder,
streaming_encoding=cfg.dataset.streaming_encoding, streaming_encoding=cfg.dataset.streaming_encoding,
encoder_queue_maxsize=cfg.dataset.encoder_queue_maxsize, encoder_queue_maxsize=cfg.dataset.encoder_queue_maxsize,
encoder_threads=cfg.dataset.encoder_threads, encoder_threads=cfg.dataset.encoder_threads,
+37 -7
View File
@@ -77,15 +77,28 @@ from lerobot.utils.constants import ACTION, DONE, OBS_STATE, REWARD
from lerobot.utils.utils import init_logging from lerobot.utils.utils import init_logging
def check_chw_float32(frame: torch.Tensor) -> None:
"""
Check if a frame is a channel-first, float32 tensor.
"""
assert frame.dtype == torch.float32
assert frame.ndim == 3
c, h, w = frame.shape
assert c < h and c < w, f"expect channel first images, but instead {frame.shape}"
def to_hwc_uint8_numpy(chw_float32_torch: torch.Tensor) -> np.ndarray: def to_hwc_uint8_numpy(chw_float32_torch: torch.Tensor) -> np.ndarray:
assert chw_float32_torch.dtype == torch.float32 check_chw_float32(chw_float32_torch)
assert chw_float32_torch.ndim == 3
c, h, w = chw_float32_torch.shape
assert c < h and c < w, f"expect channel first images, but instead {chw_float32_torch.shape}"
hwc_uint8_numpy = (chw_float32_torch * 255).type(torch.uint8).permute(1, 2, 0).numpy() hwc_uint8_numpy = (chw_float32_torch * 255).type(torch.uint8).permute(1, 2, 0).numpy()
return hwc_uint8_numpy return hwc_uint8_numpy
def to_hwc_uint16_numpy(chw_float32_torch: torch.Tensor) -> np.ndarray:
check_chw_float32(chw_float32_torch)
hwc_uint16_numpy = chw_float32_torch.round().type(torch.uint16).permute(1, 2, 0).numpy()
return hwc_uint16_numpy
def visualize_dataset( def visualize_dataset(
dataset: LeRobotDataset, dataset: LeRobotDataset,
episode_index: int, episode_index: int,
@@ -138,6 +151,14 @@ def visualize_dataset(
logging.info("Logging to Rerun") logging.info("Logging to Rerun")
# Use the dataset's q01/q99 depth statistics for robust depth range bounds
depth_ranges = {}
for key in dataset.meta.depth_keys:
stats = dataset.meta.stats[key]
lo = stats["q01"] if "q01" in stats else stats["min"]
hi = stats["q99"] if "q99" in stats else stats["max"]
depth_ranges[key] = (float(np.asarray(lo).item()), float(np.asarray(hi).item()))
first_index = None first_index = None
for batch in tqdm.tqdm(dataloader, total=len(dataloader)): for batch in tqdm.tqdm(dataloader, total=len(dataloader)):
if first_index is None: if first_index is None:
@@ -149,9 +170,18 @@ def visualize_dataset(
# display each camera image # display each camera image
for key in dataset.meta.camera_keys: for key in dataset.meta.camera_keys:
img = to_hwc_uint8_numpy(batch[key][i]) if key in dataset.meta.depth_keys:
img_entity = rr.Image(img).compress() if display_compressed_images else rr.Image(img) depth = to_hwc_uint16_numpy(batch[key][i])
rr.log(key, entity=img_entity) depth_entity = rr.DepthImage(
depth,
colormap=rr.components.Colormap.Viridis,
depth_range=depth_ranges[key],
)
rr.log(key, entity=depth_entity)
else:
img = to_hwc_uint8_numpy(batch[key][i])
img_entity = rr.Image(img).compress() if display_compressed_images else rr.Image(img)
rr.log(key, entity=img_entity)
# display each dimension of action space (e.g. actuators command) # display each dimension of action space (e.g. actuators command)
if ACTION in batch: if ACTION in batch:
+41 -12
View File
@@ -133,6 +133,15 @@ Convert image dataset to video format and save locally:
--new_root /path/to/output/pusht_video \ --new_root /path/to/output/pusht_video \
--operation.type convert_image_to_video --operation.type convert_image_to_video
Convert image dataset (with depth maps) to video format, customizing the depth encoder:
lerobot-edit-dataset \
--repo_id lerobot/pusht_image \
--new_root /path/to/output/pusht_video \
--operation.type convert_image_to_video \
--operation.depth_encoder.depth_min 0.01 \
--operation.depth_encoder.depth_max 10.0 \
--operation.depth_encoder.use_log true
Convert image dataset to video format and save with new repo_id: Convert image dataset to video format and save with new repo_id:
lerobot-edit-dataset \ lerobot-edit-dataset \
--repo_id lerobot/pusht_image \ --repo_id lerobot/pusht_image \
@@ -190,17 +199,17 @@ Re-encode all videos in a dataset (saves to lerobot/pusht_reencoded by default):
lerobot-edit-dataset \ lerobot-edit-dataset \
--repo_id lerobot/pusht \ --repo_id lerobot/pusht \
--operation.type reencode_videos \ --operation.type reencode_videos \
--operation.camera_encoder.vcodec h264 \ --operation.rgb_encoder.vcodec h264 \
--operation.camera_encoder.pix_fmt yuv420p \ --operation.rgb_encoder.pix_fmt yuv420p \
--operation.camera_encoder.crf 23 --operation.rgb_encoder.crf 23
Re-encode videos into a new dataset using 4 parallel processes: Re-encode videos into a new dataset using 4 parallel processes:
lerobot-edit-dataset \ lerobot-edit-dataset \
--repo_id lerobot/pusht \ --repo_id lerobot/pusht \
--new_repo_id lerobot/pusht_h264 \ --new_repo_id lerobot/pusht_h264 \
--operation.type reencode_videos \ --operation.type reencode_videos \
--operation.camera_encoder.vcodec h264 \ --operation.rgb_encoder.vcodec h264 \
--operation.camera_encoder.crf 23 \ --operation.rgb_encoder.crf 23 \
--operation.num_workers 4 --operation.num_workers 4
Re-encode videos in-place (overwrites original dataset): Re-encode videos in-place (overwrites original dataset):
@@ -208,9 +217,16 @@ Re-encode videos in-place (overwrites original dataset):
--repo_id lerobot/pusht \ --repo_id lerobot/pusht \
--new_repo_id lerobot/pusht \ --new_repo_id lerobot/pusht \
--operation.type reencode_videos \ --operation.type reencode_videos \
--operation.camera_encoder.vcodec h264 \ --operation.rgb_encoder.vcodec h264 \
--operation.overwrite true --operation.overwrite true
Re-encode both RGB and depth videos in a dataset (depth quantization params are preserved):
lerobot-edit-dataset \
--repo_id lerobot/pusht_depth \
--operation.type reencode_videos \
--operation.rgb_encoder.vcodec h264 \
--operation.depth_encoder.extra_options '{"x265-params": "lossless=1"}'
Using JSON config file: Using JSON config file:
lerobot-edit-dataset \ lerobot-edit-dataset \
--config_path path/to/edit_config.json --config_path path/to/edit_config.json
@@ -225,7 +241,13 @@ from pathlib import Path
import draccus import draccus
from lerobot.configs import VideoEncoderConfig, camera_encoder_defaults, parser from lerobot.configs import (
DepthEncoderConfig,
RGBEncoderConfig,
depth_encoder_defaults,
parser,
rgb_encoder_defaults,
)
from lerobot.datasets import ( from lerobot.datasets import (
LeRobotDataset, LeRobotDataset,
convert_image_to_video_dataset, convert_image_to_video_dataset,
@@ -287,7 +309,8 @@ class ModifyTasksConfig(OperationConfig):
@dataclass @dataclass
class ConvertImageToVideoConfig(OperationConfig): class ConvertImageToVideoConfig(OperationConfig):
output_dir: str | None = None output_dir: str | None = None
camera_encoder: VideoEncoderConfig = field(default_factory=camera_encoder_defaults) rgb_encoder: RGBEncoderConfig = field(default_factory=rgb_encoder_defaults)
depth_encoder: DepthEncoderConfig = field(default_factory=depth_encoder_defaults)
episode_indices: list[int] | None = None episode_indices: list[int] | None = None
num_workers: int = 4 num_workers: int = 4
max_episodes_per_batch: int | None = None max_episodes_per_batch: int | None = None
@@ -308,7 +331,8 @@ class RecomputeStatsConfig(OperationConfig):
@OperationConfig.register_subclass("reencode_videos") @OperationConfig.register_subclass("reencode_videos")
@dataclass @dataclass
class ReencodeVideosConfig(OperationConfig): class ReencodeVideosConfig(OperationConfig):
camera_encoder: VideoEncoderConfig = field(default_factory=camera_encoder_defaults) rgb_encoder: RGBEncoderConfig = field(default_factory=rgb_encoder_defaults)
depth_encoder: DepthEncoderConfig = field(default_factory=depth_encoder_defaults)
num_workers: int = 0 num_workers: int = 0
encoder_threads: int | None = None encoder_threads: int | None = None
overwrite: bool = False overwrite: bool = False
@@ -601,7 +625,8 @@ def handle_convert_image_to_video(cfg: EditDatasetConfig) -> None:
dataset=dataset, dataset=dataset,
output_dir=output_dir, output_dir=output_dir,
repo_id=output_repo_id, repo_id=output_repo_id,
camera_encoder=getattr(cfg.operation, "camera_encoder", None) or camera_encoder_defaults(), rgb_encoder=getattr(cfg.operation, "rgb_encoder", None) or rgb_encoder_defaults(),
depth_encoder=getattr(cfg.operation, "depth_encoder", None) or depth_encoder_defaults(),
episode_indices=getattr(cfg.operation, "episode_indices", None), episode_indices=getattr(cfg.operation, "episode_indices", None),
num_workers=getattr(cfg.operation, "num_workers", 4), num_workers=getattr(cfg.operation, "num_workers", 4),
max_episodes_per_batch=getattr(cfg.operation, "max_episodes_per_batch", None), max_episodes_per_batch=getattr(cfg.operation, "max_episodes_per_batch", None),
@@ -719,10 +744,14 @@ def handle_reencode_videos(cfg: EditDatasetConfig) -> None:
shutil.copytree(input_root, output_root) shutil.copytree(input_root, output_root)
dataset = LeRobotDataset(output_repo_id, root=output_root) dataset = LeRobotDataset(output_repo_id, root=output_root)
logging.info(f"Re-encoding videos in {output_repo_id} with {cfg.operation.camera_encoder}") logging.info(
f"Re-encoding videos in {output_repo_id} with RGB encoder {cfg.operation.rgb_encoder} "
f"and depth encoder {cfg.operation.depth_encoder}"
)
reencode_dataset( reencode_dataset(
dataset, dataset,
camera_encoder=cfg.operation.camera_encoder, rgb_encoder=cfg.operation.rgb_encoder,
depth_encoder=cfg.operation.depth_encoder,
encoder_threads=cfg.operation.encoder_threads, encoder_threads=cfg.operation.encoder_threads,
num_workers=cfg.operation.num_workers, num_workers=cfg.operation.num_workers,
) )
+8 -6
View File
@@ -79,9 +79,9 @@ lerobot-record \\
--dataset.single_task="Grab the cube" \\ --dataset.single_task="Grab the cube" \\
--dataset.streaming_encoding=true \\ --dataset.streaming_encoding=true \\
--dataset.encoder_threads=2 \\ --dataset.encoder_threads=2 \\
--dataset.camera_encoder.vcodec=h264 \\ --dataset.rgb_encoder.vcodec=h264 \\
--dataset.camera_encoder.preset=fast \\ --dataset.rgb_encoder.preset=fast \\
--dataset.camera_encoder.extra_options={"tune": "film", "profile:v": "high", "bf": 2} \\ --dataset.rgb_encoder.extra_options={"tune": "film", "profile:v": "high", "bf": 2} \\
--display_data=true --display_data=true
``` ```
""" """
@@ -400,7 +400,8 @@ def record(
cfg.dataset.repo_id, cfg.dataset.repo_id,
root=cfg.dataset.root, root=cfg.dataset.root,
batch_encoding_size=cfg.dataset.video_encoding_batch_size, batch_encoding_size=cfg.dataset.video_encoding_batch_size,
camera_encoder=cfg.dataset.camera_encoder, rgb_encoder=cfg.dataset.rgb_encoder,
depth_encoder=cfg.dataset.depth_encoder,
encoder_threads=cfg.dataset.encoder_threads, encoder_threads=cfg.dataset.encoder_threads,
streaming_encoding=cfg.dataset.streaming_encoding, streaming_encoding=cfg.dataset.streaming_encoding,
encoder_queue_maxsize=cfg.dataset.encoder_queue_maxsize, encoder_queue_maxsize=cfg.dataset.encoder_queue_maxsize,
@@ -429,7 +430,8 @@ def record(
image_writer_processes=cfg.dataset.num_image_writer_processes, image_writer_processes=cfg.dataset.num_image_writer_processes,
image_writer_threads=cfg.dataset.num_image_writer_threads_per_camera * len(robot.cameras), image_writer_threads=cfg.dataset.num_image_writer_threads_per_camera * len(robot.cameras),
batch_encoding_size=cfg.dataset.video_encoding_batch_size, batch_encoding_size=cfg.dataset.video_encoding_batch_size,
camera_encoder=cfg.dataset.camera_encoder, rgb_encoder=cfg.dataset.rgb_encoder,
depth_encoder=cfg.dataset.depth_encoder,
encoder_threads=cfg.dataset.encoder_threads, encoder_threads=cfg.dataset.encoder_threads,
streaming_encoding=cfg.dataset.streaming_encoding, streaming_encoding=cfg.dataset.streaming_encoding,
encoder_queue_maxsize=cfg.dataset.encoder_queue_maxsize, encoder_queue_maxsize=cfg.dataset.encoder_queue_maxsize,
@@ -443,7 +445,7 @@ def record(
if not cfg.dataset.streaming_encoding: if not cfg.dataset.streaming_encoding:
logging.info( logging.info(
"Streaming encoding is disabled. If you have capable hardware, consider enabling it for way faster episode saving. --dataset.streaming_encoding=true --dataset.encoder_threads=2 # --dataset.camera_encoder.vcodec=auto. More info in the documentation: https://huggingface.co/docs/lerobot/streaming_video_encoding" "Streaming encoding is disabled. If you have capable hardware, consider enabling it for way faster episode saving. --dataset.streaming_encoding=true --dataset.encoder_threads=2 # --dataset.rgb_encoder.vcodec=auto. More info in the documentation: https://huggingface.co/docs/lerobot/streaming_video_encoding"
) )
with VideoEncodingManager(dataset): with VideoEncodingManager(dataset):
+3 -3
View File
@@ -142,9 +142,9 @@ Usage examples
--robot.port=/dev/ttyACM0 \\ --robot.port=/dev/ttyACM0 \\
--task="pick up cube" --duration=60 \\ --task="pick up cube" --duration=60 \\
--display_data=true \\ --display_data=true \\
--dataset.camera_encoder.vcodec=h264 \\ --dataset.rgb_encoder.vcodec=h264 \\
--dataset.camera_encoder.preset=fast \\ --dataset.rgb_encoder.preset=fast \\
--dataset.camera_encoder.extra_options={"tune": "film", "profile:v": "high", "bf": 2} --dataset.rgb_encoder.extra_options={"tune": "film", "profile:v": "high", "bf": 2}
""" """
import logging import logging
+23 -12
View File
@@ -51,7 +51,9 @@ def hw_to_dataset_features(
This function takes a dictionary describing hardware outputs (like joint states This function takes a dictionary describing hardware outputs (like joint states
or camera image shapes) and formats it into the standard LeRobot feature or camera image shapes) and formats it into the standard LeRobot feature
specification. specification. Single-channel cameras (shape ``(H, W, 1)``) are flagged as depth
maps via ``info["is_depth_map"] = True``; three-channel cameras ``(H, W, 3)`` are
treated as RGB.
Args: Args:
hw_features (dict): Dictionary mapping feature names to their type (float for hw_features (dict): Dictionary mapping feature names to their type (float for
@@ -61,7 +63,7 @@ def hw_to_dataset_features(
use_video (bool): If True, image features are marked as "video", otherwise "image". use_video (bool): If True, image features are marked as "video", otherwise "image".
Returns: Returns:
dict: A LeRobot features dictionary. dict: A LeRobot features dictionary. Depth cameras carry ``info["is_depth_map"] = True``.
""" """
features = {} features = {}
joint_fts = { joint_fts = {
@@ -69,6 +71,7 @@ def hw_to_dataset_features(
for key, ftype in hw_features.items() for key, ftype in hw_features.items()
if ftype is float or (isinstance(ftype, PolicyFeature) and ftype.type != FeatureType.VISUAL) if ftype is float or (isinstance(ftype, PolicyFeature) and ftype.type != FeatureType.VISUAL)
} }
# TODO(CarolinePascal): we should not rely on the shape to determine if a feature is a camera !
cam_fts = {key: shape for key, shape in hw_features.items() if isinstance(shape, tuple)} cam_fts = {key: shape for key, shape in hw_features.items() if isinstance(shape, tuple)}
if joint_fts and prefix == ACTION: if joint_fts and prefix == ACTION:
@@ -86,11 +89,19 @@ def hw_to_dataset_features(
} }
for key, shape in cam_fts.items(): for key, shape in cam_fts.items():
features[f"{prefix}.images.{key}"] = { dtype = "video" if use_video else "image"
"dtype": "video" if use_video else "image", if len(shape) == 3 and shape[2] in (1, 3):
"shape": shape, features[f"{prefix}.images.{key}"] = {
"names": ["height", "width", "channels"], "dtype": dtype,
} "shape": shape,
"names": ["height", "width", "channels"],
"info": {"is_depth_map": shape[2] == 1},
}
else:
raise ValueError(
f"Camera feature '{key}' has shape {shape}. "
f"Expected a 3-tuple (H, W, C), e.g. (480, 640, 3) for RGB or (480, 640, 1) for depth."
)
_validate_feature_names(features) _validate_feature_names(features)
return features return features
@@ -149,11 +160,11 @@ def dataset_to_policy_features(features: dict[str, dict]) -> dict[str, PolicyFea
type = FeatureType.VISUAL type = FeatureType.VISUAL
if len(shape) != 3: if len(shape) != 3:
raise ValueError(f"Number of dimensions of {key} != 3 (shape={shape})") raise ValueError(f"Number of dimensions of {key} != 3 (shape={shape})")
else:
names = ft["names"] names = ft["names"]
# Backward compatibility for "channel" which is an error introduced in LeRobotDataset v2.0 for ported datasets. # Backward compatibility for "channel" which is an error introduced in LeRobotDataset v2.0 for ported datasets.
if names[2] in ["channel", "channels"]: # (h, w, c) -> (c, h, w) if names[2] in ["channel", "channels"]: # (h, w, c) -> (c, h, w)
shape = (shape[2], shape[0], shape[1]) shape = (shape[2], shape[0], shape[1])
elif key == OBS_ENV_STATE: elif key == OBS_ENV_STATE:
type = FeatureType.ENV type = FeatureType.ENV
elif key.startswith(OBS_STR): elif key.startswith(OBS_STR):
+4 -1
View File
@@ -107,7 +107,10 @@ def log_rerun_data(
for i, vi in enumerate(arr): for i, vi in enumerate(arr):
rr.log(f"{key}_{i}", rr.Scalars(float(vi))) rr.log(f"{key}_{i}", rr.Scalars(float(vi)))
else: else:
img_entity = rr.Image(arr).compress() if compress_images else rr.Image(arr) if arr.shape[-1] == 1:
img_entity = rr.DepthImage(arr, colormap=rr.components.Colormap.Viridis)
else:
img_entity = rr.Image(arr).compress() if compress_images else rr.Image(arr)
rr.log(key, entity=img_entity, static=True) rr.log(key, entity=img_entity, static=True)
if action: if action:
+3 -2
View File
@@ -47,6 +47,7 @@ class _FakeMeta:
def __init__(self, video_keys: list[str], image_keys: list[str], video_path: Path | None = None) -> None: def __init__(self, video_keys: list[str], image_keys: list[str], video_path: Path | None = None) -> None:
self.video_keys = video_keys self.video_keys = video_keys
self.camera_keys = [*video_keys, *image_keys] self.camera_keys = [*video_keys, *image_keys]
self.depth_keys = []
self._video_path = video_path self._video_path = video_path
self.episodes = {0: {f"videos/{key}/from_timestamp": 0.0 for key in video_keys}} self.episodes = {0: {f"videos/{key}/from_timestamp": 0.0 for key in video_keys}}
@@ -208,14 +209,14 @@ def test_episode_clip_path_trims_via_reencode_video(tmp_path: Path, monkeypatch)
def fake_reencode( def fake_reencode(
input_video_path, input_video_path,
output_video_path, output_video_path,
camera_encoder=None, video_encoder=None,
overwrite=False, overwrite=False,
start_time_s=None, start_time_s=None,
end_time_s=None, end_time_s=None,
): ):
captured.update( captured.update(
src=Path(input_video_path), src=Path(input_video_path),
encoder=camera_encoder, encoder=video_encoder,
start_time_s=start_time_s, start_time_s=start_time_s,
end_time_s=end_time_s, end_time_s=end_time_s,
) )
+74 -8
View File
@@ -29,7 +29,10 @@ from lerobot.configs import VIDEO_ENCODER_INFO_KEYS
from lerobot.datasets.aggregate import aggregate_datasets from lerobot.datasets.aggregate import aggregate_datasets
from lerobot.datasets.feature_utils import features_equal_for_merge from lerobot.datasets.feature_utils import features_equal_for_merge
from lerobot.datasets.lerobot_dataset import LeRobotDataset from lerobot.datasets.lerobot_dataset import LeRobotDataset
from tests.fixtures.constants import DUMMY_REPO_ID from tests.fixtures.constants import (
DUMMY_CAMERA_FEATURES_WITH_DEPTH,
DUMMY_REPO_ID,
)
def assert_data_shards_one_row_group_per_episode(root): def assert_data_shards_one_row_group_per_episode(root):
@@ -211,6 +214,26 @@ def assert_dataset_iteration_works(aggr_ds):
pass pass
def assert_depth_keys_preserved(aggr_ds, ds_0, ds_1):
"""Test that depth keys are correctly preserved after aggregation.
Ensures that the ``is_depth_map`` marker on visual features survives
aggregation, so that downstream consumers (e.g. the dataset reader's
depth decoding path) keep working on the merged dataset.
"""
expected_depth_keys = set(ds_0.meta.depth_keys)
assert expected_depth_keys == set(ds_1.meta.depth_keys), (
"Source datasets disagree on depth_keys; test setup is inconsistent"
)
actual_depth_keys = set(aggr_ds.meta.depth_keys)
assert actual_depth_keys == expected_depth_keys, (
f"Expected depth_keys {expected_depth_keys}, got {actual_depth_keys}"
)
for key in expected_depth_keys:
info = aggr_ds.meta.info.features[key].get("info") or {}
assert info.get("is_depth_map") is True, f"Depth marker lost on feature {key!r} after aggregation"
def assert_video_timestamps_within_bounds(aggr_ds): def assert_video_timestamps_within_bounds(aggr_ds):
"""Test that all video timestamps are within valid bounds for their respective video files. """Test that all video timestamps are within valid bounds for their respective video files.
@@ -260,7 +283,11 @@ def assert_video_timestamps_within_bounds(aggr_ds):
def test_aggregate_datasets(tmp_path, lerobot_dataset_factory): def test_aggregate_datasets(tmp_path, lerobot_dataset_factory):
"""Test basic aggregation functionality with standard parameters.""" """Test basic aggregation functionality with standard parameters.
Source datasets include both RGB and depth video features so the same
aggregation flow is exercised on the ``is_depth_map`` branch.
"""
ds_0_num_frames = 400 ds_0_num_frames = 400
ds_1_num_frames = 800 ds_1_num_frames = 800
ds_0_num_episodes = 10 ds_0_num_episodes = 10
@@ -272,14 +299,21 @@ def test_aggregate_datasets(tmp_path, lerobot_dataset_factory):
repo_id=f"{DUMMY_REPO_ID}_0", repo_id=f"{DUMMY_REPO_ID}_0",
total_episodes=ds_0_num_episodes, total_episodes=ds_0_num_episodes,
total_frames=ds_0_num_frames, total_frames=ds_0_num_frames,
camera_features=DUMMY_CAMERA_FEATURES_WITH_DEPTH,
) )
ds_1 = lerobot_dataset_factory( ds_1 = lerobot_dataset_factory(
root=tmp_path / "test_1", root=tmp_path / "test_1",
repo_id=f"{DUMMY_REPO_ID}_1", repo_id=f"{DUMMY_REPO_ID}_1",
total_episodes=ds_1_num_episodes, total_episodes=ds_1_num_episodes,
total_frames=ds_1_num_frames, total_frames=ds_1_num_frames,
camera_features=DUMMY_CAMERA_FEATURES_WITH_DEPTH,
) )
# Confirm depth was actually wired into the source datasets so the
# rest of the assertions exercise the depth aggregation path.
assert len(ds_0.meta.depth_keys) > 0, "ds_0 should expose at least one depth key"
assert len(ds_1.meta.depth_keys) > 0, "ds_1 should expose at least one depth key"
aggregate_datasets( aggregate_datasets(
repo_ids=[ds_0.repo_id, ds_1.repo_id], repo_ids=[ds_0.repo_id, ds_1.repo_id],
roots=[ds_0.root, ds_1.root], roots=[ds_0.root, ds_1.root],
@@ -306,6 +340,7 @@ def test_aggregate_datasets(tmp_path, lerobot_dataset_factory):
assert_episode_indices_updated_correctly(aggr_ds, ds_0, ds_1) assert_episode_indices_updated_correctly(aggr_ds, ds_0, ds_1)
assert_video_frames_integrity(aggr_ds, ds_0, ds_1) assert_video_frames_integrity(aggr_ds, ds_0, ds_1)
assert_video_timestamps_within_bounds(aggr_ds) assert_video_timestamps_within_bounds(aggr_ds)
assert_depth_keys_preserved(aggr_ds, ds_0, ds_1)
assert_dataset_iteration_works(aggr_ds) assert_dataset_iteration_works(aggr_ds)
@@ -423,7 +458,11 @@ def test_aggregate_incomplete_video_encoder_info_warns_and_nuls_encoders(
def test_aggregate_with_low_threshold(tmp_path, lerobot_dataset_factory): def test_aggregate_with_low_threshold(tmp_path, lerobot_dataset_factory):
"""Test aggregation with small file size limits to force file rotation/sharding.""" """Test aggregation with small file size limits to force file rotation/sharding.
Depth video features are included to verify that file rotation/concat
correctly handles depth-marked features alongside regular RGB ones.
"""
ds_0_num_episodes = ds_1_num_episodes = 10 ds_0_num_episodes = ds_1_num_episodes = 10
ds_0_num_frames = ds_1_num_frames = 400 ds_0_num_frames = ds_1_num_frames = 400
@@ -432,14 +471,19 @@ def test_aggregate_with_low_threshold(tmp_path, lerobot_dataset_factory):
repo_id=f"{DUMMY_REPO_ID}_small_0", repo_id=f"{DUMMY_REPO_ID}_small_0",
total_episodes=ds_0_num_episodes, total_episodes=ds_0_num_episodes,
total_frames=ds_0_num_frames, total_frames=ds_0_num_frames,
camera_features=DUMMY_CAMERA_FEATURES_WITH_DEPTH,
) )
ds_1 = lerobot_dataset_factory( ds_1 = lerobot_dataset_factory(
root=tmp_path / "small_1", root=tmp_path / "small_1",
repo_id=f"{DUMMY_REPO_ID}_small_1", repo_id=f"{DUMMY_REPO_ID}_small_1",
total_episodes=ds_1_num_episodes, total_episodes=ds_1_num_episodes,
total_frames=ds_1_num_frames, total_frames=ds_1_num_frames,
camera_features=DUMMY_CAMERA_FEATURES_WITH_DEPTH,
) )
assert len(ds_0.meta.depth_keys) > 0, "ds_0 should expose at least one depth key"
assert len(ds_1.meta.depth_keys) > 0, "ds_1 should expose at least one depth key"
# Use the new configurable parameters to force file rotation # Use the new configurable parameters to force file rotation
aggregate_datasets( aggregate_datasets(
repo_ids=[ds_0.repo_id, ds_1.repo_id], repo_ids=[ds_0.repo_id, ds_1.repo_id],
@@ -470,6 +514,7 @@ def test_aggregate_with_low_threshold(tmp_path, lerobot_dataset_factory):
assert_episode_indices_updated_correctly(aggr_ds, ds_0, ds_1) assert_episode_indices_updated_correctly(aggr_ds, ds_0, ds_1)
assert_video_frames_integrity(aggr_ds, ds_0, ds_1) assert_video_frames_integrity(aggr_ds, ds_0, ds_1)
assert_video_timestamps_within_bounds(aggr_ds) assert_video_timestamps_within_bounds(aggr_ds)
assert_depth_keys_preserved(aggr_ds, ds_0, ds_1)
assert_dataset_iteration_works(aggr_ds) assert_dataset_iteration_works(aggr_ds)
# Check that multiple files were actually created due to small size limits # Check that multiple files were actually created due to small size limits
@@ -489,7 +534,8 @@ def test_video_timestamps_regression(tmp_path, lerobot_dataset_factory):
"""Regression test for video timestamp bug when merging datasets. """Regression test for video timestamp bug when merging datasets.
This test specifically checks that video timestamps are correctly calculated This test specifically checks that video timestamps are correctly calculated
and accumulated when merging multiple datasets. and accumulated when merging multiple datasets. Depth video features are
included so depth timestamps are also covered by the regression.
""" """
datasets = [] datasets = []
for i in range(3): for i in range(3):
@@ -498,9 +544,13 @@ def test_video_timestamps_regression(tmp_path, lerobot_dataset_factory):
repo_id=f"{DUMMY_REPO_ID}_regression_{i}", repo_id=f"{DUMMY_REPO_ID}_regression_{i}",
total_episodes=2, total_episodes=2,
total_frames=100, total_frames=100,
camera_features=DUMMY_CAMERA_FEATURES_WITH_DEPTH,
) )
datasets.append(ds) datasets.append(ds)
for i, ds in enumerate(datasets):
assert len(ds.meta.depth_keys) > 0, f"Dataset {i} should expose at least one depth key"
aggregate_datasets( aggregate_datasets(
repo_ids=[ds.repo_id for ds in datasets], repo_ids=[ds.repo_id for ds in datasets],
roots=[ds.root for ds in datasets], roots=[ds.root for ds in datasets],
@@ -517,12 +567,21 @@ def test_video_timestamps_regression(tmp_path, lerobot_dataset_factory):
aggr_ds = LeRobotDataset(f"{DUMMY_REPO_ID}_regression_aggr", root=tmp_path / "regression_aggr") aggr_ds = LeRobotDataset(f"{DUMMY_REPO_ID}_regression_aggr", root=tmp_path / "regression_aggr")
assert_video_timestamps_within_bounds(aggr_ds) assert_video_timestamps_within_bounds(aggr_ds)
# Depth keys must survive the merge for the regression to cover the
# ``is_depth_map`` decoding branch.
assert set(aggr_ds.meta.depth_keys) == set(datasets[0].meta.depth_keys)
depth_keys = set(aggr_ds.meta.depth_keys)
for i in range(len(aggr_ds)): for i in range(len(aggr_ds)):
item = aggr_ds[i] item = aggr_ds[i]
for key in aggr_ds.meta.video_keys: for key in aggr_ds.meta.video_keys:
assert key in item, f"Video key {key} missing from item {i}" assert key in item, f"Video key {key} missing from item {i}"
assert item[key].shape[0] == 3, f"Expected 3 channels for video key {key}" # Depth frames are single-channel (1, H, W) after dequantization;
# standard RGB frames keep the 3-channel layout.
expected_channels = 1 if key in depth_keys else 3
assert item[key].shape[0] == expected_channels, (
f"Expected {expected_channels} channels for video key {key}, got {item[key].shape}"
)
def assert_image_schema_preserved(aggr_ds): def assert_image_schema_preserved(aggr_ds):
@@ -639,25 +698,31 @@ def test_aggregate_image_datasets(tmp_path, lerobot_dataset_factory):
ds_0_num_episodes = 2 ds_0_num_episodes = 2
ds_1_num_episodes = 3 ds_1_num_episodes = 3
# Create two image-based datasets (use_videos=False) # Create two image-based datasets (use_videos=False) with a mix of RGB
# and depth-marked cameras so the depth path is exercised in image mode.
ds_0 = lerobot_dataset_factory( ds_0 = lerobot_dataset_factory(
root=tmp_path / "image_0", root=tmp_path / "image_0",
repo_id=f"{DUMMY_REPO_ID}_image_0", repo_id=f"{DUMMY_REPO_ID}_image_0",
total_episodes=ds_0_num_episodes, total_episodes=ds_0_num_episodes,
total_frames=ds_0_num_frames, total_frames=ds_0_num_frames,
use_videos=False, # Image-based dataset use_videos=False,
camera_features=DUMMY_CAMERA_FEATURES_WITH_DEPTH,
) )
ds_1 = lerobot_dataset_factory( ds_1 = lerobot_dataset_factory(
root=tmp_path / "image_1", root=tmp_path / "image_1",
repo_id=f"{DUMMY_REPO_ID}_image_1", repo_id=f"{DUMMY_REPO_ID}_image_1",
total_episodes=ds_1_num_episodes, total_episodes=ds_1_num_episodes,
total_frames=ds_1_num_frames, total_frames=ds_1_num_frames,
use_videos=False, # Image-based dataset use_videos=False,
camera_features=DUMMY_CAMERA_FEATURES_WITH_DEPTH,
) )
# Verify source datasets have image keys # Verify source datasets have image keys
assert len(ds_0.meta.image_keys) > 0, "ds_0 should have image keys" assert len(ds_0.meta.image_keys) > 0, "ds_0 should have image keys"
assert len(ds_1.meta.image_keys) > 0, "ds_1 should have image keys" assert len(ds_1.meta.image_keys) > 0, "ds_1 should have image keys"
# And that the depth marker actually made it onto an image feature.
assert len(ds_0.meta.depth_keys) > 0, "ds_0 should expose at least one depth key"
assert len(ds_1.meta.depth_keys) > 0, "ds_1 should expose at least one depth key"
# Aggregate the datasets # Aggregate the datasets
aggregate_datasets( aggregate_datasets(
@@ -692,6 +757,7 @@ def test_aggregate_image_datasets(tmp_path, lerobot_dataset_factory):
# Image-specific assertions # Image-specific assertions
assert_image_schema_preserved(aggr_ds) assert_image_schema_preserved(aggr_ds)
assert_image_frames_integrity(aggr_ds, ds_0, ds_1) assert_image_frames_integrity(aggr_ds, ds_0, ds_1)
assert_depth_keys_preserved(aggr_ds, ds_0, ds_1)
# Verify images can be accessed and have correct shape # Verify images can be accessed and have correct shape
sample_item = aggr_ds[0] sample_item = aggr_ds[0]
+30 -9
View File
@@ -35,7 +35,11 @@ from lerobot.utils.constants import OBS_IMAGE, OBS_STATE
def mock_load_image_as_numpy(path, dtype, channel_first): def mock_load_image_as_numpy(path, dtype, channel_first):
return np.ones((3, 32, 32), dtype=dtype) if channel_first else np.ones((32, 32, 3), dtype=dtype) is_depth = "depth" in str(path)
channels = 1 if is_depth else 3
out_dtype = np.uint16 if is_depth else dtype
arr = np.arange(channels * 32 * 32, dtype=out_dtype).reshape(channels, 32, 32)
return arr if channel_first else arr.transpose(1, 2, 0)
@pytest.fixture @pytest.fixture
@@ -168,22 +172,33 @@ def test_get_feature_stats_single_value():
def test_compute_episode_stats(): def test_compute_episode_stats():
depth_key = "observation.images.depth"
episode_data = { episode_data = {
OBS_IMAGE: [f"image_{i}.jpg" for i in range(100)], OBS_IMAGE: [f"image_{i}.jpg" for i in range(100)],
depth_key: [f"depth_{i}.tiff" for i in range(100)],
OBS_STATE: np.random.rand(100, 10), OBS_STATE: np.random.rand(100, 10),
} }
features = { features = {
OBS_IMAGE: {"dtype": "image"}, OBS_IMAGE: {"dtype": "image"},
depth_key: {"dtype": "image", "info": {"is_depth_map": True}},
OBS_STATE: {"dtype": "numeric"}, OBS_STATE: {"dtype": "numeric"},
} }
with patch("lerobot.datasets.compute_stats.load_image_as_numpy", side_effect=mock_load_image_as_numpy): with patch("lerobot.datasets.compute_stats.load_image_as_numpy", side_effect=mock_load_image_as_numpy):
stats = compute_episode_stats(episode_data, features) stats = compute_episode_stats(episode_data, features)
assert OBS_IMAGE in stats and OBS_STATE in stats assert OBS_IMAGE in stats and depth_key in stats and OBS_STATE in stats
assert stats[OBS_IMAGE]["count"].item() == 100 assert stats[OBS_IMAGE]["count"].item() == 100
assert stats[depth_key]["count"].item() == 100
assert stats[OBS_STATE]["count"].item() == 100 assert stats[OBS_STATE]["count"].item() == 100
assert stats[OBS_IMAGE]["mean"].shape == (3, 1, 1) assert stats[OBS_IMAGE]["mean"].shape == (3, 1, 1)
assert stats[depth_key]["mean"].shape == (1, 1, 1)
# Depth keeps raw values: max far exceeds 255, proving no /255 and no uint8 downcast.
assert stats[depth_key]["min"].item() == 0.0
assert stats[depth_key]["max"].item() == 1023.0
# RGB is normalized to [0, 1].
np.testing.assert_allclose(stats[OBS_IMAGE]["min"], 0.0)
np.testing.assert_allclose(stats[OBS_IMAGE]["max"], 1.0)
def test_assert_type_and_shape_valid(): def test_assert_type_and_shape_valid():
@@ -618,25 +633,31 @@ def test_compute_episode_stats_with_custom_quantiles():
def test_compute_episode_stats_with_image_data(): def test_compute_episode_stats_with_image_data():
"""Test quantile computation with image features.""" """Test quantile computation with image features."""
image_paths = [f"image_{i}.jpg" for i in range(50)] image_paths = [f"image_{i}.jpg" for i in range(50)]
depth_paths = [f"depth_{i}.tiff" for i in range(50)]
episode_data = { episode_data = {
"observation.image": image_paths, "observation.image": image_paths,
"observation.images.depth": depth_paths,
"action": np.random.normal(0, 1, (50, 5)), "action": np.random.normal(0, 1, (50, 5)),
} }
features = { features = {
"observation.image": {"dtype": "image"}, "observation.image": {"dtype": "image"},
"observation.images.depth": {"dtype": "image", "info": {"is_depth_map": True}},
"action": {"dtype": "float32", "shape": (5,)}, "action": {"dtype": "float32", "shape": (5,)},
} }
with patch("lerobot.datasets.compute_stats.load_image_as_numpy", side_effect=mock_load_image_as_numpy): with patch("lerobot.datasets.compute_stats.load_image_as_numpy", side_effect=mock_load_image_as_numpy):
stats = compute_episode_stats(episode_data, features) stats = compute_episode_stats(episode_data, features)
# Image quantiles should be normalized and have correct shape # RGB image quantiles should be normalized and per-channel.
assert "q01" in stats["observation.image"] for q in ("q01", "q50", "q99"):
assert "q50" in stats["observation.image"] assert stats["observation.image"][q].shape == (3, 1, 1)
assert "q99" in stats["observation.image"]
assert stats["observation.image"]["q01"].shape == (3, 1, 1) # Depth quantiles are single-channel and kept in raw (un-normalized) units.
assert stats["observation.image"]["q50"].shape == (3, 1, 1) for q in ("q01", "q50", "q99"):
assert stats["observation.image"]["q99"].shape == (3, 1, 1) assert stats["observation.images.depth"][q].shape == (1, 1, 1)
# Depth max stays in raw units (not /255, not uint8-capped); RGB is normalized.
assert stats["observation.images.depth"]["max"].item() == 1023.0
np.testing.assert_allclose(stats["observation.image"]["max"], 1.0)
# Action quantiles should have correct shape # Action quantiles should have correct shape
assert stats["action"]["q01"].shape == (5,) assert stats["action"]["q01"].shape == (5,)
+45 -4
View File
@@ -59,11 +59,13 @@ def _make_dummy_stats(features: dict) -> dict:
stats = {} stats = {}
for key, ft in features.items(): for key, ft in features.items():
if ft["dtype"] in ("image", "video"): if ft["dtype"] in ("image", "video"):
channels = ft["shape"][-1]
stat_shape = (channels, 1, 1)
stats[key] = { stats[key] = {
"max": np.ones((3, 1, 1), dtype=np.float32), "max": np.ones(stat_shape, dtype=np.float32),
"mean": np.full((3, 1, 1), 0.5, dtype=np.float32), "mean": np.full(stat_shape, 0.5, dtype=np.float32),
"min": np.zeros((3, 1, 1), dtype=np.float32), "min": np.zeros(stat_shape, dtype=np.float32),
"std": np.full((3, 1, 1), 0.25, dtype=np.float32), "std": np.full(stat_shape, 0.25, dtype=np.float32),
"count": np.array([5]), "count": np.array([5]),
} }
elif ft["dtype"] in ("float32", "float64", "int64"): elif ft["dtype"] in ("float32", "float64", "int64"):
@@ -142,6 +144,45 @@ def test_create_without_videos_has_no_video_path(tmp_path):
assert meta.video_keys == [] assert meta.video_keys == []
@pytest.mark.parametrize(
("marker_field", "marker_key"),
[
("info", "is_depth_map"),
("info", "video.is_depth_map"),
("video_info", "video.is_depth_map"),
],
ids=["info.is_depth_map", "info.video.is_depth_map_legacy", "video_info.video.is_depth_map_legacy"],
)
def test_depth_keys_property_filters_by_marker(tmp_path, marker_field, marker_key):
"""``depth_keys`` recognises the canonical and the two legacy marker variants."""
depth_feature = {
"dtype": "video",
"shape": (64, 96, 1),
"names": ["height", "width", "channels"],
marker_field: {marker_key: True},
}
features = {
**VIDEO_FEATURES,
"observation.images.laptop_depth": depth_feature,
}
meta = LeRobotDatasetMetadata.create(
repo_id="test/depth_keys",
fps=DEFAULT_FPS,
features=features,
root=tmp_path / f"depth_keys_{marker_field}_{marker_key.replace('.', '_')}",
)
assert set(meta.video_keys) == {"observation.images.laptop", "observation.images.laptop_depth"}
assert meta.depth_keys == ["observation.images.laptop_depth"]
def test_depth_keys_empty_when_no_marker(tmp_path):
meta = LeRobotDatasetMetadata.create(
repo_id="test/no_depth", fps=DEFAULT_FPS, features=VIDEO_FEATURES, root=tmp_path / "no_depth"
)
assert meta.depth_keys == []
def test_create_raises_on_existing_directory(tmp_path): def test_create_raises_on_existing_directory(tmp_path):
"""create() raises if root directory already exists.""" """create() raises if root directory already exists."""
root = tmp_path / "existing" root = tmp_path / "existing"
+134 -10
View File
@@ -24,7 +24,7 @@ import torch
pytest.importorskip("datasets", reason="datasets is required (install lerobot[dataset])") pytest.importorskip("datasets", reason="datasets is required (install lerobot[dataset])")
from lerobot.configs import VideoEncoderConfig from lerobot.configs import DepthEncoderConfig, RGBEncoderConfig
from lerobot.datasets.dataset_tools import ( from lerobot.datasets.dataset_tools import (
add_features, add_features,
convert_image_to_video_dataset, convert_image_to_video_dataset,
@@ -37,7 +37,9 @@ from lerobot.datasets.dataset_tools import (
split_dataset, split_dataset,
) )
from lerobot.datasets.io_utils import load_info from lerobot.datasets.io_utils import load_info
from tests.datasets.test_video_encoding import _add_frames, require_h264, require_libsvtav1 from tests.datasets.test_video_encoding import require_h264, require_hevc, require_libsvtav1
from tests.fixtures.constants import DUMMY_DEPTH_FEATURES, DUMMY_DEPTH_KEY
from tests.fixtures.dataset_factories import add_frames
@pytest.fixture @pytest.fixture
@@ -1251,7 +1253,7 @@ def test_convert_image_to_video_dataset(tmp_path):
dataset=source_dataset, dataset=source_dataset,
output_dir=output_dir, output_dir=output_dir,
repo_id="lerobot/pusht_video", repo_id="lerobot/pusht_video",
camera_encoder=VideoEncoderConfig( rgb_encoder=RGBEncoderConfig(
vcodec="libsvtav1", vcodec="libsvtav1",
pix_fmt="yuv420p", pix_fmt="yuv420p",
g=2, g=2,
@@ -1332,9 +1334,131 @@ def test_convert_image_to_video_dataset_subset_episodes(tmp_path):
shutil.rmtree(output_dir) shutil.rmtree(output_dir)
@require_libsvtav1
@require_hevc
def test_convert_image_to_video_dataset_depth(tmp_path, empty_lerobot_dataset_factory):
"""Depth image features convert to depth videos using the depth encoder.
Mirrors :func:`test_convert_image_to_video_dataset` but with a small local
image dataset that mixes an RGB camera with a depth camera, so the
``depth_keys`` ``depth_encoder`` routing and ``is_depth_map`` preservation
are exercised end-to-end.
"""
features = {
"action": {"dtype": "float32", "shape": (2,), "names": ["a", "b"]},
"observation.images.cam": {
"dtype": "image",
"shape": (64, 96, 3),
"names": ["height", "width", "channels"],
},
"observation.images.depth": {
"dtype": "image",
"shape": (64, 96, 1),
"names": ["height", "width", "channels"],
"info": {"is_depth_map": True},
},
}
source_dataset = empty_lerobot_dataset_factory(
root=tmp_path / "img_ds",
features=features,
use_videos=False,
)
add_frames(source_dataset, num_frames=4)
source_dataset.save_episode()
source_dataset.finalize()
# Source is an image dataset with the depth marker on the depth camera.
assert len(source_dataset.meta.video_keys) == 0
assert "observation.images.depth" in source_dataset.meta.depth_keys
output_dir = tmp_path / "video_ds"
with (
patch("lerobot.datasets.dataset_metadata.get_safe_version") as mock_get_safe_version,
patch("lerobot.datasets.dataset_metadata.snapshot_download") as mock_snapshot_download,
):
mock_get_safe_version.return_value = "v3.0"
mock_snapshot_download.return_value = str(output_dir)
# Use non-default quantization params so the persisted metadata must
# come from the depth encoder (not RGB encoder defaults).
depth_encoder = DepthEncoderConfig(
vcodec="hevc",
pix_fmt="gray12le",
g=2,
crf=30,
depth_min=0.05,
depth_max=8.0,
shift=2.0,
use_log=False,
)
video_dataset = convert_image_to_video_dataset(
dataset=source_dataset,
output_dir=output_dir,
repo_id="dummy/depth_video",
rgb_encoder=RGBEncoderConfig(vcodec="libsvtav1", pix_fmt="yuv420p", g=2, crf=30),
depth_encoder=depth_encoder,
num_workers=1,
)
# Both cameras are now videos, and the depth marker survived the conversion.
assert "observation.images.cam" in video_dataset.meta.video_keys
assert "observation.images.depth" in video_dataset.meta.video_keys
assert "observation.images.depth" in video_dataset.meta.depth_keys
assert "observation.images.cam" not in video_dataset.meta.depth_keys
depth_path = video_dataset.root / video_dataset.meta.get_video_file_path(0, "observation.images.depth")
assert depth_path.exists(), f"Depth video file should exist: {depth_path}"
# The persisted depth-video metadata must carry the depth quantization params
# from the depth encoder (so frames dequantize correctly on read), and the RGB
# camera must not be marked as a depth map.
persisted_info = load_info(video_dataset.root)
depth_info = persisted_info.features["observation.images.depth"]["info"]
assert depth_info["is_depth_map"] is True
assert DepthEncoderConfig.from_video_info(depth_info) == depth_encoder
cam_info = persisted_info.features["observation.images.cam"]["info"]
assert cam_info.get("is_depth_map") is False
assert "video.codec" in cam_info
# ─── reencode_dataset ───────────────────────────────────────────────── # ─── reencode_dataset ─────────────────────────────────────────────────
@require_hevc
def test_reencode_dataset_depth_uses_depth_encoder(tmp_path, empty_lerobot_dataset_factory):
"""Depth videos are re-encoded with the depth encoder and keep their depth metadata.
Depth-focused companion to :func:`test_reencode_dataset_multi_key_multiprocessing`.
"""
initial_cfg = DepthEncoderConfig(vcodec="hevc", pix_fmt="gray12le", g=2, crf=30)
dataset = empty_lerobot_dataset_factory(
root=tmp_path / "ds",
features=DUMMY_DEPTH_FEATURES,
use_videos=True,
depth_encoder=initial_cfg,
)
add_frames(dataset, num_frames=4)
dataset.save_episode()
dataset.finalize()
assert DUMMY_DEPTH_KEY in dataset.meta.depth_keys
target_cfg = DepthEncoderConfig(vcodec="hevc", pix_fmt="gray12le", g=6, crf=23)
result = reencode_dataset(dataset, depth_encoder=target_cfg, num_workers=0)
assert result is dataset
persisted_info = load_info(dataset.root)
depth_info = persisted_info.features[DUMMY_DEPTH_KEY].get("info", {})
# Re-encode applied the new codec parameters to the depth video ...
assert DepthEncoderConfig.from_video_info(depth_info) == target_cfg
# ... while preserving the depth marker.
assert depth_info["is_depth_map"] is True
@require_libsvtav1 @require_libsvtav1
@require_h264 @require_h264
def test_reencode_dataset_multi_key_multiprocessing( def test_reencode_dataset_multi_key_multiprocessing(
@@ -1342,29 +1466,29 @@ def test_reencode_dataset_multi_key_multiprocessing(
): ):
"""Re-encode a two-camera dataset with num_workers=2 and verify metadata refresh.""" """Re-encode a two-camera dataset with num_workers=2 and verify metadata refresh."""
features = features_factory(use_videos=True) features = features_factory(use_videos=True)
initial_cfg = VideoEncoderConfig(vcodec="libsvtav1", g=2, crf=30, preset=12) initial_cfg = RGBEncoderConfig(vcodec="libsvtav1", g=2, crf=30, preset=12)
dataset = empty_lerobot_dataset_factory( dataset = empty_lerobot_dataset_factory(
root=tmp_path / "ds", root=tmp_path / "ds",
features=features, features=features,
use_videos=True, use_videos=True,
camera_encoder=initial_cfg, rgb_encoder=initial_cfg,
) )
_add_frames(dataset, num_frames=4) add_frames(dataset, num_frames=4)
dataset.save_episode() dataset.save_episode()
_add_frames(dataset, num_frames=4) add_frames(dataset, num_frames=4)
dataset.save_episode() dataset.save_episode()
dataset.finalize() dataset.finalize()
assert len(dataset.meta.video_keys) == 2 assert len(dataset.meta.video_keys) == 2
target_cfg = VideoEncoderConfig(vcodec="h264", g=6, crf=23, pix_fmt="yuv420p") target_cfg = RGBEncoderConfig(vcodec="h264", g=6, crf=23, pix_fmt="yuv420p")
result = reencode_dataset(dataset, camera_encoder=target_cfg, num_workers=2) result = reencode_dataset(dataset, rgb_encoder=target_cfg, num_workers=2)
assert result is dataset assert result is dataset
persisted_info = load_info(dataset.root) persisted_info = load_info(dataset.root)
for vk in dataset.meta.video_keys: for vk in dataset.meta.video_keys:
persisted_encoder = VideoEncoderConfig.from_video_info(persisted_info.features[vk].get("info", {})) persisted_encoder = RGBEncoderConfig.from_video_info(persisted_info.features[vk].get("info", {}))
assert persisted_encoder == target_cfg assert persisted_encoder == target_cfg
+7 -7
View File
@@ -53,8 +53,8 @@ def _make_frame(features: dict, task: str = "Dummy task") -> dict:
# ── Existing encode_video_worker tests ─────────────────────────────── # ── Existing encode_video_worker tests ───────────────────────────────
def test_encode_video_worker_forwards_camera_encoder(tmp_path): def test_encode_video_worker_forwards_video_encoder(tmp_path):
"""_encode_video_worker forwards camera_encoder to encode_video_frames.""" """_encode_video_worker forwards video_encoder to encode_video_frames."""
video_key = "observation.images.laptop" video_key = "observation.images.laptop"
fpath = DEFAULT_IMAGE_PATH.format(image_key=video_key, episode_index=0, frame_index=0) fpath = DEFAULT_IMAGE_PATH.format(image_key=video_key, episode_index=0, frame_index=0)
img_dir = tmp_path / Path(fpath).parent img_dir = tmp_path / Path(fpath).parent
@@ -74,16 +74,16 @@ def test_encode_video_worker_forwards_camera_encoder(tmp_path):
0, 0,
tmp_path, tmp_path,
fps=30, fps=30,
camera_encoder=VideoEncoderConfig(vcodec="h264", preset=None), video_encoder=VideoEncoderConfig(vcodec="h264", preset=None),
encoder_threads=4, encoder_threads=4,
) )
assert captured_kwargs["camera_encoder"].vcodec == "h264" assert captured_kwargs["video_encoder"].vcodec == "h264"
assert captured_kwargs["encoder_threads"] == 4 assert captured_kwargs["encoder_threads"] == 4
def test_encode_video_worker_default_camera_encoder(tmp_path): def test_encode_video_worker_default_video_encoder(tmp_path):
"""_encode_video_worker passes None camera_encoder which encode_video_frames defaults.""" """_encode_video_worker passes None video_encoder which encode_video_frames defaults."""
video_key = "observation.images.laptop" video_key = "observation.images.laptop"
fpath = DEFAULT_IMAGE_PATH.format(image_key=video_key, episode_index=0, frame_index=0) fpath = DEFAULT_IMAGE_PATH.format(image_key=video_key, episode_index=0, frame_index=0)
img_dir = tmp_path / Path(fpath).parent img_dir = tmp_path / Path(fpath).parent
@@ -100,7 +100,7 @@ def test_encode_video_worker_default_camera_encoder(tmp_path):
with patch("lerobot.datasets.dataset_writer.encode_video_frames", side_effect=mock_encode): with patch("lerobot.datasets.dataset_writer.encode_video_frames", side_effect=mock_encode):
_encode_video_worker(video_key, 0, tmp_path, fps=30) _encode_video_worker(video_key, 0, tmp_path, fps=30)
assert captured_kwargs["camera_encoder"] is None assert captured_kwargs["video_encoder"] is None
assert captured_kwargs["encoder_threads"] is None assert captured_kwargs["encoder_threads"] is None
+4
View File
@@ -1534,6 +1534,10 @@ def test_valid_video_codecs_constant():
assert "auto" in VALID_VIDEO_CODECS assert "auto" in VALID_VIDEO_CODECS
assert "h264_videotoolbox" in VALID_VIDEO_CODECS assert "h264_videotoolbox" in VALID_VIDEO_CODECS
assert "h264_nvenc" in VALID_VIDEO_CODECS assert "h264_nvenc" in VALID_VIDEO_CODECS
assert "h264_vaapi" in VALID_VIDEO_CODECS
assert "h264_qsv" in VALID_VIDEO_CODECS
assert "hevc_videotoolbox" in VALID_VIDEO_CODECS
assert "hevc_nvenc" in VALID_VIDEO_CODECS
assert len(VALID_VIDEO_CODECS) == 10 assert len(VALID_VIDEO_CODECS) == 10
+247
View File
@@ -0,0 +1,247 @@
"""Tests for the depth-integration feature.
Covers:
- ``depth_utils`` quantize/dequantize round-trips and backend agreement.
- Image-writer support for single-channel depth.
- Hardware-feature depth flag routing.
- Feature-to-file-format routing through the dataset writer.
Depth metadata detection on ``LeRobotDatasetMetadata.depth_keys`` lives in
``test_dataset_metadata.py``. Depth video encoding/decoding lives in
``test_video_encoding.py``.
"""
from pathlib import Path
import pytest
pytest.importorskip("av", reason="av is required (install lerobot[dataset])")
import av
import numpy as np
import PIL.Image
import torch
from lerobot.configs import DepthEncoderConfig
from lerobot.configs.video import (
DEFAULT_DEPTH_MAX,
DEFAULT_DEPTH_MIN,
DEPTH_METER_UNIT,
DEPTH_MILLIMETER_UNIT,
DEPTH_QMAX,
)
from lerobot.datasets.depth_utils import dequantize_depth, quantize_depth
from lerobot.datasets.image_writer import image_array_to_pil_image, write_image
from tests.fixtures.constants import (
DEFAULT_FPS,
DUMMY_CAMERA_FEATURES,
DUMMY_CAMERA_FEATURES_WITH_DEPTH,
DUMMY_CHW,
DUMMY_DEPTH_CAMERA_FEATURES,
DUMMY_REPO_ID,
)
from tests.fixtures.dataset_factories import add_frames
_, H, W = DUMMY_CHW
def _depth_metres_ramp() -> np.ndarray:
"""Linearly-spaced float32 depth in metres covering the default range."""
return np.linspace(DEFAULT_DEPTH_MIN, DEFAULT_DEPTH_MAX, H * W, dtype=np.float32).reshape(H, W)
# ── 1. Quantize / dequantize round-trips ──────────────────────────────
class TestQuantizeDequantize:
"""Numerical contract of ``quantize_depth`` / ``dequantize_depth``."""
@pytest.mark.parametrize("use_log", [False, True])
@pytest.mark.parametrize("output_unit", [DEPTH_METER_UNIT, DEPTH_MILLIMETER_UNIT])
@pytest.mark.parametrize("output_channel_last", [False, True])
def test_roundtrip(self, use_log, output_unit, output_channel_last):
"""quantize → dequantize recovers depth; layout and unit are honored."""
depth = _depth_metres_ramp()
quantized = quantize_depth(depth, use_log=use_log, video_backend=None)
recovered = dequantize_depth(
quantized,
use_log=use_log,
output_unit=output_unit,
output_tensor=False,
output_channel_last=output_channel_last,
)
expected_shape = (H, W, 1) if output_channel_last else (1, H, W)
assert recovered.shape == expected_shape
recovered_m = recovered.astype(np.float32)
if output_unit == DEPTH_MILLIMETER_UNIT:
recovered_m = recovered_m / 1000.0
recovered_2d = recovered_m[..., 0] if output_channel_last else recovered_m[0]
if use_log:
# Log mode: tighter near-range error than far-range (the whole point).
near = depth < 1.0
far = depth > 8.0
err_near = np.abs(recovered_2d[near] - depth[near])
err_far = np.abs(recovered_2d[far] - depth[far])
assert err_near.mean() < err_far.mean()
else:
# Linear mode: bounded by quant step + 1 mm of unit-conversion rounding.
tol = (DEFAULT_DEPTH_MAX - DEFAULT_DEPTH_MIN) / DEPTH_QMAX + 1e-3
np.testing.assert_allclose(recovered_2d, depth, atol=tol)
@pytest.mark.parametrize("use_log", [False, True])
@pytest.mark.parametrize("output_unit", [DEPTH_METER_UNIT, DEPTH_MILLIMETER_UNIT])
def test_numpy_torch_agree(self, use_log, output_unit):
"""Batched torch path produces the same values as the numpy path."""
batch_size = 3
per_frame = np.linspace(0, DEPTH_QMAX, H * W, dtype=np.uint16).reshape(H, W)
batch_np = np.broadcast_to(per_frame[None, None, ...], (batch_size, 1, H, W)).copy()
batch_t = torch.from_numpy(batch_np.astype(np.int32)) # torch.uint16 support is patchy.
ref = dequantize_depth(batch_np, use_log=use_log, output_unit=output_unit, output_tensor=False)
out = dequantize_depth(batch_t, use_log=use_log, output_unit=output_unit, output_tensor=True)
assert isinstance(out, torch.Tensor)
assert out.shape == (batch_size, 1, H, W)
# ``m``: float32 noise (~10 µm in log mode, after ``exp``) — still 200× below the ~2 mm quant step.
# ``mm`` + tensor stays in float32 (no uint16 round-trip), so allow 1 mm slop.
atol = 1e-5 if output_unit == DEPTH_METER_UNIT else 1.0
np.testing.assert_allclose(out.cpu().numpy().astype(np.float64), ref.astype(np.float64), atol=atol)
@pytest.mark.parametrize(
"input_shape,output_shape",
[
((H, W), (1, H, W)),
((1, H, W), (1, H, W)),
((H, W, 1), (1, H, W)),
((3, 1, H, W), (3, 1, H, W)),
((3, H, W, 1), (3, 1, H, W)),
],
)
def test_input_layouts_accepted(self, input_shape, output_shape):
"""All documented input layouts decode to the channel-first default."""
quantized = np.full(input_shape, DEPTH_QMAX // 2, dtype=np.uint16)
out = dequantize_depth(quantized, output_unit=DEPTH_METER_UNIT, output_tensor=False)
assert out.shape == output_shape
def test_pyav_frame_roundtrip(self):
"""quantize → av.VideoFrame → dequantize works."""
depth = _depth_metres_ramp()
frame = quantize_depth(depth, use_log=False, video_backend="pyav")
assert isinstance(frame, av.VideoFrame)
recovered = dequantize_depth(frame, use_log=False, output_unit=DEPTH_METER_UNIT, output_tensor=False)
assert recovered.shape == (1, H, W)
tol = (DEFAULT_DEPTH_MAX - DEFAULT_DEPTH_MIN) / DEPTH_QMAX + 1e-3
np.testing.assert_allclose(recovered[0], depth, atol=tol)
def test_invalid_log_params_raises(self):
with pytest.raises(ValueError, match=r"depth_min \+ shift must be positive"):
quantize_depth(_depth_metres_ramp(), depth_min=1.0, shift=-2.0, use_log=True, video_backend=None)
# ── 2. Image writer depth support ─────────────────────────────────────
class TestImageWriterDepth:
"""``image_array_to_pil_image`` and ``write_image`` for depth maps."""
@pytest.mark.parametrize("dtype,expected_mode", [(np.uint16, "I;16"), (np.float32, "F")])
@pytest.mark.parametrize("shape", [(H, W), (H, W, 1), (1, H, W)])
def test_pil_depth_modes_and_squeeze(self, dtype, expected_mode, shape):
"""Single-channel depth converts to PIL with the right mode and (W, H) size."""
arr = np.zeros(shape, dtype=dtype)
img = image_array_to_pil_image(arr)
assert img.mode == expected_mode
assert img.size == (W, H)
def test_write_image_tiff_roundtrip(self, tmp_path):
"""uint16 depth round-trips through .tiff."""
arr = np.arange(H * W, dtype=np.uint16).reshape(H, W)
fpath = tmp_path / "depth.tiff"
write_image(arr, fpath)
with PIL.Image.open(fpath) as loaded:
recovered = np.array(loaded)
np.testing.assert_array_equal(recovered, arr)
# ── 3. Hardware-feature → depth flag ──────────────────────────────────
class TestHwToDatasetFeaturesDepth:
"""``hw_to_dataset_features`` flags single-channel cameras as depth."""
@pytest.mark.parametrize("channels,is_depth", [(1, True), (3, False)])
def test_depth_marker_by_channels(self, channels, is_depth):
from lerobot.utils.feature_utils import hw_to_dataset_features
features = hw_to_dataset_features({"cam": (480, 640, channels)}, prefix="observation")
assert features["observation.images.cam"]["info"]["is_depth_map"] is is_depth
def test_invalid_channel_count_raises(self):
from lerobot.utils.feature_utils import hw_to_dataset_features
with pytest.raises(ValueError, match="Expected a 3-tuple"):
hw_to_dataset_features({"cam": (480, 640, 2)}, prefix="observation")
# ── 4. Feature-to-file-format routing ────────────────────────────────
# Keys derived from DUMMY_CAMERA_FEATURES_WITH_DEPTH; pick one RGB and the depth camera.
RGB_KEY = next(iter(DUMMY_CAMERA_FEATURES))
DEPTH_KEY = next(iter(DUMMY_DEPTH_CAMERA_FEATURES))
class TestFeatureFileRouting:
"""Depth vs RGB features route to the correct file format."""
NUM_FRAMES = 5
def test_image_mode_depth_tiff_rgb_png(self, tmp_path, features_factory):
"""Without video encoding: depth → .tiff, RGB → .png."""
from lerobot.datasets.lerobot_dataset import LeRobotDataset
features = features_factory(camera_features=DUMMY_CAMERA_FEATURES_WITH_DEPTH, use_videos=False)
dataset = LeRobotDataset.create(
repo_id=DUMMY_REPO_ID,
fps=DEFAULT_FPS,
features=features,
root=tmp_path / "ds",
use_videos=False,
)
add_frames(dataset, num_frames=self.NUM_FRAMES)
buf = dataset.writer.episode_buffer
assert all(Path(p).suffix == ".tiff" for p in buf[DEPTH_KEY])
assert all(Path(p).suffix == ".png" for p in buf[RGB_KEY])
dataset.save_episode()
dataset.finalize()
def test_video_mode_depth_uses_depth_encoder(self, tmp_path, features_factory):
"""With streaming video encoding: depth → DepthEncoderConfig, RGB does not."""
from lerobot.datasets.lerobot_dataset import LeRobotDataset
features = features_factory(camera_features=DUMMY_CAMERA_FEATURES_WITH_DEPTH, use_videos=True)
dataset = LeRobotDataset.create(
repo_id=DUMMY_REPO_ID,
fps=DEFAULT_FPS,
features=features,
root=tmp_path / "ds",
use_videos=True,
streaming_encoding=True,
)
add_frames(dataset, num_frames=self.NUM_FRAMES)
encoder = dataset.writer._streaming_encoder
assert encoder is not None
assert isinstance(encoder._threads[DEPTH_KEY].video_encoder, DepthEncoderConfig)
assert not isinstance(encoder._threads[RGB_KEY].video_encoder, DepthEncoderConfig)
dataset.save_episode()
dataset.finalize()
+2 -2
View File
@@ -94,7 +94,7 @@ def test_image_array_to_pil_image_pytorch_format(img_array_factory):
def test_image_array_to_pil_image_single_channel(img_array_factory): def test_image_array_to_pil_image_single_channel(img_array_factory):
img_array = img_array_factory(channels=1) img_array = img_array_factory(channels=1)
with pytest.raises(NotImplementedError): with pytest.raises(ValueError, match="Unsupported single-channel image dtype"):
image_array_to_pil_image(img_array) image_array_to_pil_image(img_array)
@@ -344,7 +344,7 @@ def test_with_different_image_formats(tmp_path, img_array_factory):
writer = AsyncImageWriter() writer = AsyncImageWriter()
try: try:
image_array = img_array_factory() image_array = img_array_factory()
formats = ["png", "jpeg", "bmp"] formats = ["png", "tiff", "tif"]
for fmt in formats: for fmt in formats:
fpath = tmp_path / f"test_image.{fmt}" fpath = tmp_path / f"test_image.{fmt}"
write_image(image_array, fpath) write_image(image_array, fpath)
+20 -25
View File
@@ -26,7 +26,7 @@ pytest.importorskip("av", reason="av is required (install lerobot[dataset])")
import av # noqa: E402 import av # noqa: E402
from lerobot.configs import VideoEncoderConfig from lerobot.configs import RGBEncoderConfig
from lerobot.datasets.pyav_utils import get_codec from lerobot.datasets.pyav_utils import get_codec
from lerobot.datasets.video_utils import ( from lerobot.datasets.video_utils import (
StreamingVideoEncoder, StreamingVideoEncoder,
@@ -57,13 +57,11 @@ class TestCameraEncoderThread:
result_queue: queue.Queue = queue.Queue(maxsize=1) result_queue: queue.Queue = queue.Queue(maxsize=1)
stop_event = threading.Event() stop_event = threading.Event()
enc_cfg = VideoEncoderConfig(vcodec="libsvtav1", pix_fmt="yuv420p", g=2, crf=30, preset=13) enc_cfg = RGBEncoderConfig(vcodec="libsvtav1", pix_fmt="yuv420p", g=2, crf=30, preset=13)
encoder_thread = _CameraEncoderThread( encoder_thread = _CameraEncoderThread(
video_path=video_path, video_path=video_path,
fps=fps, fps=fps,
vcodec=enc_cfg.vcodec, video_encoder=enc_cfg,
pix_fmt=enc_cfg.pix_fmt,
codec_options=enc_cfg.get_codec_options(as_strings=True),
frame_queue=frame_queue, frame_queue=frame_queue,
result_queue=result_queue, result_queue=result_queue,
stop_event=stop_event, stop_event=stop_event,
@@ -108,13 +106,11 @@ class TestCameraEncoderThread:
result_queue: queue.Queue = queue.Queue(maxsize=1) result_queue: queue.Queue = queue.Queue(maxsize=1)
stop_event = threading.Event() stop_event = threading.Event()
enc_cfg = VideoEncoderConfig(vcodec="libsvtav1", pix_fmt="yuv420p", g=2, crf=30, preset=13) enc_cfg = RGBEncoderConfig(vcodec="libsvtav1", pix_fmt="yuv420p", g=2, crf=30, preset=13)
encoder_thread = _CameraEncoderThread( encoder_thread = _CameraEncoderThread(
video_path=video_path, video_path=video_path,
fps=fps, fps=fps,
vcodec=enc_cfg.vcodec, video_encoder=enc_cfg,
pix_fmt=enc_cfg.pix_fmt,
codec_options=enc_cfg.get_codec_options(as_strings=True),
frame_queue=frame_queue, frame_queue=frame_queue,
result_queue=result_queue, result_queue=result_queue,
stop_event=stop_event, stop_event=stop_event,
@@ -142,13 +138,11 @@ class TestCameraEncoderThread:
result_queue: queue.Queue = queue.Queue(maxsize=1) result_queue: queue.Queue = queue.Queue(maxsize=1)
stop_event = threading.Event() stop_event = threading.Event()
enc_cfg = VideoEncoderConfig(vcodec="libsvtav1", pix_fmt="yuv420p", g=2, crf=30, preset=13) enc_cfg = RGBEncoderConfig(vcodec="libsvtav1", pix_fmt="yuv420p", g=2, crf=30, preset=13)
encoder_thread = _CameraEncoderThread( encoder_thread = _CameraEncoderThread(
video_path=video_path, video_path=video_path,
fps=fps, fps=fps,
vcodec=enc_cfg.vcodec, video_encoder=enc_cfg,
pix_fmt=enc_cfg.pix_fmt,
codec_options=enc_cfg.get_codec_options(as_strings=True),
frame_queue=frame_queue, frame_queue=frame_queue,
result_queue=result_queue, result_queue=result_queue,
stop_event=stop_event, stop_event=stop_event,
@@ -171,15 +165,15 @@ class TestCameraEncoderThread:
class TestStreamingVideoEncoder: class TestStreamingVideoEncoder:
def _make_encoder_config(self, **kwargs): def _make_encoder_config(self, **kwargs):
"""Helper to build a VideoEncoderConfig.""" """Helper to build an RGBEncoderConfig."""
return VideoEncoderConfig(**kwargs) return RGBEncoderConfig(**kwargs)
def test_single_camera_episode(self, tmp_path): def test_single_camera_episode(self, tmp_path):
"""Test encoding a single camera episode.""" """Test encoding a single camera episode."""
video_keys = [f"{OBS_IMAGES}.laptop"] video_keys = [f"{OBS_IMAGES}.laptop"]
encoder = StreamingVideoEncoder( encoder = StreamingVideoEncoder(
fps=30, fps=30,
camera_encoder=self._make_encoder_config( rgb_encoder=self._make_encoder_config(
vcodec="libsvtav1", pix_fmt="yuv420p", g=2, crf=30, preset=13 vcodec="libsvtav1", pix_fmt="yuv420p", g=2, crf=30, preset=13
), ),
) )
@@ -211,7 +205,7 @@ class TestStreamingVideoEncoder:
video_keys = [f"{OBS_IMAGES}.laptop", f"{OBS_IMAGES}.phone"] video_keys = [f"{OBS_IMAGES}.laptop", f"{OBS_IMAGES}.phone"]
encoder = StreamingVideoEncoder( encoder = StreamingVideoEncoder(
fps=30, fps=30,
camera_encoder=self._make_encoder_config(vcodec="libsvtav1", pix_fmt="yuv420p", g=2, crf=30), rgb_encoder=self._make_encoder_config(vcodec="libsvtav1", pix_fmt="yuv420p", g=2, crf=30),
) )
encoder.start_episode(video_keys, tmp_path) encoder.start_episode(video_keys, tmp_path)
@@ -237,7 +231,7 @@ class TestStreamingVideoEncoder:
video_keys = [f"{OBS_IMAGES}.cam"] video_keys = [f"{OBS_IMAGES}.cam"]
encoder = StreamingVideoEncoder( encoder = StreamingVideoEncoder(
fps=30, fps=30,
camera_encoder=self._make_encoder_config(vcodec="libsvtav1", pix_fmt="yuv420p", g=2, crf=30), rgb_encoder=self._make_encoder_config(vcodec="libsvtav1", pix_fmt="yuv420p", g=2, crf=30),
) )
for ep in range(3): for ep in range(3):
@@ -263,7 +257,7 @@ class TestStreamingVideoEncoder:
video_keys = [f"{OBS_IMAGES}.cam"] video_keys = [f"{OBS_IMAGES}.cam"]
encoder = StreamingVideoEncoder( encoder = StreamingVideoEncoder(
fps=30, fps=30,
camera_encoder=self._make_encoder_config(vcodec="libsvtav1", pix_fmt="yuv420p", g=2, crf=30), rgb_encoder=self._make_encoder_config(vcodec="libsvtav1", pix_fmt="yuv420p", g=2, crf=30),
) )
encoder.start_episode(video_keys, tmp_path) encoder.start_episode(video_keys, tmp_path)
@@ -309,7 +303,7 @@ class TestStreamingVideoEncoder:
video_keys = [f"{OBS_IMAGES}.cam"] video_keys = [f"{OBS_IMAGES}.cam"]
encoder = StreamingVideoEncoder( encoder = StreamingVideoEncoder(
fps=30, fps=30,
camera_encoder=self._make_encoder_config( rgb_encoder=self._make_encoder_config(
vcodec="libsvtav1", pix_fmt="yuv420p", g=2, crf=30, preset=13 vcodec="libsvtav1", pix_fmt="yuv420p", g=2, crf=30, preset=13
), ),
) )
@@ -346,7 +340,7 @@ class TestStreamingVideoEncoder:
video_keys = [f"{OBS_IMAGES}.cam1", f"{OBS_IMAGES}.cam2"] video_keys = [f"{OBS_IMAGES}.cam1", f"{OBS_IMAGES}.cam2"]
encoder = StreamingVideoEncoder( encoder = StreamingVideoEncoder(
fps=30, fps=30,
camera_encoder=self._make_encoder_config(vcodec="libsvtav1", pix_fmt="yuv420p", g=2, crf=30), rgb_encoder=self._make_encoder_config(vcodec="libsvtav1", pix_fmt="yuv420p", g=2, crf=30),
) )
encoder.start_episode(video_keys, tmp_path) encoder.start_episode(video_keys, tmp_path)
@@ -375,7 +369,7 @@ class TestStreamingVideoEncoder:
def test_encoder_threads_passed_to_thread(self, tmp_path): def test_encoder_threads_passed_to_thread(self, tmp_path):
"""Test that encoder_threads is stored and passed through to encoder threads.""" """Test that encoder_threads is stored and passed through to encoder threads."""
video_keys = [f"{OBS_IMAGES}.cam"] video_keys = [f"{OBS_IMAGES}.cam"]
cfg = VideoEncoderConfig( cfg = RGBEncoderConfig(
vcodec="libsvtav1", vcodec="libsvtav1",
pix_fmt="yuv420p", pix_fmt="yuv420p",
g=2, g=2,
@@ -383,7 +377,7 @@ class TestStreamingVideoEncoder:
) )
encoder = StreamingVideoEncoder( encoder = StreamingVideoEncoder(
fps=30, fps=30,
camera_encoder=cfg, rgb_encoder=cfg,
encoder_threads=2, encoder_threads=2,
) )
assert encoder._encoder_threads == 2 assert encoder._encoder_threads == 2
@@ -391,7 +385,8 @@ class TestStreamingVideoEncoder:
# Verify codec options include thread tuning for libsvtav1 (lp=…) # Verify codec options include thread tuning for libsvtav1 (lp=…)
thread = encoder._threads[f"{OBS_IMAGES}.cam"] thread = encoder._threads[f"{OBS_IMAGES}.cam"]
assert "svtav1-params" in thread.codec_options or "threads" in thread.codec_options codec_opts = thread.video_encoder.get_codec_options(encoder_threads=thread.encoder_threads)
assert "svtav1-params" in codec_opts or "threads" in codec_opts
# Feed some frames and finish to ensure it works end-to-end # Feed some frames and finish to ensure it works end-to-end
num_frames = 10 num_frames = 10
@@ -422,7 +417,7 @@ class TestStreamingVideoEncoder:
video_keys = [f"{OBS_IMAGES}.cam"] video_keys = [f"{OBS_IMAGES}.cam"]
encoder = StreamingVideoEncoder( encoder = StreamingVideoEncoder(
fps=30, fps=30,
camera_encoder=self._make_encoder_config( rgb_encoder=self._make_encoder_config(
vcodec="libsvtav1", pix_fmt="yuv420p", g=2, crf=30, preset=13 vcodec="libsvtav1", pix_fmt="yuv420p", g=2, crf=30, preset=13
), ),
queue_maxsize=1, queue_maxsize=1,
+359 -128
View File
@@ -26,7 +26,7 @@ pytest.importorskip("av", reason="av is required (install lerobot[dataset])")
import av # noqa: E402 import av # noqa: E402
from lerobot.configs import VALID_VIDEO_CODECS, VideoEncoderConfig from lerobot.configs import VALID_VIDEO_CODECS, DepthEncoderConfig, RGBEncoderConfig, VideoEncoderConfig
from lerobot.datasets.image_writer import write_image from lerobot.datasets.image_writer import write_image
from lerobot.datasets.lerobot_dataset import LeRobotDataset from lerobot.datasets.lerobot_dataset import LeRobotDataset
from lerobot.datasets.pyav_utils import get_codec from lerobot.datasets.pyav_utils import get_codec
@@ -37,7 +37,15 @@ from lerobot.datasets.video_utils import (
get_video_info, get_video_info,
reencode_video, reencode_video,
) )
from tests.fixtures.constants import DUMMY_VIDEO_INFO from tests.fixtures.constants import (
DUMMY_DEPTH_FEATURES,
DUMMY_DEPTH_KEY,
DUMMY_DEPTH_VIDEO_INFO_FULL,
DUMMY_VIDEO_FEATURES,
DUMMY_VIDEO_INFO,
DUMMY_VIDEO_KEY,
)
from tests.fixtures.dataset_factories import add_frames
# Per-codec skip markers — validation tests only fire when the codec is available # Per-codec skip markers — validation tests only fire when the codec is available
@@ -48,19 +56,74 @@ def _require_encoder(vcodec: str) -> pytest.MarkDecorator:
require_libsvtav1 = _require_encoder("libsvtav1") require_libsvtav1 = _require_encoder("libsvtav1")
require_h264 = _require_encoder("h264") require_h264 = _require_encoder("h264")
require_hevc = _require_encoder("hevc")
require_videotoolbox = _require_encoder("h264_videotoolbox") require_videotoolbox = _require_encoder("h264_videotoolbox")
require_nvenc = _require_encoder("h264_nvenc") require_nvenc = _require_encoder("h264_nvenc")
require_vaapi = _require_encoder("h264_vaapi") require_vaapi = _require_encoder("h264_vaapi")
require_qsv = _require_encoder("h264_qsv") require_qsv = _require_encoder("h264_qsv")
# ─── VideoEncoderConfig / codec options ────────────────────────────── TEST_ARTIFACTS_DIR = Path(__file__).parent.parent / "artifacts" / "encoded_videos"
def _write_color_frames(imgs_dir: Path, num_frames: int = 4, height: int = 64, width: int = 96) -> None:
imgs_dir.mkdir(parents=True, exist_ok=True)
for i in range(num_frames):
arr = np.random.randint(0, 256, (height, width, 3), dtype=np.uint8)
write_image(arr, imgs_dir / f"frame-{i:06d}.png")
def _write_depth_frames(imgs_dir: Path, num_frames: int = 4, height: int = 64, width: int = 96) -> None:
"""Write synthetic uint16 depth TIFFs (millimetres) for depth encoder tests.
Uses a smooth linear ramp + per-frame offset (not white noise) so HEVC Main 12
on ``gray12le`` compresses well. Values span ~100 mm to 10 m, covering most
of the default ``[DEPTH_MIN, DEPTH_MAX]`` metres range after
``quantize_depth(input_unit="auto"="mm")``.
"""
imgs_dir.mkdir(parents=True, exist_ok=True)
base = np.linspace(100.0, 10_000.0, height * width, dtype=np.float32).reshape(height, width)
for i in range(num_frames):
arr = (base + 50.0 * i).clip(0, 65535).astype(np.uint16)
write_image(arr, imgs_dir / f"frame-{i:06d}.tiff")
def _encode_video(
path: Path,
num_frames: int = 4,
fps: int = 30,
cfg: VideoEncoderConfig | None = None,
depth: bool = False,
) -> Path:
"""Write synthetic frames to a temp dir and encode them to ``path``.
``depth=False`` writes uint8 RGB PNG noise and encodes with ``cfg``
(defaulting to the library default). ``depth=True`` writes synthetic uint16
depth TIFFs and encodes with ``cfg`` or a default :class:`DepthEncoderConfig`
(HEVC Main 12 / ``gray12le``).
"""
imgs_dir = path.parent / f"imgs_{path.stem}"
if depth:
_write_depth_frames(imgs_dir, num_frames=num_frames)
cfg = cfg or DepthEncoderConfig()
else:
_write_color_frames(imgs_dir, num_frames=num_frames)
encode_video_frames(imgs_dir, path, fps=fps, video_encoder=cfg, overwrite=True)
return path
def _read_feature_info(dataset: LeRobotDataset, key: str = DUMMY_VIDEO_KEY) -> dict:
info = json.loads((dataset.root / INFO_PATH).read_text())
return info["features"][key]["info"]
# ─── RGBEncoderConfig / codec options ──────────────────────────────
class TestCodecOptions: class TestCodecOptions:
@require_libsvtav1 @require_libsvtav1
def test_libsvtav1_defaults(self): def test_libsvtav1_defaults(self):
cfg = VideoEncoderConfig() cfg = RGBEncoderConfig()
opts = cfg.get_codec_options() opts = cfg.get_codec_options()
assert opts["g"] == 2 assert opts["g"] == 2
assert opts["crf"] == 30 assert opts["crf"] == 30
@@ -68,12 +131,12 @@ class TestCodecOptions:
@require_libsvtav1 @require_libsvtav1
def test_libsvtav1_custom_preset(self): def test_libsvtav1_custom_preset(self):
cfg = VideoEncoderConfig(preset=8) cfg = RGBEncoderConfig(preset=8)
assert cfg.get_codec_options()["preset"] == 8 assert cfg.get_codec_options()["preset"] == 8
@require_h264 @require_h264
def test_h264_options(self): def test_h264_options(self):
cfg = VideoEncoderConfig(vcodec="h264", g=10, crf=23, preset=None) cfg = RGBEncoderConfig(vcodec="h264", g=10, crf=23, preset=None)
opts = cfg.get_codec_options() opts = cfg.get_codec_options()
assert opts["g"] == 10 assert opts["g"] == 10
assert opts["crf"] == 23 assert opts["crf"] == 23
@@ -81,120 +144,120 @@ class TestCodecOptions:
@require_videotoolbox @require_videotoolbox
def test_videotoolbox_options(self): def test_videotoolbox_options(self):
cfg = VideoEncoderConfig(vcodec="h264_videotoolbox", g=2, crf=30, preset=None) cfg = RGBEncoderConfig(vcodec="h264_videotoolbox", g=2, crf=30, preset=None)
opts = cfg.get_codec_options() opts = cfg.get_codec_options()
assert opts["g"] == 2 assert opts["g"] == 2
assert opts["q:v"] == 40 assert opts["q:v"] == 40
assert "crf" not in opts assert "crf" not in opts
@_require_encoder("h264_nvenc") @require_nvenc
def test_nvenc_options(self): def test_nvenc_options(self):
cfg = VideoEncoderConfig(vcodec="h264_nvenc", g=2, crf=25, preset=None) cfg = RGBEncoderConfig(vcodec="h264_nvenc", g=2, crf=25, preset=None)
opts = cfg.get_codec_options() opts = cfg.get_codec_options()
assert opts["rc"] == 0 assert opts["rc"] == 0
assert opts["qp"] == 25 assert opts["qp"] == 25
assert "crf" not in opts assert "crf" not in opts
assert opts["g"] == 2 assert opts["g"] == 2
@_require_encoder("h264_vaapi") @require_vaapi
def test_vaapi_options(self): def test_vaapi_options(self):
cfg = VideoEncoderConfig(vcodec="h264_vaapi", crf=28, preset=None) cfg = RGBEncoderConfig(vcodec="h264_vaapi", crf=28, preset=None)
assert cfg.get_codec_options()["qp"] == 28 assert cfg.get_codec_options()["qp"] == 28
@_require_encoder("h264_qsv") @require_qsv
def test_qsv_options(self): def test_qsv_options(self):
cfg = VideoEncoderConfig(vcodec="h264_qsv", crf=25, preset=None) cfg = RGBEncoderConfig(vcodec="h264_qsv", crf=25, preset=None)
assert cfg.get_codec_options()["global_quality"] == 25 assert cfg.get_codec_options()["global_quality"] == 25
@require_h264 @require_h264
def test_no_g_no_crf(self): def test_no_g_no_crf(self):
cfg = VideoEncoderConfig(vcodec="h264", g=None, crf=None, preset=None) cfg = RGBEncoderConfig(vcodec="h264", g=None, crf=None, preset=None)
opts = cfg.get_codec_options() opts = cfg.get_codec_options()
assert "g" not in opts assert "g" not in opts
assert "crf" not in opts assert "crf" not in opts
@require_libsvtav1 @require_libsvtav1
def test_encoder_threads_libsvtav1(self): def test_encoder_threads_libsvtav1(self):
cfg = VideoEncoderConfig(fast_decode=0) cfg = RGBEncoderConfig(fast_decode=0)
opts = cfg.get_codec_options(encoder_threads=4) opts = cfg.get_codec_options(encoder_threads=4)
assert "lp=4" in opts.get("svtav1-params", "") assert "lp=4" in opts.get("svtav1-params", "")
@require_h264 @require_h264
def test_encoder_threads_h264(self): def test_encoder_threads_h264(self):
cfg = VideoEncoderConfig(vcodec="h264", preset=None) cfg = RGBEncoderConfig(vcodec="h264", preset=None)
assert cfg.get_codec_options(encoder_threads=2)["threads"] == 2 assert cfg.get_codec_options(encoder_threads=2)["threads"] == 2
@require_libsvtav1 @require_libsvtav1
def test_fast_decode_libsvtav1(self): def test_fast_decode_libsvtav1(self):
cfg = VideoEncoderConfig(fast_decode=1) cfg = RGBEncoderConfig(fast_decode=1)
opts = cfg.get_codec_options() opts = cfg.get_codec_options()
assert "fast-decode=1" in opts.get("svtav1-params", "") assert "fast-decode=1" in opts.get("svtav1-params", "")
@require_libsvtav1 @require_libsvtav1
def test_libsvtav1_fast_decode_clamped_to_svt_range(self): def test_libsvtav1_fast_decode_clamped_to_svt_range(self):
"""Out-of-range fast_decode is clamped to [0, 2] in svtav1-params (SVT-AV1 FastDecode).""" """Out-of-range fast_decode is clamped to [0, 2] in svtav1-params (SVT-AV1 FastDecode)."""
cfg = VideoEncoderConfig(fast_decode=100) cfg = RGBEncoderConfig(fast_decode=100)
assert "fast-decode=2" in cfg.get_codec_options().get("svtav1-params", "") assert "fast-decode=2" in cfg.get_codec_options().get("svtav1-params", "")
cfg_neg = VideoEncoderConfig(fast_decode=-5) cfg_neg = RGBEncoderConfig(fast_decode=-5)
assert "fast-decode=0" in cfg_neg.get_codec_options().get("svtav1-params", "") assert "fast-decode=0" in cfg_neg.get_codec_options().get("svtav1-params", "")
@require_h264 @require_h264
def test_fast_decode_h264(self): def test_fast_decode_h264(self):
cfg = VideoEncoderConfig(vcodec="h264", fast_decode=1, preset=None) cfg = RGBEncoderConfig(vcodec="h264", fast_decode=1, preset=None)
assert cfg.get_codec_options()["tune"] == "fastdecode" assert cfg.get_codec_options()["tune"] == "fastdecode"
@require_libsvtav1 @require_libsvtav1
def test_pix_fmt_unsupported_raises(self): def test_pix_fmt_unsupported_raises(self):
"""Passing an unsupported pix_fmt is a hard error.""" """Passing an unsupported pix_fmt is a hard error."""
with pytest.raises(ValueError, match="pix_fmt"): with pytest.raises(ValueError, match="pix_fmt"):
VideoEncoderConfig(pix_fmt="yuv444p") # libsvtav1 only supports yuv420p variants RGBEncoderConfig(pix_fmt="yuv444p") # libsvtav1 only supports yuv420p variants
@require_libsvtav1 @require_libsvtav1
@require_h264 @require_h264
def test_preset_default_behaviour(self): def test_preset_default_behaviour(self):
"""Empty constructor picks preset=12 (libsvtav1 path); other codecs stay None.""" """Empty constructor picks preset=12 (libsvtav1 path); other codecs stay None."""
assert VideoEncoderConfig().preset == 12 assert RGBEncoderConfig().preset == 12
assert VideoEncoderConfig(vcodec="libsvtav1").preset == 12 assert RGBEncoderConfig(vcodec="libsvtav1").preset == 12
assert VideoEncoderConfig(vcodec="h264").preset is None assert RGBEncoderConfig(vcodec="h264").preset is None
assert VideoEncoderConfig(vcodec="h264", preset=None).preset is None assert RGBEncoderConfig(vcodec="h264", preset=None).preset is None
@require_h264 @require_h264
def test_preset_string_on_h264(self): def test_preset_string_on_h264(self):
"""h264 accepts string presets and forwards them to FFmpeg.""" """h264 accepts string presets and forwards them to FFmpeg."""
cfg = VideoEncoderConfig(vcodec="h264", preset="slow") cfg = RGBEncoderConfig(vcodec="h264", preset="slow")
assert cfg.get_codec_options()["preset"] == "slow" assert cfg.get_codec_options()["preset"] == "slow"
@require_videotoolbox @require_videotoolbox
def test_preset_on_videotoolbox_not_set(self): def test_preset_on_videotoolbox_not_set(self):
"""videotoolbox has no preset option at all.""" """videotoolbox has no preset option at all."""
cfg = VideoEncoderConfig(vcodec="h264_videotoolbox", preset="slow") cfg = RGBEncoderConfig(vcodec="h264_videotoolbox", preset="slow")
assert "preset" not in cfg.get_codec_options() assert "preset" not in cfg.get_codec_options()
@require_libsvtav1 @require_libsvtav1
def test_libsvtav1_preset_out_of_range_raises(self): def test_libsvtav1_preset_out_of_range_raises(self):
"""libsvtav1 preset must sit in [-2, 13] as exposed by PyAV.""" """libsvtav1 preset must sit in [-2, 13] as exposed by PyAV."""
with pytest.raises(ValueError, match="out of range"): with pytest.raises(ValueError, match="out of range"):
VideoEncoderConfig(vcodec="libsvtav1", preset=100) RGBEncoderConfig(vcodec="libsvtav1", preset=100)
with pytest.raises(ValueError, match="out of range"): with pytest.raises(ValueError, match="out of range"):
VideoEncoderConfig(vcodec="libsvtav1", preset=-3) RGBEncoderConfig(vcodec="libsvtav1", preset=-3)
@require_libsvtav1 @require_libsvtav1
def test_libsvtav1_crf_out_of_range_raises(self): def test_libsvtav1_crf_out_of_range_raises(self):
"""libsvtav1 crf must sit in [0, 63].""" """libsvtav1 crf must sit in [0, 63]."""
with pytest.raises(ValueError, match="crf.*out of range"): with pytest.raises(ValueError, match="crf.*out of range"):
VideoEncoderConfig(vcodec="libsvtav1", crf=64) RGBEncoderConfig(vcodec="libsvtav1", crf=64)
@require_libsvtav1 @require_libsvtav1
def test_libsvtav1_crf_rejects_python_float(self): def test_libsvtav1_crf_rejects_python_float(self):
"""libsvtav1 exposes ``crf`` as an INT AVOption; Python float must not pass validation.""" """libsvtav1 exposes ``crf`` as an INT AVOption; Python float must not pass validation."""
with pytest.raises(ValueError, match="float values are not allowed"): with pytest.raises(ValueError, match="float values are not allowed"):
VideoEncoderConfig(vcodec="libsvtav1", crf=2.5) RGBEncoderConfig(vcodec="libsvtav1", crf=2.5)
@require_libsvtav1 @require_libsvtav1
def test_libsvtav1_extra_crf_rejects_fractional_string(self): def test_libsvtav1_extra_crf_rejects_fractional_string(self):
"""INT options reject fractional values even when supplied only via ``extra_options``.""" """INT options reject fractional values even when supplied only via ``extra_options``."""
with pytest.raises(ValueError, match="float values are not allowed"): with pytest.raises(ValueError, match="float values are not allowed"):
VideoEncoderConfig( RGBEncoderConfig(
vcodec="libsvtav1", vcodec="libsvtav1",
crf=None, crf=None,
extra_options={"crf": "2.5"}, extra_options={"crf": "2.5"},
@@ -203,7 +266,7 @@ class TestCodecOptions:
@require_libsvtav1 @require_libsvtav1
def test_libsvtav1_extra_crf_rejects_float(self): def test_libsvtav1_extra_crf_rejects_float(self):
with pytest.raises(ValueError, match="float values are not allowed"): with pytest.raises(ValueError, match="float values are not allowed"):
VideoEncoderConfig( RGBEncoderConfig(
vcodec="libsvtav1", vcodec="libsvtav1",
crf=None, crf=None,
extra_options={"crf": 2.5}, extra_options={"crf": 2.5},
@@ -212,13 +275,13 @@ class TestCodecOptions:
@require_h264 @require_h264
def test_h264_crf_accepts_float_and_int(self): def test_h264_crf_accepts_float_and_int(self):
"""x264 exposes crf as a FLOAT option, so both int and float are accepted.""" """x264 exposes crf as a FLOAT option, so both int and float are accepted."""
assert VideoEncoderConfig(vcodec="h264", crf=23).get_codec_options()["crf"] == 23 assert RGBEncoderConfig(vcodec="h264", crf=23).get_codec_options()["crf"] == 23
assert VideoEncoderConfig(vcodec="h264", crf=23.5).get_codec_options()["crf"] == 23.5 assert RGBEncoderConfig(vcodec="h264", crf=23.5).get_codec_options()["crf"] == 23.5
@require_libsvtav1 @require_libsvtav1
def test_validate_is_rerunnable(self): def test_validate_is_rerunnable(self):
"""After mutating a field, validate() re-checks and surfaces new issues.""" """After mutating a field, validate() re-checks and surfaces new issues."""
cfg = VideoEncoderConfig(vcodec="libsvtav1") cfg = RGBEncoderConfig(vcodec="libsvtav1")
cfg.preset = 100 # now out of range cfg.preset = 100 # now out of range
with pytest.raises(ValueError, match="out of range"): with pytest.raises(ValueError, match="out of range"):
cfg.validate() cfg.validate()
@@ -227,58 +290,58 @@ class TestCodecOptions:
class TestExtraOptions: class TestExtraOptions:
@require_libsvtav1 @require_libsvtav1
def test_default_is_empty_dict(self): def test_default_is_empty_dict(self):
cfg = VideoEncoderConfig() cfg = RGBEncoderConfig()
assert cfg.extra_options == {} assert cfg.extra_options == {}
@require_libsvtav1 @require_libsvtav1
def test_unknown_key_passes_through(self): def test_unknown_key_passes_through(self):
"""Keys not published as AVOptions are forwarded to FFmpeg.""" """Keys not published as AVOptions are forwarded to FFmpeg."""
cfg = VideoEncoderConfig(extra_options={"totally_made_up_option": "value"}) cfg = RGBEncoderConfig(extra_options={"totally_made_up_option": "value"})
assert cfg.extra_options == {"totally_made_up_option": "value"} assert cfg.extra_options == {"totally_made_up_option": "value"}
@require_libsvtav1 @require_libsvtav1
def test_numeric_value_in_range_ok(self): def test_numeric_value_in_range_ok(self):
"""libsvtav1 exposes ``qp`` as INT in [0, 63].""" """libsvtav1 exposes ``qp`` as INT in [0, 63]."""
cfg = VideoEncoderConfig(extra_options={"qp": 30}) cfg = RGBEncoderConfig(extra_options={"qp": 30})
assert cfg.extra_options == {"qp": 30} assert cfg.extra_options == {"qp": 30}
@require_libsvtav1 @require_libsvtav1
def test_numeric_out_of_range_raises(self): def test_numeric_out_of_range_raises(self):
with pytest.raises(ValueError, match=r"qp=.*out of range"): with pytest.raises(ValueError, match=r"qp=.*out of range"):
VideoEncoderConfig(extra_options={"qp": 999}) RGBEncoderConfig(extra_options={"qp": 999})
@require_libsvtav1 @require_libsvtav1
def test_numeric_string_accepted_in_range(self): def test_numeric_string_accepted_in_range(self):
"""Numeric strings are accepted for numeric options (mirrors FFmpeg).""" """Numeric strings are accepted for numeric options (mirrors FFmpeg)."""
cfg = VideoEncoderConfig(extra_options={"qp": "18"}) cfg = RGBEncoderConfig(extra_options={"qp": "18"})
assert cfg.extra_options == {"qp": "18"} assert cfg.extra_options == {"qp": "18"}
@require_libsvtav1 @require_libsvtav1
def test_numeric_string_out_of_range_raises(self): def test_numeric_string_out_of_range_raises(self):
with pytest.raises(ValueError, match=r"qp=.*out of range"): with pytest.raises(ValueError, match=r"qp=.*out of range"):
VideoEncoderConfig(extra_options={"qp": "999"}) RGBEncoderConfig(extra_options={"qp": "999"})
@require_libsvtav1 @require_libsvtav1
def test_non_numeric_string_on_numeric_option_raises(self): def test_non_numeric_string_on_numeric_option_raises(self):
with pytest.raises(ValueError, match=r"qp=.*not numeric"): with pytest.raises(ValueError, match=r"qp=.*not numeric"):
VideoEncoderConfig(extra_options={"qp": "medium"}) RGBEncoderConfig(extra_options={"qp": "medium"})
@require_libsvtav1 @require_libsvtav1
def test_bool_on_numeric_option_raises(self): def test_bool_on_numeric_option_raises(self):
"""``bool`` is explicitly rejected for numeric options.""" """``bool`` is explicitly rejected for numeric options."""
with pytest.raises(ValueError, match=r"qp=.*not numeric"): with pytest.raises(ValueError, match=r"qp=.*not numeric"):
VideoEncoderConfig(extra_options={"qp": True}) RGBEncoderConfig(extra_options={"qp": True})
@require_h264 @require_h264
def test_string_option_passes_through_unchecked(self): def test_string_option_passes_through_unchecked(self):
"""String-typed AVOptions are NOT enum-checked (too many accept freeform).""" """String-typed AVOptions are NOT enum-checked (too many accept freeform)."""
cfg = VideoEncoderConfig(vcodec="h264", preset=None, extra_options={"tune": "some-future-tune"}) cfg = RGBEncoderConfig(vcodec="h264", preset=None, extra_options={"tune": "some-future-tune"})
assert cfg.extra_options == {"tune": "some-future-tune"} assert cfg.extra_options == {"tune": "some-future-tune"}
@require_libsvtav1 @require_libsvtav1
def test_merged_into_codec_options_and_stringified(self): def test_merged_into_codec_options_and_stringified(self):
"""Typed merge by default; ``as_strings=True`` matches FFmpeg option dict.""" """Typed merge by default; ``as_strings=True`` matches FFmpeg option dict."""
cfg = VideoEncoderConfig(extra_options={"qp": 20}) cfg = RGBEncoderConfig(extra_options={"qp": 20})
opts = cfg.get_codec_options() opts = cfg.get_codec_options()
assert opts["qp"] == 20 assert opts["qp"] == 20
assert isinstance(opts["qp"], int) assert isinstance(opts["qp"], int)
@@ -287,25 +350,25 @@ class TestExtraOptions:
@require_libsvtav1 @require_libsvtav1
def test_structured_fields_win_on_collision(self): def test_structured_fields_win_on_collision(self):
"""A colliding extra_options key is discarded; the structured field wins.""" """A colliding extra_options key is discarded; the structured field wins."""
cfg = VideoEncoderConfig(crf=30, extra_options={"crf": 18}) cfg = RGBEncoderConfig(crf=30, extra_options={"crf": 18})
assert cfg.get_codec_options()["crf"] == 30 assert cfg.get_codec_options()["crf"] == 30
class TestEncoderDetection: class TestEncoderDetection:
@require_h264 @require_h264
def test_explicit_codec_kept_when_available(self): def test_explicit_codec_kept_when_available(self):
cfg = VideoEncoderConfig(vcodec="h264") cfg = RGBEncoderConfig(vcodec="h264")
assert cfg.vcodec == "h264" assert cfg.vcodec == "h264"
@require_videotoolbox @require_videotoolbox
def test_auto_picks_videotoolbox_when_available(self): def test_auto_picks_videotoolbox_when_available(self):
"""``h264_videotoolbox`` sits at the top of ``HW_VIDEO_CODECS`` so it wins when present.""" """``h264_videotoolbox`` sits at the top of ``HW_VIDEO_CODECS`` so it wins when present."""
cfg = VideoEncoderConfig(vcodec="auto") cfg = RGBEncoderConfig(vcodec="auto")
assert cfg.vcodec == "h264_videotoolbox" assert cfg.vcodec == "h264_videotoolbox"
def test_invalid_codec_raises(self): def test_invalid_codec_raises(self):
with pytest.raises(ValueError, match="Invalid vcodec"): with pytest.raises(ValueError, match="Invalid vcodec"):
VideoEncoderConfig(vcodec="not_a_real_codec") RGBEncoderConfig(vcodec="not_a_real_codec")
def test_hw_encoder_names_listed_as_valid(self): def test_hw_encoder_names_listed_as_valid(self):
assert "auto" in VALID_VIDEO_CODECS assert "auto" in VALID_VIDEO_CODECS
@@ -313,59 +376,6 @@ class TestEncoderDetection:
assert "h264_nvenc" in VALID_VIDEO_CODECS assert "h264_nvenc" in VALID_VIDEO_CODECS
TEST_ARTIFACTS_DIR = Path(__file__).parent.parent / "artifacts" / "encoded_videos"
# Default video feature set used by persistence tests.
VIDEO_FEATURES = {
"observation.images.cam": {
"dtype": "video",
"shape": (64, 96, 3),
"names": ["height", "width", "channels"],
},
"action": {"dtype": "float32", "shape": (2,), "names": ["a", "b"]},
}
VIDEO_KEY = "observation.images.cam"
def _write_frames(imgs_dir: Path, num_frames: int = 4, height: int = 64, width: int = 96) -> None:
imgs_dir.mkdir(parents=True, exist_ok=True)
for i in range(num_frames):
arr = np.random.randint(0, 256, (height, width, 3), dtype=np.uint8)
write_image(arr, imgs_dir / f"frame-{i:06d}.png")
def _encode_video(
path: Path, num_frames: int = 4, fps: int = 30, cfg: VideoEncoderConfig | None = None
) -> Path:
imgs_dir = path.parent / f"imgs_{path.stem}"
_write_frames(imgs_dir, num_frames=num_frames)
encode_video_frames(imgs_dir, path, fps=fps, camera_encoder=cfg, overwrite=True)
return path
def _read_feature_info(dataset: LeRobotDataset) -> dict:
info = json.loads((dataset.root / INFO_PATH).read_text())
return info["features"][VIDEO_KEY]["info"]
def _add_frames(dataset: LeRobotDataset, num_frames: int, video_keys: list[str] | None = None) -> None:
from lerobot.utils.constants import DEFAULT_FEATURES
if video_keys is None:
video_keys = dataset.meta.video_keys
for _ in range(num_frames):
frame: dict = {"task": "test"}
for key, ft in dataset.meta.features.items():
if key in DEFAULT_FEATURES:
continue
shape = ft["shape"]
if key in video_keys:
frame[key] = np.random.randint(0, 256, shape, dtype=np.uint8)
else:
frame[key] = np.zeros(shape, dtype=np.float32)
dataset.add_frame(frame)
class TestGetVideoInfo: class TestGetVideoInfo:
def test_returns_all_stream_fields(self): def test_returns_all_stream_fields(self):
info = get_video_info(TEST_ARTIFACTS_DIR / "clip_4frames.mp4") info = get_video_info(TEST_ARTIFACTS_DIR / "clip_4frames.mp4")
@@ -375,7 +385,7 @@ class TestGetVideoInfo:
assert info["video.pix_fmt"] == "yuv420p" assert info["video.pix_fmt"] == "yuv420p"
assert info["video.fps"] == 30 assert info["video.fps"] == 30
assert info["video.channels"] == 3 assert info["video.channels"] == 3
assert info["video.is_depth_map"] is False assert info["is_depth_map"] is False
assert info["has_audio"] is False assert info["has_audio"] is False
assert "video.g" not in info assert "video.g" not in info
assert "video.crf" not in info assert "video.crf" not in info
@@ -383,9 +393,9 @@ class TestGetVideoInfo:
@require_libsvtav1 @require_libsvtav1
def test_merges_encoder_config_as_video_prefixed_entries(self): def test_merges_encoder_config_as_video_prefixed_entries(self):
cfg = VideoEncoderConfig(vcodec="libsvtav1", g=2, crf=30, preset=12) cfg = RGBEncoderConfig(vcodec="libsvtav1", g=2, crf=30, preset=12)
info = get_video_info(TEST_ARTIFACTS_DIR / "clip_4frames.mp4", camera_encoder=cfg) info = get_video_info(TEST_ARTIFACTS_DIR / "clip_4frames.mp4", video_encoder=cfg)
assert info["video.g"] == 2 assert info["video.g"] == 2
assert info["video.crf"] == 30 assert info["video.crf"] == 30
@@ -396,13 +406,18 @@ class TestGetVideoInfo:
@require_libsvtav1 @require_libsvtav1
def test_stream_derived_keys_take_precedence_over_config(self): def test_stream_derived_keys_take_precedence_over_config(self):
cfg = VideoEncoderConfig(vcodec="libsvtav1", pix_fmt="yuv420p") cfg = RGBEncoderConfig(vcodec="libsvtav1", pix_fmt="yuv420p")
info = get_video_info(TEST_ARTIFACTS_DIR / "clip_4frames.mp4", camera_encoder=cfg) info = get_video_info(TEST_ARTIFACTS_DIR / "clip_4frames.mp4", video_encoder=cfg)
assert info["video.codec"] # populated from stream, not from config's vcodec assert info["video.codec"] # populated from stream, not from config's vcodec
assert info["video.pix_fmt"] == "yuv420p" assert info["video.pix_fmt"] == "yuv420p"
def test_depth_encoder_config_sets_is_depth_map_true(self):
"""A ``DepthEncoderConfig`` causes ``get_video_info`` to mark the stream as depth."""
info = get_video_info(TEST_ARTIFACTS_DIR / "clip_4frames.mp4", video_encoder=DepthEncoderConfig())
assert info["is_depth_map"] is True
class TestEncodeVideoFrames: class TestEncodeVideoFrames:
@require_libsvtav1 @require_libsvtav1
@@ -434,7 +449,7 @@ class TestEncodeVideoFrames:
def test_overwrite_false_skips_existing_file(self, tmp_path): def test_overwrite_false_skips_existing_file(self, tmp_path):
imgs_dir = tmp_path / "imgs" imgs_dir = tmp_path / "imgs"
_write_frames(imgs_dir) _write_color_frames(imgs_dir)
video_path = tmp_path / "out.mp4" video_path = tmp_path / "out.mp4"
sentinel = b"pre-existing content" sentinel = b"pre-existing content"
video_path.write_bytes(sentinel) video_path.write_bytes(sentinel)
@@ -446,7 +461,7 @@ class TestEncodeVideoFrames:
@require_libsvtav1 @require_libsvtav1
def test_overwrite_true_replaces_existing_file(self, tmp_path): def test_overwrite_true_replaces_existing_file(self, tmp_path):
imgs_dir = tmp_path / "imgs" imgs_dir = tmp_path / "imgs"
_write_frames(imgs_dir) _write_color_frames(imgs_dir)
video_path = tmp_path / "out.mp4" video_path = tmp_path / "out.mp4"
video_path.write_bytes(b"stale content") video_path.write_bytes(b"stale content")
@@ -458,10 +473,10 @@ class TestEncodeVideoFrames:
@require_libsvtav1 @require_libsvtav1
def test_custom_encoder_config_fields_stored_in_info(self, tmp_path): def test_custom_encoder_config_fields_stored_in_info(self, tmp_path):
"""All stream-derived and encoder config fields are present after encoding.""" """All stream-derived and encoder config fields are present after encoding."""
cfg = VideoEncoderConfig(vcodec="libsvtav1", g=4, crf=25, preset=10) cfg = RGBEncoderConfig(vcodec="libsvtav1", g=4, crf=25, preset=10)
video_path = _encode_video(tmp_path / "out.mp4", num_frames=4, fps=30, cfg=cfg) video_path = _encode_video(tmp_path / "out.mp4", num_frames=4, fps=30, cfg=cfg)
info = get_video_info(video_path, camera_encoder=cfg) info = get_video_info(video_path, video_encoder=cfg)
# Stream-derived # Stream-derived
assert info["video.height"] == 64 assert info["video.height"] == 64
@@ -470,7 +485,7 @@ class TestEncodeVideoFrames:
assert info["video.codec"] == "av1" assert info["video.codec"] == "av1"
assert info["video.pix_fmt"] == "yuv420p" assert info["video.pix_fmt"] == "yuv420p"
assert info["video.fps"] == 30 assert info["video.fps"] == 30
assert info["video.is_depth_map"] is False assert info["is_depth_map"] is False
assert info["has_audio"] is False assert info["has_audio"] is False
# Encoder config # Encoder config
assert info["video.g"] == 4 assert info["video.g"] == 4
@@ -487,15 +502,15 @@ class TestReencodeVideo:
def test_reencode_video(self, tmp_path): def test_reencode_video(self, tmp_path):
src = TEST_ARTIFACTS_DIR / "clip_4frames.mp4" src = TEST_ARTIFACTS_DIR / "clip_4frames.mp4"
out = tmp_path / "reencoded.mp4" out = tmp_path / "reencoded.mp4"
cfg = VideoEncoderConfig(vcodec="h264", g=6, crf=23, pix_fmt="yuv444p") cfg = RGBEncoderConfig(vcodec="h264", g=6, crf=23, pix_fmt="yuv444p")
reencode_video(src, out, camera_encoder=cfg, overwrite=True) reencode_video(src, out, video_encoder=cfg, overwrite=True)
assert out.exists() assert out.exists()
with av.open(str(out)) as container: with av.open(str(out)) as container:
n_frames = sum(1 for _ in container.decode(video=0)) n_frames = sum(1 for _ in container.decode(video=0))
assert n_frames == 4 assert n_frames == 4
info = get_video_info(out, camera_encoder=cfg) info = get_video_info(out, video_encoder=cfg)
assert info["video.codec"] == "h264" assert info["video.codec"] == "h264"
assert info["video.pix_fmt"] == "yuv444p" assert info["video.pix_fmt"] == "yuv444p"
assert info["video.height"] == 64 assert info["video.height"] == 64
@@ -508,8 +523,8 @@ class TestReencodeVideo:
def test_reencode_video_trim_window(self, tmp_path): def test_reencode_video_trim_window(self, tmp_path):
src = TEST_ARTIFACTS_DIR / "clip_6frames.mp4" src = TEST_ARTIFACTS_DIR / "clip_6frames.mp4"
out = tmp_path / "trim_window.mp4" out = tmp_path / "trim_window.mp4"
cfg = VideoEncoderConfig(vcodec="h264") cfg = RGBEncoderConfig(vcodec="h264")
reencode_video(src, out, camera_encoder=cfg, start_time_s=0.05, end_time_s=0.12, overwrite=True) reencode_video(src, out, video_encoder=cfg, start_time_s=0.05, end_time_s=0.12, overwrite=True)
with av.open(str(out)) as container: with av.open(str(out)) as container:
frames = list(container.decode(video=0)) frames = list(container.decode(video=0))
@@ -578,12 +593,12 @@ class TestEncoderConfigPersistence:
@require_libsvtav1 @require_libsvtav1
def test_first_episode_save_persists_encoder_config(self, tmp_path, empty_lerobot_dataset_factory): def test_first_episode_save_persists_encoder_config(self, tmp_path, empty_lerobot_dataset_factory):
cfg = VideoEncoderConfig(vcodec="libsvtav1", g=2, crf=30, preset=12) cfg = RGBEncoderConfig(vcodec="libsvtav1", g=2, crf=30, preset=12)
dataset = empty_lerobot_dataset_factory( dataset = empty_lerobot_dataset_factory(
root=tmp_path / "ds", features=VIDEO_FEATURES, use_videos=True, camera_encoder=cfg root=tmp_path / "ds", features=DUMMY_VIDEO_FEATURES, use_videos=True, rgb_encoder=cfg
) )
_add_frames(dataset, num_frames=4) add_frames(dataset, num_frames=4)
dataset.save_episode() dataset.save_episode()
dataset.finalize() dataset.finalize()
@@ -601,16 +616,16 @@ class TestEncoderConfigPersistence:
@require_libsvtav1 @require_libsvtav1
def test_second_episode_does_not_overwrite_encoder_fields(self, tmp_path, empty_lerobot_dataset_factory): def test_second_episode_does_not_overwrite_encoder_fields(self, tmp_path, empty_lerobot_dataset_factory):
cfg = VideoEncoderConfig(vcodec="libsvtav1", g=2, crf=30, preset=12) cfg = RGBEncoderConfig(vcodec="libsvtav1", g=2, crf=30, preset=12)
dataset = empty_lerobot_dataset_factory( dataset = empty_lerobot_dataset_factory(
root=tmp_path / "ds", features=VIDEO_FEATURES, use_videos=True, camera_encoder=cfg root=tmp_path / "ds", features=DUMMY_VIDEO_FEATURES, use_videos=True, rgb_encoder=cfg
) )
_add_frames(dataset, num_frames=4) add_frames(dataset, num_frames=4)
dataset.save_episode() dataset.save_episode()
first_info = dict(_read_feature_info(dataset)) first_info = dict(_read_feature_info(dataset))
_add_frames(dataset, num_frames=4) add_frames(dataset, num_frames=4)
dataset.save_episode() dataset.save_episode()
dataset.finalize() dataset.finalize()
@@ -618,13 +633,13 @@ class TestEncoderConfigPersistence:
class TestFromVideoInfo: class TestFromVideoInfo:
"""``VideoEncoderConfig.from_video_info`` reconstructs an encoder config """``RGBEncoderConfig.from_video_info`` reconstructs an encoder config
from the ``video.*`` keys persisted in a dataset's ``info.json``. from the ``video.*`` keys persisted in a dataset's ``info.json``.
""" """
@require_libsvtav1 @require_libsvtav1
def test_reconstructs_from_dummy_video_info(self): def test_reconstructs_from_dummy_video_info(self):
cfg = VideoEncoderConfig.from_video_info(DUMMY_VIDEO_INFO) cfg = RGBEncoderConfig.from_video_info(DUMMY_VIDEO_INFO)
# Canonical stream codec ``"av1"`` is aliased to the encoder name. # Canonical stream codec ``"av1"`` is aliased to the encoder name.
assert cfg.vcodec == "libsvtav1" assert cfg.vcodec == "libsvtav1"
@@ -636,4 +651,220 @@ class TestFromVideoInfo:
assert cfg.video_backend == DUMMY_VIDEO_INFO["video.video_backend"] assert cfg.video_backend == DUMMY_VIDEO_INFO["video.video_backend"]
# ``{}`` placeholder (typical after a merge with disagreeing sources) # ``{}`` placeholder (typical after a merge with disagreeing sources)
# must not leak into the reconstructed config. # must not leak into the reconstructed config.
assert cfg.extra_options == VideoEncoderConfig().extra_options assert cfg.extra_options == RGBEncoderConfig().extra_options
# ─── Depth-specific encoding tests ────────────────────────────────────
class TestEncodeDepthVideoFrames:
"""Depth mirror of :class:`TestEncodeVideoFrames`.
Exercises ``encode_video_frames`` end-to-end through
:class:`DepthEncoderConfig` (HEVC Main 12 / ``gray12le``) on synthetic
uint16 depth TIFFs.
"""
@require_hevc
def test_produces_readable_file(self, tmp_path):
video_path = _encode_video(tmp_path / "out.mp4", depth=True)
assert video_path.exists()
info = get_video_info(video_path, video_encoder=DepthEncoderConfig())
assert info["video.height"] == 64
assert info["video.width"] == 96
assert info["video.codec"] == "hevc"
assert info["video.pix_fmt"] == "gray12le"
assert info["video.channels"] == 1
assert info["is_depth_map"] is True
@require_hevc
def test_frame_count_and_duration_match_input(self, tmp_path):
num_frames = 10
fps = 30
video_path = _encode_video(tmp_path / "out.mp4", num_frames=num_frames, fps=fps, depth=True)
with av.open(str(video_path)) as container:
stream = container.streams.video[0]
actual_frames = sum(1 for _ in container.decode(stream))
duration = (
float(stream.duration * stream.time_base)
if stream.duration is not None
else float(container.duration / av.time_base)
)
assert actual_frames == num_frames
assert abs(duration - num_frames / fps) < 0.1
def test_overwrite_false_skips_existing_file(self, tmp_path):
"""Codec-agnostic: file-system semantics must hold even without an HEVC encoder."""
imgs_dir = tmp_path / "imgs"
_write_depth_frames(imgs_dir)
video_path = tmp_path / "out.mp4"
sentinel = b"pre-existing depth content"
video_path.write_bytes(sentinel)
encode_video_frames(imgs_dir, video_path, fps=30, video_encoder=DepthEncoderConfig(), overwrite=False)
assert video_path.read_bytes() == sentinel
@require_hevc
def test_overwrite_true_replaces_existing_file(self, tmp_path):
imgs_dir = tmp_path / "imgs"
_write_depth_frames(imgs_dir)
video_path = tmp_path / "out.mp4"
video_path.write_bytes(b"stale content")
encode_video_frames(imgs_dir, video_path, fps=30, video_encoder=DepthEncoderConfig(), overwrite=True)
info = get_video_info(video_path, video_encoder=DepthEncoderConfig())
assert info["video.height"] == 64
assert info["video.pix_fmt"] == "gray12le"
assert info["is_depth_map"] is True
@require_hevc
def test_custom_encoder_config_fields_stored_in_info(self, tmp_path):
"""All stream-derived and depth-encoder config fields are present after encoding."""
cfg = DepthEncoderConfig(
vcodec="hevc",
pix_fmt="gray12le",
g=4,
crf=25,
extra_options={},
depth_min=0.05,
depth_max=8.0,
shift=2.5,
use_log=False,
)
video_path = _encode_video(tmp_path / "out.mp4", num_frames=4, fps=30, cfg=cfg, depth=True)
info = get_video_info(video_path, video_encoder=cfg)
# Stream-derived
assert info["video.height"] == 64
assert info["video.width"] == 96
assert info["video.channels"] == 1
assert info["video.codec"] == "hevc"
assert info["video.pix_fmt"] == "gray12le"
assert info["video.fps"] == 30
assert info["is_depth_map"] is True
assert info["has_audio"] is False
# Base encoder config
assert info["video.g"] == 4
assert info["video.crf"] == 25
assert info["video.fast_decode"] == 0
assert info["video.video_backend"] == "pyav"
assert info["video.extra_options"] == {}
# Depth-specific tuning
assert info["video.depth_min"] == 0.05
assert info["video.depth_max"] == 8.0
assert info["video.shift"] == 2.5
assert info["video.use_log"] is False
class TestDepthEncoderConfigPersistence:
"""Depth mirror of :class:`TestEncoderConfigPersistence`.
``DepthEncoderConfig`` must be stored as ``video.<field>`` entries
(including the depth-specific ``depth_min`` / ``depth_max`` / ``shift`` /
``use_log``) under ``info["features"][<depth_key>]["info"]`` when the
first episode is saved.
"""
@require_hevc
def test_first_episode_save_persists_depth_encoder_config(self, tmp_path, empty_lerobot_dataset_factory):
cfg = DepthEncoderConfig(
vcodec="hevc",
pix_fmt="gray12le",
g=2,
crf=30,
extra_options={},
depth_min=0.05,
depth_max=8.0,
shift=2.5,
use_log=False,
)
dataset = empty_lerobot_dataset_factory(
root=tmp_path / "ds", features=DUMMY_DEPTH_FEATURES, use_videos=True, depth_encoder=cfg
)
add_frames(dataset, num_frames=4)
dataset.save_episode()
dataset.finalize()
info = _read_feature_info(dataset, key=DUMMY_DEPTH_KEY)
# Stream-derived
assert info["video.height"] == 64
assert info["video.width"] == 96
assert info["video.fps"] == 30
assert info["video.codec"] == "hevc"
assert info["video.pix_fmt"] == "gray12le"
assert info["is_depth_map"] is True
# Base encoder config
assert info["video.g"] == 2
assert info["video.crf"] == 30
assert info["video.fast_decode"] == 0
assert info["video.video_backend"] == "pyav"
assert info["video.extra_options"] == {}
# Depth-specific tuning
assert info["video.depth_min"] == 0.05
assert info["video.depth_max"] == 8.0
assert info["video.shift"] == 2.5
assert info["video.use_log"] is False
@require_hevc
def test_second_episode_does_not_overwrite_depth_encoder_fields(
self, tmp_path, empty_lerobot_dataset_factory
):
cfg = DepthEncoderConfig(
vcodec="hevc",
pix_fmt="gray12le",
g=2,
crf=30,
depth_min=0.05,
depth_max=8.0,
shift=2.5,
use_log=False,
)
dataset = empty_lerobot_dataset_factory(
root=tmp_path / "ds", features=DUMMY_DEPTH_FEATURES, use_videos=True, depth_encoder=cfg
)
add_frames(dataset, num_frames=4)
dataset.save_episode()
first_info = dict(_read_feature_info(dataset, key=DUMMY_DEPTH_KEY))
add_frames(dataset, num_frames=4)
dataset.save_episode()
dataset.finalize()
assert _read_feature_info(dataset, key=DUMMY_DEPTH_KEY) == first_info
class TestDepthFromVideoInfo:
"""``DepthEncoderConfig.from_video_info`` reconstructs a depth encoder
config from the ``video.*`` keys persisted in a dataset's ``info.json``.
Depth mirror of :class:`TestFromVideoInfo`.
"""
@require_hevc
def test_reconstructs_from_dummy_depth_video_info(self):
cfg = DepthEncoderConfig.from_video_info(DUMMY_DEPTH_VIDEO_INFO_FULL)
# No alias for ``"hevc"``; the canonical stream codec is reused as-is.
assert cfg.vcodec == "hevc"
assert cfg.pix_fmt == DUMMY_DEPTH_VIDEO_INFO_FULL["video.pix_fmt"]
assert cfg.g == DUMMY_DEPTH_VIDEO_INFO_FULL["video.g"]
assert cfg.crf == DUMMY_DEPTH_VIDEO_INFO_FULL["video.crf"]
assert cfg.fast_decode == DUMMY_DEPTH_VIDEO_INFO_FULL["video.fast_decode"]
assert cfg.video_backend == DUMMY_DEPTH_VIDEO_INFO_FULL["video.video_backend"]
# ``{}`` placeholder (typical after a merge with disagreeing sources)
# must not leak into the reconstructed config.
assert cfg.extra_options == DepthEncoderConfig().extra_options
# Depth-specific tuning round-trips through ``info.json``.
assert cfg.depth_min == DUMMY_DEPTH_VIDEO_INFO_FULL["video.depth_min"]
assert cfg.depth_max == DUMMY_DEPTH_VIDEO_INFO_FULL["video.depth_max"]
assert cfg.shift == DUMMY_DEPTH_VIDEO_INFO_FULL["video.shift"]
assert cfg.use_log == DUMMY_DEPTH_VIDEO_INFO_FULL["video.use_log"]
+45 -1
View File
@@ -39,12 +39,56 @@ DUMMY_VIDEO_INFO = {
"video.crf": 30, "video.crf": 30,
"video.preset": 12, "video.preset": 12,
"video.fast_decode": 0, "video.fast_decode": 0,
"video.is_depth_map": False, "is_depth_map": False,
"has_audio": False, "has_audio": False,
} }
DUMMY_CAMERA_FEATURES = { DUMMY_CAMERA_FEATURES = {
"laptop": {"shape": (64, 96, 3), "names": ["height", "width", "channels"], "info": DUMMY_VIDEO_INFO}, "laptop": {"shape": (64, 96, 3), "names": ["height", "width", "channels"], "info": DUMMY_VIDEO_INFO},
"phone": {"shape": (64, 96, 3), "names": ["height", "width", "channels"], "info": DUMMY_VIDEO_INFO}, "phone": {"shape": (64, 96, 3), "names": ["height", "width", "channels"], "info": DUMMY_VIDEO_INFO},
} }
DUMMY_DEPTH_VIDEO_INFO = {
**DUMMY_VIDEO_INFO,
"is_depth_map": True,
}
DUMMY_DEPTH_VIDEO_INFO_FULL = {
**{k: v for k, v in DUMMY_VIDEO_INFO.items() if k != "video.preset"},
"video.codec": "hevc",
"video.pix_fmt": "gray12le",
"is_depth_map": True,
"video.depth_min": 0.05,
"video.depth_max": 8.0,
"video.shift": 2.5,
"video.use_log": True,
}
DUMMY_DEPTH_CAMERA_FEATURES = {
"laptop_depth": {
"shape": (64, 96, 1),
"names": ["height", "width", "channels"],
"info": DUMMY_DEPTH_VIDEO_INFO,
},
}
DUMMY_CAMERA_FEATURES_WITH_DEPTH = {**DUMMY_CAMERA_FEATURES, **DUMMY_DEPTH_CAMERA_FEATURES}
DUMMY_CHW = (3, 96, 128) DUMMY_CHW = (3, 96, 128)
DUMMY_HWC = (96, 128, 3) DUMMY_HWC = (96, 128, 3)
# Default video feature set used by video-encoding persistence tests.
DUMMY_VIDEO_FEATURES = {
"observation.images.cam": {
"dtype": "video",
"shape": (64, 96, 3),
"names": ["height", "width", "channels"],
},
"action": {"dtype": "float32", "shape": (2,), "names": ["a", "b"]},
}
DUMMY_VIDEO_KEY = "observation.images.cam"
DUMMY_DEPTH_FEATURES = {
"observation.images.depth": {
"dtype": "video",
"shape": (64, 96, 1),
"names": ["height", "width", "channels"],
"info": {"is_depth_map": True},
},
"action": {"dtype": "float32", "shape": (2,), "names": ["a", "b"]},
}
DUMMY_DEPTH_KEY = "observation.images.depth"
+38
View File
@@ -49,6 +49,39 @@ from tests.fixtures.constants import (
) )
def add_frames(dataset: LeRobotDataset, num_frames: int) -> None:
"""Append ``num_frames`` synthetic frames to ``dataset``.
Generates per-feature payloads from ``dataset.meta``: uint16 depth ramps for
keys in ``dataset.meta.depth_keys``, uint8 random noise for video/image keys,
and float32 zeros for everything else. ``DEFAULT_FEATURES`` (timestamp,
frame_index, ...) are auto-populated by ``add_frame`` and skipped here.
"""
video_keys = dataset.meta.video_keys
depth_keys = dataset.meta.depth_keys
# Smooth gradient base reused per (H, W) to keep depth frames cheap to
# encode (HEVC Main 12 hates white noise).
_depth_base_cache: dict[tuple[int, int], np.ndarray] = {}
for i in range(num_frames):
frame: dict = {"task": "test"}
for key, ft in dataset.meta.features.items():
if key in DEFAULT_FEATURES:
continue
shape = ft["shape"]
if key in depth_keys:
h, w, _ = shape
base = _depth_base_cache.setdefault(
(h, w),
np.linspace(100.0, 10_000.0, h * w, dtype=np.float32).reshape(h, w, 1),
)
frame[key] = (base + 50.0 * i).clip(0, 65535).astype(np.uint16)
elif key in video_keys:
frame[key] = np.random.randint(0, 256, shape, dtype=np.uint8)
else:
frame[key] = np.zeros(shape, dtype=np.float32)
dataset.add_frame(frame)
class LeRobotDatasetFactory(Protocol): class LeRobotDatasetFactory(Protocol):
def __call__(self, *args, **kwargs) -> LeRobotDataset: ... def __call__(self, *args, **kwargs) -> LeRobotDataset: ...
@@ -485,10 +518,14 @@ def lerobot_dataset_factory(
hf_dataset: datasets.Dataset | None = None, hf_dataset: datasets.Dataset | None = None,
data_files_size_in_mb: float = DEFAULT_DATA_FILE_SIZE_IN_MB, data_files_size_in_mb: float = DEFAULT_DATA_FILE_SIZE_IN_MB,
chunks_size: int = DEFAULT_CHUNK_SIZE, chunks_size: int = DEFAULT_CHUNK_SIZE,
camera_features: dict | None = None,
**kwargs, **kwargs,
) -> LeRobotDataset: ) -> LeRobotDataset:
# Instantiate objects # Instantiate objects
if info is None: if info is None:
info_kwargs = {}
if camera_features is not None:
info_kwargs["camera_features"] = camera_features
info = info_factory( info = info_factory(
total_episodes=total_episodes, total_episodes=total_episodes,
total_frames=total_frames, total_frames=total_frames,
@@ -496,6 +533,7 @@ def lerobot_dataset_factory(
use_videos=use_videos, use_videos=use_videos,
data_files_size_in_mb=data_files_size_in_mb, data_files_size_in_mb=data_files_size_in_mb,
chunks_size=chunks_size, chunks_size=chunks_size,
**info_kwargs,
) )
if stats is None: if stats is None:
stats = stats_factory(features=info.features) stats = stats_factory(features=info.features)
@@ -27,6 +27,7 @@ from lerobot.scripts.lerobot_edit_dataset import (
MergeConfig, MergeConfig,
ModifyTasksConfig, ModifyTasksConfig,
OperationConfig, OperationConfig,
ReencodeVideosConfig,
RemoveFeatureConfig, RemoveFeatureConfig,
SplitConfig, SplitConfig,
_validate_config, _validate_config,
@@ -103,3 +104,47 @@ class TestOperationTypeParsing:
) )
resolved_name = OperationConfig.get_choice_name(type(cfg.operation)) resolved_name = OperationConfig.get_choice_name(type(cfg.operation))
assert resolved_name == type_name assert resolved_name == type_name
class TestDepthEncoderParsing:
"""Test that the depth encoder is exposed and parsed for video operations."""
def test_reencode_has_default_depth_encoder(self):
cfg = parse_cfg(["--repo_id", "test/repo", "--operation.type", "reencode_videos"])
assert isinstance(cfg.operation, ReencodeVideosConfig)
# A depth encoder is configured by default so depth videos are re-encoded too.
assert cfg.operation.depth_encoder is not None
assert hasattr(cfg.operation.depth_encoder, "depth_min")
def test_reencode_parses_depth_encoder_overrides(self):
cfg = parse_cfg(
[
"--repo_id",
"test/repo",
"--operation.type",
"reencode_videos",
"--operation.depth_encoder.extra_options",
'{"x265-params": "lossless=1"}',
"--operation.depth_encoder.depth_max",
"12.0",
"--operation.depth_encoder.use_log",
"false",
]
)
assert cfg.operation.depth_encoder.extra_options == {"x265-params": "lossless=1"}
assert cfg.operation.depth_encoder.depth_max == 12.0
assert cfg.operation.depth_encoder.use_log is False
def test_convert_image_to_video_parses_depth_encoder_overrides(self):
cfg = parse_cfg(
[
"--repo_id",
"test/repo",
"--operation.type",
"convert_image_to_video",
"--operation.depth_encoder.depth_min",
"0.05",
]
)
assert isinstance(cfg.operation, ConvertImageToVideoConfig)
assert cfg.operation.depth_encoder.depth_min == 0.05
+8 -1
View File
@@ -43,6 +43,11 @@ def mock_rerun(monkeypatch):
def __init__(self, arr): def __init__(self, arr):
self.arr = arr self.arr = arr
class DummyDepthImage:
def __init__(self, arr, colormap=None):
self.arr = arr
self.colormap = colormap
def dummy_log(key, obj=None, **kwargs): def dummy_log(key, obj=None, **kwargs):
# Accept either positional `obj` or keyword `entity` and record remaining kwargs. # Accept either positional `obj` or keyword `entity` and record remaining kwargs.
if obj is None and "entity" in kwargs: if obj is None and "entity" in kwargs:
@@ -55,6 +60,8 @@ def mock_rerun(monkeypatch):
__spec__=SimpleNamespace(name="rerun", submodule_search_locations=None), __spec__=SimpleNamespace(name="rerun", submodule_search_locations=None),
Scalars=DummyScalar, Scalars=DummyScalar,
Image=DummyImage, Image=DummyImage,
DepthImage=DummyDepthImage,
components=SimpleNamespace(Colormap=SimpleNamespace(Viridis="viridis")),
log=dummy_log, log=dummy_log,
init=lambda *a, **k: None, init=lambda *a, **k: None,
spawn=lambda *a, **k: None, spawn=lambda *a, **k: None,
@@ -225,7 +232,7 @@ def test_log_rerun_data_kwargs_only(mock_rerun):
assert temp.value == pytest.approx(10.0) assert temp.value == pytest.approx(10.0)
img = _obj_for(calls, "observation.gray") img = _obj_for(calls, "observation.gray")
assert type(img).__name__ == "DummyImage" assert type(img).__name__ == "DummyDepthImage" # single-channel -> DepthImage
assert img.arr.shape == (8, 8, 1) # remains HWC assert img.arr.shape == (8, 8, 1) # remains HWC
assert _kwargs_for(calls, "observation.gray").get("static", False) is True assert _kwargs_for(calls, "observation.gray").get("static", False) is True