annotation-studio / test_subset.py
pepijn223's picture
pepijn223 HF Staff
Add Human demo videos (robot → human, fal MiniMax-H3 480P) as a Video generation option (#1)
f19f993
Raw History Blame Contribute Delete
14.3 kB
"""Real shared-file video/data fixtures: subset correctness and ordinary LeRobot loading."""
import json
import sys
from pathlib import Path
import av
import numpy as np
import pandas as pd
import pyarrow as pa
import pyarrow.parquet as pq
import pytest
sys.path.insert(0, str(Path(__file__).parent))
from subset import EpisodeSubset, trim_video # noqa: E402
from worker import Executor, PreservingWriter # noqa: E402
from lerobot.annotations.steerable_pipeline.reader import iter_episodes
from lerobot.annotations.steerable_pipeline.staging import EpisodeStaging
from lerobot.datasets.feature_utils import create_empty_dataset_info
from lerobot.datasets.io_utils import write_info, write_tasks
from lerobot.datasets.lerobot_dataset import LeRobotDataset
from lerobot.utils.constants import DEFAULT_FEATURES
@pytest.fixture
def shared_source(tmp_path):
root = tmp_path / "source"
cameras = ["observation.images.top", "observation.images.wrist"]
features = {
**DEFAULT_FEATURES,
"action": {"dtype": "float32", "shape": [2], "names": ["joint", "joint2"]},
**{
c: {
"dtype": "video",
"shape": [32, 32, 3],
"names": ["height", "width", "channels"],
"info": {
"video.height": 32,
"video.width": 32,
"video.codec": "h264",
"video.fps": 10,
"video.pix_fmt": "yuv420p",
"video.is_depth_map": False,
"has_audio": False,
},
}
for c in cameras
},
}
info = create_empty_dataset_info("v3.0", 10, features, True)
info.total_episodes, info.total_frames, info.total_tasks = 3, 30, 3
info.splits = {"train": "0:3"}
write_info(info, root)
write_tasks(
pd.DataFrame(
{"task_index": [0, 1, 2]}, index=pd.Index(["task zero", "excluded task", "task two"], name="task")
),
root,
)
for ci, camera in enumerate(cameras):
video = root / info.video_path.format(video_key=camera, chunk_index=0, file_index=0)
video.parent.mkdir(parents=True, exist_ok=True)
with av.open(str(video), "w") as container:
stream = container.add_stream("libx264", rate=10, options={"crf": "0", "g": "30"})
stream.width = stream.height = 32
stream.pix_fmt = "yuv420p"
for idx in range(30):
frame = av.VideoFrame.from_ndarray(
np.full((32, 32, 3), 30 + (idx // 10) * 80 + ci * 10, dtype=np.uint8), format="rgb24"
)
for packet in stream.encode(frame):
container.mux(packet)
for packet in stream.encode():
container.mux(packet)
table = pa.table(
{
"episode_index": pa.array([i // 10 for i in range(30)], type=pa.int64()),
"frame_index": pa.array([i % 10 for i in range(30)], type=pa.int64()),
"timestamp": pa.array([(i % 10) / 10 for i in range(30)], type=pa.float32()),
"index": pa.array(range(30), type=pa.int64()),
"task_index": pa.array([i // 10 for i in range(30)], type=pa.int64()),
"action": pa.array([[float(i), float(i)] for i in range(30)], type=pa.list_(pa.float32(), 2)),
}
)
path = root / info.data_path.format(chunk_index=0, file_index=0)
path.parent.mkdir(parents=True, exist_ok=True)
pq.write_table(table, path, row_group_size=10)
rows = []
for ep in range(3):
row = {
"episode_index": ep,
"length": 10,
"tasks": [["task zero", "excluded task", "task two"][ep]],
"data/chunk_index": 0,
"data/file_index": 0,
"dataset_from_index": ep * 10,
"dataset_to_index": (ep + 1) * 10,
"meta/episodes/chunk_index": 0,
"meta/episodes/file_index": 0,
"stats/action/min": [float(ep * 10)],
"stats/action/max": [float(ep * 10 + 9)],
"stats/action/mean": [float(ep * 10 + 4.5)],
"stats/action/std": [float(np.std(np.arange(10)))],
"stats/action/count": [10],
}
for camera in cameras:
row.update(
{
f"videos/{camera}/chunk_index": 0,
f"videos/{camera}/file_index": 0,
f"videos/{camera}/from_timestamp": float(ep),
f"videos/{camera}/to_timestamp": float(ep + 1),
}
)
rows.append(row)
meta_path = root / "meta/episodes/chunk-000/file-000.parquet"
meta_path.parent.mkdir(parents=True, exist_ok=True)
pq.write_table(pa.Table.from_pylist(rows), meta_path)
return root, cameras, path
def test_successful_episodes_only_all_cameras_and_real_lerobot_load(shared_source, tmp_path):
root, cameras, data_path = shared_source
records = [r for r in iter_episodes(root) if r.episode_index != 1]
staging = tmp_path / "staging"
for r in records:
EpisodeStaging(staging, r.episode_index).write(
"plan",
[
{
"role": "assistant",
"content": f"annotated source {r.episode_index}",
"style": "subtask",
"timestamp": 0.0,
"camera": None,
"tool_calls": None,
}
],
)
PreservingWriter([0, 2], {"subtask"}).write_all(records, staging, root)
Executor._ensure_annotation_metadata_in_info(root)
output = tmp_path / "subset"
builder = EpisodeSubset(root, output, [0, 1, 2])
table = pq.read_table(data_path)
builder.add(0, table)
# Loading a checkpoint works before all completed episodes have been copied.
first = LeRobotDataset("test/subset", root=output, video_backend="pyav", token=False)
assert first.num_episodes == 1 and len(first) == 10
assert "annotated source 0" in str(first[0]["language_persistent"])
builder.add(2, table)
result = LeRobotDataset("test/subset", root=output, video_backend="pyav", token=False)
assert result.num_episodes == 2 and len(result) == 20
assert result.meta.info.splits == {"train": "0:2"}
assert result.meta.total_tasks == 2
assert list(result.meta.tasks.index) == ["task zero", "task two"]
assert builder.mapping == [
{"episode_index": 0, "source_episode_index": 0},
{"episode_index": 1, "source_episode_index": 2},
]
data = pq.read_table(output / "data").to_pylist()
assert [r["index"] for r in data] == list(range(20))
assert {r["episode_index"] for r in data} == {0, 1}
assert {r["task_index"] for r in data} == {0, 1}
assert [r["action"][0] for r in data] == list(range(10)) + list(range(20, 30))
for episode in [0, 1]:
for ci, camera in enumerate(cameras):
path = output / builder.info["video_path"].format(
video_key=camera, chunk_index=0, file_index=episode
)
with av.open(str(path)) as video:
frames = list(video.decode(video=0))
assert len(frames) == 10
assert float(frames[0].pts * frames[0].time_base) == 0
assert abs(frames[0].to_ndarray(format="rgb24").mean() - (30 + episode * 160 + ci * 10)) < 5
loaded = result[episode * 10]
assert f"annotated source {episode * 2}" in str(loaded["language_persistent"])
assert tuple(loaded[cameras[0]].shape) == (3, 32, 32)
stats = json.loads((output / "meta/stats.json").read_text())
assert stats["action"]["mean"] == [14.5]
assert stats["index"]["mean"] == [9.5]
assert stats["episode_index"]["max"] == [1]
assert len(list((output / "videos").rglob("*.mp4"))) == 4
with pytest.raises(ValueError, match="already added"):
builder.add(2, table)
def test_truncated_video_fails_instead_of_copying_wrong_frames(shared_source, tmp_path):
root, cameras, _ = shared_source
source = root / f"videos/{cameras[0]}/chunk-000/file-000.mp4"
with pytest.raises(ValueError, match="ended early"):
trim_video(source, tmp_path / "too-long.mp4", 2.0, 20, 10)
def test_subset_visualizer_link_uses_new_index_not_source_id():
from reporting import annotation_url
from studio import COPY
report = {
"manifest": {"mode": COPY, "copy_scope": "annotated_episodes", "output": "test/subset"},
"status": "completed",
"episodes": [{"episode_index": 900, "status": "completed"}],
"written_episode_ids": [],
}
assert annotation_url(report) is None
report["written_episode_ids"] = [0]
assert annotation_url(report).endswith("/test/subset/0?tab=annotations")
def test_worker_publishes_only_successful_subset_without_hub_duplication(
shared_source, tmp_path, monkeypatch
):
import hashlib
import shutil
from types import SimpleNamespace
from unittest.mock import Mock
import worker
from reporting import annotation_url
from studio import COPY
source, cameras, _ = shared_source
(source / "README.md").write_text("Original dataset attribution")
(source / "LICENSE").write_text("Original license")
published = tmp_path / "published"
published.mkdir()
api = Mock()
api.list_repo_refs.return_value = SimpleNamespace(tags=[])
commits = []
def write_file(path, data):
target = published / path
target.parent.mkdir(parents=True, exist_ok=True)
target.write_bytes(data if isinstance(data, bytes) else Path(data).read_bytes())
def commit(*args, operations, **kwargs):
for op in operations:
write_file(op.path_in_repo, op.path_or_fileobj)
if any(op.path_in_repo == "meta/info.json" for op in operations):
snapshot = LeRobotDataset("test/output", root=published, video_backend="pyav", token=False)
commits.append(snapshot.num_episodes)
return SimpleNamespace(oid="output-commit")
api.create_commit.side_effect = commit
api.upload_file.side_effect = lambda **kw: write_file(kw["path_in_repo"], kw["path_or_fileobj"])
api.get_paths_info.side_effect = lambda repo, files, **kw: [
SimpleNamespace(size=(source / f).stat().st_size) for f in files
]
monkeypatch.setattr(worker, "HfApi", lambda **kw: api)
monkeypatch.setenv("HF_TOKEN", "local-test-token")
monkeypatch.delenv("ANNOTATION_DEADLINE_EPOCH", raising=False)
rows = pq.read_table(source / "meta/episodes").to_pylist()
info = json.loads((source / "meta/info.json").read_text())
monkeypatch.setattr(worker, "read_selection", lambda *args, **kw: {"info": info, "episodes": rows})
def snapshot(*args, local_dir, **kw):
shutil.copytree(source / "meta", local_dir / "meta")
for filename in ["README.md", "LICENSE"]:
shutil.copyfile(source / filename, local_dir / filename)
def download(repo, filename, *, local_dir, **kw):
target = local_dir / filename
target.parent.mkdir(parents=True, exist_ok=True)
shutil.copyfile(source / filename, target)
return str(target)
monkeypatch.setattr(worker, "snapshot_download", snapshot)
monkeypatch.setattr(worker, "hf_hub_download", download)
monkeypatch.setattr(worker, "Provider", lambda *args, **kw: SimpleNamespace(concurrency=2))
def annotate(record, row, info, root, manifest, provider, staging, media=None):
if record.episode_index == 1:
raise ValueError("Simulated annotation validation failure")
EpisodeStaging(staging, record.episode_index).write(
"plan",
[
{
"role": "assistant",
"content": f"annotated source {record.episode_index}",
"style": "subtask",
"timestamp": 0.0,
"camera": None,
"tool_calls": None,
}
],
)
return {"episode_index": record.episode_index, "status": "completed", "rows": 1}
monkeypatch.setattr(worker, "annotate_episode", annotate)
manifest = {
"dataset": "test/source",
"dataset_revision": "pinned",
"output": "test/output",
"mode": COPY,
"visibility": "public",
"copy_scope": "annotated_episodes",
"source_license": "apache-2.0",
"start": 0,
"count": 3,
"episode_ids_sha256": hashlib.sha256(json.dumps([0, 1, 2]).encode()).hexdigest(),
"inference_budget": 1,
"model": "mock",
"camera": cameras[0],
"features": ["Subtasks"],
"costs": {"input_rate": 0.1, "output_rate": 0.2},
}
worker.run(manifest, tmp_path / "work")
api.duplicate_repo.assert_not_called()
api.create_repo.assert_called_once_with("test/output", repo_type="dataset", private=False, exist_ok=False)
assert commits == [1, 2]
final = LeRobotDataset("test/output", root=published, video_backend="pyav", token=False)
assert final.num_episodes == 2 and len(final) == 20
assert "annotated source 0" in str(final[0]["language_persistent"])
assert "annotated source 2" in str(final[10]["language_persistent"])
assert not list(published.rglob("*.jsonl"))
assert not (published / "annotation_studio/episodes").exists()
report = json.loads((published / "annotation_studio/run.json").read_text())
assert report["status"] == "partial"
assert report["written_episode_ids"] == [0, 1]
assert [m["source_episode_index"] for m in report["source_episode_map"]] == [0, 2]
assert annotation_url(report).endswith("/test/output/0?tab=annotations")
assert (published / "annotation_studio/source_README.md").read_text() == "Original dataset attribution"
assert (published / "LICENSE").read_text() == "Original license"
assert 'license: "apache-2.0"' in (published / "README.md").read_text()
assert not list((tmp_path / "work/subset/videos").rglob("*.mp4"))
assert not list((tmp_path / "work/subset/data").rglob("*.parquet"))
api.create_tag.assert_called_once_with(
"test/output", tag="v3.0", repo_type="dataset", revision="output-commit"
)