"""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())