From a7179c873d25ccdcb54d82e01aeeed492aff73cf Mon Sep 17 00:00:00 2001 From: Brandon Thomas Date: Thu, 6 Feb 2025 00:39:20 -0500 Subject: [PATCH] result filename --- .../animatex/inference_animate_x_entrance.py | 19 +++++++++++++------ animation/animate-x/inference_cli.py | 16 ++++++++++++++++ 2 files changed, 29 insertions(+), 6 deletions(-) diff --git a/animation/animate-x/animatex/inference_animate_x_entrance.py b/animation/animate-x/animatex/inference_animate_x_entrance.py index ba9cea2..96f8647 100644 --- a/animation/animate-x/animatex/inference_animate_x_entrance.py +++ b/animation/animate-x/animatex/inference_animate_x_entrance.py @@ -606,6 +606,7 @@ def worker(gpu, cfg, cfg_update): file_name = f'{cap_name}_{pose_name}_{name}_rank_{cfg.world_size:02d}_{cfg.rank:02d}_{idx:02d}_{cfg.resolution[1]}x{cfg.resolution[0]}.mp4' local_path = os.path.join(cfg.log_dir, f'{file_name}') local_path_1col = os.path.join(cfg.log_dir, f'{file_name[:-4]}_results_1col.mp4') + os.makedirs(os.path.dirname(local_path), exist_ok=True) captions = "human" del model_kwargs_one_vis[0][list(model_kwargs_one_vis[0].keys())[0]] @@ -614,14 +615,20 @@ def worker(gpu, cfg, cfg_update): del model_kwargs_one_vis[0]["pose_embeddings"] del model_kwargs_one_vis[1]["pose_embeddings"] + result_filename = local_path_1col + + if cfg.result_filename: + result_filename = cfg.result_filename - - save_video_multiple_conditions_not_gif_horizontal_3col(local_path, video_data.cpu(), model_kwargs_one_vis, misc_backups, - cfg.mean, cfg.std, nrow=1, save_fps=cfg.save_fps) - - save_video_multiple_conditions_not_gif_horizontal_1col(local_path_1col, video_data.cpu(), model_kwargs_one_vis, misc_backups, + save_video_multiple_conditions_not_gif_horizontal_1col(result_filename, video_data.cpu(), model_kwargs_one_vis, misc_backups, cfg.mean, cfg.std, nrow=1, save_fps=cfg.save_fps) - logging.info(f'video saved in {local_path}!') + + if cfg.generate_comparison_video: + save_video_multiple_conditions_not_gif_horizontal_3col(local_path, video_data.cpu(), model_kwargs_one_vis, misc_backups, + cfg.mean, cfg.std, nrow=1, save_fps=cfg.save_fps) + + + logging.info(f'video saved in {result_filename}!') logging.info('Congratulations! The inference is completed!') diff --git a/animation/animate-x/inference_cli.py b/animation/animate-x/inference_cli.py index 8311ada..5b3b014 100755 --- a/animation/animate-x/inference_cli.py +++ b/animation/animate-x/inference_cli.py @@ -41,6 +41,12 @@ def parse_args(): default=None, required=True ) + parser.add_argument( + "--result_filename", + type=str, + default=None, + required=False + ) parser.add_argument( "--max_frames", type=int, @@ -71,6 +77,12 @@ def parse_args(): default=1, required=False ) + parser.add_argument( + "--generate_comparison_video", + default=False, + action='store_true', + required=False + ) # TODO(bt): Means to control output filename @@ -113,6 +125,9 @@ def main(): args.seed, ]] + if args.result_filename: + cfg_update.cfg_dict['result_filename'] = args.result_filename + # NB(bt): I think this controls extra inference rounds. cfg_update.cfg_dict['round'] = args.round @@ -128,6 +143,7 @@ def main(): cfg_update.cfg_dict['embedder']['pretrained'] = os.path.join(model_dir, 'open_clip_pytorch_model.bin') cfg_update.cfg_dict['auto_encoder']['pretrained'] = os.path.join(model_dir, 'v2-1_512-ema-pruned.ckpt') + cfg_update.cfg_dict['generate_comparison_video'] = args.generate_comparison_video print("Configurations:\n\n", cfg_update.cfg_dict, "\n\n")