feat(dataset): integrate episode streaming into training

This commit is contained in:
Pepijn
2026-07-23 16:10:39 +02:00
parent b85620657f
commit 3d70b21aac
34 changed files with 2771 additions and 942 deletions
+9 -10
View File
@@ -323,14 +323,13 @@ class TestDepthUnitMetadata:
np.testing.assert_allclose(float(np.asarray(stats["mean"]).reshape(-1)[0]), expected, rtol=0.05)
np.testing.assert_allclose(float(np.asarray(stats["count"]).reshape(-1)[0]), count)
if not use_videos:
depth = read_dataset[0][DEPTH_KEY]
assert torch.allclose(depth, torch.full_like(depth, expected))
from lerobot.datasets.streaming_dataset import StreamingLeRobotDataset
from lerobot.datasets.streaming_dataset import StreamingLeRobotDataset
stream_dataset = StreamingLeRobotDataset(
repo_id=DUMMY_REPO_ID, root=tmp_path / "ds", depth_output_unit=output_unit
)
stream_depth = next(iter(stream_dataset))[DEPTH_KEY]
assert torch.allclose(stream_depth, torch.full_like(stream_depth, expected))
stream_dataset = StreamingLeRobotDataset(
repo_id=DUMMY_REPO_ID, root=tmp_path / "ds", depth_output_unit=output_unit
)
stream_item = next(iter(stream_dataset))
stream_depth = stream_item[DEPTH_KEY]
reference_depth = read_dataset[int(stream_item["index"])][DEPTH_KEY]
assert torch.allclose(stream_depth, reference_depth)
assert torch.allclose(stream_depth, torch.full_like(stream_depth, expected), rtol=0.05)
@@ -0,0 +1,110 @@
#!/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
from __future__ import annotations
from pathlib import Path
import fsspec
import pytest
pytest.importorskip("pyarrow", reason="pyarrow is required (install lerobot[dataset])")
import pyarrow as pa
import pyarrow.parquet as pq
from lerobot.datasets.episode_parquet import EpisodeParquetReader
def _table(episodes: list[int]) -> pa.Table:
frame_counts: dict[int, int] = {}
frame_indices = []
values = []
ignored = []
for episode in episodes:
frame_index = frame_counts.get(episode, 0)
frame_counts[episode] = frame_index + 1
frame_indices.append(frame_index)
values.append(episode * 10 + frame_index)
ignored.append(f"ignored-{episode}-{frame_index}")
return pa.table(
{
"episode_index": episodes,
"frame_index": frame_indices,
"value": values,
"ignored": ignored,
}
)
def _write_episode_row_groups(path: Path) -> None:
path.parent.mkdir(parents=True, exist_ok=True)
writer = pq.ParquetWriter(path, _table([0]).schema)
try:
writer.write_table(_table([0, 0]))
writer.write_table(_table([1, 1, 1]))
finally:
writer.close()
def test_reader_projects_columns_and_reads_matching_row_group(tmp_path: Path) -> None:
path = tmp_path / "data/chunk-000/file-000.parquet"
_write_episode_row_groups(path)
reader = EpisodeParquetReader(tmp_path, columns=("episode_index", "frame_index", "value"))
table = reader.read_episode(path.relative_to(tmp_path), episode_index=1, expected_rows=3)
assert table.column_names == ["episode_index", "frame_index", "value"]
assert table.column("value").to_pylist() == [10, 11, 12]
def test_reader_filters_legacy_mixed_row_group(tmp_path: Path) -> None:
path = tmp_path / "data/chunk-000/file-000.parquet"
path.parent.mkdir(parents=True)
pq.write_table(_table([0, 0, 1, 1, 1]), path)
reader = EpisodeParquetReader(tmp_path, columns=("episode_index", "frame_index", "value"))
table = reader.read_episode(path.relative_to(tmp_path), episode_index=1, expected_rows=3)
assert table.column("episode_index").to_pylist() == [1, 1, 1]
assert table.column("frame_index").to_pylist() == [0, 1, 2]
def test_reader_rejects_partial_episode(tmp_path: Path) -> None:
path = tmp_path / "data/chunk-000/file-000.parquet"
path.parent.mkdir(parents=True)
pq.write_table(_table([2, 2]), path)
reader = EpisodeParquetReader(tmp_path, columns=("episode_index", "frame_index"))
with pytest.raises(ValueError, match="expected 3 rows, found 2"):
reader.read_episode(path.relative_to(tmp_path), episode_index=2, expected_rows=3)
def test_reader_rejects_missing_episode(tmp_path: Path) -> None:
path = tmp_path / "data/chunk-000/file-000.parquet"
path.parent.mkdir(parents=True)
pq.write_table(_table([0, 0]), path)
reader = EpisodeParquetReader(tmp_path, columns=("episode_index", "frame_index"))
with pytest.raises(ValueError, match="episode 4"):
reader.read_episode(path.relative_to(tmp_path), episode_index=4, expected_rows=1)
def test_reader_supports_fsspec_remote_root() -> None:
filesystem = fsspec.filesystem("memory")
root = "memory://episode-reader"
path = "episode-reader/data/chunk-000/file-000.parquet"
with filesystem.open(path, "wb") as output:
pq.write_table(_table([0, 0, 0]), output)
reader = EpisodeParquetReader(root, columns=("episode_index", "value"))
table = reader.read_episode("data/chunk-000/file-000.parquet", episode_index=0, expected_rows=3)
assert table.column("value").to_pylist() == [0, 1, 2]
+118 -29
View File
@@ -10,12 +10,19 @@
import json
import struct
import threading
from concurrent.futures import ThreadPoolExecutor
import numpy as np
import pytest
from lerobot.datasets.episode_video_streaming import assert_hf_hub_range_cache_branch
from lerobot.datasets.mp4 import (
from lerobot.streaming.episode_video import (
EpisodeByteCache,
EpisodeVideoManifest,
ThreadLocalRangeFetcher,
_log_http_failure,
)
from lerobot.streaming.mp4 import (
_box,
_co64,
_dinf,
@@ -89,33 +96,115 @@ def test_parser_accepts_co64_chunk_offsets():
np.testing.assert_array_equal(mp4.sample_offsets, np.array([10_000, 10_050, 10_025]))
def test_hf_hub_branch_assertion_accepts_requested_revision(monkeypatch):
class FakeDist:
def read_text(self, name):
assert name == "direct_url.json"
return json.dumps(
{
"url": "https://github.com/huggingface/huggingface_hub.git",
"vcs_info": {"requested_revision": "feat/hffs-cache-cdn-range-reads"},
}
)
def _fake_cache(monkeypatch, tmp_path, *, byte_budget=8, max_open_decoders=1):
manifest = EpisodeVideoManifest(video_keys=["camera"], files=[], spans={})
cache = EpisodeByteCache(
manifest,
tmp_path,
byte_budget=byte_budget,
workers=1,
open_decoders=False,
max_open_decoders=max_open_decoders,
)
monkeypatch.setattr(
"lerobot.datasets.episode_video_streaming.metadata.distribution", lambda _: FakeDist()
cache,
"_fetch_and_synthesize",
lambda episode_index, _camera_key: {"bytes": bytes([episode_index]) * 5, "_timings": None},
)
return cache
def test_byte_cache_does_not_evict_retained_episode(monkeypatch, tmp_path):
with _fake_cache(monkeypatch, tmp_path, byte_budget=10) as cache:
cache.retain_episode(0)
cache.ensure_ready(0)
cache.ensure_ready(1)
cache.ensure_ready(2)
assert (0, "camera") in cache._cache
assert (1, "camera") not in cache._cache
assert cache.resident_bytes <= cache.byte_budget
def test_byte_cache_rejects_retained_set_larger_than_budget(monkeypatch, tmp_path):
with _fake_cache(monkeypatch, tmp_path, byte_budget=4) as cache:
cache.retain_episode(0)
with pytest.raises(MemoryError, match="byte budget"):
cache.ensure_ready(0)
def test_decoder_count_has_independent_limit(monkeypatch, tmp_path):
opened = []
class FakeDecoder:
pass
def open_decoder(_data):
decoder = FakeDecoder()
opened.append(decoder)
return decoder
monkeypatch.setattr("lerobot.streaming.episode_video.open_video_decoder", open_decoder)
with _fake_cache(monkeypatch, tmp_path, byte_budget=20, max_open_decoders=1) as cache:
first = cache.get_decoder(0, "camera")
second = cache.get_decoder(1, "camera")
assert first is not second
assert cache.open_decoder_count == 1
def test_releasing_episode_allows_immediate_eviction(monkeypatch, tmp_path):
with _fake_cache(monkeypatch, tmp_path, byte_budget=5) as cache:
cache.retain_episode(0)
cache.ensure_ready(0)
cache.release_episode(0)
cache.ensure_ready(1)
assert (0, "camera") not in cache._cache
assert (1, "camera") in cache._cache
def test_range_fetcher_closes_handles_from_all_worker_threads(tmp_path):
(tmp_path / "video.mp4").write_bytes(b"0123456789")
fetcher = ThreadLocalRangeFetcher(tmp_path)
barrier = threading.Barrier(2)
def read_from_worker(offset):
barrier.wait()
return fetcher.read_range("video.mp4", offset, 1)
with ThreadPoolExecutor(max_workers=2) as pool:
futures = [pool.submit(read_from_worker, offset) for offset in range(2)]
assert [future.result() for future in futures] == [b"0", b"1"]
handles = list(fetcher._all_handles.values())
assert len(handles) == 2
fetcher.close()
assert not fetcher._all_handles
assert all(handle.closed for handle in handles)
def test_http_failure_log_does_not_write_credentials(tmp_path, monkeypatch):
log_path = tmp_path / "http-failures.jsonl"
monkeypatch.setenv("LEROBOT_HTTP_FAILURE_LOG", str(log_path))
_log_http_failure(
backend="native-http",
method="GET",
url="https://cdn.example/private/video.mp4?token=url-secret",
headers={
"Authorization": "Bearer header-secret",
"Range": "bytes=0-10",
"X-Request-Id": "safe-request-id",
},
elapsed_s=0.1,
status_code=403,
)
assert_hf_hub_range_cache_branch()
def test_hf_hub_branch_assertion_rejects_plain_install(monkeypatch):
class FakeDist:
def read_text(self, name):
assert name == "direct_url.json"
return json.dumps({"url": "https://github.com/huggingface/huggingface_hub.git"})
monkeypatch.setattr(
"lerobot.datasets.episode_video_streaming.metadata.distribution", lambda _: FakeDist()
)
with pytest.raises(AssertionError):
assert_hf_hub_range_cache_branch()
record = json.loads(log_path.read_text())
assert record["host"] == "cdn.example"
assert record["path"] == "/private/video.mp4"
assert record["request_id"] == "safe-request-id"
assert "secret" not in log_path.read_text()
+1 -1
View File
@@ -16,7 +16,7 @@
from collections import Counter
from lerobot.datasets.episode_video_streaming import ExactCoveragePool
from lerobot.streaming.episode_video import ExactCoveragePool
EPISODES = [(0, 5), (1, 3), (2, 8), (3, 1), (4, 6), (5, 4), (6, 7), (7, 2)]
TOTAL = sum(n for _, n in EPISODES)
+14
View File
@@ -25,6 +25,7 @@ from lerobot.datasets.language import ( # noqa: E402
language_persistent_arrow_type,
validate_camera_field,
)
from lerobot.datasets.streaming_dataset import StreamingLeRobotDataset # noqa: E402
from lerobot.datasets.utils import DEFAULT_DATA_PATH # noqa: E402
@@ -171,3 +172,16 @@ def test_lerobot_dataset_passes_language_columns_through(tmp_path, empty_lerobot
assert first[LANGUAGE_EVENTS] == [event]
assert second[LANGUAGE_PERSISTENT] == persistent
assert second[LANGUAGE_EVENTS] == []
streamed = {
int(item["index"]): item
for item in StreamingLeRobotDataset(
repo_id=dataset.repo_id,
root=root,
shuffle=False,
)
}
assert streamed[0][LANGUAGE_PERSISTENT] == persistent
assert streamed[0][LANGUAGE_EVENTS] == [event]
assert streamed[1][LANGUAGE_PERSISTENT] == persistent
assert streamed[1][LANGUAGE_EVENTS] == []
+15 -84
View File
@@ -13,64 +13,16 @@
# 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.
import numpy as np
import pytest
import torch
pytest.importorskip("datasets", reason="datasets is required (install lerobot[dataset])")
from lerobot.datasets.streaming_dataset import StreamingLeRobotDataset
from lerobot.datasets.utils import safe_shard
from lerobot.utils.constants import ACTION
from tests.fixtures.constants import DUMMY_REPO_ID
def get_frames_expected_order(streaming_ds: StreamingLeRobotDataset) -> list[int]:
"""Replicates the shuffling logic of StreamingLeRobotDataset to get the expected order of indices."""
rng = np.random.default_rng(streaming_ds.seed)
buffer_size = streaming_ds.buffer_size
num_shards = streaming_ds.num_shards
shards_indices = []
for shard_idx in range(num_shards):
shard = streaming_ds.hf_dataset.shard(num_shards, index=shard_idx)
shard_indices = [item["index"] for item in shard]
shards_indices.append(shard_indices)
shard_iterators = {i: iter(s) for i, s in enumerate(shards_indices)}
buffer_indices_generator = streaming_ds._iter_random_indices(rng, buffer_size)
frames_buffer = []
expected_indices = []
while shard_iterators: # While there are still available shards
available_shard_keys = list(shard_iterators.keys())
if not available_shard_keys:
break
# Call _infinite_generator_over_elements with current available shards (key difference!)
shard_key = next(streaming_ds._infinite_generator_over_elements(rng, available_shard_keys))
try:
frame_index = next(shard_iterators[shard_key])
if len(frames_buffer) == buffer_size:
i = next(buffer_indices_generator)
expected_indices.append(frames_buffer[i])
frames_buffer[i] = frame_index
else:
frames_buffer.append(frame_index)
except StopIteration:
del shard_iterators[shard_key] # Remove exhausted shard
rng.shuffle(frames_buffer)
expected_indices.extend(frames_buffer)
return expected_indices
def test_single_frame_consistency(tmp_path, lerobot_dataset_factory):
"""Test if are correctly accessed"""
ds_num_frames = 400
@@ -92,7 +44,7 @@ def test_single_frame_consistency(tmp_path, lerobot_dataset_factory):
key_checks = []
for _ in range(ds_num_frames):
streaming_frame = next(streaming_ds)
frame_idx = streaming_frame["index"]
frame_idx = int(streaming_frame["index"])
target_frame = ds[frame_idx]
for key in streaming_frame:
@@ -141,22 +93,16 @@ def test_frames_order_over_epochs(tmp_path, lerobot_dataset_factory, shuffle):
repo_id=repo_id, root=local_path, buffer_size=buffer_size, seed=seed, shuffle=shuffle
)
first_epoch_indices = [frame["index"] for frame in streaming_ds]
expected_indices = get_frames_expected_order(streaming_ds)
assert first_epoch_indices == expected_indices, "First epoch indices do not match expected indices"
expected_indices = get_frames_expected_order(streaming_ds)
first_epoch_indices = [int(frame["index"]) for frame in streaming_ds]
assert sorted(first_epoch_indices) == list(range(ds_num_frames))
for _ in range(n_epochs):
streaming_indices = [frame["index"] for frame in streaming_ds]
frames_match = all(
s_index == e_index for s_index, e_index in zip(streaming_indices, expected_indices, strict=True)
)
streaming_indices = [int(frame["index"]) for frame in streaming_ds]
assert sorted(streaming_indices) == list(range(ds_num_frames))
if shuffle:
assert not frames_match
assert streaming_indices != first_epoch_indices
else:
assert frames_match
assert streaming_indices == first_epoch_indices
@pytest.mark.parametrize(
@@ -196,22 +142,16 @@ def test_frames_order_with_shards(tmp_path, lerobot_dataset_factory, shuffle):
max_num_shards=4,
)
first_epoch_indices = [frame["index"] for frame in streaming_ds]
expected_indices = get_frames_expected_order(streaming_ds)
assert first_epoch_indices == expected_indices, "First epoch indices do not match expected indices"
first_epoch_indices = [int(frame["index"]) for frame in streaming_ds]
assert sorted(first_epoch_indices) == list(range(ds_num_frames))
for _ in range(n_epochs):
streaming_indices = [
frame["index"] for frame in streaming_ds
] # NOTE: this is the same as first_epoch_indices
frames_match = all(
s_index == e_index for s_index, e_index in zip(streaming_indices, expected_indices, strict=True)
)
streaming_indices = [int(frame["index"]) for frame in streaming_ds]
assert sorted(streaming_indices) == list(range(ds_num_frames))
if shuffle:
assert not frames_match
assert streaming_indices != first_epoch_indices
else:
assert frames_match
assert streaming_indices == first_epoch_indices
@pytest.mark.parametrize(
@@ -261,7 +201,7 @@ def test_frames_with_delta_consistency(tmp_path, lerobot_dataset_factory, state_
for i in range(ds_num_frames):
streaming_frame = next(streaming_ds)
frame_idx = streaming_frame["index"]
frame_idx = int(streaming_frame["index"])
target_frame = ds[frame_idx]
assert set(streaming_frame.keys()) == set(target_frame.keys()), (
@@ -344,20 +284,11 @@ def test_frames_with_delta_consistency_with_shards(
max_num_shards=4,
)
iter(streaming_ds)
num_shards = 4
shards_indices = []
for shard_idx in range(num_shards):
shard = safe_shard(streaming_ds.hf_dataset, shard_idx, num_shards)
shard_indices = [item["index"] for item in shard]
shards_indices.append(shard_indices)
streaming_ds = iter(streaming_ds)
for i in range(ds_num_frames):
streaming_frame = next(streaming_ds)
frame_idx = streaming_frame["index"]
frame_idx = int(streaming_frame["index"])
target_frame = ds[frame_idx]
assert set(streaming_frame.keys()) == set(target_frame.keys()), (
+58
View File
@@ -0,0 +1,58 @@
#!/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
from types import SimpleNamespace
import pytest
pytest.importorskip("datasets", reason="datasets is required (install lerobot[dataset])")
from lerobot.configs.default import DatasetConfig
from lerobot.datasets import factory
def test_factory_wires_production_streaming_settings(monkeypatch):
captured = {}
class DummyStreamingDataset:
def __init__(self, *args, **kwargs):
captured["args"] = args
captured["kwargs"] = kwargs
self.meta = SimpleNamespace(camera_keys=[], depth_keys=[], stats={})
monkeypatch.setattr(factory, "LeRobotDatasetMetadata", lambda *args, **kwargs: object())
monkeypatch.setattr(factory, "resolve_delta_timestamps", lambda *args, **kwargs: {"action": [0.0]})
monkeypatch.setattr(factory, "StreamingLeRobotDataset", DummyStreamingDataset)
dataset_config = DatasetConfig(
repo_id="owner/dataset",
streaming=True,
streaming_data_root="memory://payload",
streaming_episode_pool_size=7,
streaming_prefetch_episodes=3,
streaming_byte_budget_gb=2.5,
)
cfg = SimpleNamespace(
dataset=dataset_config,
trainable_config=object(),
num_workers=0,
tolerance_s=1e-4,
)
dataset = factory.make_dataset(cfg)
assert isinstance(dataset, DummyStreamingDataset)
assert captured["args"] == ("owner/dataset",)
assert captured["kwargs"]["data_root"] == "memory://payload"
assert captured["kwargs"]["episode_pool_size"] == 7
assert captured["kwargs"]["prefetch_episodes"] == 3
assert captured["kwargs"]["byte_budget_gb"] == 2.5
assert captured["kwargs"]["max_num_shards"] == 1
assert captured["kwargs"]["return_uint8"] is True
assert captured["kwargs"]["repeat"] is True
+519
View File
@@ -0,0 +1,519 @@
#!/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
from __future__ import annotations
from itertools import islice
from pathlib import Path
import fsspec
import pytest
import torch
pytest.importorskip("datasets", reason="datasets is required (install lerobot[dataset])")
from lerobot.datasets.streaming_dataset import StreamingLeRobotDataset
from lerobot.utils.utils import cycle
from tests.fixtures.constants import DUMMY_REPO_ID
def _indices(dataset: StreamingLeRobotDataset) -> list[int]:
return [int(item["index"]) for item in dataset]
def _assert_item_equal(left: dict, right: dict) -> None:
assert left.keys() == right.keys()
for key in left:
if isinstance(left[key], torch.Tensor):
assert torch.equal(left[key], right[key]), key
else:
assert left[key] == right[key], key
def test_streaming_matches_map_style_with_exact_coverage(tmp_path: Path, lerobot_dataset_factory) -> None:
root = tmp_path / "dataset"
map_dataset = lerobot_dataset_factory(
root=root,
repo_id=DUMMY_REPO_ID,
total_episodes=4,
total_frames=40,
use_videos=False,
)
streaming = StreamingLeRobotDataset(
DUMMY_REPO_ID,
root=root,
shuffle=False,
buffer_size=3,
)
samples = list(streaming)
assert len(samples) == len(map_dataset)
assert sorted(int(sample["index"]) for sample in samples) == list(range(len(map_dataset)))
for sample in samples:
_assert_item_equal(sample, map_dataset[int(sample["index"])])
def test_streaming_rgb_video_matches_map_style(tmp_path: Path, lerobot_dataset_factory) -> None:
root = tmp_path / "dataset"
map_dataset = lerobot_dataset_factory(
root=root,
repo_id=DUMMY_REPO_ID,
total_episodes=2,
total_frames=20,
)
streaming = StreamingLeRobotDataset(
DUMMY_REPO_ID,
root=root,
shuffle=False,
buffer_size=2,
)
for sample in streaming:
reference = map_dataset[int(sample["index"])]
assert sample.keys() == reference.keys()
for camera_key in map_dataset.meta.camera_keys:
assert torch.equal(sample[camera_key], reference[camera_key]), (
camera_key,
int(sample["index"]),
float((sample[camera_key] - reference[camera_key]).abs().max()),
)
def test_streaming_applies_rgb_transforms_and_preserves_uint8(
tmp_path: Path, lerobot_dataset_factory
) -> None:
root = tmp_path / "dataset"
def flip_width(image: torch.Tensor) -> torch.Tensor:
return image.flip(-1)
map_dataset = lerobot_dataset_factory(
root=root,
repo_id=DUMMY_REPO_ID,
total_episodes=2,
total_frames=10,
image_transforms=flip_width,
return_uint8=True,
)
streaming = StreamingLeRobotDataset(
DUMMY_REPO_ID,
root=root,
shuffle=False,
buffer_size=2,
image_transforms=flip_width,
return_uint8=True,
)
sample = next(iter(streaming))
reference = map_dataset[int(sample["index"])]
for camera_key in map_dataset.meta.camera_keys:
assert sample[camera_key].dtype == torch.uint8
assert torch.equal(sample[camera_key], reference[camera_key])
def test_streaming_honors_episode_subset(tmp_path: Path, lerobot_dataset_factory) -> None:
root = tmp_path / "dataset"
map_dataset = lerobot_dataset_factory(
root=root,
repo_id=DUMMY_REPO_ID,
total_episodes=5,
total_frames=50,
use_videos=False,
)
selected = [1, 3]
streaming = StreamingLeRobotDataset(
DUMMY_REPO_ID,
root=root,
episodes=selected,
shuffle=False,
buffer_size=2,
)
indices = _indices(streaming)
expected = [
index
for episode in selected
for index in range(
map_dataset.meta.episodes[episode]["dataset_from_index"],
map_dataset.meta.episodes[episode]["dataset_to_index"],
)
]
assert sorted(indices) == sorted(expected)
def test_streaming_reads_episode_parquet_from_configured_fsspec_root(
tmp_path: Path, lerobot_dataset_factory
) -> None:
root = tmp_path / "metadata"
map_dataset = lerobot_dataset_factory(
root=root,
repo_id=DUMMY_REPO_ID,
total_episodes=3,
total_frames=30,
use_videos=False,
)
remote_root = "memory://streaming-production"
filesystem = fsspec.filesystem("memory")
for path in (root / "data").glob("*/*.parquet"):
relative = path.relative_to(root).as_posix()
filesystem.put(str(path), f"streaming-production/{relative}")
streaming = StreamingLeRobotDataset(
DUMMY_REPO_ID,
root=root,
data_root=remote_root,
shuffle=False,
buffer_size=2,
)
assert sorted(_indices(streaming)) == list(range(len(map_dataset)))
def test_streaming_reads_video_bytes_from_configured_fsspec_root(
tmp_path: Path, lerobot_dataset_factory
) -> None:
root = tmp_path / "metadata"
map_dataset = lerobot_dataset_factory(
root=root,
repo_id=DUMMY_REPO_ID,
total_episodes=2,
total_frames=10,
)
namespace = f"streaming-video-{tmp_path.name}"
remote_root = f"memory://{namespace}"
filesystem = fsspec.filesystem("memory")
for path in [*(root / "data").glob("*/*.parquet"), *(root / "videos").glob("*/*/*.mp4")]:
relative = path.relative_to(root).as_posix()
filesystem.put(str(path), f"{namespace}/{relative}")
streaming = StreamingLeRobotDataset(
DUMMY_REPO_ID,
root=root,
data_root=remote_root,
shuffle=False,
buffer_size=2,
)
sample = next(iter(streaming))
reference = map_dataset[int(sample["index"])]
for camera_key in map_dataset.meta.camera_keys:
assert torch.equal(sample[camera_key], reference[camera_key])
def test_streaming_rank_shards_are_disjoint(tmp_path: Path, lerobot_dataset_factory, monkeypatch) -> None:
root = tmp_path / "dataset"
map_dataset = lerobot_dataset_factory(
root=root,
repo_id=DUMMY_REPO_ID,
total_episodes=8,
total_frames=80,
use_videos=False,
)
per_rank = []
for rank in range(2):
monkeypatch.setenv("RANK", str(rank))
monkeypatch.setenv("WORLD_SIZE", "2")
per_rank.append(
set(
_indices(
StreamingLeRobotDataset(
DUMMY_REPO_ID,
root=root,
shuffle=False,
buffer_size=2,
)
)
)
)
assert per_rank[0].isdisjoint(per_rank[1])
assert per_rank[0] | per_rank[1] == set(range(len(map_dataset)))
def test_streaming_workers_do_not_duplicate_frames(tmp_path: Path, lerobot_dataset_factory) -> None:
root = tmp_path / "dataset"
map_dataset = lerobot_dataset_factory(
root=root,
repo_id=DUMMY_REPO_ID,
total_episodes=8,
total_frames=80,
use_videos=False,
)
streaming = StreamingLeRobotDataset(
DUMMY_REPO_ID,
root=root,
shuffle=False,
buffer_size=2,
)
loader = torch.utils.data.DataLoader(streaming, batch_size=None, num_workers=2)
indices = [int(item["index"]) for item in loader]
assert len(indices) == len(map_dataset)
assert set(indices) == set(range(len(map_dataset)))
def test_streaming_persistent_workers_advance_epochs(tmp_path: Path, lerobot_dataset_factory) -> None:
root = tmp_path / "dataset"
map_dataset = lerobot_dataset_factory(
root=root,
repo_id=DUMMY_REPO_ID,
total_episodes=8,
total_frames=80,
use_videos=False,
)
streaming = StreamingLeRobotDataset(
DUMMY_REPO_ID,
root=root,
seed=23,
shuffle=True,
buffer_size=2,
)
loader = torch.utils.data.DataLoader(
streaming,
batch_size=None,
num_workers=2,
persistent_workers=True,
)
try:
first = [int(item["index"]) for item in loader]
second = [int(item["index"]) for item in loader]
finally:
if loader._iterator is not None:
loader._iterator._shutdown_workers()
assert sorted(first) == list(range(len(map_dataset)))
assert sorted(second) == list(range(len(map_dataset)))
assert first != second
def test_streaming_worker_exception_propagates_and_workers_stop(
tmp_path: Path, lerobot_dataset_factory
) -> None:
root = tmp_path / "dataset"
lerobot_dataset_factory(
root=root,
repo_id=DUMMY_REPO_ID,
total_episodes=4,
total_frames=40,
use_videos=False,
)
streaming = StreamingLeRobotDataset(
DUMMY_REPO_ID,
root=root,
shuffle=False,
buffer_size=2,
)
next((root / "data").glob("*/*.parquet")).write_bytes(b"corrupt parquet")
loader = torch.utils.data.DataLoader(
streaming,
batch_size=None,
num_workers=2,
persistent_workers=True,
)
try:
with pytest.raises(Exception, match="Parquet"):
list(loader)
finally:
if loader._iterator is not None:
loader._iterator._shutdown_workers()
assert not any(worker.is_alive() for worker in loader._iterator._workers)
def test_streaming_resume_reproduces_remaining_stream(tmp_path: Path, lerobot_dataset_factory) -> None:
root = tmp_path / "dataset"
lerobot_dataset_factory(
root=root,
repo_id=DUMMY_REPO_ID,
total_episodes=5,
total_frames=50,
use_videos=False,
)
full = _indices(
StreamingLeRobotDataset(
DUMMY_REPO_ID,
root=root,
seed=7,
shuffle=True,
buffer_size=3,
)
)
resumed = StreamingLeRobotDataset(
DUMMY_REPO_ID,
root=root,
seed=7,
shuffle=True,
buffer_size=3,
)
resumed.load_state_dict({"epoch": 0, "offset": 11})
assert _indices(resumed) == full[11:]
@pytest.mark.parametrize(("batch_size", "offset"), [(None, 17), (4, 20)])
def test_streaming_worker_resume_reproduces_remaining_stream(
tmp_path: Path,
lerobot_dataset_factory,
batch_size: int | None,
offset: int,
) -> None:
root = tmp_path / "dataset"
lerobot_dataset_factory(
root=root,
repo_id=DUMMY_REPO_ID,
total_episodes=8,
total_frames=80,
use_videos=False,
)
def load(dataset: StreamingLeRobotDataset) -> list[int]:
loader = torch.utils.data.DataLoader(dataset, batch_size=batch_size, num_workers=2)
if batch_size is None:
return [int(item["index"]) for item in loader]
return [int(index) for batch in loader for index in batch["index"]]
full = load(
StreamingLeRobotDataset(
DUMMY_REPO_ID,
root=root,
seed=31,
shuffle=True,
buffer_size=2,
)
)
resumed = StreamingLeRobotDataset(
DUMMY_REPO_ID,
root=root,
seed=31,
shuffle=True,
buffer_size=2,
)
resumed.load_state_dict({"epoch": 0, "offset": offset, "batch_size": batch_size or 1})
assert load(resumed) == full[offset:]
def test_streaming_state_dict_round_trip_mid_epoch(tmp_path: Path, lerobot_dataset_factory) -> None:
root = tmp_path / "dataset"
lerobot_dataset_factory(
root=root,
repo_id=DUMMY_REPO_ID,
total_episodes=5,
total_frames=50,
use_videos=False,
)
source = StreamingLeRobotDataset(
DUMMY_REPO_ID,
root=root,
seed=17,
shuffle=True,
buffer_size=3,
)
iterator = iter(source)
consumed = [int(next(iterator)["index"]) for _ in range(13)]
state = source.state_dict()
remaining = [int(item["index"]) for item in iterator]
restored = StreamingLeRobotDataset(
DUMMY_REPO_ID,
root=root,
seed=17,
shuffle=True,
buffer_size=3,
)
restored.load_state_dict(state)
assert len(consumed) == state["offset"]
assert _indices(restored) == remaining
def test_streaming_worker_resume_after_epoch_boundary(tmp_path: Path, lerobot_dataset_factory) -> None:
root = tmp_path / "dataset"
lerobot_dataset_factory(
root=root,
repo_id=DUMMY_REPO_ID,
total_episodes=4,
total_frames=24,
use_videos=False,
)
def infinite_indices(dataset: StreamingLeRobotDataset, count: int) -> list[int]:
loader = torch.utils.data.DataLoader(
dataset,
batch_size=4,
num_workers=2,
persistent_workers=True,
)
try:
return [
int(index) for batch in islice(cycle(loader), (count + 3) // 4) for index in batch["index"]
][:count]
finally:
if loader._iterator is not None:
loader._iterator._shutdown_workers()
full = infinite_indices(
StreamingLeRobotDataset(
DUMMY_REPO_ID,
root=root,
seed=47,
shuffle=True,
buffer_size=2,
repeat=True,
),
56,
)
offset = 32
resumed = StreamingLeRobotDataset(
DUMMY_REPO_ID,
root=root,
seed=47,
shuffle=True,
buffer_size=2,
repeat=True,
)
resumed.load_state_dict({"epoch": 0, "offset": offset, "batch_size": 4})
assert infinite_indices(resumed, 24) == full[offset : offset + 24]
def test_streaming_local_training_step_smoke(tmp_path: Path, lerobot_dataset_factory) -> None:
root = tmp_path / "dataset"
lerobot_dataset_factory(
root=root,
repo_id=DUMMY_REPO_ID,
total_episodes=4,
total_frames=24,
use_videos=False,
)
dataset = StreamingLeRobotDataset(
DUMMY_REPO_ID,
root=root,
seed=53,
buffer_size=2,
repeat=True,
)
loader = torch.utils.data.DataLoader(dataset, batch_size=4, num_workers=2)
iterator = iter(loader)
try:
batch = next(iterator)
finally:
iterator._shutdown_workers()
model = torch.nn.Linear(batch["action"].shape[-1], batch["action"].shape[-1])
optimizer = torch.optim.Adam(model.parameters(), lr=1e-3)
loss = torch.nn.functional.mse_loss(model(batch["action"]), batch["action"])
loss.backward()
optimizer.step()
assert torch.isfinite(loss)
assert batch["index"].shape == (4,)
@@ -0,0 +1,47 @@
#!/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
from pathlib import Path
from types import SimpleNamespace
import pytest
pytest.importorskip("datasets", reason="datasets is required (install lerobot[dataset])")
from lerobot.datasets.streaming_sidecar import range_backend_for_root, streaming_data_root
def test_hub_data_root_is_revision_qualified() -> None:
meta = SimpleNamespace(repo_id="owner/dataset", revision="commit-sha")
root = streaming_data_root(meta, requested_root=None, configured_data_root=None)
assert root == "hf://datasets/owner/dataset@commit-sha"
assert range_backend_for_root(root) == "native-http"
def test_explicit_bucket_root_is_preserved() -> None:
meta = SimpleNamespace(repo_id="owner/dataset", revision="commit-sha")
bucket = "hf://buckets/owner/dataset-bucket/prefix/"
root = streaming_data_root(meta, requested_root=None, configured_data_root=bucket)
assert root == bucket.rstrip("/")
assert range_backend_for_root(root) == "native-http"
def test_local_and_generic_remote_roots_use_fsspec(tmp_path: Path) -> None:
meta = SimpleNamespace(repo_id="owner/dataset", revision="commit-sha")
local = streaming_data_root(meta, requested_root=tmp_path, configured_data_root=None)
assert local == str(tmp_path)
assert range_backend_for_root(local) == "fsspec"
assert range_backend_for_root("memory://dataset") == "fsspec"