Download scripts/compute_norm_stats.py from Dengliming/StreamPIReal6_3w: direct link, hf CLI and curl.
- Browser
- Download file 4.68 kB
-
https://huggingface.co/Dengliming/StreamPIReal6_3w/resolve/main/scripts/compute_norm_stats.py
- Command line
-
hf download hf://Dengliming/StreamPIReal6_3w/scripts/compute_norm_stats.py
-
curl -L -o compute_norm_stats.py https://huggingface.co/Dengliming/StreamPIReal6_3w/resolve/main/scripts/compute_norm_stats.py
4.68 kB
| """Compute normalization statistics for a config. | |
| This script is used to compute the normalization statistics for a given config. It | |
| will compute the mean and standard deviation of the data in the dataset and save it | |
| to the config assets directory. | |
| """ | |
| import pathlib | |
| import numpy as np | |
| import tqdm | |
| import tyro | |
| import openpi.models.model as _model | |
| import openpi.shared.normalize as normalize | |
| import openpi.training.config as _config | |
| import openpi.training.data_loader as _data_loader | |
| import openpi.transforms as transforms | |
| class RemoveStrings(transforms.DataTransformFn): | |
| def __call__(self, x: dict) -> dict: | |
| return {k: v for k, v in x.items() if not np.issubdtype(np.asarray(v).dtype, np.str_)} | |
| def create_torch_dataloader( | |
| data_config: _config.DataConfig, | |
| action_horizon: int, | |
| batch_size: int, | |
| model_config: _model.BaseModelConfig, | |
| num_workers: int, | |
| max_frames: int | None = None, | |
| ) -> tuple[_data_loader.Dataset, int]: | |
| if data_config.repo_id is None: | |
| raise ValueError("Data config must have a repo_id") | |
| # This script only reads `state`/`actions` from each batch (see `keys` below), never the raw | |
| # image/video keys, so we can skip mp4 decoding entirely for the v3.0 RoboDojo adapter dataset. | |
| # This is a large speedup (hours -> minutes) with no effect on the computed statistics. | |
| dataset = _data_loader.create_torch_dataset( | |
| data_config, action_horizon, model_config, skip_video_decode=True | |
| ) | |
| dataset = _data_loader.TransformedDataset( | |
| dataset, | |
| [ | |
| *data_config.repack_transforms.inputs, | |
| *data_config.data_transforms.inputs, | |
| # Remove strings since they are not supported by JAX and are not needed to compute norm stats. | |
| RemoveStrings(), | |
| ], | |
| ) | |
| if max_frames is not None and max_frames < len(dataset): | |
| num_batches = max_frames // batch_size | |
| shuffle = False | |
| else: | |
| num_batches = len(dataset) // batch_size | |
| shuffle = False | |
| data_loader = _data_loader.TorchDataLoader( | |
| dataset, | |
| local_batch_size=batch_size, | |
| num_workers=num_workers, | |
| shuffle=shuffle, | |
| num_batches=num_batches, | |
| ) | |
| return data_loader, num_batches | |
| def create_rlds_dataloader( | |
| data_config: _config.DataConfig, | |
| action_horizon: int, | |
| batch_size: int, | |
| max_frames: int | None = None, | |
| ) -> tuple[_data_loader.Dataset, int]: | |
| dataset = _data_loader.create_rlds_dataset(data_config, action_horizon, batch_size, shuffle=False) | |
| dataset = _data_loader.IterableTransformedDataset( | |
| dataset, | |
| [ | |
| *data_config.repack_transforms.inputs, | |
| *data_config.data_transforms.inputs, | |
| # Remove strings since they are not supported by JAX and are not needed to compute norm stats. | |
| RemoveStrings(), | |
| ], | |
| is_batched=True, | |
| ) | |
| if max_frames is not None and max_frames < len(dataset): | |
| num_batches = max_frames // batch_size | |
| else: | |
| # NOTE: this length is currently hard-coded for DROID. | |
| num_batches = len(dataset) // batch_size | |
| data_loader = _data_loader.RLDSDataLoader( | |
| dataset, | |
| num_batches=num_batches, | |
| ) | |
| return data_loader, num_batches | |
| def main(config_name: str, max_frames: int | None = None, num_workers: int | None = None): | |
| config = _config.get_config(config_name) | |
| data_config = config.data.create(config.assets_dirs, config.model) | |
| workers = config.num_workers if num_workers is None else num_workers | |
| if data_config.rlds_data_dir is not None: | |
| data_loader, num_batches = create_rlds_dataloader( | |
| data_config, config.model.action_horizon, config.batch_size, max_frames | |
| ) | |
| else: | |
| data_loader, num_batches = create_torch_dataloader( | |
| data_config, config.model.action_horizon, config.batch_size, config.model, workers, max_frames | |
| ) | |
| keys = ["state", "actions"] | |
| stats = {key: normalize.RunningStats() for key in keys} | |
| for batch in tqdm.tqdm(data_loader, total=num_batches, desc="Computing stats"): | |
| for key in keys: | |
| stats[key].update(np.asarray(batch[key])) | |
| norm_stats = {key: stats.get_statistics() for key, stats in stats.items()} | |
| if data_config.asset_id is None: | |
| raise ValueError("Data config must have an asset_id for normalization stats.") | |
| assets_dir = config.data.assets.assets_dir or config.assets_dirs | |
| output_path = pathlib.Path(assets_dir) / data_config.asset_id | |
| print(f"Writing stats to: {output_path}") | |
| normalize.save(output_path, norm_stats) | |
| if __name__ == "__main__": | |
| tyro.cli(main) |