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