| 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 |
|
|
| |
| 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() |
|
|