File size: 3,680 Bytes
44aecac
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
"""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())