CoolFace
Apppublic

ALSv/self-forcing

sourceHugging Faceupdated 1y agoView on Hugging Face
0likes
train.py48 linesDownload Raw Back to root
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