openpi RoboCasa365 fork + pi0.5 port, single-task PickPlaceCounterToCabinet configs
44aecac verified Download scripts/_test_data_loader.py from Ronaldo-GOAT/pi_training: direct link, hf CLI and curl.
- Browser
- Download file 3.68 kB
-
https://huggingface.co/Ronaldo-GOAT/pi_training/resolve/main/scripts/_test_data_loader.py
- Command line
-
hf download hf://Ronaldo-GOAT/pi_training/scripts/_test_data_loader.py
-
curl -L -o _test_data_loader.py https://huggingface.co/Ronaldo-GOAT/pi_training/resolve/main/scripts/_test_data_loader.py
3.68 kB
| """Ad-hoc data loading sanity check for pi0_fast_robocasa_pretrain_human300. | |
| Loads the first dataset from the config's data_dirs only, builds a tiny | |
| GrootOpenpiSingleDataset, and pulls one sample + one transformed batch. | |
| """ | |
| import dataclasses | |
| import os | |
| import sys | |
| import time | |
| os.environ.setdefault("JAX_PLATFORMS", "cpu") # keep this probe CPU-only; model isn't loaded | |
| import numpy as np | |
| from openpi.training import config as _config | |
| from openpi.training import data_loader as _dl | |
| def main() -> int: | |
| cfg = _config.get_config("pi0_fast_robocasa_pretrain_human300") | |
| # Pre-filter data_dirs to only those that exist so norm-stats loading doesn't crash | |
| # on a missing dataset (e.g. RecycleBottlesBySize). | |
| full = list(cfg.data.data_dirs) | |
| print(f"[info] total registered datasets: {len(full)}") | |
| available = [e for e in full if os.path.exists(e["path"])] | |
| print(f"[info] available on disk: {len(available)}") | |
| if not available: | |
| print("[err ] no dataset paths exist on disk") | |
| return 1 | |
| pick = available[0] | |
| print(f"[info] using first available:\n task={pick['task']}\n path={pick['path']}") | |
| # Build a LeRobotRobocasaDataConfig that points at just this one dataset, | |
| # then let it produce a DataConfig (which will try to load norm_stats for this dir only). | |
| narrow_factory = dataclasses.replace(cfg.data, data_dirs=[pick]) | |
| single_data_config = narrow_factory.create(cfg.assets_dirs, cfg.model) | |
| t0 = time.time() | |
| raw_ds = _dl.create_torch_dataset( | |
| single_data_config, action_horizon=cfg.model.action_horizon, model_config=cfg.model | |
| ) | |
| print(f"[info] raw dataset class: {type(raw_ds).__name__}") | |
| print(f"[info] len(raw_ds) : {len(raw_ds)}") | |
| print(f"[time] build raw_ds : {time.time() - t0:.2f}s") | |
| t0 = time.time() | |
| sample = raw_ds[0] | |
| print(f"[time] raw_ds[0] : {time.time() - t0:.2f}s") | |
| print("[info] raw sample keys :", list(sample.keys())) | |
| for k, v in sample.items(): | |
| if isinstance(v, np.ndarray): | |
| print(f" {k}: ndarray shape={v.shape} dtype={v.dtype}") | |
| elif isinstance(v, (list, tuple)): | |
| print(f" {k}: {type(v).__name__} len={len(v)}") | |
| else: | |
| print(f" {k}: {type(v).__name__}") | |
| # Now put the transforms on top and pull a mini-batch. | |
| ds = _dl.transform_dataset(raw_ds, single_data_config, skip_norm_stats=True) | |
| print(f"[info] transformed cls : {type(ds).__name__}") | |
| loader = _dl.TorchDataLoader( | |
| ds, | |
| local_batch_size=2, | |
| shuffle=False, | |
| num_batches=1, | |
| num_workers=0, | |
| ) | |
| t0 = time.time() | |
| batch = next(iter(loader)) | |
| print(f"[time] pull 1 batch : {time.time() - t0:.2f}s") | |
| print(f"[info] batch container : {type(batch).__name__}") | |
| def _describe(x, prefix): | |
| if isinstance(x, dict): | |
| for k, v in x.items(): | |
| _describe(v, f"{prefix}.{k}") | |
| elif hasattr(x, "shape") or hasattr(x, "dtype"): | |
| print(f" {prefix}: shape={getattr(x, 'shape', None)} dtype={getattr(x, 'dtype', None)}") | |
| elif isinstance(x, (list, tuple)): | |
| print(f" {prefix}: {type(x).__name__}(len={len(x)})") | |
| for i, e in enumerate(x): | |
| _describe(e, f"{prefix}[{i}]") | |
| elif dataclasses.is_dataclass(x): | |
| _describe(dataclasses.asdict(x), prefix) | |
| else: | |
| print(f" {prefix}: {type(x).__name__}") | |
| _describe(batch, "batch") | |
| print("[ok ] data loading smoke test passed") | |
| return 0 | |
| if __name__ == "__main__": | |
| sys.exit(main()) | |