CoolFace
Apppublic

raaraya/AnimateDiff

sourceHugging Faceapache-2.0updated 3y agoView on Hugging Face
0likes
animate.py160 linesDownload Raw Back to scripts
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