mirror of
https://github.com/huggingface/lerobot.git
synced 2026-07-26 11:16:00 +00:00
feat(dataset): integrate episode streaming into training
This commit is contained in:
@@ -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]
|
||||
@@ -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()
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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] == []
|
||||
|
||||
@@ -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()), (
|
||||
|
||||
@@ -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
|
||||
@@ -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"
|
||||
Reference in New Issue
Block a user