Self-Forcing / train.py
Cccccz's picture
Add files using upload-large-folder tool
d1d122e verified
Raw
History Blame Contribute Delete
2.77 kB
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()