PascalLiu/FNeVR_demo
1
1import os2import gradio as gr3import yaml4from argparse import ArgumentParser5from tqdm import tqdm6 7import numpy as np8import imageio9from skimage.transform import resize10from skimage import img_as_ubyte11from scipy.spatial import ConvexHull12import torch13from sync_batchnorm import DataParallelWithCallback14import face_alignment15 16from modules.generator import OcclusionAwareGenerator_SPADE17from modules.keypoint_detector import KPDetector18 19 20def normalize_kp(kp_source, kp_driving, kp_driving_initial, adapt_movement_scale=False,21 use_relative_movement=False, use_relative_jacobian=False):22 kp_new = {k: v for k, v in kp_driving.items()}23 24 if adapt_movement_scale:25 source_area = ConvexHull(kp_source['value'][0].data.cpu().numpy()).volume26 driving_area = ConvexHull(kp_driving_initial['value'][0].data.cpu().numpy()).volume27 adapt_movement_scale = np.sqrt(source_area) / np.sqrt(driving_area)28 kp_new['value'] = kp_driving['value'] * adapt_movement_scale # for reenactment demo29 else:30 adapt_movement_scale = 131 32 if use_relative_movement:33 kp_value_diff = (kp_driving['value'] - kp_driving_initial['value'])34 kp_value_diff *= adapt_movement_scale35 kp_new['value'] = kp_value_diff + kp_source['value']36 37 if use_relative_jacobian:38 jacobian_diff = torch.matmul(kp_driving['jacobian'], torch.inverse(kp_driving_initial['jacobian']))39 kp_new['jacobian'] = torch.matmul(jacobian_diff, kp_source['jacobian'])40 41 return kp_new42 43 44def load_checkpoints(config_path, checkpoint_path, cpu=False):45 with open(config_path) as f:46 # config = yaml.load(f)47 config = yaml.load(f, Loader=yaml.FullLoader)48 49 generator = OcclusionAwareGenerator_SPADE(**config['model_params']['generator_params'],50 **config['model_params']['common_params'])51 if not cpu:52 generator.cuda()53 54 kp_detector = KPDetector(**config['model_params']['kp_detector_params'],55 **config['model_params']['common_params'])56 if not cpu:57 kp_detector.cuda()58 59 if cpu:60 checkpoint = torch.load(checkpoint_path, map_location=torch.device('cpu'))61 else:62 checkpoint = torch.load(checkpoint_path)63 64 generator.load_state_dict(checkpoint['generator'])65 kp_detector.load_state_dict(checkpoint['kp_detector'])66 67 if not cpu:68 generator = DataParallelWithCallback(generator)69 kp_detector = DataParallelWithCallback(kp_detector)70 71 generator.eval()72 kp_detector.eval()73 74 return generator, kp_detector75 76 77def make_animation(source_image, driving_video, generator, kp_detector, relative=True, adapt_movement_scale=True,78 cpu=False):79 with torch.no_grad():80 predictions = []81 source = torch.tensor(source_image[np.newaxis].astype(np.float32)).permute(0, 3, 1, 2)82 if not cpu:83 source = source.cuda()84 driving = torch.tensor(np.array(driving_video)[np.newaxis].astype(np.float32)).permute(0, 4, 1, 2, 3)85 kp_source = kp_detector(source)86 kp_driving_initial = kp_detector(driving[:, :, 0])87 88 for frame_idx in tqdm(range(driving.shape[2])):89 driving_frame = driving[:, :, frame_idx]90 if not cpu:91 driving_frame = driving_frame.cuda()92 kp_driving = kp_detector(driving_frame)93 kp_norm = normalize_kp(kp_source=kp_source, kp_driving=kp_driving,94 kp_driving_initial=kp_driving_initial, use_relative_movement=relative,95 use_relative_jacobian=relative, adapt_movement_scale=adapt_movement_scale)96 out = generator(source, kp_source=kp_source, kp_driving=kp_norm)97 98 predictions.append(np.transpose(out['prediction'].data.cpu().numpy(), [0, 2, 3, 1])[0])99 return predictions100 101 102def find_best_frame_func(source, driving, cpu=False):103 def normalize_kp_infunc(kp):104 kp = kp - kp.mean(axis=0, keepdims=True)105 area = ConvexHull(kp[:, :2]).volume106 area = np.sqrt(area)107 kp[:, :2] = kp[:, :2] / area108 return kp109 110 fa = face_alignment.FaceAlignment(face_alignment.LandmarksType._2D, flip_input=True,111 device='cpu' if cpu else 'cuda')112 kp_source = fa.get_landmarks(255 * source)[0]113 kp_source = normalize_kp_infunc(kp_source)114 norm = float('inf')115 frame_num = 0116 for i, image in tqdm(enumerate(driving)):117 kp_driving = fa.get_landmarks(255 * image)[0]118 kp_driving = normalize_kp_infunc(kp_driving)119 new_norm = (np.abs(kp_source - kp_driving) ** 2).sum()120 if new_norm < norm:121 norm = new_norm122 frame_num = i123 return frame_num124 125 126def drive_im(source_image, driving_image, adapt_scale):127 source_image = resize(source_image, (256, 256))[..., :3]128 driving_image = [resize(driving_image, (256, 256))[..., :3]]129 130 prediction = make_animation(source_image, driving_image, generator, kp_detector, relative=False,131 adapt_movement_scale=adapt_scale, cpu=cpu)132 return img_as_ubyte(prediction[0])133 134 135def drive_vi(source_image, driving_video, mode, find_best_frame, best_frame, relative, adapt_scale):136 reader = imageio.get_reader(driving_video)137 fps = reader.get_meta_data()['fps']138 driving_video = []139 try:140 for im in reader:141 driving_video.append(im)142 except RuntimeError:143 pass144 reader.close()145 146 147 if mode == 'reconstruction':148 source_image = driving_video[0]149 150 source_image = resize(source_image, (256, 256))[..., :3]151 driving_video = [resize(frame, (256, 256))[..., :3] for frame in driving_video]152 153 i = 0154 if find_best_frame != "specific ref frame" or best_frame > 0:155 i = best_frame if find_best_frame == "specific ref frame" else find_best_frame_func(source_image, driving_video, cpu=cpu)156 print("Best frame: " + str(i))157 driving_forward = driving_video[i:]158 driving_backward = driving_video[:(i + 1)][::-1]159 predictions_forward = make_animation(source_image, driving_forward, generator, kp_detector,160 relative=relative, adapt_movement_scale=adapt_scale, cpu=cpu)161 predictions_backward = make_animation(source_image, driving_backward, generator, kp_detector,162 relative=relative, adapt_movement_scale=adapt_scale, cpu=cpu)163 predictions = predictions_backward[::-1] + predictions_forward[1:]164 else:165 predictions = make_animation(source_image, driving_video, generator, kp_detector, relative=relative,166 adapt_movement_scale=adapt_scale, cpu=cpu)167 result_video_path = "result_video.mp4"168 imageio.mimsave(result_video_path, [img_as_ubyte(frame) for frame in predictions], fps=fps)169 return result_video_path, i170 171 172config = "config/vox-256.yaml"173checkpoint = "00000099-checkpoint.pth.tar"174cpu = True # decided by the deploying environment175 176description = "We propose a Face Neural Volume Rendering (FNeVR) network for more realistic face animation, by taking the merits of 2D motion warping on facial expression transformation and 3D volume rendering on high-quality image synthesis in a unified framework.<br>[Paper](https://arxiv.org/abs/2209.10340) and [Code](https://github.com/zengbohan0217/FNeVR)"177im_description = "We can animate a face portrait by a single image in this tab.<br>Please input the origin face and the driving face which provides pose and expression information, then we can obtain the virtual generated face.<br>We can select \"adaptive scale\" parameter for better optic flow estimation using adaptive movement scale based on convex hull of keypoints."178vi_description = "We can animate a face portrait by a video in this tab.<br>Please input the origin face and the driving video which provides pose and expression information, then we can obtain the virtual generated video.<br>Please select inference mode (reenactment for different identities and reconstruction for the same identities).<br>We can select \"relative motion\" paramter to use relative keypoint coordinates for preserving global object geometry, select \"adaptive scale\" parameter for better optic flow estimation using adaptive movement scale based on convex hull of keypoints, and select \"find best ref frame\" parameter to generate video from the frame that is the most alligned with source image."179acknowledgements = "This work was supported by “the Fundamental Research Funds for the Central Universities”, and the National Natural Science Foundation of China under Grant 62076016, Beijing Natural Science Foundation-Xiaomi Innovation Joint Fund L223024. Besides, we gratefully acknowledge the support of [MindSpore](https://www.mindspore.cn), CANN (Compute Architecture for Neural Networks) and Ascend AI processor used for this research.<br>Our FNeVR implementation is inspired by [FOMM](https://github.com/AliaksandrSiarohin/first-order-model) and [DECA](https://github.com/YadiraF/DECA). We appreciate the authors of these papers for making their codes available to the public."180 181generator, kp_detector = load_checkpoints(config_path=config, checkpoint_path=checkpoint, cpu=cpu)182 183# iface = gr.Interface(fn=drive_im, 184# inputs=[gr.Image(label="Origin face"), 185# gr.Image(label="Driving face"), 186# gr.CheckboxGroup(label="adapt scale")], 187# outputs=gr.Image(label="Generated face"), examples=[["sup-mat/source.png"], ["sup-mat/driving.png"]], 188# title="Demostration of FNeVR", description=description)189 190with gr.Blocks(title="Demostration of FNeVR") as demo:191 gr.Markdown("# <center> Demostration of FNeVR")192 gr.Markdown(description)193 194 with gr.Tab("Driving by image"):195 gr.Markdown(im_description)196 197 with gr.Row():198 with gr.Column():199 gr.Markdown("#### Inputs")200 inp2 = gr.Image(label="Driving face")201 inp1 = gr.Image(label="Origin face")202 203 gr.Markdown("#### Parameter")204 inp3 = gr.Checkbox(value=True, label="adaptive scale")205 206 btn1 = gr.Button(value="Animate")207 with gr.Column():208 gr.Markdown("#### Output")209 outp = gr.Image(label="Generated face")210 211 with gr.Row():212 with gr.Column():213 btn2 = gr.Button(value="Reset")214 with gr.Column():215 btn3 = gr.Button(value="Cancel")216 217 gr.Examples([["sup-mat/driving.png", "sup-mat/source.png"]], [inp2, inp1])218 219 def reset_output():220 return outp.update(value=None)221 222 def reset_all():223 return inp1.update(value=None), inp2.update(value=None), inp3.update(value=True), outp.update(value=None)224 225 run = btn1.click(fn=drive_im, inputs=[inp1, inp2, inp3], outputs=outp)226 btn2.click(fn=reset_all, outputs=[inp1, inp2, inp3, outp])227 btn3.click(fn=reset_output, outputs=[outp], cancels=[run])228 with gr.Tab("Driving by video"):229 gr.Markdown(vi_description)230 231 with gr.Row():232 with gr.Column():233 gr.Markdown("#### Inputs")234 inp2 = gr.Video(label="Driving video")235 inp1 = gr.Image(label="Origin face")236 237 gr.Markdown("#### Parameters")238 inp3 = gr.Radio(choices=["reenactment", "reconstruction"], value="reenactment", label="mode (if \"reconstruction\" selected, origin face is the first frame of driving video)")239 inp6 = gr.Checkbox(value=True, label="relative motion")240 inp7 = gr.Checkbox(value=True, label="adaptive scale")241 inp4 = gr.Radio(choices=["find best ref frame (more time consumed)", "specific ref frame"], value="find best ref frame (more time consumed)", label="set ref frame (used by relative motion and adaptive scale)")242 inp5 = gr.Number(label="specific ref frame (default: 0)", value=0, precision=0, visible=False)243 244 def reset_ref(inp4):245 return inp5.update(visible=True) if inp4 == "specific ref frame" else inp5.update(value=0, visible=False)246 247 inp4.change(fn=reset_ref, inputs=inp4, outputs=inp5)248 249 btn1 = gr.Button(value="Animate")250 with gr.Column():251 gr.Markdown("#### Output")252 outp1 = gr.Video(label="Generated video")253 outp2 = gr.Number(label="Ref frame", value=0, precision=0)254 255 # file = gr.File(value="result_video.mp4", visible=False)256 257 with gr.Row():258 with gr.Column():259 btn2 = gr.Button(value="Reset")260 with gr.Column():261 btn3 = gr.Button(value="Cancel")262 263 gr.Examples([["sup-mat/driving.mp4", "sup-mat/source_for_video.png", "specific ref frame", 53]], [inp2, inp1, inp4, inp5])264 265 def reset_output():266 return outp1.update(value=None), outp2.update(value=0)267 268 def reset_all():269 return inp1.update(value=None), inp2.update(value=None), inp3.update(value="reenactment"), inp4.update(value="find best ref frame (more time consumed)"), inp5.update(value=0), inp6.update(value=True), inp7.update(value=True), outp1.update(value=None), outp2.update(value=0)270 271 run = btn1.click(fn=drive_vi, inputs=[inp1, inp2, inp3, inp4, inp5, inp6, inp7], outputs=[outp1, outp2])272 btn2.click(fn=reset_all, outputs=[inp1, inp2, inp3, inp4, inp5, inp6, inp7, outp1, outp2])273 btn3.click(fn=reset_output, outputs=[outp1, outp2], cancels=[run])274 with gr.Tab("Real time animation"):275 gr.Markdown("Real time animation is coming soon.")276 277 gr.Markdown("## Acknowledgements")278 gr.Markdown(acknowledgements)279 280demo.queue()281demo.launch()282 