CoolFace
Apppublic

batkovdev/i2v-vtk

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
app-i2v.py402 linesDownload Raw Back to root
1 2import os3import json4import torch5import random6 7import gradio as gr8from glob import glob9from omegaconf import OmegaConf10from datetime import datetime11from safetensors import safe_open12 13from diffusers import AutoencoderKL14from diffusers.utils.import_utils import is_xformers_available15from transformers import CLIPTextModel, CLIPTokenizer16 17from animatelcm.scheduler.lcm_scheduler import LCMScheduler18from animatelcm.models.unet import UNet3DConditionModel19from animatelcm.pipelines.pipeline_animation import AnimationPipeline20from animatelcm.utils.util import save_videos_grid21from animatelcm.utils.convert_from_ckpt import convert_ldm_unet_checkpoint, convert_ldm_clip_checkpoint, convert_ldm_vae_checkpoint22from animatelcm.utils.convert_lora_safetensor_to_diffusers import convert_lora23from animatelcm.utils.lcm_utils import convert_lcm_lora24import copy25 26sample_idx = 027scheduler_dict = {28    "LCM": LCMScheduler,29}30 31css = """32.toolbutton {33    margin-buttom: 0em 0em 0em 0em;34    max-width: 2.5em;35    min-width: 2.5em !important;36    height: 2.5em;37}38"""39 40 41class AnimateController:42    def __init__(self):43 44        # config dirs45        self.basedir = os.getcwd()46        self.stable_diffusion_dir = os.path.join(47            self.basedir, "models", "StableDiffusion")48        self.motion_module_dir = os.path.join(49            self.basedir, "models", "Motion_Module")50        self.personalized_model_dir = os.path.join(51            self.basedir, "models", "DreamBooth_LoRA")52        self.savedir = os.path.join(53            self.basedir, "samples", datetime.now().strftime("Gradio-%Y-%m-%dT%H-%M-%S"))54        self.savedir_sample = os.path.join(self.savedir, "sample")55        self.lcm_lora_path = "models/LCM_LoRA/AnimateLCM_sd15_i2v_lora.safetensors"56        os.makedirs(self.savedir, exist_ok=True)57 58        self.stable_diffusion_list = []59        self.motion_module_list = []60        self.personalized_model_list = []61 62        self.refresh_stable_diffusion()63        self.refresh_motion_module()64        self.refresh_personalized_model()65 66        # config models67        self.tokenizer = None68        self.text_encoder = None69        self.vae = None70        self.unet = None71        self.pipeline = None72        self.lora_model_state_dict = {}73 74        self.inference_config = OmegaConf.load("configs/inference-i2v.yaml")75 76    def refresh_stable_diffusion(self):77        self.stable_diffusion_list = glob(78            os.path.join(self.stable_diffusion_dir, "*/"))79 80    def refresh_motion_module(self):81        motion_module_list = glob(os.path.join(82            self.motion_module_dir, "*.ckpt"))83        self.motion_module_list = [84            os.path.basename(p) for p in motion_module_list]85 86    def refresh_personalized_model(self):87        personalized_model_list = glob(os.path.join(88            self.personalized_model_dir, "*.safetensors"))89        self.personalized_model_list = [90            os.path.basename(p) for p in personalized_model_list]91 92    def update_stable_diffusion(self, stable_diffusion_dropdown):93        stable_diffusion_dropdown = os.path.join(self.stable_diffusion_dir,stable_diffusion_dropdown)94        self.tokenizer = CLIPTokenizer.from_pretrained(95            stable_diffusion_dropdown, subfolder="tokenizer")96        self.text_encoder = CLIPTextModel.from_pretrained(97            stable_diffusion_dropdown, subfolder="text_encoder").cuda()98        self.vae = AutoencoderKL.from_pretrained(99            stable_diffusion_dropdown, subfolder="vae").cuda()100        self.unet = UNet3DConditionModel.from_pretrained_2d(101            stable_diffusion_dropdown, subfolder="unet", unet_additional_kwargs=OmegaConf.to_container(self.inference_config.unet_additional_kwargs)).cuda()102        return gr.Dropdown.update()103 104    def update_motion_module(self, motion_module_dropdown):105        if self.unet is None:106            gr.Info(f"Please select a pretrained model path.")107            return gr.Dropdown.update(value=None)108        else:109            motion_module_dropdown = os.path.join(110                self.motion_module_dir, motion_module_dropdown)111            motion_module_state_dict = torch.load(112                motion_module_dropdown, map_location="cpu")113            missing, unexpected = self.unet.load_state_dict(114                motion_module_state_dict, strict=False)115            assert len(unexpected) == 0116            return gr.Dropdown.update()117 118    def update_base_model(self, base_model_dropdown):119        if self.unet is None:120            gr.Info(f"Please select a pretrained model path.")121            return gr.Dropdown.update(value=None)122        else:123            base_model_dropdown = os.path.join(124                self.personalized_model_dir, base_model_dropdown)125            base_model_state_dict = {}126            with safe_open(base_model_dropdown, framework="pt", device="cpu") as f:127                for key in f.keys():128                    base_model_state_dict[key] = f.get_tensor(key)129 130            converted_vae_checkpoint = convert_ldm_vae_checkpoint(131                base_model_state_dict, self.vae.config)132            self.vae.load_state_dict(converted_vae_checkpoint)133 134            converted_unet_checkpoint = convert_ldm_unet_checkpoint(135                base_model_state_dict, self.unet.config)136            self.unet.load_state_dict(converted_unet_checkpoint, strict=False)137 138            self.text_encoder = convert_ldm_clip_checkpoint(base_model_state_dict)139            return gr.Dropdown.update()140 141    def update_lora_model(self, lora_model_dropdown):142        lora_model_dropdown = os.path.join(143            self.personalized_model_dir, lora_model_dropdown)144        self.lora_model_state_dict = {}145        if lora_model_dropdown == "none":146            pass147        else:148            with safe_open(lora_model_dropdown, framework="pt", device="cpu") as f:149                for key in f.keys():150                    self.lora_model_state_dict[key] = f.get_tensor(key)151        return gr.Dropdown.update()152 153    def animate(154        self,155        lora_alpha_slider,156        spatial_lora_slider,157        prompt_textbox,158        negative_prompt_textbox,159        sampler_dropdown,160        sample_step_slider,161        width_slider,162        length_slider,163        height_slider,164        cfg_scale_slider,165        seed_textbox,166        image_upload,167        beta_end_slider,168        motion_scale_slider,169    ):170        print(image_upload)171 172        if is_xformers_available():173            self.unet.enable_xformers_memory_efficient_attention()174 175        print(self.inference_config.noise_scheduler_kwargs["beta_end"])176        self.inference_config.noise_scheduler_kwargs["beta_end"] = beta_end_slider177        178        self.unet.img_encoder.motion_scale = motion_scale_slider179        pipeline = AnimationPipeline(180            vae=self.vae, text_encoder=self.text_encoder, tokenizer=self.tokenizer, unet=self.unet,181            scheduler=scheduler_dict[sampler_dropdown](182                **OmegaConf.to_container(self.inference_config.noise_scheduler_kwargs))183        ).to("cuda")184 185        if self.lora_model_state_dict != {}:186            pipeline = convert_lora(187                pipeline, self.lora_model_state_dict, alpha=lora_alpha_slider)188 189        pipeline.unet = convert_lcm_lora(copy.deepcopy(190            self.unet), self.lcm_lora_path, spatial_lora_slider)191 192        pipeline.to("cuda")193 194        if seed_textbox != -1 and seed_textbox != "":195            torch.manual_seed(int(seed_textbox))196        else:197            torch.seed()198        seed = torch.initial_seed()199 200        with torch.autocast("cuda",torch.float16):201            sample = pipeline(202                prompt_textbox,203                negative_prompt=negative_prompt_textbox,204                num_inference_steps=sample_step_slider,205                guidance_scale=cfg_scale_slider,206                width=width_slider,207                height=height_slider,208                video_length=length_slider,209                image_path=image_upload210            ).videos211 212        save_sample_path = os.path.join(213            self.savedir_sample, f"{sample_idx}.mp4")214        save_videos_grid(sample, save_sample_path)215 216        sample_config = {217            "prompt": prompt_textbox,218            "n_prompt": negative_prompt_textbox,219            "sampler": sampler_dropdown,220            "num_inference_steps": sample_step_slider,221            "guidance_scale": cfg_scale_slider,222            "width": width_slider,223            "height": height_slider,224            "video_length": length_slider,225            "seed": seed,226        }227        json_str = json.dumps(sample_config, indent=4)228        with open(os.path.join(self.savedir, "logs.json"), "a") as f:229            f.write(json_str)230            f.write("\n\n")231        return gr.Video.update(value=save_sample_path)232 233 234controller = AnimateController()235 236controller.update_stable_diffusion("stable-diffusion-v1-5")237controller.update_motion_module("AnimateLCM_sd15_i2v.ckpt")238controller.update_base_model("realistic1.safetensors")239 240 241def ui():242    with gr.Blocks(css=css) as demo:243        gr.Markdown(244            """245            # [AnimateLCM: Accelerating the Animation of Personalized Diffusion Models and Adapters with Decoupled Consistency Learning](https://arxiv.org/abs/2402.00769)246            Fu-Yun Wang, Zhaoyang Huang (*Corresponding Author), Xiaoyu Shi, Weikang Bian, Guanglu Song, Yu Liu, Hongsheng Li (*Corresponding Author)<br>247            [arXiv Report](https://arxiv.org/abs/2402.00769) | [Project Page](https://animatelcm.github.io/) | [Github](https://github.com/G-U-N/AnimateLCM) | [Civitai](https://civitai.com/models/290375/animatelcm-fast-video-generation) | [Replicate](https://replicate.com/camenduru/animate-lcm)248            """249            250            '''251            Important Notes: 252            1. The generation speed is around few seconds. 253            2. Increase the sampling step for better sample quality.254            '''255        )256        with gr.Column(variant="panel"):257            with gr.Row():258                image_upload = gr.Image(label="Upload Image", tool="select", type="filepath")259 260                base_model_dropdown = gr.Dropdown(261                    label="Select base Dreambooth model (required)",262                    choices=controller.personalized_model_list,263                    interactive=True,264                    value="realistic1.safetensors"265                )266                base_model_dropdown.change(fn=controller.update_base_model, inputs=[267                                           base_model_dropdown], outputs=[base_model_dropdown])268 269                lora_model_dropdown = gr.Dropdown(270                    label="Select LoRA model (optional)",271                    choices=["none"],272                    value="none",273                    interactive=True,274                )275                lora_model_dropdown.change(fn=controller.update_lora_model, inputs=[276                                           lora_model_dropdown], outputs=[lora_model_dropdown])277 278                lora_alpha_slider = gr.Slider(279                    label="LoRA alpha", value=0.8, minimum=0, maximum=2, interactive=True)280                spatial_lora_slider = gr.Slider(281                    label="LCM LoRA alpha", value=0.8, minimum=0.0, maximum=1.0, interactive=True)282 283                personalized_refresh_button = gr.Button(284                    value="\U0001F503", elem_classes="toolbutton")285 286                def update_personalized_model():287                    controller.refresh_personalized_model()288                    return [289                        gr.Dropdown.update(290                            choices=controller.personalized_model_list),291                        gr.Dropdown.update(292                            choices=["none"] + controller.personalized_model_list)293                    ]294                personalized_refresh_button.click(fn=update_personalized_model, inputs=[], outputs=[295                                                  base_model_dropdown, lora_model_dropdown])296 297        with gr.Column(variant="panel"):298            gr.Markdown(299                """300                ### 2. Configs for AnimateLCM.301                """302            )303 304            prompt_textbox = gr.Textbox(label="Prompt", lines=2, value="best quality")305            negative_prompt_textbox = gr.Textbox(306                label="Negative prompt", lines=2, value="bad quality")307 308            with gr.Row().style(equal_height=False):309                with gr.Column():310                    with gr.Row():311                        sampler_dropdown = gr.Dropdown(label="Sampling method", choices=list(312                            scheduler_dict.keys()), value=list(scheduler_dict.keys())[0])313                        sample_step_slider = gr.Slider(314                            label="Sampling steps", value=4, minimum=1, maximum=25, step=1)315 316                    motion_scale_slider = gr.Slider(317                        label="motion scale (better identity with smaller scale)", value=0.8, minimum=0.0, maximum=1.5, step=0.05318                    )319                    beta_end_slider = gr.Slider(320                        label="beta end (a tricky way for selecting noisy steps)", value=0.014, minimum=0.012, maximum=0.016, step=0.001)321                    width_slider = gr.Slider(322                        label="Width",            value=512, minimum=256, maximum=1024, step=64)323                    height_slider = gr.Slider(324                        label="Height",           value=512, minimum=256, maximum=1024, step=64)325                    length_slider = gr.Slider(326                        label="Animation length", value=16,  minimum=12,   maximum=20,   step=1)327                    cfg_scale_slider = gr.Slider(328                        label="CFG Scale",        value=1, minimum=1,   maximum=2)329 330                    with gr.Row():331                        seed_textbox = gr.Textbox(label="Seed", value=-1)332                        seed_button = gr.Button(333                            value="\U0001F3B2", elem_classes="toolbutton")334                        seed_button.click(fn=lambda: gr.Textbox.update(335                            value=random.randint(1, 1e8)), inputs=[], outputs=[seed_textbox])336 337                    generate_button = gr.Button(338                        value="Generate", variant='primary')339 340                result_video = gr.Video(341                    label="Generated Animation", interactive=False)342            343            344            generate_button.click(345                fn=controller.animate,346                inputs=[347                    lora_alpha_slider,348                    spatial_lora_slider,349                    prompt_textbox,350                    negative_prompt_textbox,351                    sampler_dropdown,352                    sample_step_slider,353                    width_slider,354                    length_slider,355                    height_slider,356                    cfg_scale_slider,357                    seed_textbox,358                    image_upload,359                    beta_end_slider,360                    motion_scale_slider,361                ],362                outputs=[result_video]363            )364            examples = [365                [0.8, 0.8, "good quality, cloud", "bad quality", "LCM", 4, 768, 16, 512, 1, 1234, "test_imgs/cloud.jpeg", 0.014, 0.8],366                [0.8, 0.8, "good quality, dog", "bad quality", "LCM", 4, 768, 16, 512, 1, 1234, "test_imgs/dog.jpg", 0.014, 0.8],367                [0.8, 0.8, "good quality, fire", "bad quality", "LCM", 4, 768, 16, 512, 1, 1234, "test_imgs/fire.jpg", 0.014, 0.8],368                [0.8, 0.8, "good quality, fox, snow", "bad quality", "LCM", 4, 768, 16, 512, 1, 1234, "test_imgs/fox.jpg", 0.014, 0.8],369                [0.8, 0.8, "good quality, girl, wind, flower", "bad quality", "LCM", 4, 768, 16, 512, 1, 1235, "test_imgs/girl_flower.jpg", 0.014, 0.8],370                [0.8, 0.8, "good quality, lighter, fire", "bad quality", "LCM", 4, 768, 16, 512, 1, 1235, "test_imgs/lighter.jpg", 0.014, 0.8],371                [0.8, 0.8, "good quality, snowman, fire", "bad quality", "LCM", 4, 768, 16, 512, 1, 1234, "test_imgs/snow_man_fire.jpg", 0.014, 0.8],372            ]373            gr.Examples(374                examples = examples,375                inputs=[376                    lora_alpha_slider,377                    spatial_lora_slider,378                    prompt_textbox,379                    negative_prompt_textbox,380                    sampler_dropdown,381                    sample_step_slider,382                    width_slider,383                    length_slider,384                    height_slider,385                    cfg_scale_slider,386                    seed_textbox,387                    image_upload,388                    beta_end_slider,389                    motion_scale_slider,390                ],391                outputs=[result_video],392                fn=controller.animate,393                cache_examples=True,394            )395    return demo396 397 398if __name__ == "__main__":399    demo = ui()400    demo.queue(concurrency_count=3, max_size=20)401    demo.launch(share=True, server_name="127.0.0.1")402