pusht-flashwam / training_code /preflight.py
SleepMastger's picture
add model card, conditioning, and training-time processing
2622c40 verified
Raw History Blame Contribute Delete
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()