mirror of
https://github.com/huggingface/lerobot.git
synced 2026-08-08 17:39:44 +00:00
fix(dataset): use conservative bounds for quantile aggregation instead of incorrect weighted mean (#3804)
* fix(stats): use conservative bounds for quantile aggregation instead of incorrect weighted mean * docs: add --overwrite/--skip-images/--root options to augment_dataset_quantile_stats usage * fix(dataset): clarify quantile aggregation semantics * fix(augment): handle quantile stats edge cases
This commit is contained in:
@@ -242,6 +242,17 @@ python src/lerobot/scripts/augment_dataset_quantile_stats.py \
|
|||||||
--repo-id=your_dataset
|
--repo-id=your_dataset
|
||||||
```
|
```
|
||||||
|
|
||||||
|
Recording, resuming, and merging aggregate quantiles from per-episode summaries, so `meta/stats.json` ends up holding a conservative envelope (`min` for `q <= 50`, `max` for `q > 50`) rather than whole-dataset quantiles. To estimate the latter, scan every episode with a running histogram:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
python src/lerobot/scripts/augment_dataset_quantile_stats.py \
|
||||||
|
--repo-id=your_dataset \
|
||||||
|
--overwrite \
|
||||||
|
--skip-images
|
||||||
|
```
|
||||||
|
|
||||||
|
`--skip-images` keeps the existing image statistics and avoids video decoding when only `STATE`/`ACTION` need recomputing, and `--root` reads a local dataset instead of the Hub. These values are histogram estimates, subject to discretization and rebinning error, so they can differ from the conservative ones — which changes MolmoAct2's normalized targets and therefore its loss scale. Statistics already saved inside an existing checkpoint are not affected.
|
||||||
|
|
||||||
Alternatively, train MolmoAct2 with mean/std normalization:
|
Alternatively, train MolmoAct2 with mean/std normalization:
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
|
|||||||
@@ -127,6 +127,17 @@ lerobot-edit-dataset \
|
|||||||
|
|
||||||
Or keep the dataset as-is and pass `--policy.normalization_mapping='{"ACTION": "MEAN_STD", "STATE": "MEAN_STD", "VISUAL": "IDENTITY"}'`.
|
Or keep the dataset as-is and pass `--policy.normalization_mapping='{"ACTION": "MEAN_STD", "STATE": "MEAN_STD", "VISUAL": "IDENTITY"}'`.
|
||||||
|
|
||||||
|
Recording, resuming, and merging aggregate quantiles from per-episode summaries, so `meta/stats.json` ends up holding a conservative envelope (`min` for `q <= 50`, `max` for `q > 50`) rather than whole-dataset quantiles. To estimate the latter, scan every episode with a running histogram:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
python src/lerobot/scripts/augment_dataset_quantile_stats.py \
|
||||||
|
--repo-id=your_dataset \
|
||||||
|
--overwrite \
|
||||||
|
--skip-images
|
||||||
|
```
|
||||||
|
|
||||||
|
`--skip-images` keeps the existing image statistics and avoids video decoding when only `STATE`/`ACTION` need recomputing, and `--root` reads a local dataset instead of the Hub. These values are histogram estimates, subject to discretization and rebinning error, so they can differ from the conservative ones — which changes π₀.₅'s normalized targets and therefore its loss scale. Statistics already saved inside an existing checkpoint are not affected.
|
||||||
|
|
||||||
### Training Command Example
|
### Training Command Example
|
||||||
|
|
||||||
The same finetune with the VLM frozen: less memory, at some cost in success rate. Swap `--dataset.repo_id` for your own dataset.
|
The same finetune with the VLM frozen: less memory, at some cost in success rate. Swap `--dataset.repo_id` for your own dataset.
|
||||||
|
|||||||
@@ -613,8 +613,15 @@ def aggregate_feature_stats(stats_ft_list: list[dict[str, dict]]) -> dict[str, d
|
|||||||
for q_key in quantile_keys:
|
for q_key in quantile_keys:
|
||||||
if all(q_key in s for s in stats_ft_list):
|
if all(q_key in s for s in stats_ft_list):
|
||||||
quantile_values = np.stack([s[q_key] for s in stats_ft_list])
|
quantile_values = np.stack([s[q_key] for s in stats_ft_list])
|
||||||
weighted_quantiles = quantile_values * counts
|
# Exact global quantiles cannot be recovered from quantile summaries.
|
||||||
aggregated[q_key] = weighted_quantiles.sum(axis=0) / total_count
|
# Keep a conservative envelope of the available estimates: min
|
||||||
|
# for lower quantiles and max for upper quantiles. The resulting
|
||||||
|
# values are bounds across the inputs, not global quantile estimates.
|
||||||
|
q_percent = int(q_key[1:])
|
||||||
|
if q_percent <= 50:
|
||||||
|
aggregated[q_key] = np.min(quantile_values, axis=0)
|
||||||
|
else:
|
||||||
|
aggregated[q_key] = np.max(quantile_values, axis=0)
|
||||||
|
|
||||||
return aggregated
|
return aggregated
|
||||||
|
|
||||||
|
|||||||
@@ -25,6 +25,11 @@ quantile statistics (q01, q10, q50, q90, q99) in their metadata. This script:
|
|||||||
3. If missing, computes quantile statistics for all features
|
3. If missing, computes quantile statistics for all features
|
||||||
4. Updates the dataset metadata with the new quantile statistics
|
4. Updates the dataset metadata with the new quantile statistics
|
||||||
|
|
||||||
|
Statistics are accumulated into a single running histogram per feature across
|
||||||
|
all episodes rather than aggregating per-episode quantile summaries. The
|
||||||
|
resulting quantiles are histogram approximations, subject to discretization and
|
||||||
|
range-rebinning error; image/video frames are sampled by default.
|
||||||
|
|
||||||
Usage:
|
Usage:
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
@@ -34,9 +39,7 @@ python src/lerobot/scripts/augment_dataset_quantile_stats.py \
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
import argparse
|
import argparse
|
||||||
import concurrent.futures
|
|
||||||
import logging
|
import logging
|
||||||
import os
|
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
@@ -49,11 +52,10 @@ from lerobot.datasets import (
|
|||||||
CODEBASE_VERSION,
|
CODEBASE_VERSION,
|
||||||
DEFAULT_QUANTILES,
|
DEFAULT_QUANTILES,
|
||||||
LeRobotDataset,
|
LeRobotDataset,
|
||||||
aggregate_stats,
|
|
||||||
get_feature_stats,
|
get_feature_stats,
|
||||||
write_stats,
|
write_stats,
|
||||||
)
|
)
|
||||||
from lerobot.datasets.compute_stats import sample_indices
|
from lerobot.datasets.compute_stats import RunningQuantileStats, sample_indices
|
||||||
from lerobot.utils.utils import init_logging
|
from lerobot.utils.utils import init_logging
|
||||||
|
|
||||||
|
|
||||||
@@ -79,20 +81,25 @@ def has_quantile_stats(stats: dict[str, dict] | None, quantile_list_keys: list[s
|
|||||||
return False
|
return False
|
||||||
|
|
||||||
|
|
||||||
def process_single_episode(dataset: LeRobotDataset, episode_idx: int, use_sampling: bool = True) -> dict:
|
def collect_episode_arrays(
|
||||||
"""Process a single episode and return its statistics.
|
dataset: LeRobotDataset,
|
||||||
|
episode_idx: int,
|
||||||
|
use_sampling: bool = True,
|
||||||
|
skip_images: bool = False,
|
||||||
|
) -> dict[str, tuple[np.ndarray, int]]:
|
||||||
|
"""Collect one episode's frames per feature, flattened to (num_samples, dim).
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
dataset: The LeRobot dataset
|
dataset: The LeRobot dataset
|
||||||
episode_idx: Index of the episode to process
|
episode_idx: Index of the episode to read
|
||||||
use_sampling: If True, sub-sample image/video frames per episode to bound
|
use_sampling: If True, sub-sample image/video frames to bound memory.
|
||||||
memory. If False, use every frame (exact, higher memory).
|
If False, use every frame (higher memory).
|
||||||
|
skip_images: If True, skip image/video features entirely.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
Dictionary containing episode statistics
|
Mapping of feature name to that episode's values and the number of frames
|
||||||
|
they came from (which differs from the row count for image features).
|
||||||
"""
|
"""
|
||||||
logging.info(f"Computing stats for episode {episode_idx}")
|
|
||||||
|
|
||||||
start_idx = dataset.meta.episodes[episode_idx]["dataset_from_index"]
|
start_idx = dataset.meta.episodes[episode_idx]["dataset_from_index"]
|
||||||
end_idx = dataset.meta.episodes[episode_idx]["dataset_to_index"]
|
end_idx = dataset.meta.episodes[episode_idx]["dataset_to_index"]
|
||||||
|
|
||||||
@@ -102,7 +109,9 @@ def process_single_episode(dataset: LeRobotDataset, episode_idx: int, use_sampli
|
|||||||
# numeric columns are cheap, so read them in full (exact).
|
# numeric columns are cheap, so read them in full (exact).
|
||||||
image_keys = [k for k in dataset.features if dataset.features[k]["dtype"] in ("image", "video")]
|
image_keys = [k for k in dataset.features if dataset.features[k]["dtype"] in ("image", "video")]
|
||||||
numeric_keys = [
|
numeric_keys = [
|
||||||
k for k in dataset.features if dataset.features[k]["dtype"] not in ("image", "video", "string")
|
k
|
||||||
|
for k in dataset.features
|
||||||
|
if dataset.features[k]["dtype"] not in ("image", "video", "string", "language")
|
||||||
]
|
]
|
||||||
|
|
||||||
collected_data: dict[str, list] = {}
|
collected_data: dict[str, list] = {}
|
||||||
@@ -114,7 +123,7 @@ def process_single_episode(dataset: LeRobotDataset, episode_idx: int, use_sampli
|
|||||||
collected_data[key] = [torch.as_tensor(v) for v in numeric_cols[key]]
|
collected_data[key] = [torch.as_tensor(v) for v in numeric_cols[key]]
|
||||||
|
|
||||||
# Image/video features: decode only a sampled subset of frames.
|
# Image/video features: decode only a sampled subset of frames.
|
||||||
if image_keys:
|
if image_keys and not skip_images:
|
||||||
sampled_offsets = sample_indices(episode_len) if use_sampling else list(range(episode_len))
|
sampled_offsets = sample_indices(episode_len) if use_sampling else list(range(episode_len))
|
||||||
for offset in sampled_offsets:
|
for offset in sampled_offsets:
|
||||||
item = dataset[start_idx + offset]
|
item = dataset[start_idx + offset]
|
||||||
@@ -122,87 +131,82 @@ def process_single_episode(dataset: LeRobotDataset, episode_idx: int, use_sampli
|
|||||||
if key in item:
|
if key in item:
|
||||||
collected_data.setdefault(key, []).append(item[key])
|
collected_data.setdefault(key, []).append(item[key])
|
||||||
|
|
||||||
ep_stats = {}
|
episode_arrays: dict[str, tuple[np.ndarray, int]] = {}
|
||||||
for key, data_list in collected_data.items():
|
for key, data_list in collected_data.items():
|
||||||
if dataset.features[key]["dtype"] == "string":
|
|
||||||
continue
|
|
||||||
|
|
||||||
data = torch.stack(data_list).cpu().numpy()
|
data = torch.stack(data_list).cpu().numpy()
|
||||||
if dataset.features[key]["dtype"] in ["image", "video"]:
|
if dataset.features[key]["dtype"] in ["image", "video"]:
|
||||||
if data.dtype == np.uint8:
|
if data.dtype == np.uint8:
|
||||||
data = data.astype(np.float32) / 255.0
|
data = data.astype(np.float32) / 255.0
|
||||||
|
# (N, C, H, W) -> (N * H * W, C) so quantiles are computed per channel.
|
||||||
axes_to_reduce = (0, 2, 3)
|
channels = data.shape[1]
|
||||||
keepdims = True
|
values = data.transpose(0, 2, 3, 1).reshape(-1, channels)
|
||||||
else:
|
else:
|
||||||
axes_to_reduce = 0
|
values = data.reshape(-1, data.shape[-1]) if data.ndim > 1 else data.reshape(-1, 1)
|
||||||
keepdims = data.ndim == 1
|
episode_arrays[key] = (values, len(data_list))
|
||||||
|
|
||||||
ep_stats[key] = get_feature_stats(
|
return episode_arrays
|
||||||
data, axis=axes_to_reduce, keepdims=keepdims, quantile_list=DEFAULT_QUANTILES
|
|
||||||
)
|
|
||||||
|
|
||||||
if dataset.features[key]["dtype"] in ["image", "video"]:
|
|
||||||
ep_stats[key] = {
|
|
||||||
k: v if k == "count" else np.squeeze(v, axis=0) for k, v in ep_stats[key].items()
|
|
||||||
}
|
|
||||||
|
|
||||||
return ep_stats
|
|
||||||
|
|
||||||
|
|
||||||
def compute_quantile_stats_for_dataset(dataset: LeRobotDataset, use_sampling: bool = True) -> dict[str, dict]:
|
def compute_quantile_stats_for_dataset(
|
||||||
"""Compute quantile statistics for all episodes in the dataset.
|
dataset: LeRobotDataset,
|
||||||
|
use_sampling: bool = True,
|
||||||
|
skip_images: bool = False,
|
||||||
|
) -> dict[str, dict]:
|
||||||
|
"""Compute whole-dataset statistics with one running histogram per feature.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
dataset: The LeRobot dataset to compute statistics for
|
dataset: The LeRobot dataset to compute statistics for
|
||||||
use_sampling: If True, sub-sample image/video frames per episode to bound
|
use_sampling: If True, sub-sample image/video frames per episode to bound
|
||||||
memory. If False, use every frame (exact, higher memory).
|
memory. If False, use every frame (higher memory).
|
||||||
|
skip_images: If True, skip image/video features and leave their stats untouched.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
Dictionary containing aggregated statistics with quantiles
|
Dictionary containing statistics with histogram-based global quantile estimates
|
||||||
|
|
||||||
Note:
|
Note:
|
||||||
Video decoding operations are not thread-safe, so we process episodes sequentially
|
Episodes are accumulated sequentially because the running accumulators are
|
||||||
when video keys are present. For datasets without videos, we use parallel processing
|
shared across all of them.
|
||||||
with ThreadPoolExecutor for better performance.
|
|
||||||
"""
|
"""
|
||||||
logging.info(f"Computing quantile statistics for dataset with {dataset.num_episodes} episodes")
|
logging.info(f"Computing quantile statistics for dataset with {dataset.num_episodes} episodes")
|
||||||
|
|
||||||
episode_stats_list = []
|
running_stats: dict[str, RunningQuantileStats] = {}
|
||||||
has_videos = len(dataset.meta.video_keys) > 0
|
frame_counts: dict[str, int] = {}
|
||||||
|
row_counts: dict[str, int] = {}
|
||||||
|
# Kept only while a feature has a single row, so it can still be finalized.
|
||||||
|
single_row_arrays: dict[str, np.ndarray] = {}
|
||||||
|
|
||||||
if has_videos:
|
for episode_idx in tqdm(range(dataset.num_episodes), desc="Processing episodes"):
|
||||||
logging.info("Dataset contains video keys - using sequential processing for thread safety")
|
episode_arrays = collect_episode_arrays(
|
||||||
for episode_idx in tqdm(range(dataset.num_episodes), desc="Processing episodes"):
|
dataset, episode_idx, use_sampling=use_sampling, skip_images=skip_images
|
||||||
ep_stats = process_single_episode(dataset, episode_idx, use_sampling)
|
)
|
||||||
episode_stats_list.append(ep_stats)
|
for key, (array, num_frames) in episode_arrays.items():
|
||||||
else:
|
running_stats.setdefault(key, RunningQuantileStats()).update(array)
|
||||||
logging.info("Dataset has no video keys - using parallel processing for better performance")
|
frame_counts[key] = frame_counts.get(key, 0) + num_frames
|
||||||
max_workers = min(dataset.num_episodes, int(os.environ.get("LEROBOT_STATS_MAX_WORKERS", 16)))
|
row_counts[key] = row_counts.get(key, 0) + len(array)
|
||||||
|
if row_counts[key] < 2:
|
||||||
|
single_row_arrays[key] = array
|
||||||
|
else:
|
||||||
|
single_row_arrays.pop(key, None)
|
||||||
|
|
||||||
with concurrent.futures.ThreadPoolExecutor(max_workers=max_workers) as executor:
|
if not running_stats:
|
||||||
future_to_episode = {
|
|
||||||
executor.submit(process_single_episode, dataset, episode_idx, use_sampling): episode_idx
|
|
||||||
for episode_idx in range(dataset.num_episodes)
|
|
||||||
}
|
|
||||||
|
|
||||||
episode_results = {}
|
|
||||||
with tqdm(total=dataset.num_episodes, desc="Processing episodes") as pbar:
|
|
||||||
for future in concurrent.futures.as_completed(future_to_episode):
|
|
||||||
episode_idx = future_to_episode[future]
|
|
||||||
ep_stats = future.result()
|
|
||||||
episode_results[episode_idx] = ep_stats
|
|
||||||
pbar.update(1)
|
|
||||||
|
|
||||||
for episode_idx in range(dataset.num_episodes):
|
|
||||||
if episode_idx in episode_results:
|
|
||||||
episode_stats_list.append(episode_results[episode_idx])
|
|
||||||
|
|
||||||
if not episode_stats_list:
|
|
||||||
raise ValueError("No episode data found for computing statistics")
|
raise ValueError("No episode data found for computing statistics")
|
||||||
|
|
||||||
logging.info(f"Aggregating statistics from {len(episode_stats_list)} episodes")
|
aggregated_stats: dict[str, dict] = {}
|
||||||
return aggregate_stats(episode_stats_list)
|
for key, accumulator in running_stats.items():
|
||||||
|
if row_counts[key] < 2:
|
||||||
|
# Histograms need at least two samples; mirror get_feature_stats' basic-stats path.
|
||||||
|
stats = get_feature_stats(single_row_arrays[key], axis=0, keepdims=False)
|
||||||
|
else:
|
||||||
|
stats = accumulator.get_statistics()
|
||||||
|
if dataset.features[key]["dtype"] in ["image", "video"]:
|
||||||
|
# Image stats are stored as (C, 1, 1) to broadcast over height and width.
|
||||||
|
stats = {k: v if k == "count" else v[:, np.newaxis, np.newaxis] for k, v in stats.items()}
|
||||||
|
# `get_feature_stats` counts frames, not the per-channel rows the accumulator sees.
|
||||||
|
stats["count"] = np.array([frame_counts[key]])
|
||||||
|
aggregated_stats[key] = stats
|
||||||
|
|
||||||
|
logging.info(f"Computed global histogram statistics for {len(aggregated_stats)} features")
|
||||||
|
return aggregated_stats
|
||||||
|
|
||||||
|
|
||||||
def augment_dataset_with_quantile_stats(
|
def augment_dataset_with_quantile_stats(
|
||||||
@@ -210,6 +214,7 @@ def augment_dataset_with_quantile_stats(
|
|||||||
root: str | Path | None = None,
|
root: str | Path | None = None,
|
||||||
overwrite: bool = False,
|
overwrite: bool = False,
|
||||||
use_sampling: bool = True,
|
use_sampling: bool = True,
|
||||||
|
skip_images: bool = False,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""Augment a dataset with quantile statistics if they are missing.
|
"""Augment a dataset with quantile statistics if they are missing.
|
||||||
|
|
||||||
@@ -218,7 +223,8 @@ def augment_dataset_with_quantile_stats(
|
|||||||
root: Local root directory for the dataset
|
root: Local root directory for the dataset
|
||||||
overwrite: Overwrite existing quantile statistics if they already exist
|
overwrite: Overwrite existing quantile statistics if they already exist
|
||||||
use_sampling: If True, sub-sample image/video frames per episode to bound
|
use_sampling: If True, sub-sample image/video frames per episode to bound
|
||||||
memory. If False, use every frame (exact, higher memory).
|
memory. If False, use every frame (higher memory).
|
||||||
|
skip_images: If True, skip image/video features and keep their existing stats
|
||||||
"""
|
"""
|
||||||
logging.info(f"Loading dataset: {repo_id}")
|
logging.info(f"Loading dataset: {repo_id}")
|
||||||
dataset = LeRobotDataset(
|
dataset = LeRobotDataset(
|
||||||
@@ -232,7 +238,13 @@ def augment_dataset_with_quantile_stats(
|
|||||||
|
|
||||||
logging.info("Dataset does not contain quantile statistics. Computing them now...")
|
logging.info("Dataset does not contain quantile statistics. Computing them now...")
|
||||||
|
|
||||||
new_stats = compute_quantile_stats_for_dataset(dataset, use_sampling=use_sampling)
|
new_stats = compute_quantile_stats_for_dataset(
|
||||||
|
dataset, use_sampling=use_sampling, skip_images=skip_images
|
||||||
|
)
|
||||||
|
|
||||||
|
if skip_images and dataset.meta.stats:
|
||||||
|
for key, feature_stats in dataset.meta.stats.items():
|
||||||
|
new_stats.setdefault(key, feature_stats)
|
||||||
|
|
||||||
logging.info("Updating dataset metadata with new quantile statistics")
|
logging.info("Updating dataset metadata with new quantile statistics")
|
||||||
dataset.meta.stats = new_stats
|
dataset.meta.stats = new_stats
|
||||||
@@ -276,10 +288,15 @@ def main():
|
|||||||
"--no-sampling",
|
"--no-sampling",
|
||||||
action="store_true",
|
action="store_true",
|
||||||
help=(
|
help=(
|
||||||
"Compute stats over every frame (exact, higher memory). By default, "
|
"Compute stats over every frame (higher memory). By default, "
|
||||||
"image/video frames are sub-sampled per episode to bound memory."
|
"image/video frames are sub-sampled per episode to bound memory."
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--skip-images",
|
||||||
|
action="store_true",
|
||||||
|
help="Skip image/video features and preserve their existing stats",
|
||||||
|
)
|
||||||
|
|
||||||
args = parser.parse_args()
|
args = parser.parse_args()
|
||||||
root = Path(args.root) if args.root else None
|
root = Path(args.root) if args.root else None
|
||||||
@@ -291,6 +308,7 @@ def main():
|
|||||||
root=root,
|
root=root,
|
||||||
overwrite=args.overwrite,
|
overwrite=args.overwrite,
|
||||||
use_sampling=not args.no_sampling,
|
use_sampling=not args.no_sampling,
|
||||||
|
skip_images=args.skip_images,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -12,8 +12,11 @@
|
|||||||
# See the License for the specific language governing permissions and
|
# See the License for the specific language governing permissions and
|
||||||
# limitations under the License.
|
# limitations under the License.
|
||||||
|
|
||||||
|
from types import SimpleNamespace
|
||||||
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import pytest
|
import pytest
|
||||||
|
import torch
|
||||||
|
|
||||||
pytest.importorskip("datasets", reason="datasets is required (install lerobot[dataset])")
|
pytest.importorskip("datasets", reason="datasets is required (install lerobot[dataset])")
|
||||||
|
|
||||||
@@ -24,7 +27,9 @@ from lerobot.scripts.augment_dataset_quantile_stats import (
|
|||||||
|
|
||||||
|
|
||||||
def _numeric_keys(dataset):
|
def _numeric_keys(dataset):
|
||||||
return [k for k, v in dataset.features.items() if v["dtype"] not in ("image", "video", "string")]
|
return [
|
||||||
|
k for k, v in dataset.features.items() if v["dtype"] not in ("image", "video", "string", "language")
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
def _image_keys(dataset):
|
def _image_keys(dataset):
|
||||||
@@ -102,3 +107,112 @@ def test_quantile_stats_present_after_compute(tmp_path, lerobot_dataset_factory)
|
|||||||
)
|
)
|
||||||
stats = compute_quantile_stats_for_dataset(dataset, use_sampling=True)
|
stats = compute_quantile_stats_for_dataset(dataset, use_sampling=True)
|
||||||
assert has_quantile_stats(stats)
|
assert has_quantile_stats(stats)
|
||||||
|
|
||||||
|
|
||||||
|
class FakeHFDataset:
|
||||||
|
"""Minimal stand-in exposing the column slicing used by the augment script."""
|
||||||
|
|
||||||
|
def __init__(self, columns: dict[str, list]):
|
||||||
|
self._columns = columns
|
||||||
|
|
||||||
|
def select_columns(self, keys):
|
||||||
|
return FakeHFDataset({key: self._columns[key] for key in keys})
|
||||||
|
|
||||||
|
def __getitem__(self, index):
|
||||||
|
return {key: values[index] for key, values in self._columns.items()}
|
||||||
|
|
||||||
|
|
||||||
|
def test_compute_quantile_stats_skips_language_features():
|
||||||
|
class FakeDataset:
|
||||||
|
num_episodes = 1
|
||||||
|
features = {
|
||||||
|
"action": {"dtype": "float32"},
|
||||||
|
"observation.language": {"dtype": "language"},
|
||||||
|
}
|
||||||
|
meta = SimpleNamespace(episodes=[{"dataset_from_index": 0, "dataset_to_index": 2}])
|
||||||
|
hf_dataset = FakeHFDataset(
|
||||||
|
{
|
||||||
|
"action": [[0.0], [1.0]],
|
||||||
|
"observation.language": [
|
||||||
|
[{"role": "user", "content": "pick"}],
|
||||||
|
[{"role": "assistant", "content": "done"}],
|
||||||
|
],
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
stats = compute_quantile_stats_for_dataset(FakeDataset())
|
||||||
|
|
||||||
|
assert set(stats) == {"action"}
|
||||||
|
|
||||||
|
|
||||||
|
def test_compute_quantile_stats_skip_images_avoids_decoding():
|
||||||
|
class FakeDataset:
|
||||||
|
num_episodes = 1
|
||||||
|
features = {
|
||||||
|
"action": {"dtype": "float32"},
|
||||||
|
"observation.images.cam": {"dtype": "video"},
|
||||||
|
}
|
||||||
|
meta = SimpleNamespace(episodes=[{"dataset_from_index": 0, "dataset_to_index": 2}])
|
||||||
|
hf_dataset = FakeHFDataset({"action": [[0.0], [1.0]]})
|
||||||
|
|
||||||
|
def __getitem__(self, index):
|
||||||
|
raise AssertionError(f"video frame {index} was decoded despite skip_images=True")
|
||||||
|
|
||||||
|
stats = compute_quantile_stats_for_dataset(FakeDataset(), skip_images=True)
|
||||||
|
|
||||||
|
assert set(stats) == {"action"}
|
||||||
|
|
||||||
|
|
||||||
|
def test_compute_quantile_stats_handles_single_frame():
|
||||||
|
class FakeDataset:
|
||||||
|
num_episodes = 1
|
||||||
|
features = {"action": {"dtype": "float32"}}
|
||||||
|
meta = SimpleNamespace(episodes=[{"dataset_from_index": 0, "dataset_to_index": 1}])
|
||||||
|
hf_dataset = FakeHFDataset({"action": [[5.0, 7.0]]})
|
||||||
|
|
||||||
|
stats = compute_quantile_stats_for_dataset(FakeDataset())
|
||||||
|
|
||||||
|
np.testing.assert_array_equal(stats["action"]["count"], np.array([1]))
|
||||||
|
for key in ("min", "max", "mean", "q01", "q10", "q50", "q90", "q99"):
|
||||||
|
np.testing.assert_allclose(stats["action"][key], np.array([5.0, 7.0]))
|
||||||
|
|
||||||
|
|
||||||
|
def test_compute_quantile_stats_image_count_uses_frames():
|
||||||
|
frames = [torch.zeros(3, 2, 2), torch.ones(3, 2, 2)]
|
||||||
|
|
||||||
|
class FakeDataset:
|
||||||
|
num_episodes = 1
|
||||||
|
features = {"observation.images.cam": {"dtype": "video"}}
|
||||||
|
meta = SimpleNamespace(episodes=[{"dataset_from_index": 0, "dataset_to_index": 2}])
|
||||||
|
hf_dataset = FakeHFDataset({})
|
||||||
|
|
||||||
|
def __getitem__(self, index):
|
||||||
|
return {"observation.images.cam": frames[index]}
|
||||||
|
|
||||||
|
stats = compute_quantile_stats_for_dataset(FakeDataset(), use_sampling=False)
|
||||||
|
image_stats = stats["observation.images.cam"]
|
||||||
|
|
||||||
|
np.testing.assert_array_equal(image_stats["count"], np.array([2]))
|
||||||
|
assert image_stats["mean"].shape == (3, 1, 1)
|
||||||
|
np.testing.assert_allclose(image_stats["mean"], np.full((3, 1, 1), 0.5))
|
||||||
|
|
||||||
|
|
||||||
|
def test_compute_quantile_stats_accumulates_across_episodes():
|
||||||
|
values = [[float(value)] for value in range(100)] + [[float(value)] for value in range(1000, 1010)]
|
||||||
|
|
||||||
|
class FakeDataset:
|
||||||
|
num_episodes = 2
|
||||||
|
features = {"action": {"dtype": "float32"}}
|
||||||
|
meta = SimpleNamespace(
|
||||||
|
episodes=[
|
||||||
|
{"dataset_from_index": 0, "dataset_to_index": 100},
|
||||||
|
{"dataset_from_index": 100, "dataset_to_index": 110},
|
||||||
|
]
|
||||||
|
)
|
||||||
|
hf_dataset = FakeHFDataset({"action": values})
|
||||||
|
|
||||||
|
stats = compute_quantile_stats_for_dataset(FakeDataset())
|
||||||
|
|
||||||
|
np.testing.assert_array_equal(stats["action"]["count"], np.array([110]))
|
||||||
|
expected_q90 = np.percentile(np.asarray(values), 90, axis=0)
|
||||||
|
np.testing.assert_allclose(stats["action"]["q90"], expected_q90, atol=0.1)
|
||||||
|
|||||||
@@ -688,7 +688,7 @@ def test_compute_episode_stats_string_features_skipped():
|
|||||||
|
|
||||||
|
|
||||||
def test_aggregate_feature_stats_with_quantiles():
|
def test_aggregate_feature_stats_with_quantiles():
|
||||||
"""Test aggregating feature stats that include quantiles."""
|
"""Test aggregating feature stats that include quantiles uses conservative bounds."""
|
||||||
stats_ft_list = [
|
stats_ft_list = [
|
||||||
{
|
{
|
||||||
"min": np.array([1.0]),
|
"min": np.array([1.0]),
|
||||||
@@ -697,6 +697,9 @@ def test_aggregate_feature_stats_with_quantiles():
|
|||||||
"std": np.array([2.0]),
|
"std": np.array([2.0]),
|
||||||
"count": np.array([100]),
|
"count": np.array([100]),
|
||||||
"q01": np.array([1.5]),
|
"q01": np.array([1.5]),
|
||||||
|
"q10": np.array([2.0]),
|
||||||
|
"q50": np.array([5.0]),
|
||||||
|
"q90": np.array([9.0]),
|
||||||
"q99": np.array([9.5]),
|
"q99": np.array([9.5]),
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
@@ -706,22 +709,21 @@ def test_aggregate_feature_stats_with_quantiles():
|
|||||||
"std": np.array([2.5]),
|
"std": np.array([2.5]),
|
||||||
"count": np.array([150]),
|
"count": np.array([150]),
|
||||||
"q01": np.array([2.5]),
|
"q01": np.array([2.5]),
|
||||||
|
"q10": np.array([3.0]),
|
||||||
|
"q50": np.array([6.0]),
|
||||||
|
"q90": np.array([11.0]),
|
||||||
"q99": np.array([11.5]),
|
"q99": np.array([11.5]),
|
||||||
},
|
},
|
||||||
]
|
]
|
||||||
|
|
||||||
result = aggregate_feature_stats(stats_ft_list)
|
result = aggregate_feature_stats(stats_ft_list)
|
||||||
|
|
||||||
# Should preserve quantiles
|
# Lower quantiles use min; upper quantiles use max, regardless of counts.
|
||||||
assert "q01" in result
|
np.testing.assert_allclose(result["q01"], np.array([1.5]), atol=1e-6)
|
||||||
assert "q99" in result
|
np.testing.assert_allclose(result["q10"], np.array([2.0]), atol=1e-6)
|
||||||
|
np.testing.assert_allclose(result["q50"], np.array([5.0]), atol=1e-6)
|
||||||
# Verify quantile aggregation (weighted average)
|
np.testing.assert_allclose(result["q90"], np.array([11.0]), atol=1e-6)
|
||||||
expected_q01 = (1.5 * 100 + 2.5 * 150) / 250 # ≈ 2.1
|
np.testing.assert_allclose(result["q99"], np.array([11.5]), atol=1e-6)
|
||||||
expected_q99 = (9.5 * 100 + 11.5 * 150) / 250 # ≈ 10.7
|
|
||||||
|
|
||||||
np.testing.assert_allclose(result["q01"], np.array([expected_q01]), atol=1e-6)
|
|
||||||
np.testing.assert_allclose(result["q99"], np.array([expected_q99]), atol=1e-6)
|
|
||||||
|
|
||||||
|
|
||||||
def test_aggregate_stats_mixed_quantiles():
|
def test_aggregate_stats_mixed_quantiles():
|
||||||
@@ -878,3 +880,60 @@ def test_fixed_quantiles_always_computed():
|
|||||||
for q_key in expected_quantiles:
|
for q_key in expected_quantiles:
|
||||||
assert q_key in episode_stats[key]
|
assert q_key in episode_stats[key]
|
||||||
assert episode_stats[key][q_key].shape == (features[key]["shape"][0],)
|
assert episode_stats[key][q_key].shape == (features[key]["shape"][0],)
|
||||||
|
|
||||||
|
|
||||||
|
def test_aggregate_stats_incremental_resume():
|
||||||
|
"""Verify conservative bounds remain associative across incremental additions."""
|
||||||
|
# Start with episode 1 stats (narrow distribution)
|
||||||
|
ep1_stats = {
|
||||||
|
"action": {
|
||||||
|
"min": np.array([-10.0, -5.0]),
|
||||||
|
"max": np.array([10.0, 5.0]),
|
||||||
|
"mean": np.array([0.0, 0.0]),
|
||||||
|
"std": np.array([3.0, 1.5]),
|
||||||
|
"count": np.array([500]),
|
||||||
|
"q01": np.array([-9.0, -4.5]),
|
||||||
|
"q99": np.array([9.0, 4.5]),
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
# Episode 2: wider distribution on dim 0
|
||||||
|
ep2_stats = {
|
||||||
|
"action": {
|
||||||
|
"min": np.array([-30.0, -5.0]),
|
||||||
|
"max": np.array([40.0, 6.0]),
|
||||||
|
"mean": np.array([5.0, 0.5]),
|
||||||
|
"std": np.array([15.0, 2.0]),
|
||||||
|
"count": np.array([100]),
|
||||||
|
"q01": np.array([-25.0, -4.0]),
|
||||||
|
"q99": np.array([35.0, 5.5]),
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
# First aggregation: ep1 + ep2 (simulates save_episode for ep2)
|
||||||
|
cumulative = aggregate_stats([ep1_stats, ep2_stats])
|
||||||
|
|
||||||
|
# q01 should take min (conservative lower bound)
|
||||||
|
np.testing.assert_allclose(cumulative["action"]["q01"], np.array([-25.0, -4.5]))
|
||||||
|
# q99 should take max (conservative upper bound)
|
||||||
|
np.testing.assert_allclose(cumulative["action"]["q99"], np.array([35.0, 5.5]))
|
||||||
|
|
||||||
|
# Episode 3: even wider on dim 1
|
||||||
|
ep3_stats = {
|
||||||
|
"action": {
|
||||||
|
"min": np.array([-8.0, -20.0]),
|
||||||
|
"max": np.array([8.0, 25.0]),
|
||||||
|
"mean": np.array([0.0, 3.0]),
|
||||||
|
"std": np.array([2.0, 8.0]),
|
||||||
|
"count": np.array([50]),
|
||||||
|
"q01": np.array([-7.0, -18.0]),
|
||||||
|
"q99": np.array([7.0, 22.0]),
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
# Second aggregation: cumulative + ep3 (simulates save_episode for ep3)
|
||||||
|
cumulative2 = aggregate_stats([cumulative, ep3_stats])
|
||||||
|
|
||||||
|
# Bounds should widen monotonically
|
||||||
|
np.testing.assert_allclose(cumulative2["action"]["q01"], np.array([-25.0, -18.0]))
|
||||||
|
np.testing.assert_allclose(cumulative2["action"]["q99"], np.array([35.0, 22.0]))
|
||||||
|
|||||||
Reference in New Issue
Block a user