#!/usr/bin/env python """Train LingBot-VA on George's real-world lift_new Franka dataset. This repo carries its own independent copy of the lingbot-va code (./lingbot-va -- a snapshot of vla_or_wam/lingbot-va's working tree, `git remote` removed, no relation to vla_or_wam from here on) and its own .venv. Rather than forking wan_va/train.py itself, this script imports wan_va.train from that local clone as a library and registers our own config (lift_new_configs/va_lift_new_cfg.py) into its VA_CONFIGS dict at runtime. vla_or_wam is never read from or written to by this repo. Usage (single process, for a quick check): .venv/bin/python train_lift_new.py --config-name lift_new_scratch Usage (multi-GPU, what slurm/train_lift_new.sbatch actually runs): PYTORCH_CUDA_ALLOC_CONF="expandable_segments:True" \ .venv/bin/python -m torch.distributed.run --nproc_per_node "$NGPU" \ train_lift_new.py --config-name lift_new_scratch """ import os import sys REPO_ROOT = os.path.dirname(os.path.abspath(__file__)) LINGBOT_VA_ROOT = os.environ.get( "LINGBOT_VA_ROOT", os.path.join(REPO_ROOT, "lingbot-va") ) # Our own independent lingbot-va clone, used as a library: needed for `import wan_va...`. sys.path.insert(0, LINGBOT_VA_ROOT) # our own config module (named lift_new_configs, NOT `configs` -- lingbot-va/wan_va/train.py # does `from configs import VA_CONFIGS` relying on wan_va/'s own dir being on sys.path, so a # same-named top-level `configs` package here would shadow it and break that import). sys.path.insert(0, REPO_ROOT) from wan_va import train as wan_va_train # noqa: E402 (import after sys.path setup) from lift_new_configs.va_lift_new_cfg import LIFT_NEW_CONFIGS # noqa: E402 from lift_new_configs.va_place_cube_bowl_cfg import PLACE_CUBE_BOWL_CONFIGS # noqa: E402 from lift_new_configs.va_fruit_pick_cfg import FRUIT_PICK_CONFIGS # noqa: E402 wan_va_train.VA_CONFIGS.update(LIFT_NEW_CONFIGS) wan_va_train.VA_CONFIGS.update(PLACE_CUBE_BOWL_CONFIGS) wan_va_train.VA_CONFIGS.update(FRUIT_PICK_CONFIGS) if __name__ == "__main__": wan_va_train.init_logger() wan_va_train.main()