()
| 6 | DEFAULT_TRAINING_CONFIG = os.path.join(os.getcwd(), 'configs', 'training.yaml') |
| 7 | |
| 8 | def main(): |
| 9 | # load training config |
| 10 | argv = sys.argv |
| 11 | try: |
| 12 | idx = argv.index('--config') |
| 13 | training_cfg_path = argv[idx + 1] if idx + 1 < len(argv) else DEFAULT_TRAINING_CONFIG |
| 14 | except ValueError: |
| 15 | print(f"No training config path provided, using default: {DEFAULT_TRAINING_CONFIG}") |
| 16 | training_cfg_path = DEFAULT_TRAINING_CONFIG |
| 17 | |
| 18 | train_config: TrainingConfig = load_config(TrainingConfig, default_config_path=training_cfg_path) |
| 19 | |
| 20 | # torchrun if distributed, else python |
| 21 | if train_config.nproc_per_node > 1: |
| 22 | launcher = [ |
| 23 | "torchrun", |
| 24 | "--nproc_per_node", str(train_config.nproc_per_node), |
| 25 | ] |
| 26 | if train_config.standalone: |
| 27 | launcher += ["--standalone"] |
| 28 | else: |
| 29 | launcher = [sys.executable] |
| 30 | |
| 31 | # top-level run root and export child processes can use |
| 32 | run_root, run_name = prepare_pipeline_run_root(base_cwd=os.getcwd()) |
| 33 | os.environ['NG_RUN_ROOT_DIR'] = run_root |
| 34 | |
| 35 | if train_config.run_video_tokenizer: |
| 36 | v_cmd = launcher + [ |
| 37 | "scripts/train_video_tokenizer.py", |
| 38 | "--config", train_config.video_tokenizer_config, |
| 39 | "--training_config", training_cfg_path, |
| 40 | ] |
| 41 | if not run_command(v_cmd, "Video Tokenizer Training"): |
| 42 | return |
| 43 | |
| 44 | if train_config.run_latent_actions: |
| 45 | latent_actions_cmd = launcher + [ |
| 46 | "scripts/train_latent_actions.py", |
| 47 | "--config", train_config.latent_actions_config, |
| 48 | "--training_config", training_cfg_path, |
| 49 | ] |
| 50 | if not run_command(latent_actions_cmd, "Latent Actions Training"): |
| 51 | return |
| 52 | |
| 53 | # need to get above checkpoints and pass in to dynamics |
| 54 | video_tokenizer_checkpoint = find_latest_checkpoint(".", "video_tokenizer") |
| 55 | latent_actions_checkpoint = find_latest_checkpoint(".", "latent_actions") |
| 56 | |
| 57 | if train_config.run_dynamics: |
| 58 | dyn_cmd = launcher + [ |
| 59 | "scripts/train_dynamics.py", |
| 60 | "--config", train_config.dynamics_config, |
| 61 | "--training_config", training_cfg_path, |
| 62 | f"video_tokenizer_path={video_tokenizer_checkpoint}", |
| 63 | f"latent_actions_path={latent_actions_checkpoint}", |
| 64 | ] |
| 65 | if not run_command(dyn_cmd, "Dynamics Model Training"): |
no test coverage detected