mirror of
https://github.com/huggingface/lerobot.git
synced 2026-07-26 03:06:01 +00:00
fix(rollout): properly propagating video_files_size_in_mb to lerobot_dataset (#3470)
This commit is contained in:
@@ -630,6 +630,8 @@ class LeRobotDataset(torch.utils.data.Dataset):
|
|||||||
streaming_encoding: bool = False,
|
streaming_encoding: bool = False,
|
||||||
encoder_queue_maxsize: int = 30,
|
encoder_queue_maxsize: int = 30,
|
||||||
encoder_threads: int | None = None,
|
encoder_threads: int | None = None,
|
||||||
|
video_files_size_in_mb: int | None = None,
|
||||||
|
data_files_size_in_mb: int | None = None,
|
||||||
) -> "LeRobotDataset":
|
) -> "LeRobotDataset":
|
||||||
"""Create a new LeRobotDataset from scratch for recording data.
|
"""Create a new LeRobotDataset from scratch for recording data.
|
||||||
|
|
||||||
@@ -677,6 +679,8 @@ class LeRobotDataset(torch.utils.data.Dataset):
|
|||||||
root=root,
|
root=root,
|
||||||
use_videos=use_videos,
|
use_videos=use_videos,
|
||||||
metadata_buffer_size=metadata_buffer_size,
|
metadata_buffer_size=metadata_buffer_size,
|
||||||
|
video_files_size_in_mb=video_files_size_in_mb,
|
||||||
|
data_files_size_in_mb=data_files_size_in_mb,
|
||||||
)
|
)
|
||||||
obj.repo_id = obj.meta.repo_id
|
obj.repo_id = obj.meta.repo_id
|
||||||
obj._requested_root = obj.meta.root
|
obj._requested_root = obj.meta.root
|
||||||
|
|||||||
@@ -75,7 +75,7 @@ class SentryStrategyConfig(RolloutStrategyConfig):
|
|||||||
# Target video file size in MB for episode rotation. Episodes are
|
# Target video file size in MB for episode rotation. Episodes are
|
||||||
# saved once the estimated video duration would exceed this limit.
|
# saved once the estimated video duration would exceed this limit.
|
||||||
# Defaults to DEFAULT_VIDEO_FILE_SIZE_IN_MB when set to None.
|
# Defaults to DEFAULT_VIDEO_FILE_SIZE_IN_MB when set to None.
|
||||||
target_video_file_size_mb: float | None = None
|
target_video_file_size_mb: int | None = None
|
||||||
|
|
||||||
|
|
||||||
@RolloutStrategyConfig.register_subclass("highlight")
|
@RolloutStrategyConfig.register_subclass("highlight")
|
||||||
@@ -90,7 +90,7 @@ class HighlightStrategyConfig(RolloutStrategyConfig):
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
ring_buffer_seconds: float = 10.0
|
ring_buffer_seconds: float = 10.0
|
||||||
ring_buffer_max_memory_mb: float = 1024.0
|
ring_buffer_max_memory_mb: int = 1024
|
||||||
save_key: str = "s"
|
save_key: str = "s"
|
||||||
push_key: str = "h"
|
push_key: str = "h"
|
||||||
|
|
||||||
@@ -150,7 +150,7 @@ class DAggerStrategyConfig(RolloutStrategyConfig):
|
|||||||
upload_every_n_episodes: int = 5
|
upload_every_n_episodes: int = 5
|
||||||
# Target video file size in MB for episode rotation (record_autonomous
|
# Target video file size in MB for episode rotation (record_autonomous
|
||||||
# mode only). Defaults to DEFAULT_VIDEO_FILE_SIZE_IN_MB when None.
|
# mode only). Defaults to DEFAULT_VIDEO_FILE_SIZE_IN_MB when None.
|
||||||
target_video_file_size_mb: float | None = None
|
target_video_file_size_mb: int | None = None
|
||||||
input_device: str = "keyboard"
|
input_device: str = "keyboard"
|
||||||
keyboard: DAggerKeyboardConfig = field(default_factory=DAggerKeyboardConfig)
|
keyboard: DAggerKeyboardConfig = field(default_factory=DAggerKeyboardConfig)
|
||||||
pedal: DAggerPedalConfig = field(default_factory=DAggerPedalConfig)
|
pedal: DAggerPedalConfig = field(default_factory=DAggerPedalConfig)
|
||||||
|
|||||||
@@ -355,6 +355,7 @@ def build_rollout_context(
|
|||||||
"Use --dataset.repo_id=<user>/rollout_<name> for policy deployment datasets."
|
"Use --dataset.repo_id=<user>/rollout_<name> for policy deployment datasets."
|
||||||
)
|
)
|
||||||
cfg.dataset.stamp_repo_id()
|
cfg.dataset.stamp_repo_id()
|
||||||
|
target_video_mb = getattr(cfg.strategy, "target_video_file_size_mb", None)
|
||||||
dataset = LeRobotDataset.create(
|
dataset = LeRobotDataset.create(
|
||||||
cfg.dataset.repo_id,
|
cfg.dataset.repo_id,
|
||||||
cfg.dataset.fps,
|
cfg.dataset.fps,
|
||||||
@@ -370,6 +371,7 @@ def build_rollout_context(
|
|||||||
streaming_encoding=cfg.dataset.streaming_encoding,
|
streaming_encoding=cfg.dataset.streaming_encoding,
|
||||||
encoder_queue_maxsize=cfg.dataset.encoder_queue_maxsize,
|
encoder_queue_maxsize=cfg.dataset.encoder_queue_maxsize,
|
||||||
encoder_threads=cfg.dataset.encoder_threads,
|
encoder_threads=cfg.dataset.encoder_threads,
|
||||||
|
video_files_size_in_mb=target_video_mb,
|
||||||
)
|
)
|
||||||
|
|
||||||
if dataset is not None:
|
if dataset is not None:
|
||||||
|
|||||||
@@ -47,7 +47,7 @@ class RolloutRingBuffer:
|
|||||||
count.
|
count.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(self, max_seconds: float = 30.0, max_memory_mb: float = 2048.0, fps: float = 30.0) -> None:
|
def __init__(self, max_seconds: float = 30.0, max_memory_mb: int = 2048, fps: float = 30.0) -> None:
|
||||||
self._max_frames = int(max_seconds * fps)
|
self._max_frames = int(max_seconds * fps)
|
||||||
self._max_bytes = int(max_memory_mb * 1024 * 1024)
|
self._max_bytes = int(max_memory_mb * 1024 * 1024)
|
||||||
self._buffer: deque[dict] = deque(maxlen=self._max_frames)
|
self._buffer: deque[dict] = deque(maxlen=self._max_frames)
|
||||||
|
|||||||
@@ -416,6 +416,18 @@ def test_create_initial_counts_zero(tmp_path):
|
|||||||
assert dataset.num_frames == 0
|
assert dataset.num_frames == 0
|
||||||
|
|
||||||
|
|
||||||
|
def test_create_propagates_video_files_size_in_mb(tmp_path):
|
||||||
|
"""video_files_size_in_mb passed to create() is reflected in the dataset metadata."""
|
||||||
|
dataset = LeRobotDataset.create(
|
||||||
|
repo_id=DUMMY_REPO_ID,
|
||||||
|
fps=DEFAULT_FPS,
|
||||||
|
features=SIMPLE_FEATURES,
|
||||||
|
root=tmp_path / "ds",
|
||||||
|
video_files_size_in_mb=42.0,
|
||||||
|
)
|
||||||
|
assert dataset.meta.video_files_size_in_mb == 42.0
|
||||||
|
|
||||||
|
|
||||||
def test_add_frame_works_in_write_mode(tmp_path):
|
def test_add_frame_works_in_write_mode(tmp_path):
|
||||||
"""add_frame() succeeds on a dataset created via create()."""
|
"""add_frame() succeeds on a dataset created via create()."""
|
||||||
dataset = LeRobotDataset.create(
|
dataset = LeRobotDataset.create(
|
||||||
|
|||||||
Reference in New Issue
Block a user