File size: 10,430 Bytes
2622c40
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
"""Preflight checks for the two pusht training runs.

CPU-only, no GPU, no Slurm job needed. Run from anywhere:

    /shared_work/environments/miniconda3/envs/fastwam-robotwin-rui/bin/python \
        /shared_work/george/real_world_wam/pusht_train/preflight.py \
        --task pusht_flashwam_scratch [--dataset-dir <path>]

Checks:
1. Hydra config composes (our --config-dir overlaid on the checkout configs)
   and the model / batch / epoch / resume wiring is what we intend.
2. Text embedding cache hit for the exact pusht prompt.
3. Train dataset builds with the REAL config; one sample has the exact tensor
   contract (shapes, value ranges, prompt, context).
4. THE DEGENERATE-CHANNEL CONTRACT. This dataset has 6 constant channels
   (action drx/dry/drz/gripper, proprio gripL/gripR) that the source README
   warns will divide by zero under min/max normalization. The normalizer's
   `ignore_dim = input_range < range_tol` guard (range_tol=1e-4) is supposed to
   catch all six and emit a constant 0. This asserts that it actually does:
     - every constant channel is finite and stays within +/-range_tol after
       normalization (the guard sets scale=1.0 / offset=-min, so an ignored
       dim maps to x-min: exactly 0 for the 4 truly-constant action dims, and
       [0, 3.2e-05] for the two gripper state dims)
     - the live channels (action dx/dy/dz, proprio x/y/z/rx/ry/rz) actually
       vary and reach the normalizer's endpoints
   A NaN or a nonzero-but-constant channel here means the guard did not fire,
   and training would silently produce garbage on those dims.

--dataset-dir defaults to the shared-FS copy so this runs on the login node;
the configs point at the node-local /tmp copy that only exists at job time.
"""

import argparse
import hashlib
import sys
import tempfile
from pathlib import Path

FASTWAM_ROOT = Path("/shared_work/physical_intelligence/policies/Fast-WAM/fastwam")
CONFIG_DIR = Path("/shared_work/george/real_world_wam/pusht_train/configs")
SHARED_DATASET = "/shared_work/george/real_world_wam/datasets/pusht_lerobot_v21"
CACHE_DIR = "/shared_work/george/real_world_wam/pusht_train/text_embeds_cache"
TASK_STR = "push the T block to the target outline"

# Measured on the raw HDF5 across all 100 episodes / 32,131 frames.
# index -> (name, is_constant)
ACTION_DIMS = [(0, "dx", False), (1, "dy", False), (2, "dz", False),
               (3, "drx", True), (4, "dry", True), (5, "drz", True),
               (6, "gripper", True)]
PROPRIO_DIMS = [(0, "x", False), (1, "y", False), (2, "z", False),
                (3, "rx", False), (4, "ry", False), (5, "rz", False),
                (6, "gripL", True), (7, "gripR", True)]

EXPECT = {
    "pusht_flashwam_scratch": {
        "target": "fastwam.models.wan22.fasterwam_decoupled.create_fasterwam_decoupled",
        "action_layers": 1,
        "kv_source_mode": "fused_kv",
        "fixed_rope": True,
        "action_dit_pretrained_none": True,
    },
    "pusht_fastwam_scratch": {
        "target": "fastwam.runtime.create_fastwam",
        "action_layers": 30,
        "kv_source_mode": None,
        "fixed_rope": None,
        "action_dit_pretrained_none": False,
    },
}

sys.path.insert(0, str(FASTWAM_ROOT / "src"))


def compose_cfg(task, exp, dataset_dir):
    from hydra import compose, initialize_config_dir

    from fastwam.utils.config_resolvers import register_default_resolvers

    register_default_resolvers()
    with initialize_config_dir(config_dir=str(FASTWAM_ROOT / "configs"), version_base="1.3"):
        cfg = compose(
            config_name="train",
            overrides=[f"task={task}", f"hydra.searchpath=[{CONFIG_DIR}]"],
        )
    assert cfg.data.train.dataset_dirs == ["/tmp/george_pusht/pusht_lerobot_v21"], \
        cfg.data.train.dataset_dirs
    assert cfg.model._target_ == exp["target"], cfg.model._target_
    assert cfg.model.action_dit_config.num_layers == exp["action_layers"]
    assert cfg.model.video_dit_config.num_layers == 30
    assert cfg.num_epochs == 30, cfg.num_epochs
    assert cfg.batch_size == 8, cfg.batch_size
    assert cfg.gradient_accumulation_steps == 1, cfg.gradient_accumulation_steps
    assert cfg.learning_rate == 1e-4, cfg.learning_rate
    assert not cfg.resume, f"expected scratch (resume null), got {cfg.resume}"
    assert cfg.data.train.processor.norm_default_mode == "min/max"
    assert cfg.data.train.val_set_proportion == 0.0
    if exp["kv_source_mode"] is not None:
        assert cfg.model.kv_source_mode == exp["kv_source_mode"], cfg.model.kv_source_mode
    if exp["fixed_rope"] is not None:
        assert cfg.model.fixed_rope is exp["fixed_rope"]
    if exp["action_dit_pretrained_none"]:
        assert cfg.model.action_dit_pretrained_path is None, cfg.model.action_dit_pretrained_path
    else:
        assert cfg.model.action_dit_pretrained_path is not None

    # Re-point at the shared-FS copy + shared cache so this runs off-node.
    cfg.data.train.dataset_dirs = [dataset_dir]
    cfg.data.train.text_embedding_cache_dir = CACHE_DIR
    print(f"[1/4] hydra config composes OK "
          f"({exp['action_layers']}-layer action expert, scratch, global batch "
          f"{cfg.batch_size * 4}, {cfg.num_epochs} epochs)")
    return cfg


def check_text_cache(cfg):
    from fastwam.datasets.lerobot.robot_video_dataset import DEFAULT_PROMPT

    prompt = DEFAULT_PROMPT.format(task=TASK_STR)
    hashed = hashlib.sha256(prompt.encode("utf-8")).hexdigest()
    cache = Path(cfg.data.train.text_embedding_cache_dir) / f"{hashed}.t5_len128.wan22ti2v5b.pt"
    assert cache.exists(), f"Missing text embedding cache {cache}"
    print(f"[2/4] text embedding cache hit: {cache.name[:16]}… (task {TASK_STR!r})")


def check_dataset(cfg):
    import torch
    from hydra.utils import instantiate

    from fastwam.utils import misc

    with tempfile.TemporaryDirectory(prefix="pusht_preflight_") as tmp:
        misc.register_work_dir(tmp)  # dataset_stats.json goes here, not into ./runs
        ds = instantiate(cfg.data.train)
        print(f"    dataset: {len(ds)} samples")
        sample = ds[0]

        video = sample["video"]
        assert tuple(video.shape) == (3, 9, 224, 448), video.shape
        # tolerance: (2/255)*x - 1 arithmetic can land a float epsilon above 1.0
        assert video.min() >= -1.0 - 1e-5 and video.max() <= 1.0 + 1e-5
        assert tuple(sample["action"].shape) == (32, 7), sample["action"].shape
        assert tuple(sample["proprio"].shape) == (32, 8), sample["proprio"].shape
        assert sample["prompt"].endswith(TASK_STR), sample["prompt"]
        assert tuple(sample["context"].shape) == (128, 4096), sample["context"].shape
        print("[3/4] dataset contract OK (video 3x9x224x448, action 32x7, "
              "proprio 32x8, context 128x4096, prompt matches)")

        # ---- degenerate-channel contract -----------------------------------
        n = len(ds)
        idxs = range(0, n, max(1, n // 60))
        alo = torch.full((7,), float("inf"))
        ahi = torch.full((7,), float("-inf"))
        plo = torch.full((8,), float("inf"))
        phi = torch.full((8,), float("-inf"))
        TOL = 1e-4  # normalizer's range_tol; bounds an ignored dim's output
        nan_hits = []
        for i in idxs:
            s = ds[i]
            a, p = s["action"], s["proprio"]
            if not torch.isfinite(a).all():
                nan_hits.append(f"action sample {i}")
            if not torch.isfinite(p).all():
                nan_hits.append(f"proprio sample {i}")
            pad = s["action_is_pad"]
            a_valid = a[~pad] if (~pad).any() else a
            alo = torch.minimum(alo, a_valid.min(0).values)
            ahi = torch.maximum(ahi, a_valid.max(0).values)
            plo = torch.minimum(plo, p.min(0).values)
            phi = torch.maximum(phi, p.max(0).values)
        assert not nan_hits, f"NON-FINITE normalized values — the range_tol guard did NOT fire: {nan_hits[:5]}"

        problems = []
        print("    normalized action ranges:")
        for i, name, is_const in ACTION_DIMS:
            lo, hi = alo[i].item(), ahi[i].item()
            tag = "const" if is_const else "live "
            print(f"      [{i}] {name:8s} {tag} [{lo:+.4f}, {hi:+.4f}]")
            if is_const:
                # The guard sets scale=1.0 and offset=-min, so an ignored dim
                # normalizes to (x - min): identically 0 only when the raw
                # channel is EXACTLY constant. The guarantee that matters is
                # that it stays finite and bounded by its raw range (< TOL),
                # i.e. the guard fired instead of dividing by ~0.
                if not (abs(lo) <= TOL and abs(hi) <= TOL):
                    problems.append(f"action[{i}] {name} should stay within +/-{TOL} after "
                                    f"the range_tol guard, got [{lo}, {hi}]")
            elif hi - lo < 1.0:
                problems.append(f"action[{i}] {name} should span the normalized range, got [{lo}, {hi}]")
        print("    normalized proprio ranges:")
        for i, name, is_const in PROPRIO_DIMS:
            lo, hi = plo[i].item(), phi[i].item()
            tag = "const" if is_const else "live "
            print(f"      [{i}] {name:8s} {tag} [{lo:+.4f}, {hi:+.4f}]")
            if is_const:
                if not (abs(lo) <= TOL and abs(hi) <= TOL):
                    problems.append(f"proprio[{i}] {name} should stay within +/-{TOL} after "
                                    f"the range_tol guard, got [{lo}, {hi}]")
            elif hi - lo < 0.5:
                problems.append(f"proprio[{i}] {name} looks near-constant, got [{lo}, {hi}]")
        assert not problems, "degenerate-channel contract violated:\n  " + "\n  ".join(problems)
    print("[4/4] degenerate-channel contract OK: all 6 constant channels are "
          "finite and within +/-1e-4; all 9 live channels vary")


def main():
    ap = argparse.ArgumentParser(description=__doc__)
    ap.add_argument("--task", required=True, choices=sorted(EXPECT))
    ap.add_argument("--dataset-dir", default=SHARED_DATASET)
    args = ap.parse_args()

    exp = EXPECT[args.task]
    print(f"=== preflight: {args.task} (dataset {args.dataset_dir}) ===")
    cfg = compose_cfg(args.task, exp, args.dataset_dir)
    check_text_cache(cfg)
    check_dataset(cfg)
    print(f"=== {args.task}: ALL CHECKS PASSED ===")


if __name__ == "__main__":
    main()