CoolFace
Apppublic

batkovdev/i2v-vtk

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
app.py390 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", "Personalized")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_t2v_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-t2v.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_textbox166    ):167 168        if is_xformers_available():169            self.unet.enable_xformers_memory_efficient_attention()170 171        pipeline = AnimationPipeline(172            vae=self.vae, text_encoder=self.text_encoder, tokenizer=self.tokenizer, unet=self.unet,173            scheduler=scheduler_dict[sampler_dropdown](174                **OmegaConf.to_container(self.inference_config.noise_scheduler_kwargs))175        ).to("cuda")176 177        if self.lora_model_state_dict != {}:178            pipeline = convert_lora(179                pipeline, self.lora_model_state_dict, alpha=lora_alpha_slider)180 181        pipeline.unet = convert_lcm_lora(copy.deepcopy(182            self.unet), self.lcm_lora_path, spatial_lora_slider)183 184        pipeline.to("cuda")185 186        if seed_textbox != -1 and seed_textbox != "":187            torch.manual_seed(int(seed_textbox))188        else:189            torch.seed()190        seed = torch.initial_seed()191 192        sample = pipeline(193            prompt_textbox,194            negative_prompt=negative_prompt_textbox,195            num_inference_steps=sample_step_slider,196            guidance_scale=cfg_scale_slider,197            width=width_slider,198            height=height_slider,199            video_length=length_slider,200        ).videos201 202        save_sample_path = os.path.join(203            self.savedir_sample, f"{sample_idx}.mp4")204        save_videos_grid(sample, save_sample_path)205 206        sample_config = {207            "prompt": prompt_textbox,208            "n_prompt": negative_prompt_textbox,209            "sampler": sampler_dropdown,210            "num_inference_steps": sample_step_slider,211            "guidance_scale": cfg_scale_slider,212            "width": width_slider,213            "height": height_slider,214            "video_length": length_slider,215            "seed": seed216        }217        json_str = json.dumps(sample_config, indent=4)218        with open(os.path.join(self.savedir, "logs.json"), "a") as f:219            f.write(json_str)220            f.write("\n\n")221        return gr.Video.update(value=save_sample_path)222 223 224controller = AnimateController()225 226controller.update_stable_diffusion("stable-diffusion-v1-5")227controller.update_motion_module("AnimateLCM_sd15_t2v.ckpt")228controller.update_base_model("realistic2.safetensors")229 230 231def ui():232    with gr.Blocks(css=css) as demo:233        gr.Markdown(234            """235            # [AnimateLCM: Accelerating the Animation of Personalized Diffusion Models and Adapters with Decoupled Consistency Learning](https://arxiv.org/abs/2402.00769)236            Fu-Yun Wang, Zhaoyang Huang (*Corresponding Author), Xiaoyu Shi, Weikang Bian, Guanglu Song, Yu Liu, Hongsheng Li (*Corresponding Author)<br>237            [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)238            """239            240            '''241            Important Notes: 242            1. The generation speed is around 1~2 seconds. There is delay in the space.243            2. Increase the sampling step and cfg if you want more fancy videos.244            '''245        )246        with gr.Column(variant="panel"):247            with gr.Row():248 249                base_model_dropdown = gr.Dropdown(250                    label="Select base Dreambooth model (required)",251                    choices=controller.personalized_model_list,252                    interactive=True,253                    value="realistic2.safetensors"254                )255                256                motion_module_dropdown = gr.Dropdown(257                    label="Select motion modules",258                    choices=controller.motion_module_list,259                    interactive=True,260                    value="sd15_t2v_beta_motion.ckpt"261                )262                base_model_dropdown.change(fn=controller.update_base_model, inputs=[263                                           base_model_dropdown], outputs=[base_model_dropdown])264                265                motion_module_dropdown.change(fn=controller.update_motion_module, inputs=[motion_module_dropdown],outputs=[motion_module_dropdown])266 267                lora_model_dropdown = gr.Dropdown(268                    label="Select LoRA model (optional)",269                    choices=["none"],270                    value="none",271                    interactive=True,272                )273                lora_model_dropdown.change(fn=controller.update_lora_model, inputs=[274                                           lora_model_dropdown], outputs=[lora_model_dropdown])275 276                lora_alpha_slider = gr.Slider(277                    label="LoRA alpha", value=0.8, minimum=0, maximum=2, interactive=True)278                spatial_lora_slider = gr.Slider(279                    label="LCM LoRA alpha", value=0.8, minimum=0.0, maximum=1.0, interactive=True)280 281                personalized_refresh_button = gr.Button(282                    value="\U0001F503", elem_classes="toolbutton")283 284                def update_personalized_model():285                    controller.refresh_personalized_model()286                    return [287                        gr.Dropdown.update(288                            choices=controller.personalized_model_list),289                        gr.Dropdown.update(290                            choices=["none"] + controller.personalized_model_list)291                    ]292                personalized_refresh_button.click(fn=update_personalized_model, inputs=[], outputs=[293                                                  base_model_dropdown, lora_model_dropdown])294 295        with gr.Column(variant="panel"):296            gr.Markdown(297                """298                ### 2. Configs for AnimateLCM.299                """300            )301 302            prompt_textbox = gr.Textbox(label="Prompt", lines=2, value="a boy holding a rabbit")303            negative_prompt_textbox = gr.Textbox(304                label="Negative prompt", lines=2, value="bad quality")305 306            with gr.Row().style(equal_height=False):307                with gr.Column():308                    with gr.Row():309                        sampler_dropdown = gr.Dropdown(label="Sampling method", choices=list(310                            scheduler_dict.keys()), value=list(scheduler_dict.keys())[0])311                        sample_step_slider = gr.Slider(312                            label="Sampling steps", value=6, minimum=1, maximum=25, step=1)313 314                    width_slider = gr.Slider(315                        label="Width",            value=512, minimum=256, maximum=1024, step=64)316                    height_slider = gr.Slider(317                        label="Height",           value=512, minimum=256, maximum=1024, step=64)318                    length_slider = gr.Slider(319                        label="Animation length", value=16,  minimum=12,   maximum=20,   step=1)320                    cfg_scale_slider = gr.Slider(321                        label="CFG Scale",        value=1.5, minimum=1,   maximum=2)322 323                    with gr.Row():324                        seed_textbox = gr.Textbox(label="Seed", value=-1)325                        seed_button = gr.Button(326                            value="\U0001F3B2", elem_classes="toolbutton")327                        seed_button.click(fn=lambda: gr.Textbox.update(328                            value=random.randint(1, 1e8)), inputs=[], outputs=[seed_textbox])329 330                    generate_button = gr.Button(331                        value="Generate", variant='primary')332 333                result_video = gr.Video(334                    label="Generated Animation", interactive=False)335            336            337            generate_button.click(338                fn=controller.animate,339                inputs=[340                    lora_alpha_slider,341                    spatial_lora_slider,342                    prompt_textbox,343                    negative_prompt_textbox,344                    sampler_dropdown,345                    sample_step_slider,346                    width_slider,347                    length_slider,348                    height_slider,349                    cfg_scale_slider,350                    seed_textbox,351                ],352                outputs=[result_video]353            )354            examples = [355                [0.8, 0.8, "a boy is holding a rabbit", "bad quality", "LCM", 8, 512, 16, 512, 1.5, 1234],356                [0.8, 0.8, "1girl smiling", "bad quality", "LCM", 4, 512, 16, 512, 1.5, 1233],357                [0.8, 0.8, "1girl,face,white background,", "bad quality", "LCM", 6, 512, 16, 512, 1.5, 1234],358                [0.8, 0.8, "clouds in the sky, best quality", "bad quality", "LCM", 4, 512, 16, 512, 1.5, 1234],359                360                361            ]362            gr.Examples(363                examples = examples,364                inputs=[365                    lora_alpha_slider,366                    spatial_lora_slider,367                    prompt_textbox,368                    negative_prompt_textbox,369                    sampler_dropdown,370                    sample_step_slider,371                    width_slider,372                    length_slider,373                    height_slider,374                    cfg_scale_slider,375                    seed_textbox,376                ],377                outputs=[result_video],378                fn=controller.animate,379                cache_examples=True,380            )381 382    return demo383 384 385if __name__ == "__main__":386    demo = ui()387    # gr.close_all()388    demo.queue(concurrency_count=3, max_size=20)389    demo.launch(share=True, server_name="127.0.0.1")390