fruit-picking-lingbot / training_code /train_entrypoint.py
SleepMastger's picture
add model card, conditioning, and training-time processing
0613ebe verified
Raw History Blame Contribute Delete
2.12 kB
#!/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()