mirror of
https://github.com/storytold/storyteller-ml.git
synced 2026-10-09 00:09:55 +00:00
result filename
This commit is contained in:
@@ -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!')
|
||||
|
||||
@@ -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")
|
||||
|
||||
|
||||
Reference in New Issue
Block a user