(args)
| 7 | |
| 8 | @hydra.main(config_path='.', config_name='config_crema', version_base='1.1') |
| 9 | def main(args): |
| 10 | device = 'cuda' if args.gpu and torch.cuda.is_available() else 'cpu' |
| 11 | |
| 12 | print('Loading model...') |
| 13 | unet = torch.jit.load(args.checkpoint) |
| 14 | diffusion = Diffusion(unet, device, **args.diffusion).to(device) |
| 15 | diffusion.space(args.inference_steps) |
| 16 | |
| 17 | id_frame = get_id_frame(args.id_frame, random=args.id_frame_random, resize=args.diffusion.image_size).to(device) |
| 18 | audio, audio_emb = get_audio_emb(args.audio, args.encoder_checkpoint, device) |
| 19 | |
| 20 | samples = diffusion.sample(id_frame, audio_emb.unsqueeze(0), **args.unet) |
| 21 | |
| 22 | save_video(args.output, samples, audio=audio, fps=25, audio_rate=16000) |
| 23 | print(f'Results saved at {args.output}') |
| 24 | |
| 25 | |
| 26 | if __name__ == '__main__': |
no test coverage detected