MCPcopy Create free account
hub / github.com/AlmondGod/tinyworlds / main

Function main

scripts/full_train.py:8–72  ·  view source on GitHub ↗
()

Source from the content-addressed store, hash-verified

6DEFAULT_TRAINING_CONFIG = os.path.join(os.getcwd(), 'configs', 'training.yaml')
7
8def 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"):

Callers 1

full_train.pyFile · 0.70

Calls 4

load_configFunction · 0.90
run_commandFunction · 0.90
find_latest_checkpointFunction · 0.90

Tested by

no test coverage detected