Team Ai
Apppublic

PascalLiu/FNeVR_demo

sourceHugging Faceapache-2.0updated 4y agoView on Hugging Face
1likes
app.py282 linesDownload Raw Back to root
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