cocktailpeanut/MotionDirector
0
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 