From 265abe6c7948fe1096e5f47a9058be20d16b8bdc Mon Sep 17 00:00:00 2001 From: Steven Palma Date: Wed, 29 Jul 2026 17:07:34 +0200 Subject: [PATCH] chore(datasets): add typing to aggregate helpers (#4211) * chore(datasets): add typing to aggregate helpers Signed-off-by: nathon-lee * chore(dataset): add more typing aggregate * chore(test): remove panda test --------- Signed-off-by: nathon-lee Co-authored-by: nathon-lee --- src/lerobot/datasets/aggregate.py | 172 ++++++++++++++++++++---------- 1 file changed, 114 insertions(+), 58 deletions(-) diff --git a/src/lerobot/datasets/aggregate.py b/src/lerobot/datasets/aggregate.py index f5bf70eba..45d8d82e1 100644 --- a/src/lerobot/datasets/aggregate.py +++ b/src/lerobot/datasets/aggregate.py @@ -19,6 +19,7 @@ import copy import logging import shutil from pathlib import Path +from typing import Any, NotRequired, TypedDict import datasets import pandas as pd @@ -49,8 +50,32 @@ from .utils import ( ) from .video_utils import concatenate_video_files, get_video_duration_in_s +logger = logging.getLogger(__name__) -def merge_video_feature_info_for_aggregate(all_metadata: list[LeRobotDatasetMetadata]) -> dict[str, dict]: +type FeatureDict = dict[str, dict[str, Any]] +type ChunkFile = tuple[int, int] + + +class IndexState(TypedDict): + chunk: int + file: int + src_to_dst: NotRequired[dict[ChunkFile, ChunkFile]] + + +class VideoIndex(TypedDict): + chunk: int + file: int + latest_duration: float + episode_duration: float + src_to_offset: NotRequired[dict[ChunkFile, float]] + src_to_dst: NotRequired[dict[ChunkFile, ChunkFile]] + dst_file_durations: NotRequired[dict[ChunkFile, float]] + + +type VideoIndexState = dict[str, VideoIndex] + + +def merge_video_feature_info_for_aggregate(all_metadata: list[LeRobotDatasetMetadata]) -> FeatureDict: """Create a merged video feature info dictionary for aggregation. The video encoder info is merged field-by-field: each key is kept only when every source agrees; otherwise that key is set to ``null`` (or ``{}`` for ``video.extra_options``) and a warning is logged. Args: @@ -59,14 +84,14 @@ def merge_video_feature_info_for_aggregate(all_metadata: list[LeRobotDatasetMeta Returns: dict: A dictionary of merged video feature info. """ - merged_info = copy.deepcopy(all_metadata[0].features) + merged_info: FeatureDict = copy.deepcopy(all_metadata[0].features) video_keys = [k for k in merged_info if merged_info[k].get("dtype") == "video"] for vk in video_keys: video_infos = [m.features.get(vk, {}).get("info") or {} for m in all_metadata] base_video_info = video_infos[0] - merged_encoder_info: dict = {} + merged_encoder_info: dict[str, Any] = {} fallback_keys: list[str] = [] for info_key in VIDEO_ENCODER_INFO_KEYS: values = [info.get(info_key, None) for info in video_infos] @@ -80,7 +105,7 @@ def merge_video_feature_info_for_aggregate(all_metadata: list[LeRobotDatasetMeta merged_encoder_info[info_key] = {} if info_key == "video.extra_options" else None if fallback_keys: - logging.warning( + logger.warning( f"Merging heterogeneous or incomplete video encoder metadata for feature {vk}. " f"Setting these keys to null: {fallback_keys}.", ) @@ -92,7 +117,7 @@ def merge_video_feature_info_for_aggregate(all_metadata: list[LeRobotDatasetMeta return merged_info -def validate_all_metadata(all_metadata: list[LeRobotDatasetMetadata]): +def validate_all_metadata(all_metadata: list[LeRobotDatasetMetadata]) -> tuple[int, str | None, FeatureDict]: """Validates that all dataset metadata have consistent properties. Ensures all datasets have the same fps, robot_type, and features to guarantee @@ -129,7 +154,9 @@ def validate_all_metadata(all_metadata: list[LeRobotDatasetMetadata]): return fps, robot_type, features -def update_data_df(df, src_meta, dst_meta): +def update_data_df( + df: pd.DataFrame, src_meta: LeRobotDatasetMetadata, dst_meta: LeRobotDatasetMetadata +) -> pd.DataFrame: """Updates a data DataFrame with new indices and task mappings for aggregation. Adjusts episode indices, frame indices, and task indices to account for @@ -154,12 +181,12 @@ def update_data_df(df, src_meta, dst_meta): def update_meta_data( - df, - dst_meta, - meta_idx, - data_idx, - videos_idx, -): + df: pd.DataFrame, + dst_meta: LeRobotDatasetMetadata, + meta_idx: IndexState, + data_idx: IndexState, + videos_idx: VideoIndexState, +) -> pd.DataFrame: """Updates metadata DataFrame with new chunk, file, and timestamp indices. Adjusts all indices and timestamps to account for previously aggregated @@ -289,7 +316,7 @@ def aggregate_datasets( chunk_size: int | None = None, concatenate_videos: bool = True, concatenate_data: bool = True, -): +) -> None: """Aggregates multiple LeRobot datasets into a single unified dataset. This is the main function that orchestrates the aggregation process by: @@ -309,7 +336,7 @@ def aggregate_datasets( concatenate_videos: When False, keep one mp4 per source file instead of packing into shards. concatenate_data: When False, keep one parquet per source file instead of packing into shards. """ - logging.info("Start aggregate_datasets") + logger.info("Start aggregate_datasets") if data_files_size_in_mb is None: data_files_size_in_mb = DEFAULT_DATA_FILE_SIZE_IN_MB @@ -341,15 +368,15 @@ def aggregate_datasets( video_files_size_in_mb=video_files_size_in_mb, ) - logging.info("Find all tasks") + logger.info("Find all tasks") unique_tasks = pd.concat([m.tasks for m in all_metadata]).index.unique() dst_meta.tasks = pd.DataFrame( {"task_index": range(len(unique_tasks))}, index=pd.Index(unique_tasks, name="task") ) - meta_idx = {"chunk": 0, "file": 0} - data_idx = {"chunk": 0, "file": 0} - videos_idx = { + meta_idx: IndexState = {"chunk": 0, "file": 0} + data_idx: IndexState = {"chunk": 0, "file": 0} + videos_idx: VideoIndexState = { key: {"chunk": 0, "file": 0, "latest_duration": 0, "episode_duration": 0} for key in video_keys } @@ -373,12 +400,17 @@ def aggregate_datasets( dst_meta.info.total_frames += src_meta.total_frames finalize_aggregation(dst_meta, all_metadata) - logging.info("Aggregation complete.") + logger.info("Aggregation complete.") def aggregate_videos( - src_meta, dst_meta, videos_idx, video_files_size_in_mb, chunk_size, concatenate_videos=True -): + src_meta: LeRobotDatasetMetadata, + dst_meta: LeRobotDatasetMetadata, + videos_idx: VideoIndexState, + video_files_size_in_mb: float, + chunk_size: int, + concatenate_videos: bool = True, +) -> VideoIndexState: """Aggregates video chunks from a source dataset into the destination dataset. Handles video file concatenation and rotation based on file size limits. @@ -406,15 +438,16 @@ def aggregate_videos( videos_idx[key]["dst_file_durations"] = {} for key, video_idx in videos_idx.items(): - unique_chunk_file_pairs = { - (chunk, file) - for chunk, file in zip( - src_meta.episodes[f"videos/{key}/chunk_index"], - src_meta.episodes[f"videos/{key}/file_index"], - strict=False, - ) - } - unique_chunk_file_pairs = sorted(unique_chunk_file_pairs) + unique_chunk_file_pairs: list[ChunkFile] = sorted( + { + (chunk, file) + for chunk, file in zip( + src_meta.episodes[f"videos/{key}/chunk_index"], + src_meta.episodes[f"videos/{key}/file_index"], + strict=False, + ) + } + ) chunk_idx = video_idx["chunk"] file_idx = video_idx["file"] @@ -489,7 +522,14 @@ def aggregate_videos( return videos_idx -def aggregate_data(src_meta, dst_meta, data_idx, data_files_size_in_mb, chunk_size, concatenate_data=True): +def aggregate_data( + src_meta: LeRobotDatasetMetadata, + dst_meta: LeRobotDatasetMetadata, + data_idx: IndexState, + data_files_size_in_mb: float, + chunk_size: int, + concatenate_data: bool = True, +) -> IndexState: """Aggregates data chunks from a source dataset into the destination dataset. Reads source data files, updates indices to match the aggregated dataset, @@ -510,14 +550,16 @@ def aggregate_data(src_meta, dst_meta, data_idx, data_files_size_in_mb, chunk_si Returns: dict: Updated data_idx with current chunk and file indices. """ - unique_chunk_file_ids = { - (c, f) - for c, f in zip( - src_meta.episodes["data/chunk_index"], src_meta.episodes["data/file_index"], strict=False - ) - } - - unique_chunk_file_ids = sorted(unique_chunk_file_ids) + unique_chunk_file_ids: list[ChunkFile] = sorted( + { + (c, f) + for c, f in zip( + src_meta.episodes["data/chunk_index"], + src_meta.episodes["data/file_index"], + strict=False, + ) + } + ) contains_images = len(dst_meta.image_keys) > 0 # retrieve features schema for proper image typing in parquet @@ -525,7 +567,7 @@ def aggregate_data(src_meta, dst_meta, data_idx, data_files_size_in_mb, chunk_si # Track source to destination file mapping for metadata update # This is critical for handling datasets that are already results of a merge - src_to_dst: dict[tuple[int, int], tuple[int, int]] = {} + src_to_dst: dict[ChunkFile, ChunkFile] = {} for src_chunk_idx, src_file_idx in unique_chunk_file_ids: src_path = src_meta.root / DEFAULT_DATA_PATH.format( @@ -564,7 +606,13 @@ def aggregate_data(src_meta, dst_meta, data_idx, data_files_size_in_mb, chunk_si return data_idx -def aggregate_metadata(src_meta, dst_meta, meta_idx, data_idx, videos_idx): +def aggregate_metadata( + src_meta: LeRobotDatasetMetadata, + dst_meta: LeRobotDatasetMetadata, + meta_idx: IndexState, + data_idx: IndexState, + videos_idx: VideoIndexState, +) -> IndexState: """Aggregates metadata from a source dataset into the destination dataset. Reads source metadata files, updates all indices and timestamps, @@ -580,16 +628,16 @@ def aggregate_metadata(src_meta, dst_meta, meta_idx, data_idx, videos_idx): Returns: dict: Updated meta_idx with current chunk and file indices. """ - chunk_file_ids = { - (c, f) - for c, f in zip( - src_meta.episodes["meta/episodes/chunk_index"], - src_meta.episodes["meta/episodes/file_index"], - strict=False, - ) - } - - chunk_file_ids = sorted(chunk_file_ids) + chunk_file_ids: list[ChunkFile] = sorted( + { + (c, f) + for c, f in zip( + src_meta.episodes["meta/episodes/chunk_index"], + src_meta.episodes["meta/episodes/file_index"], + strict=False, + ) + } + ) for chunk_idx, file_idx in chunk_file_ids: src_path = src_meta.root / DEFAULT_EPISODES_PATH.format(chunk_index=chunk_idx, file_index=file_idx) df = pd.read_parquet(src_path) @@ -622,16 +670,16 @@ def aggregate_metadata(src_meta, dst_meta, meta_idx, data_idx, videos_idx): def append_or_create_parquet_file( df: pd.DataFrame, src_path: Path, - idx: dict[str, int], + idx: IndexState, max_mb: float, chunk_size: int, default_path: str, contains_images: bool = False, - aggr_root: Path = None, + aggr_root: Path | None = None, hf_features: datasets.Features | None = None, concatenate: bool = True, one_row_group_per_episode: bool = False, -) -> tuple[dict[str, int], tuple[int, int]]: +) -> tuple[IndexState, ChunkFile]: """Appends data to an existing parquet file or creates a new one based on size constraints. Manages file rotation when size limits are exceeded to prevent individual files @@ -654,7 +702,13 @@ def append_or_create_parquet_file( Returns: tuple: (updated_idx, (dst_chunk, dst_file)) where updated_idx is the index dict and (dst_chunk, dst_file) is the actual destination file the data was written to. + + Raises: + ValueError: If aggr_root is not provided. """ + if aggr_root is None: + raise ValueError("aggr_root must be provided.") + dst_chunk, dst_file = idx["chunk"], idx["file"] dst_path = aggr_root / default_path.format(chunk_index=dst_chunk, file_index=dst_file) @@ -698,7 +752,9 @@ def append_or_create_parquet_file( return idx, (dst_chunk, dst_file) -def finalize_aggregation(aggr_meta, all_metadata): +def finalize_aggregation( + aggr_meta: LeRobotDatasetMetadata, all_metadata: list[LeRobotDatasetMetadata] +) -> None: """Finalizes the dataset aggregation by writing summary files and statistics. Writes the tasks file, info file with total counts and splits, and @@ -708,16 +764,16 @@ def finalize_aggregation(aggr_meta, all_metadata): aggr_meta: Aggregated dataset metadata. all_metadata: List of all source dataset metadata objects. """ - logging.info("write tasks") + logger.info("write tasks") write_tasks(aggr_meta.tasks, aggr_meta.root) - logging.info("write info") + logger.info("write info") aggr_meta.info.total_tasks = len(aggr_meta.tasks) aggr_meta.info.total_episodes = sum(m.total_episodes for m in all_metadata) aggr_meta.info.total_frames = sum(m.total_frames for m in all_metadata) aggr_meta.info.splits = {"train": f"0:{sum(m.total_episodes for m in all_metadata)}"} write_info(aggr_meta.info, aggr_meta.root) - logging.info("write stats") + logger.info("write stats") aggr_meta.stats = aggregate_stats([m.stats for m in all_metadata]) write_stats(aggr_meta.stats, aggr_meta.root)