File size: 6,560 Bytes
b152c62
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
#!/usr/bin/env python
"""Validate a packed wds dataset dir (or several) against the loader contract.

Checks, per dataset:
  * manifest.json lists every shard-*.tar that exists, tasks.json parses
  * samples decode: cam0 JPEG opens at the declared size, meta.npz has all keys
  * dtypes/shapes: state f32[S], action_chunk f32[chunk,A], action_is_pad bool
  * normalized sanity: |mean| < 5, std in (1e-6, 50) over a sample of frames
  * task_index maps into tasks.json
Then, jointly: a few batches through wds_mix.make_mixture_loader and the batch
schema the trainer expects (incl. the "task" string -> present and str).

Usage:
    python validate_wds.py /data/wds/SetA /data/wds/SetB --batches 4
"""

from __future__ import annotations

import argparse
import io
import json
import sys
import tarfile
from pathlib import Path

import numpy as np


def check_dataset(d: Path, n_samples: int = 300) -> dict:
    errs = []
    mf = json.loads((d / "manifest.json").read_text())
    tasks = json.loads((d / "tasks.json").read_text())
    tars = sorted(d.glob("shard-*.tar"))
    if sorted(mf["shards"]) != [t.name for t in tars]:
        errs.append(f"manifest shards != on-disk tars ({len(mf['shards'])} vs {len(tars)})")
    chunk = int(mf["chunk"])

    seen, s_all, a_all = 0, [], []
    from PIL import Image

    for t in tars:
        with tarfile.open(t) as tf:
            members = tf.getmembers()
            by_key = {}
            for m in members:
                k, _, ext = m.name.partition(".")
                by_key.setdefault(k, {})[ext] = m
            for k in by_key:
                if seen >= n_samples:
                    break
                need = {"cam0.jpg", "cam1.jpg", "meta.npz"}
                raw = {m.name[len(k) + 1:]: m for m in members if m.name.startswith(k + ".")}
                if not need.issubset(raw):
                    errs.append(f"{t.name}:{k}: missing {need - set(raw)}")
                    seen += 1
                    continue
                c0 = tf.extractfile(raw["cam0.jpg"]).read()
                img = Image.open(io.BytesIO(c0))
                img.load()
                meta = np.load(io.BytesIO(tf.extractfile(raw["meta.npz"]).read()))
                for kk, dt in (("state", np.float32), ("action_chunk", np.float32),
                               ("action_is_pad", np.bool_)):
                    if kk not in meta:
                        errs.append(f"{k}: meta missing {kk}")
                    elif meta[kk].dtype != dt:
                        errs.append(f"{k}: {kk} dtype {meta[kk].dtype} != {dt}")
                ac = meta["action_chunk"]
                if ac.shape[0] != chunk:
                    errs.append(f"{k}: chunk {ac.shape[0]} != manifest {chunk}")
                if int(meta["action_dim"]) != ac.shape[1] or int(meta["state_dim"]) != meta["state"].shape[0]:
                    errs.append(f"{k}: dim fields disagree with arrays")
                if str(int(meta["task_index"])) not in tasks:
                    errs.append(f"{k}: task_index {int(meta['task_index'])} not in tasks.json")
                s_all.append(meta["state"]); a_all.append(ac[0])
                seen += 1
        if seen >= n_samples:
            break

    if s_all:
        s = np.stack(s_all); a = np.stack(a_all)
        for nm, arr in (("state", s), ("action", a)):
            m, sd = np.abs(arr.mean(0)).max(), arr.std(0)
            if m > 5:
                errs.append(f"{nm}: |mean| up to {m:.2f} — normalization looks wrong")
            if sd.max() > 50:
                errs.append(f"{nm}: std up to {sd.max():.1f} — normalization looks wrong")
    return {"dir": str(d), "samples_checked": seen, "declared": mf.get("samples"),
            "episodes": mf.get("episodes"), "errors": errs}


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("dirs", nargs="+", type=Path)
    ap.add_argument("--batches", type=int, default=2)
    ap.add_argument("--batch-size", type=int, default=8)
    ap.add_argument("--max-state-dim", type=int, default=256)
    ap.add_argument("--max-action-dim", type=int, default=64)
    ap.add_argument("--samples", type=int, default=300)
    args = ap.parse_args()

    bad = False
    for d in args.dirs:
        r = check_dataset(d, args.samples)
        status = "OK " if not r["errors"] else "FAIL"
        print(f"[{status}] {r['dir']}: {r['samples_checked']} samples checked "
              f"(declared {r['declared']}, {r['episodes']} eps)")
        for e in r["errors"][:20]:
            print("   -", e)
        bad = bad or bool(r["errors"])

    if args.batches > 0:
        sys.path.insert(0, str(Path(__file__).resolve().parents[1] / "src"))
        from tinyvla.data.wds_mix import make_mixture_loader

        loader, mfs = make_mixture_loader(
            [{"dir": str(d)} for d in args.dirs],
            batch_size=args.batch_size, num_workers=0,
            max_state_dim=args.max_state_dim, max_action_dim=args.max_action_dim,
            shuffle_buffer=64, steps_per_epoch=args.batches + 1,
        )
        import torch

        want = {
            "observation.images.cam0": (torch.float32, (3, 256, 256)),
            "observation.images.cam1": (torch.float32, (3, 256, 256)),
            "camera_mask": (torch.bool, (2,)),
            "observation.state": (torch.float32, (args.max_state_dim,)),
            "action": (torch.float32, (None, args.max_action_dim)),
            "action_dim_mask": (torch.bool, (args.max_action_dim,)),
            "action_is_pad": (torch.bool, None),
            "embodiment_id": (torch.int64, ()),
            "task_index": (torch.int64, ()),
        }
        it = iter(loader)
        for b in range(args.batches):
            batch = next(it)
            for k, (dt, shp) in want.items():
                assert k in batch, f"batch missing {k}"
                assert batch[k].dtype == dt, f"{k}: dtype {batch[k].dtype} != {dt}"
                if shp:
                    tail = batch[k].shape[1:]
                    for i, s in enumerate(shp):
                        if s is not None:
                            assert tail[i] == s, f"{k}: shape {tuple(batch[k].shape)}"
            assert "task" in batch and isinstance(batch["task"][0], str), "task strings missing"
            assert not any(k.startswith("__") and torch.is_tensor(v) for k, v in batch.items())
        print(f"[OK ] mixture loader: {args.batches} batches, schema verified, "
              f"tasks e.g. {batch['task'][:2]}")

    sys.exit(1 if bad else 0)


if __name__ == "__main__":
    main()