pi_training / scripts /_test_data_loader.py
Ronaldo-GOAT's picture
openpi RoboCasa365 fork + pi0.5 port, single-task PickPlaceCounterToCabinet configs
44aecac verified
Raw History Blame Contribute Delete
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())