mirror of
https://github.com/huggingface/lerobot.git
synced 2026-07-26 11:16:00 +00:00
refactor(dataset): make streaming pool rank-owned
This commit is contained in:
@@ -38,6 +38,7 @@ from lerobot.streaming.mp4 import (
|
||||
_vmhd,
|
||||
parse_mp4_index,
|
||||
synthesize_mp4,
|
||||
synthesized_mp4_size,
|
||||
)
|
||||
|
||||
|
||||
@@ -90,6 +91,15 @@ def test_synthesized_mp4_rebases_one_chunk_per_sample_offsets():
|
||||
np.testing.assert_array_equal(mini_index.sample_sizes, np.array([10, 10, 10]))
|
||||
|
||||
|
||||
def test_synthesized_mp4_size_matches_materialized_bytes():
|
||||
mp4 = parse_mp4_index("test.mp4", _minimal_mp4([10_000, 10_050, 10_025]))
|
||||
sample_slice = mp4.sample_slice(0.0, 2.0, keyframe_pad_s=0, keyframe_pad_fraction=0)
|
||||
|
||||
mini = synthesize_mp4(mp4, sample_slice, b"x" * sample_slice.byte_length)
|
||||
|
||||
assert synthesized_mp4_size(mp4, sample_slice) == len(mini)
|
||||
|
||||
|
||||
def test_parser_accepts_co64_chunk_offsets():
|
||||
mp4 = parse_mp4_index("test.mp4", _minimal_mp4([10_000, 10_050, 10_025], use_co64=True))
|
||||
|
||||
|
||||
@@ -16,6 +16,8 @@
|
||||
|
||||
from collections import Counter
|
||||
|
||||
import pytest
|
||||
|
||||
from lerobot.streaming.episode_video import ExactCoveragePool
|
||||
|
||||
EPISODES = [(0, 5), (1, 3), (2, 8), (3, 1), (4, 6), (5, 4), (6, 7), (7, 2)]
|
||||
@@ -95,3 +97,50 @@ def test_zero_length_episodes_skipped():
|
||||
pool = ExactCoveragePool([(0, 3), (1, 0), (2, 2)], pool_size=8, seed=0)
|
||||
out, _ = _drain(pool)
|
||||
assert Counter(out) == Counter({(0, 0): 1, (0, 1): 1, (0, 2): 1, (2, 0): 1, (2, 1): 1})
|
||||
|
||||
|
||||
def test_byte_aware_admission_never_exceeds_budget():
|
||||
sizes = {0: 7, 1: 6, 2: 4, 3: 3, 4: 2}
|
||||
pool = ExactCoveragePool(
|
||||
EPISODES[:5],
|
||||
pool_size=4,
|
||||
seed=9,
|
||||
episode_byte_sizes=sizes,
|
||||
byte_budget=10,
|
||||
)
|
||||
out = []
|
||||
max_resident_bytes = pool.resident_bytes
|
||||
while pool.remaining_total:
|
||||
out.append(next(pool))
|
||||
max_resident_bytes = max(max_resident_bytes, pool.resident_bytes)
|
||||
|
||||
assert Counter(out) == Counter((ep, frame) for ep, count in EPISODES[:5] for frame in range(count))
|
||||
assert max_resident_bytes <= 10
|
||||
assert len(pool.admission_order) == len(EPISODES[:5])
|
||||
|
||||
|
||||
def test_byte_aware_admission_rejects_one_oversized_episode():
|
||||
with pytest.raises(ValueError, match="Episode 1.*byte budget"):
|
||||
ExactCoveragePool(
|
||||
[(0, 2), (1, 3)],
|
||||
pool_size=2,
|
||||
seed=0,
|
||||
episode_byte_sizes={0: 4, 1: 11},
|
||||
byte_budget=10,
|
||||
)
|
||||
|
||||
|
||||
def test_prefetch_candidates_follow_deterministic_pending_frontier():
|
||||
pool = ExactCoveragePool(
|
||||
EPISODES,
|
||||
pool_size=2,
|
||||
seed=17,
|
||||
episode_byte_sizes={episode: 1 for episode, _ in EPISODES},
|
||||
byte_budget=2,
|
||||
)
|
||||
|
||||
candidates = pool.prefetch_candidates(3)
|
||||
|
||||
assert len(candidates) == 3
|
||||
assert not set(candidates) & set(pool.resident)
|
||||
assert candidates == pool.prefetch_candidates(3)
|
||||
|
||||
@@ -19,7 +19,7 @@ import torch
|
||||
|
||||
pytest.importorskip("datasets", reason="datasets is required (install lerobot[dataset])")
|
||||
|
||||
from lerobot.datasets.streaming_dataset import StreamingLeRobotDataset
|
||||
from lerobot.datasets.streaming_dataset import StreamingLeRobotDataset, _balanced_episode_shards
|
||||
from lerobot.utils.utils import cycle
|
||||
from tests.fixtures.constants import DUMMY_REPO_ID
|
||||
|
||||
@@ -209,6 +209,28 @@ def test_streaming_reads_video_bytes_from_configured_fsspec_root(
|
||||
assert torch.equal(sample[camera_key], reference[camera_key])
|
||||
|
||||
|
||||
def test_streaming_rejects_episode_larger_than_rank_byte_budget(
|
||||
tmp_path: Path, lerobot_dataset_factory
|
||||
) -> None:
|
||||
root = tmp_path / "dataset"
|
||||
lerobot_dataset_factory(
|
||||
root=root,
|
||||
repo_id=DUMMY_REPO_ID,
|
||||
total_episodes=2,
|
||||
total_frames=10,
|
||||
)
|
||||
streaming = StreamingLeRobotDataset(
|
||||
DUMMY_REPO_ID,
|
||||
root=root,
|
||||
shuffle=False,
|
||||
buffer_size=2,
|
||||
byte_budget_gb=1 / 1024**3,
|
||||
)
|
||||
|
||||
with pytest.raises(ValueError, match="Episode .*byte budget"):
|
||||
next(iter(streaming))
|
||||
|
||||
|
||||
def test_streaming_rank_shards_are_disjoint(tmp_path: Path, lerobot_dataset_factory, monkeypatch) -> None:
|
||||
root = tmp_path / "dataset"
|
||||
map_dataset = lerobot_dataset_factory(
|
||||
@@ -239,9 +261,22 @@ def test_streaming_rank_shards_are_disjoint(tmp_path: Path, lerobot_dataset_fact
|
||||
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:
|
||||
def test_rank_shards_are_greedily_balanced_by_frame_count() -> None:
|
||||
shards = _balanced_episode_shards(
|
||||
[0, 1, 2, 3, 4],
|
||||
{0: 100, 1: 90, 2: 20, 3: 10, 4: 5},
|
||||
world_size=2,
|
||||
)
|
||||
|
||||
assert {episode for shard in shards for episode in shard} == {0, 1, 2, 3, 4}
|
||||
assert set(shards[0]).isdisjoint(shards[1])
|
||||
totals = [sum({0: 100, 1: 90, 2: 20, 3: 10, 4: 5}[episode] for episode in shard) for shard in shards]
|
||||
assert max(totals) - min(totals) <= 15
|
||||
|
||||
|
||||
def test_streaming_rejects_multiple_sampling_workers(tmp_path: Path, lerobot_dataset_factory) -> None:
|
||||
root = tmp_path / "dataset"
|
||||
map_dataset = lerobot_dataset_factory(
|
||||
lerobot_dataset_factory(
|
||||
root=root,
|
||||
repo_id=DUMMY_REPO_ID,
|
||||
total_episodes=8,
|
||||
@@ -256,10 +291,8 @@ def test_streaming_workers_do_not_duplicate_frames(tmp_path: Path, lerobot_datas
|
||||
)
|
||||
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)))
|
||||
with pytest.raises(RuntimeError, match="one DataLoader worker per rank"):
|
||||
list(loader)
|
||||
|
||||
|
||||
def test_streaming_persistent_workers_advance_epochs(tmp_path: Path, lerobot_dataset_factory) -> None:
|
||||
@@ -281,7 +314,7 @@ def test_streaming_persistent_workers_advance_epochs(tmp_path: Path, lerobot_dat
|
||||
loader = torch.utils.data.DataLoader(
|
||||
streaming,
|
||||
batch_size=None,
|
||||
num_workers=2,
|
||||
num_workers=1,
|
||||
persistent_workers=True,
|
||||
)
|
||||
try:
|
||||
@@ -317,7 +350,7 @@ def test_streaming_worker_exception_propagates_and_workers_stop(
|
||||
loader = torch.utils.data.DataLoader(
|
||||
streaming,
|
||||
batch_size=None,
|
||||
num_workers=2,
|
||||
num_workers=1,
|
||||
persistent_workers=True,
|
||||
)
|
||||
try:
|
||||
@@ -376,7 +409,7 @@ def test_streaming_worker_resume_reproduces_remaining_stream(
|
||||
)
|
||||
|
||||
def load(dataset: StreamingLeRobotDataset) -> list[int]:
|
||||
loader = torch.utils.data.DataLoader(dataset, batch_size=batch_size, num_workers=2)
|
||||
loader = torch.utils.data.DataLoader(dataset, batch_size=batch_size, num_workers=1)
|
||||
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"]]
|
||||
@@ -450,7 +483,7 @@ def test_streaming_worker_resume_after_epoch_boundary(tmp_path: Path, lerobot_da
|
||||
loader = torch.utils.data.DataLoader(
|
||||
dataset,
|
||||
batch_size=4,
|
||||
num_workers=2,
|
||||
num_workers=1,
|
||||
persistent_workers=True,
|
||||
)
|
||||
try:
|
||||
@@ -502,7 +535,7 @@ def test_streaming_local_training_step_smoke(tmp_path: Path, lerobot_dataset_fac
|
||||
buffer_size=2,
|
||||
repeat=True,
|
||||
)
|
||||
loader = torch.utils.data.DataLoader(dataset, batch_size=4, num_workers=2)
|
||||
loader = torch.utils.data.DataLoader(dataset, batch_size=4, num_workers=1)
|
||||
iterator = iter(loader)
|
||||
try:
|
||||
batch = next(iterator)
|
||||
|
||||
Reference in New Issue
Block a user