CoolFace
Apppublic

cocktailpeanut/MotionDirector

sourceHugging Faceapache-2.0updated 3y agoView on Hugging Face
0likes
MotionDirector_inference.py285 linesDownload Raw Back to root
1import argparse2import os3import platform4import re5import warnings6from typing import Optional7 8import torch9from diffusers import DDIMScheduler, TextToVideoSDPipeline10from einops import rearrange11from torch import Tensor12from torch.nn.functional import interpolate13from tqdm import trange14import random15 16from MotionDirector_train import export_to_video, handle_memory_attention, load_primary_models, unet_and_text_g_c, freeze_models17from utils.lora_handler import LoraHandler18from utils.ddim_utils import ddim_inversion19import imageio20 21 22def initialize_pipeline(23    model: str,24    device: str = "cuda",25    xformers: bool = False,26    sdp: bool = False,27    lora_path: str = "",28    lora_rank: int = 64,29    lora_scale: float = 1.0,30):31    with warnings.catch_warnings():32        warnings.simplefilter("ignore")33 34        scheduler, tokenizer, text_encoder, vae, unet = load_primary_models(model)35 36    # Freeze any necessary models37    freeze_models([vae, text_encoder, unet])38 39    # Enable xformers if available40    handle_memory_attention(xformers, sdp, unet)41 42    lora_manager_temporal = LoraHandler(43        version="cloneofsimo",44        use_unet_lora=True,45        use_text_lora=False,46        save_for_webui=False,47        only_for_webui=False,48        unet_replace_modules=["TransformerTemporalModel"],49        text_encoder_replace_modules=None,50        lora_bias=None51    )52 53    unet_lora_params, unet_negation = lora_manager_temporal.add_lora_to_model(54        True, unet, lora_manager_temporal.unet_replace_modules, 0, lora_path, r=lora_rank, scale=lora_scale)55 56    unet.eval()57    text_encoder.eval()58    unet_and_text_g_c(unet, text_encoder, False, False)59 60    pipe = TextToVideoSDPipeline.from_pretrained(61        pretrained_model_name_or_path=model,62        scheduler=scheduler,63        tokenizer=tokenizer,64        text_encoder=text_encoder.to(device=device, dtype=torch.half),65        vae=vae.to(device=device, dtype=torch.half),66        unet=unet.to(device=device, dtype=torch.half),67    )68    pipe.scheduler = DDIMScheduler.from_config(pipe.scheduler.config)69 70    return pipe71 72 73def inverse_video(pipe, latents, num_steps):74    ddim_inv_scheduler = DDIMScheduler.from_config(pipe.scheduler.config)75    ddim_inv_scheduler.set_timesteps(num_steps)76 77    ddim_inv_latent = ddim_inversion(78        pipe, ddim_inv_scheduler, video_latent=latents.to(pipe.device),79        num_inv_steps=num_steps, prompt="")[-1]80    return ddim_inv_latent81 82 83def prepare_input_latents(84    pipe: TextToVideoSDPipeline,85    batch_size: int,86    num_frames: int,87    height: int,88    width: int,89    latents_path:str,90    noise_prior: float91):92    # initialize with random gaussian noise93    scale = pipe.vae_scale_factor94    shape = (batch_size, pipe.unet.config.in_channels, num_frames, height // scale, width // scale)95    if noise_prior > 0.:96        cached_latents = torch.load(latents_path)97        if 'inversion_noise' not in cached_latents:98            latents = inverse_video(pipe, cached_latents['latents'].unsqueeze(0), 50).squeeze(0)99        else:100            latents = torch.load(latents_path)['inversion_noise'].unsqueeze(0)101        if latents.shape[0] != batch_size:102            latents = latents.repeat(batch_size, 1, 1, 1, 1)103        if latents.shape != shape:104            latents = interpolate(rearrange(latents, "b c f h w -> (b f) c h w", b=batch_size), (height // scale, width // scale), mode='bilinear')105            latents = rearrange(latents, "(b f) c h w -> b c f h w", b=batch_size)106        noise = torch.randn_like(latents, dtype=torch.half)107        latents = (noise_prior) ** 0.5 * latents + (1 - noise_prior) ** 0.5 * noise108    else:109        latents = torch.randn(shape, dtype=torch.half)110 111    return latents112 113 114def encode(pipe: TextToVideoSDPipeline, pixels: Tensor, batch_size: int = 8):115    nf = pixels.shape[2]116    pixels = rearrange(pixels, "b c f h w -> (b f) c h w")117 118    latents = []119    for idx in trange(120        0, pixels.shape[0], batch_size, desc="Encoding to latents...", unit_scale=batch_size, unit="frame"121    ):122        pixels_batch = pixels[idx : idx + batch_size].to(pipe.device, dtype=torch.half)123        latents_batch = pipe.vae.encode(pixels_batch).latent_dist.sample()124        latents_batch = latents_batch.mul(pipe.vae.config.scaling_factor).cpu()125        latents.append(latents_batch)126    latents = torch.cat(latents)127 128    latents = rearrange(latents, "(b f) c h w -> b c f h w", f=nf)129 130    return latents131 132 133@torch.inference_mode()134def inference(135    model: str,136    prompt: str,137    negative_prompt: Optional[str] = None,138    width: int = 256,139    height: int = 256,140    num_frames: int = 24,141    num_steps: int = 50,142    guidance_scale: float = 15,143    device: str = "cuda",144    xformers: bool = False,145    sdp: bool = False,146    lora_path: str = "",147    lora_rank: int = 64,148    lora_scale: float = 1.0,149    seed: Optional[int] = None,150    latents_path: str="",151    noise_prior: float = 0.,152    repeat_num: int = 1,153):154    if seed is not None:155        random_seed = seed156        torch.manual_seed(seed)157 158    with torch.autocast(device, dtype=torch.half):159        # prepare models160        pipe = initialize_pipeline(model, device, xformers, sdp, lora_path, lora_rank, lora_scale)161 162        for i in range(repeat_num):163            if seed is None:164                random_seed = random.randint(100, 10000000)165                torch.manual_seed(random_seed)166 167            # prepare input latents168            init_latents = prepare_input_latents(169                pipe=pipe,170                batch_size=len(prompt),171                num_frames=num_frames,172                height=height,173                width=width,174                latents_path=latents_path,175                noise_prior=noise_prior176            )177 178            with torch.no_grad():179                video_frames = pipe(180                    prompt=prompt,181                    negative_prompt=negative_prompt,182                    width=width,183                    height=height,184                    num_frames=num_frames,185                    num_inference_steps=num_steps,186                    guidance_scale=guidance_scale,187                    latents=init_latents188                ).frames189 190            # =========================================191            # ========= write outputs to file =========192            # =========================================193            os.makedirs(args.output_dir, exist_ok=True)194 195            # save to mp4196            export_to_video(video_frames, f"{out_name}_{random_seed}.mp4", args.fps)197 198            # # save to gif199            file_name = f"{out_name}_{random_seed}.gif"200            imageio.mimsave(file_name, video_frames, 'GIF', duration=1000 * 1 / args.fps, loop=0)201 202    return video_frames203 204 205if __name__ == "__main__":206    import decord207 208    decord.bridge.set_bridge("torch")209 210    # fmt: off211    parser = argparse.ArgumentParser()212    parser.add_argument("-m", "--model", type=str, required=True,213                        help="HuggingFace repository or path to model checkpoint directory")214    parser.add_argument("-p", "--prompt", type=str, required=True, help="Text prompt to condition on")215    parser.add_argument("-n", "--negative-prompt", type=str, default=None, help="Text prompt to condition against")216    parser.add_argument("-o", "--output_dir", type=str, default="./outputs/inference", help="Directory to save output video to")217    parser.add_argument("-B", "--batch-size", type=int, default=1, help="Batch size for inference")218    parser.add_argument("-W", "--width", type=int, default=384, help="Width of output video")219    parser.add_argument("-H", "--height", type=int, default=384, help="Height of output video")220    parser.add_argument("-T", "--num-frames", type=int, default=16, help="Total number of frames to generate")221    parser.add_argument("-s", "--num-steps", type=int, default=30, help="Number of diffusion steps to run per frame.")222    parser.add_argument("-g", "--guidance-scale", type=float, default=12, help="Scale for guidance loss (higher values = more guidance, but possibly more artifacts).")223    parser.add_argument("-f", "--fps", type=int, default=8, help="FPS of output video")224    parser.add_argument("-d", "--device", type=str, default="cuda", help="Device to run inference on (defaults to cuda).")225    parser.add_argument("-x", "--xformers", action="store_true", help="Use XFormers attnetion, a memory-efficient attention implementation (requires `pip install xformers`).")226    parser.add_argument("-S", "--sdp", action="store_true", help="Use SDP attention, PyTorch's built-in memory-efficient attention implementation.")227    parser.add_argument("-cf", "--checkpoint_folder", type=str, default=None, help="Path to Low Rank Adaptation checkpoint file (defaults to empty string, which uses no LoRA).")228    parser.add_argument("-lr", "--lora_rank", type=int, default=32, help="Size of the LoRA checkpoint's projection matrix (defaults to 32).")229    parser.add_argument("-ls", "--lora_scale", type=float, default=1.0, help="Scale of LoRAs.")230    parser.add_argument("-r", "--seed", type=int, default=None, help="Random seed to make generations reproducible.")231    parser.add_argument("-np", "--noise_prior", type=float, default=0., help="Scale of the influence of inversion noise.")232    parser.add_argument("-ci", "--checkpoint_index", type=int, required=True,233                        help="The index of checkpoint, such as 300.")234    parser.add_argument("-rn", "--repeat_num", type=int, default=1,235                        help="How many results to generate with the same prompt.")236 237    args = parser.parse_args()238    # fmt: on239 240    # =========================================241    # ====== validate and prepare inputs ======242    # =========================================243 244    out_name = f"{args.output_dir}/"245    prompt = re.sub(r'[<>:"/\\|?*\x00-\x1F]', "_", args.prompt) if platform.system() == "Windows" else args.prompt246    out_name += f"{prompt}".replace(' ','_').replace(',', '').replace('.', '')247 248    args.prompt = [prompt] * args.batch_size249    if args.negative_prompt is not None:250        args.negative_prompt = [args.negative_prompt] * args.batch_size251 252    # =========================================253    # ============= sample videos =============254    # =========================================255    if args.checkpoint_index is not None:256        lora_path = f"{args.checkpoint_folder}/checkpoint-{args.checkpoint_index}/temporal/lora"257    else:258        lora_path = f"{args.checkpoint_folder}/checkpoint-default/temporal/lora"259    latents_folder = f"{args.checkpoint_folder}/cached_latents"260    latents_path = f"{latents_folder}/{random.choice(os.listdir(latents_folder))}"261    assert os.path.exists(lora_path)262    video_frames = inference(263        model=args.model,264        prompt=args.prompt,265        negative_prompt=args.negative_prompt,266        width=args.width,267        height=args.height,268        num_frames=args.num_frames,269        num_steps=args.num_steps,270        guidance_scale=args.guidance_scale,271        device=args.device,272        xformers=args.xformers,273        sdp=args.sdp,274        lora_path=lora_path,275        lora_rank=args.lora_rank,276        lora_scale = args.lora_scale,277        seed=args.seed,278        latents_path=latents_path,279        noise_prior=args.noise_prior,280        repeat_num=args.repeat_num281    )282 283 284 285