File size: 2,774 Bytes
d1d122e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
import argparse
import os
import sys
from omegaconf import OmegaConf


def main():
    parser = argparse.ArgumentParser()
    parser.add_argument("--config_path", type=str, required=True)
    parser.add_argument("--no_save", action="store_true")
    parser.add_argument("--no_visualize", action="store_true")
    parser.add_argument("--logdir", type=str, default="", help="Path to the directory to save logs")
    parser.add_argument("--wandb-save-dir", type=str, default="", help="Path to the directory to save wandb logs")
    parser.add_argument("--disable-wandb", action="store_true")
    parser.add_argument(
        "--config_override",
        action="append",
        default=[],
        help="OmegaConf dot-list override, for example predictor_v4.batch_size=16",
    )

    args = parser.parse_args()

    config = OmegaConf.load(args.config_path)
    default_config = OmegaConf.load("configs/default_config.yaml")
    config = OmegaConf.merge(default_config, config)
    if args.config_override:
        config = OmegaConf.merge(
            config,
            OmegaConf.from_dotlist(args.config_override),
        )
    config.no_save = args.no_save
    config.no_visualize = args.no_visualize

    # get the filename of config_path
    config_name = os.path.basename(args.config_path).split(".")[0]
    config.config_name = config_name
    config.logdir = args.logdir
    config.wandb_save_dir = args.wandb_save_dir
    config.disable_wandb = args.disable_wandb

    if config.trainer == "diffusion":
        from trainer.diffusion import Trainer as DiffusionTrainer
        trainer = DiffusionTrainer(config)
    elif config.trainer == "gan":
        from trainer.gan import Trainer as GANTrainer
        trainer = GANTrainer(config)
    elif config.trainer == "ode":
        from trainer.ode import Trainer as ODETrainer
        trainer = ODETrainer(config)
    elif config.trainer == "score_distillation":
        from trainer.distillation import Trainer as ScoreDistillationTrainer
        trainer = ScoreDistillationTrainer(config)
    elif config.trainer == "predictor_v4":
        from trainer.predictor_v4 import Trainer as PredictorV4Trainer
        trainer = PredictorV4Trainer(config)
    elif config.trainer == "predictor_v4_rollout":
        from trainer.predictor_v4_rollout import Trainer as PredictorV4RolloutTrainer
        trainer = PredictorV4RolloutTrainer(config)
    elif config.trainer == "predictor_v4_dmd":
        from trainer.predictor_v4_dmd import Trainer as PredictorV4DMDTrainer
        trainer = PredictorV4DMDTrainer(config)
    else:
        raise ValueError(f"Unknown trainer: {config.trainer}")
    trainer.train()

    if "wandb" in sys.modules:
        sys.modules["wandb"].finish()


if __name__ == "__main__":
    main()