result filename

This commit is contained in:
Brandon Thomas
2025-02-06 00:39:20 -05:00
parent d75e712d6a
commit a7179c873d
2 changed files with 29 additions and 6 deletions
@@ -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,
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)
save_video_multiple_conditions_not_gif_horizontal_1col(local_path_1col, 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!')
+16
View File
@@ -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")