Download training_code/preflight.py from SleepMastger/pusht-flashwam: direct link, hf CLI and curl.
- Browser
- Download file 10.4 kB
-
https://huggingface.co/SleepMastger/pusht-flashwam/resolve/main/training_code/preflight.py
- Command line
-
hf download hf://SleepMastger/pusht-flashwam/training_code/preflight.py
-
curl -L -o preflight.py https://huggingface.co/SleepMastger/pusht-flashwam/resolve/main/training_code/preflight.py
10.4 kB
| """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() | |