ALSv/self-forcing
0
1import argparse2import os3from omegaconf import OmegaConf4import wandb5 6from trainer import DiffusionTrainer, GANTrainer, ODETrainer, ScoreDistillationTrainer7 8 9def main():10 parser = argparse.ArgumentParser()11 parser.add_argument("--config_path", type=str, required=True)12 parser.add_argument("--no_save", action="store_true")13 parser.add_argument("--no_visualize", action="store_true")14 parser.add_argument("--logdir", type=str, default="", help="Path to the directory to save logs")15 parser.add_argument("--wandb-save-dir", type=str, default="", help="Path to the directory to save wandb logs")16 parser.add_argument("--disable-wandb", action="store_true")17 18 args = parser.parse_args()19 20 config = OmegaConf.load(args.config_path)21 default_config = OmegaConf.load("configs/default_config.yaml")22 config = OmegaConf.merge(default_config, config)23 config.no_save = args.no_save24 config.no_visualize = args.no_visualize25 26 # get the filename of config_path27 config_name = os.path.basename(args.config_path).split(".")[0]28 config.config_name = config_name29 config.logdir = args.logdir30 config.wandb_save_dir = args.wandb_save_dir31 config.disable_wandb = args.disable_wandb32 33 if config.trainer == "diffusion":34 trainer = DiffusionTrainer(config)35 elif config.trainer == "gan":36 trainer = GANTrainer(config)37 elif config.trainer == "ode":38 trainer = ODETrainer(config)39 elif config.trainer == "score_distillation":40 trainer = ScoreDistillationTrainer(config)41 trainer.train()42 43 wandb.finish()44 45 46if __name__ == "__main__":47 main()48 