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