raaraya/AnimateDiff
0
1import argparse2import datetime3import inspect4import os5from omegaconf import OmegaConf6 7import torch8 9import diffusers10from diffusers import AutoencoderKL, DDIMScheduler11 12from tqdm.auto import tqdm13from transformers import CLIPTextModel, CLIPTokenizer14 15from animatediff.models.unet import UNet3DConditionModel16from animatediff.pipelines.pipeline_animation import AnimationPipeline17from animatediff.utils.util import save_videos_grid18from animatediff.utils.convert_from_ckpt import convert_ldm_unet_checkpoint, convert_ldm_clip_checkpoint, convert_ldm_vae_checkpoint19from animatediff.utils.convert_lora_safetensor_to_diffusers import convert_lora20from diffusers.utils.import_utils import is_xformers_available21 22from einops import rearrange, repeat23 24import csv, pdb, glob25from safetensors import safe_open26import math27from pathlib import Path28 29 30def main(args):31 *_, func_args = inspect.getargvalues(inspect.currentframe())32 func_args = dict(func_args)33 34 time_str = datetime.datetime.now().strftime("%Y-%m-%dT%H-%M-%S")35 savedir = f"samples/{Path(args.config).stem}-{time_str}"36 os.makedirs(savedir)37 inference_config = OmegaConf.load(args.inference_config)38 39 config = OmegaConf.load(args.config)40 samples = []41 42 sample_idx = 043 for model_idx, (config_key, model_config) in enumerate(list(config.items())):44 45 motion_modules = model_config.motion_module46 motion_modules = [motion_modules] if isinstance(motion_modules, str) else list(motion_modules)47 for motion_module in motion_modules:48 49 ### >>> create validation pipeline >>> ###50 tokenizer = CLIPTokenizer.from_pretrained(args.pretrained_model_path, subfolder="tokenizer")51 text_encoder = CLIPTextModel.from_pretrained(args.pretrained_model_path, subfolder="text_encoder")52 vae = AutoencoderKL.from_pretrained(args.pretrained_model_path, subfolder="vae") 53 unet = UNet3DConditionModel.from_pretrained_2d(args.pretrained_model_path, subfolder="unet", unet_additional_kwargs=OmegaConf.to_container(inference_config.unet_additional_kwargs))54 55 if is_xformers_available(): unet.enable_xformers_memory_efficient_attention()56 else: assert False57 58 pipeline = AnimationPipeline(59 vae=vae, text_encoder=text_encoder, tokenizer=tokenizer, unet=unet,60 scheduler=DDIMScheduler(**OmegaConf.to_container(inference_config.noise_scheduler_kwargs)),61 ).to("cuda")62 63 # 1. unet ckpt64 # 1.1 motion module65 motion_module_state_dict = torch.load(motion_module, map_location="cpu")66 if "global_step" in motion_module_state_dict: func_args.update({"global_step": motion_module_state_dict["global_step"]})67 missing, unexpected = pipeline.unet.load_state_dict(motion_module_state_dict, strict=False)68 assert len(unexpected) == 069 70 # 1.2 T2I71 if model_config.path != "":72 if model_config.path.endswith(".ckpt"):73 state_dict = torch.load(model_config.path)74 pipeline.unet.load_state_dict(state_dict)75 76 elif model_config.path.endswith(".safetensors"):77 state_dict = {}78 with safe_open(model_config.path, framework="pt", device="cpu") as f:79 for key in f.keys():80 state_dict[key] = f.get_tensor(key)81 82 is_lora = all("lora" in k for k in state_dict.keys())83 if not is_lora:84 base_state_dict = state_dict85 else:86 base_state_dict = {}87 with safe_open(model_config.base, framework="pt", device="cpu") as f:88 for key in f.keys():89 base_state_dict[key] = f.get_tensor(key) 90 91 # vae92 converted_vae_checkpoint = convert_ldm_vae_checkpoint(base_state_dict, pipeline.vae.config)93 pipeline.vae.load_state_dict(converted_vae_checkpoint)94 # unet95 converted_unet_checkpoint = convert_ldm_unet_checkpoint(base_state_dict, pipeline.unet.config)96 pipeline.unet.load_state_dict(converted_unet_checkpoint, strict=False)97 # text_model98 pipeline.text_encoder = convert_ldm_clip_checkpoint(base_state_dict)99 100 # import pdb101 # pdb.set_trace()102 if is_lora:103 pipeline = convert_lora(pipeline, state_dict, alpha=model_config.lora_alpha)104 105 pipeline.to("cuda")106 ### <<< create validation pipeline <<< ###107 108 prompts = model_config.prompt109 n_prompts = list(model_config.n_prompt) * len(prompts) if len(model_config.n_prompt) == 1 else model_config.n_prompt110 111 random_seeds = model_config.get("seed", [-1])112 random_seeds = [random_seeds] if isinstance(random_seeds, int) else list(random_seeds)113 random_seeds = random_seeds * len(prompts) if len(random_seeds) == 1 else random_seeds114 115 config[config_key].random_seed = []116 for prompt_idx, (prompt, n_prompt, random_seed) in enumerate(zip(prompts, n_prompts, random_seeds)):117 118 # manually set random seed for reproduction119 if random_seed != -1: torch.manual_seed(random_seed)120 else: torch.seed()121 config[config_key].random_seed.append(torch.initial_seed())122 123 print(f"current seed: {torch.initial_seed()}")124 print(f"sampling {prompt} ...")125 sample = pipeline(126 prompt,127 negative_prompt = n_prompt,128 num_inference_steps = model_config.steps,129 guidance_scale = model_config.guidance_scale,130 width = args.W,131 height = args.H,132 video_length = args.L,133 ).videos134 samples.append(sample)135 136 prompt = "-".join((prompt.replace("/", "").split(" ")[:10]))137 save_videos_grid(sample, f"{savedir}/sample/{sample_idx}-{prompt}.gif")138 print(f"save to {savedir}/sample/{prompt}.gif")139 140 sample_idx += 1141 142 samples = torch.concat(samples)143 save_videos_grid(samples, f"{savedir}/sample.gif", n_rows=4)144 145 OmegaConf.save(config, f"{savedir}/config.yaml")146 147 148if __name__ == "__main__":149 parser = argparse.ArgumentParser()150 parser.add_argument("--pretrained_model_path", type=str, default="models/StableDiffusion/stable-diffusion-v1-5",)151 parser.add_argument("--inference_config", type=str, default="configs/inference/inference.yaml") 152 parser.add_argument("--config", type=str, required=True)153 154 parser.add_argument("--L", type=int, default=16 )155 parser.add_argument("--W", type=int, default=512)156 parser.add_argument("--H", type=int, default=512)157 158 args = parser.parse_args()159 main(args)160 