File size: 2,122 Bytes
0613ebe
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
#!/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()