batkovdev/i2v-vtk
0
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 