mirror of
https://github.com/huggingface/lerobot.git
synced 2026-07-24 10:16:09 +00:00
perf(datasets): bound memory of visual episode stats via batched streaming
Accumulate image/video per-episode stats into a RunningQuantileStats in frame batches instead of materialising every sampled frame at once. Peak memory is now bounded by one batch (frame_batch_size x C x H x W) regardless of episode length, preventing OOM on long, high-resolution episodes. count keeps the per-frame convention of get_feature_stats.
This commit is contained in:
@@ -53,11 +53,11 @@ from lerobot.utils.utils import flatten_dict
|
|||||||
|
|
||||||
from .aggregate import aggregate_datasets
|
from .aggregate import aggregate_datasets
|
||||||
from .compute_stats import (
|
from .compute_stats import (
|
||||||
|
RunningQuantileStats,
|
||||||
aggregate_stats,
|
aggregate_stats,
|
||||||
auto_downsample_height_width,
|
auto_downsample_height_width,
|
||||||
compute_episode_stats,
|
compute_episode_stats,
|
||||||
compute_relative_action_stats,
|
compute_relative_action_stats,
|
||||||
get_feature_stats,
|
|
||||||
sample_indices,
|
sample_indices,
|
||||||
)
|
)
|
||||||
from .dataset_metadata import LeRobotDatasetMetadata
|
from .dataset_metadata import LeRobotDatasetMetadata
|
||||||
@@ -1633,12 +1633,18 @@ def _compute_visual_episode_stats(
|
|||||||
dataset: LeRobotDataset,
|
dataset: LeRobotDataset,
|
||||||
ep_idx: int,
|
ep_idx: int,
|
||||||
visual_keys: list[str],
|
visual_keys: list[str],
|
||||||
|
frame_batch_size: int = 32,
|
||||||
) -> dict:
|
) -> dict:
|
||||||
"""Compute per-episode statistics for image/video features by sampling frames.
|
"""Compute per-episode statistics for image/video features by sampling frames.
|
||||||
|
|
||||||
Mirrors the image/video branch of :func:`compute_episode_stats`: per-channel stats
|
Mirrors the image/video branch of :func:`compute_episode_stats`: per-channel stats
|
||||||
are computed on downsampled sampled frames, then RGB stats are rescaled to [0, 1]
|
are computed on downsampled sampled frames, then RGB stats are rescaled to [0, 1]
|
||||||
(depth maps keep their native units).
|
(depth maps keep their native units).
|
||||||
|
|
||||||
|
Frames are decoded and accumulated into a :class:`RunningQuantileStats` in batches of
|
||||||
|
``frame_batch_size`` rather than materialising every sampled frame at once. Peak memory
|
||||||
|
is bounded by one batch (``frame_batch_size x C x H x W``) regardless of episode length,
|
||||||
|
which keeps long, high-resolution episodes from exhausting memory.
|
||||||
"""
|
"""
|
||||||
ep_length = dataset.meta.episodes[ep_idx]["length"]
|
ep_length = dataset.meta.episodes[ep_idx]["length"]
|
||||||
frame_offsets = sample_indices(ep_length)
|
frame_offsets = sample_indices(ep_length)
|
||||||
@@ -1646,15 +1652,27 @@ def _compute_visual_episode_stats(
|
|||||||
ep_stats = {}
|
ep_stats = {}
|
||||||
for key in visual_keys:
|
for key in visual_keys:
|
||||||
is_depth = key in dataset.meta.depth_keys
|
is_depth = key in dataset.meta.depth_keys
|
||||||
if dataset.meta.features[key]["dtype"] == "video":
|
is_video = dataset.meta.features[key]["dtype"] == "video"
|
||||||
frames = _load_episode_video_frames(dataset, key, ep_idx, frame_offsets, is_depth)
|
|
||||||
else:
|
|
||||||
frames = _load_episode_image_frames(dataset, key, ep_idx, frame_offsets, is_depth)
|
|
||||||
|
|
||||||
stats = get_feature_stats(frames, axis=(0, 2, 3), keepdims=True)
|
running = RunningQuantileStats()
|
||||||
|
for start in range(0, len(frame_offsets), frame_batch_size):
|
||||||
|
batch_offsets = frame_offsets[start : start + frame_batch_size]
|
||||||
|
if is_video:
|
||||||
|
frames = _load_episode_video_frames(dataset, key, ep_idx, batch_offsets, is_depth)
|
||||||
|
else:
|
||||||
|
frames = _load_episode_image_frames(dataset, key, ep_idx, batch_offsets, is_depth)
|
||||||
|
# (N, C, H, W) -> (N * H * W, C) so stats are accumulated per channel.
|
||||||
|
running.update(np.moveaxis(frames, 1, -1).reshape(-1, frames.shape[1]))
|
||||||
|
|
||||||
|
stats = running.get_statistics()
|
||||||
normalization_factor = 1.0 if is_depth else 255.0
|
normalization_factor = 1.0 if is_depth else 255.0
|
||||||
|
num_channels = stats["mean"].shape[0]
|
||||||
|
# ``count`` follows the per-frame convention of ``get_feature_stats`` (number of
|
||||||
|
# sampled frames), not the per-pixel count tracked internally by RunningQuantileStats.
|
||||||
ep_stats[key] = {
|
ep_stats[key] = {
|
||||||
k: v if k == "count" else np.squeeze(v / normalization_factor, axis=0)
|
k: np.array([len(frame_offsets)])
|
||||||
|
if k == "count"
|
||||||
|
else v.reshape(num_channels, 1, 1) / normalization_factor
|
||||||
for k, v in stats.items()
|
for k, v in stats.items()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user