mirror of
https://github.com/huggingface/lerobot.git
synced 2026-08-08 17:39:44 +00:00
tests(simplifications): further simplifying tests, reducing docstrings size and making them clearer.
This commit is contained in:
@@ -26,7 +26,6 @@ import lerobot.datasets.streaming_dataset as streaming_dataset_module
|
|||||||
from lerobot.datasets.dataset_metadata import LeRobotDatasetMetadata
|
from lerobot.datasets.dataset_metadata import LeRobotDatasetMetadata
|
||||||
from lerobot.datasets.lerobot_dataset import LeRobotDataset
|
from lerobot.datasets.lerobot_dataset import LeRobotDataset
|
||||||
from lerobot.datasets.streaming_dataset import StreamingLeRobotDataset
|
from lerobot.datasets.streaming_dataset import StreamingLeRobotDataset
|
||||||
from lerobot.datasets.utils import safe_shard
|
|
||||||
from lerobot.utils.constants import ACTION
|
from lerobot.utils.constants import ACTION
|
||||||
from tests.fixtures.constants import DUMMY_REPO_ID
|
from tests.fixtures.constants import DUMMY_REPO_ID
|
||||||
|
|
||||||
@@ -87,14 +86,9 @@ def get_frames_expected_order(streaming_ds: StreamingLeRobotDataset) -> list[int
|
|||||||
@pytest.mark.parametrize("from_local", [False, True])
|
@pytest.mark.parametrize("from_local", [False, True])
|
||||||
def test_streaming_dataset_forwards_hub_token_only_for_remote_data(tmp_path, monkeypatch, token, from_local):
|
def test_streaming_dataset_forwards_hub_token_only_for_remote_data(tmp_path, monkeypatch, token, from_local):
|
||||||
requested_root = tmp_path / "local" if from_local else None
|
requested_root = tmp_path / "local" if from_local else None
|
||||||
metadata = SimpleNamespace(
|
metadata = _fake_meta(
|
||||||
root=requested_root or tmp_path / "snapshot",
|
root=requested_root or tmp_path / "snapshot",
|
||||||
revision=streaming_dataset_module.CODEBASE_VERSION,
|
revision=streaming_dataset_module.CODEBASE_VERSION,
|
||||||
_version=streaming_dataset_module.CODEBASE_VERSION,
|
|
||||||
features={},
|
|
||||||
depth_keys=[],
|
|
||||||
image_keys=[],
|
|
||||||
rescale_depth_stats=Mock(),
|
|
||||||
)
|
)
|
||||||
metadata_cls = Mock(return_value=metadata)
|
metadata_cls = Mock(return_value=metadata)
|
||||||
load_dataset = Mock(return_value=SimpleNamespace(num_shards=1))
|
load_dataset = Mock(return_value=SimpleNamespace(num_shards=1))
|
||||||
@@ -119,10 +113,7 @@ def test_streaming_dataset_forwards_hub_token_only_for_remote_data(tmp_path, mon
|
|||||||
|
|
||||||
|
|
||||||
def assert_videos_roll_over(ds: LeRobotDataset) -> None:
|
def assert_videos_roll_over(ds: LeRobotDataset) -> None:
|
||||||
"""Guard the fixture: these tests are only meaningful if episodes live past ``file-000``.
|
"""Videos spanning several files must decode each frame from its own file (v3.0 rollover).
|
||||||
|
|
||||||
If ``episodes_per_video_file`` ever stops splitting the videos, the rollover tests below
|
|
||||||
would still pass while silently covering nothing.
|
|
||||||
"""
|
"""
|
||||||
for key in ds.meta.video_keys:
|
for key in ds.meta.video_keys:
|
||||||
episodes = [ds.meta.episodes[ep_idx] for ep_idx in range(ds.meta.total_episodes)]
|
episodes = [ds.meta.episodes[ep_idx] for ep_idx in range(ds.meta.total_episodes)]
|
||||||
@@ -171,8 +162,19 @@ def assert_frame_matches(streaming_frame: dict, target_frame: dict, ds: LeRobotD
|
|||||||
assert not mismatched, f"Streaming and target frames differ on {mismatched} ({context})"
|
assert not mismatched, f"Streaming and target frames differ on {mismatched} ({context})"
|
||||||
|
|
||||||
|
|
||||||
|
def assert_stream_matches_reference(
|
||||||
|
streaming_ds: StreamingLeRobotDataset, ds: LeRobotDataset, num_frames: int
|
||||||
|
) -> None:
|
||||||
|
"""Stream ``num_frames`` frames and assert each equals the same frame from the non-streaming reader."""
|
||||||
|
stream = iter(streaming_ds)
|
||||||
|
for i in range(num_frames):
|
||||||
|
streaming_frame = next(stream)
|
||||||
|
frame_idx = streaming_frame["index"]
|
||||||
|
assert_frame_matches(streaming_frame, ds[frame_idx], ds, context=f"i: {i}, frame_idx: {frame_idx}")
|
||||||
|
|
||||||
|
|
||||||
def test_single_frame_consistency(tmp_path, lerobot_dataset_factory):
|
def test_single_frame_consistency(tmp_path, lerobot_dataset_factory):
|
||||||
"""Test if are correctly accessed"""
|
"""Streaming without deltas returns the same frames as the non-streaming reader."""
|
||||||
ds_num_frames = 400
|
ds_num_frames = 400
|
||||||
ds_num_episodes = 10
|
ds_num_episodes = 10
|
||||||
buffer_size = 100
|
buffer_size = 100
|
||||||
@@ -187,32 +189,8 @@ def test_single_frame_consistency(tmp_path, lerobot_dataset_factory):
|
|||||||
total_frames=ds_num_frames,
|
total_frames=ds_num_frames,
|
||||||
)
|
)
|
||||||
|
|
||||||
streaming_ds = iter(StreamingLeRobotDataset(repo_id=repo_id, root=local_path, buffer_size=buffer_size))
|
streaming_ds = StreamingLeRobotDataset(repo_id=repo_id, root=local_path, buffer_size=buffer_size)
|
||||||
|
assert_stream_matches_reference(streaming_ds, ds, ds_num_frames)
|
||||||
key_checks = []
|
|
||||||
for _ in range(ds_num_frames):
|
|
||||||
streaming_frame = next(streaming_ds)
|
|
||||||
frame_idx = streaming_frame["index"]
|
|
||||||
target_frame = ds[frame_idx]
|
|
||||||
|
|
||||||
for key in streaming_frame:
|
|
||||||
left = streaming_frame[key]
|
|
||||||
right = target_frame[key]
|
|
||||||
|
|
||||||
if isinstance(left, str):
|
|
||||||
check = left == right
|
|
||||||
|
|
||||||
elif isinstance(left, torch.Tensor):
|
|
||||||
check = torch.allclose(left, right) and left.shape == right.shape
|
|
||||||
|
|
||||||
elif isinstance(left, float):
|
|
||||||
check = left == right.item() # right is a torch.Tensor
|
|
||||||
|
|
||||||
key_checks.append((key, check))
|
|
||||||
|
|
||||||
assert all(t[1] for t in key_checks), (
|
|
||||||
f"Checking {list(filter(lambda t: not t[1], key_checks))[0][0]} left and right were found different (frame_idx: {frame_idx})"
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.parametrize(
|
@pytest.mark.parametrize(
|
||||||
@@ -362,57 +340,7 @@ def test_iter_raises_on_nested_generator_error(tmp_path, lerobot_dataset_factory
|
|||||||
next(iter(streaming_ds))
|
next(iter(streaming_ds))
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.parametrize(
|
@pytest.mark.parametrize("sharded", [False, True])
|
||||||
"state_deltas, action_deltas",
|
|
||||||
[
|
|
||||||
([-1, -0.5, -0.20, 0], [0, 1, 2, 3]),
|
|
||||||
([-1, -0.5, -0.20, 0], [-1.5, -1, -0.5, -0.20, -0.10, 0]),
|
|
||||||
([-2, -1, -0.5, 0], [0, 1, 2, 3]),
|
|
||||||
([-2, -1, -0.5, 0], [-1.5, -1, -0.5, -0.20, -0.10, 0]),
|
|
||||||
],
|
|
||||||
)
|
|
||||||
def test_frames_with_delta_consistency(tmp_path, lerobot_dataset_factory, state_deltas, action_deltas):
|
|
||||||
ds_num_frames = 500
|
|
||||||
ds_num_episodes = 10
|
|
||||||
buffer_size = 100
|
|
||||||
|
|
||||||
seed = 42
|
|
||||||
|
|
||||||
local_path = tmp_path / "test"
|
|
||||||
repo_id = f"{DUMMY_REPO_ID}-ciao"
|
|
||||||
camera_key = "phone"
|
|
||||||
|
|
||||||
delta_timestamps = {
|
|
||||||
camera_key: state_deltas,
|
|
||||||
"state": state_deltas,
|
|
||||||
ACTION: action_deltas,
|
|
||||||
}
|
|
||||||
|
|
||||||
ds = lerobot_dataset_factory(
|
|
||||||
root=local_path,
|
|
||||||
repo_id=repo_id,
|
|
||||||
total_episodes=ds_num_episodes,
|
|
||||||
total_frames=ds_num_frames,
|
|
||||||
delta_timestamps=delta_timestamps,
|
|
||||||
)
|
|
||||||
|
|
||||||
streaming_ds = iter(
|
|
||||||
StreamingLeRobotDataset(
|
|
||||||
repo_id=repo_id,
|
|
||||||
root=local_path,
|
|
||||||
buffer_size=buffer_size,
|
|
||||||
seed=seed,
|
|
||||||
shuffle=False,
|
|
||||||
delta_timestamps=delta_timestamps,
|
|
||||||
)
|
|
||||||
)
|
|
||||||
|
|
||||||
for i in range(ds_num_frames):
|
|
||||||
streaming_frame = next(streaming_ds)
|
|
||||||
frame_idx = streaming_frame["index"]
|
|
||||||
assert_frame_matches(streaming_frame, ds[frame_idx], ds, context=f"i: {i}, frame_idx: {frame_idx}")
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.parametrize(
|
@pytest.mark.parametrize(
|
||||||
"state_deltas, action_deltas",
|
"state_deltas, action_deltas",
|
||||||
[
|
[
|
||||||
@@ -422,61 +350,40 @@ def test_frames_with_delta_consistency(tmp_path, lerobot_dataset_factory, state_
|
|||||||
([-2, -1, -0.5, 0], [-20, -1.5, -1, -0.5, -0.20, -0.10, 0]),
|
([-2, -1, -0.5, 0], [-20, -1.5, -1, -0.5, -0.20, -0.10, 0]),
|
||||||
],
|
],
|
||||||
)
|
)
|
||||||
def test_frames_with_delta_consistency_with_shards(
|
def test_frames_with_delta_consistency(
|
||||||
tmp_path, lerobot_dataset_factory, state_deltas, action_deltas
|
tmp_path, lerobot_dataset_factory, sharded, state_deltas, action_deltas
|
||||||
):
|
):
|
||||||
ds_num_frames = 100
|
"""Delta-window frames streamed match the non-streaming reader, with and without sharding."""
|
||||||
ds_num_episodes = 10
|
|
||||||
buffer_size = 10
|
|
||||||
data_file_size_mb = 0.001
|
|
||||||
chunks_size = 1
|
|
||||||
|
|
||||||
seed = 42
|
|
||||||
|
|
||||||
local_path = tmp_path / "test"
|
local_path = tmp_path / "test"
|
||||||
repo_id = f"{DUMMY_REPO_ID}-ciao"
|
repo_id = f"{DUMMY_REPO_ID}-ciao"
|
||||||
camera_key = "phone"
|
delta_timestamps = {"phone": state_deltas, "state": state_deltas, ACTION: action_deltas}
|
||||||
|
|
||||||
delta_timestamps = {
|
if sharded:
|
||||||
camera_key: state_deltas,
|
num_frames, buffer_size = 100, 10
|
||||||
"state": state_deltas,
|
factory_extra = {"data_files_size_in_mb": 0.001, "chunks_size": 1}
|
||||||
ACTION: action_deltas,
|
stream_extra = {"max_num_shards": 4}
|
||||||
}
|
else:
|
||||||
|
num_frames, buffer_size = 500, 100
|
||||||
|
factory_extra, stream_extra = {}, {}
|
||||||
|
|
||||||
ds = lerobot_dataset_factory(
|
ds = lerobot_dataset_factory(
|
||||||
root=local_path,
|
root=local_path,
|
||||||
repo_id=repo_id,
|
repo_id=repo_id,
|
||||||
total_episodes=ds_num_episodes,
|
total_episodes=10,
|
||||||
total_frames=ds_num_frames,
|
total_frames=num_frames,
|
||||||
delta_timestamps=delta_timestamps,
|
delta_timestamps=delta_timestamps,
|
||||||
data_files_size_in_mb=data_file_size_mb,
|
**factory_extra,
|
||||||
chunks_size=chunks_size,
|
|
||||||
)
|
)
|
||||||
streaming_ds = StreamingLeRobotDataset(
|
streaming_ds = StreamingLeRobotDataset(
|
||||||
repo_id=repo_id,
|
repo_id=repo_id,
|
||||||
root=local_path,
|
root=local_path,
|
||||||
buffer_size=buffer_size,
|
buffer_size=buffer_size,
|
||||||
seed=seed,
|
seed=42,
|
||||||
shuffle=False,
|
shuffle=False,
|
||||||
delta_timestamps=delta_timestamps,
|
delta_timestamps=delta_timestamps,
|
||||||
max_num_shards=4,
|
**stream_extra,
|
||||||
)
|
)
|
||||||
|
assert_stream_matches_reference(streaming_ds, ds, num_frames)
|
||||||
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"]
|
|
||||||
assert_frame_matches(streaming_frame, ds[frame_idx], ds, context=f"i: {i}, frame_idx: {frame_idx}")
|
|
||||||
|
|
||||||
|
|
||||||
class _StopConstructionError(Exception):
|
class _StopConstructionError(Exception):
|
||||||
@@ -490,7 +397,9 @@ def _fake_meta(*args, **kwargs):
|
|||||||
revision = kwargs.get("revision", args[2] if len(args) > 2 else None)
|
revision = kwargs.get("revision", args[2] if len(args) > 2 else None)
|
||||||
meta.root = root or "/tmp/_streaming_meta"
|
meta.root = root or "/tmp/_streaming_meta"
|
||||||
meta.revision = revision or "v0"
|
meta.revision = revision or "v0"
|
||||||
meta._version = "v3.0"
|
meta._version = streaming_dataset_module.CODEBASE_VERSION
|
||||||
|
meta.features = {}
|
||||||
|
meta.video_keys = []
|
||||||
meta.depth_keys = []
|
meta.depth_keys = []
|
||||||
meta.image_keys = []
|
meta.image_keys = []
|
||||||
meta.rescale_depth_stats = lambda *_a, **_k: None
|
meta.rescale_depth_stats = lambda *_a, **_k: None
|
||||||
@@ -644,71 +553,26 @@ def test_invalid_repo_type_fails_before_io():
|
|||||||
StreamingLeRobotDataset(DUMMY_REPO_ID, repo_type="space")
|
StreamingLeRobotDataset(DUMMY_REPO_ID, repo_type="space")
|
||||||
|
|
||||||
|
|
||||||
def test_single_frame_consistency_across_video_files(tmp_path, lerobot_dataset_factory):
|
|
||||||
"""Streaming a dataset whose videos span several files must decode each frame from its own file.
|
|
||||||
|
|
||||||
Regression test for decoding at a *global* timestamp (`index / fps`). That position only
|
|
||||||
exists while the whole dataset fits in one .mp4; once v3.0 rolls the video over, every
|
|
||||||
episode in a later file asked for a frame past the end of the file being read
|
|
||||||
(`IndexError: Invalid frame index=... must be less than ...`).
|
|
||||||
"""
|
|
||||||
buffer_size = 100
|
|
||||||
|
|
||||||
local_path = tmp_path / "test"
|
|
||||||
repo_id = f"{DUMMY_REPO_ID}-video-rollover"
|
|
||||||
|
|
||||||
ds = lerobot_dataset_factory(
|
|
||||||
root=local_path,
|
|
||||||
repo_id=repo_id,
|
|
||||||
total_episodes=MULTI_FILE_EPISODES,
|
|
||||||
total_frames=MULTI_FILE_FRAMES,
|
|
||||||
episodes_per_video_file=MULTI_FILE_EPISODES_PER_VIDEO_FILE,
|
|
||||||
)
|
|
||||||
assert_videos_roll_over(ds)
|
|
||||||
|
|
||||||
streaming_ds = iter(
|
|
||||||
StreamingLeRobotDataset(
|
|
||||||
repo_id=repo_id,
|
|
||||||
root=local_path,
|
|
||||||
buffer_size=buffer_size,
|
|
||||||
shuffle=False,
|
|
||||||
)
|
|
||||||
)
|
|
||||||
|
|
||||||
for _ in range(MULTI_FILE_FRAMES):
|
|
||||||
streaming_frame = next(streaming_ds)
|
|
||||||
frame_idx = streaming_frame["index"]
|
|
||||||
assert_frame_matches(streaming_frame, ds[frame_idx], ds, context=f"frame_idx: {frame_idx}")
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.parametrize(
|
@pytest.mark.parametrize(
|
||||||
"state_deltas, action_deltas",
|
"state_deltas, action_deltas",
|
||||||
[
|
[
|
||||||
|
(None, None),
|
||||||
([-1, -0.5, -0.20, 0], [0, 1, 2, 3]),
|
([-1, -0.5, -0.20, 0], [0, 1, 2, 3]),
|
||||||
([-2, -1, -0.5, 0], [-1.5, -1, -0.5, -0.20, -0.10, 0]),
|
([-2, -1, -0.5, 0], [-1.5, -1, -0.5, -0.20, -0.10, 0]),
|
||||||
],
|
],
|
||||||
)
|
)
|
||||||
def test_frames_with_delta_consistency_across_video_files(
|
def test_consistency_across_video_files(
|
||||||
tmp_path, lerobot_dataset_factory, state_deltas, action_deltas
|
tmp_path, lerobot_dataset_factory, state_deltas, action_deltas
|
||||||
):
|
):
|
||||||
"""Same rollover, on the delta path.
|
"""Videos spanning several files must decode each frame from its own file (v3.0 rollover).
|
||||||
|
|
||||||
Here the old global timestamp failed silently rather than raising: the query was clamped to
|
|
||||||
the episode's `to_timestamp`, so every frame decoded the episode's *last* frame — a frozen
|
|
||||||
video paired with advancing state/action.
|
|
||||||
"""
|
"""
|
||||||
buffer_size = 100
|
|
||||||
seed = 42
|
|
||||||
|
|
||||||
local_path = tmp_path / "test"
|
local_path = tmp_path / "test"
|
||||||
repo_id = f"{DUMMY_REPO_ID}-video-rollover-deltas"
|
repo_id = f"{DUMMY_REPO_ID}-video-rollover"
|
||||||
camera_key = "phone"
|
delta_timestamps = (
|
||||||
|
None
|
||||||
delta_timestamps = {
|
if state_deltas is None
|
||||||
camera_key: state_deltas,
|
else {"phone": state_deltas, "state": state_deltas, ACTION: action_deltas}
|
||||||
"state": state_deltas,
|
)
|
||||||
ACTION: action_deltas,
|
|
||||||
}
|
|
||||||
|
|
||||||
ds = lerobot_dataset_factory(
|
ds = lerobot_dataset_factory(
|
||||||
root=local_path,
|
root=local_path,
|
||||||
@@ -720,18 +584,12 @@ def test_frames_with_delta_consistency_across_video_files(
|
|||||||
)
|
)
|
||||||
assert_videos_roll_over(ds)
|
assert_videos_roll_over(ds)
|
||||||
|
|
||||||
streaming_ds = iter(
|
streaming_ds = StreamingLeRobotDataset(
|
||||||
StreamingLeRobotDataset(
|
|
||||||
repo_id=repo_id,
|
repo_id=repo_id,
|
||||||
root=local_path,
|
root=local_path,
|
||||||
buffer_size=buffer_size,
|
buffer_size=100,
|
||||||
seed=seed,
|
seed=42,
|
||||||
shuffle=False,
|
shuffle=False,
|
||||||
delta_timestamps=delta_timestamps,
|
delta_timestamps=delta_timestamps,
|
||||||
)
|
)
|
||||||
)
|
assert_stream_matches_reference(streaming_ds, ds, MULTI_FILE_FRAMES)
|
||||||
|
|
||||||
for i in range(MULTI_FILE_FRAMES):
|
|
||||||
streaming_frame = next(streaming_ds)
|
|
||||||
frame_idx = streaming_frame["index"]
|
|
||||||
assert_frame_matches(streaming_frame, ds[frame_idx], ds, context=f"i: {i}, frame_idx: {frame_idx}")
|
|
||||||
|
|||||||
Reference in New Issue
Block a user