From 6b5a2321aee8fd0cf5394652f99886dff36d7eb5 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Thu, 10 Oct 2024 16:57:22 +0300 Subject: [PATCH] img2vid --- nodes.py | 132 +++++++++++++----- .../pyramid_dit_for_video_gen_pipeline.py | 44 +++--- 2 files changed, 126 insertions(+), 50 deletions(-) diff --git a/nodes.py b/nodes.py index 95c45d6..91c88ad 100644 --- a/nodes.py +++ b/nodes.py @@ -115,8 +115,13 @@ class DownloadAndLoadPyramidFlowModel: # fuse_qkv_projections=True if pab_config is None else False, # ) - - return (model,) + pyramid_pipe = { + "model": model, + "dtype": model_dtype, + "text_encoder_dtype": text_encoder_dtype, + "vae_dtype": vae_dtype, + } + return (pyramid_pipe,) class CogVideoTextEncode: @@ -170,9 +175,9 @@ class PyramidFlowSampler: "keep_model_loaded": ("BOOLEAN", {"default": False}), }, - # "optional": { - # "samples": ("LATENT", ), - # } + "optional": { + "input_latent": ("LATENT", ), + } } RETURN_TYPES = ("PYRAMIDFLOWMODEL", "LATENT", ) @@ -180,38 +185,51 @@ class PyramidFlowSampler: FUNCTION = "sample" CATEGORY = "PyramidFlowWrapper" - def sample(self, model, steps, prompt_embeds, seed, height, width, video_steps, temp, guidance_scale, video_guidance_scale, keep_model_loaded): + def sample(self, model, steps, prompt_embeds, seed, height, width, video_steps, temp, guidance_scale, video_guidance_scale, + keep_model_loaded, input_latent=None): mm.soft_empty_cache() device = mm.get_torch_device() offload_device = mm.unet_offload_device() - model.vae.enable_tiling() torch.manual_seed(seed) torch.cuda.manual_seed(seed) - autocastcondition = not model.dtype == torch.float32 + autocastcondition = not model["model"].dtype == torch.float32 autocast_context = torch.autocast(mm.get_autocast_device(device)) if autocastcondition else nullcontext() - #model.dit.to(device) - #model.vae.to(device) - #model.text_encoder.to(device) - with autocast_context: - latents = model.generate( - prompt_embeds_dict = prompt_embeds, - device=device, - num_inference_steps=[steps, steps, steps], #why's this a list - video_num_inference_steps=[video_steps, video_steps, video_steps], #why's this a list - height=height, - width=width, - temp=temp, - guidance_scale=guidance_scale, # The guidance for the first frame - video_guidance_scale=video_guidance_scale, # The guidance for the other video latent - output_type="latent", - ) + if input_latent is None: + with autocast_context: + latents = model["model"].generate( + prompt_embeds_dict = prompt_embeds, + device=device, + num_inference_steps=[steps, steps, steps], #why's this a list + video_num_inference_steps=[video_steps, video_steps, video_steps], #why's this a list + height=height, + width=width, + temp=temp, + guidance_scale=guidance_scale, # The guidance for the first frame + video_guidance_scale=video_guidance_scale, # The guidance for the other video latent + output_type="latent", + ) + else: + with autocast_context: + latents = model["model"].generate_i2v( + prompt_embeds_dict = prompt_embeds, + input_image_latent=input_latent, + device=device, + num_inference_steps=[steps, steps, steps], #why's this a list + height=height, + width=width, + temp=temp, + guidance_scale=guidance_scale, # The guidance for the first frame + video_guidance_scale=video_guidance_scale, # The guidance for the other video latent + output_type="latent", + ) + if not keep_model_loaded: - model.dit.to(offload_device) + model["model"].dit.to(offload_device) return (model, {"samples": latents},) @@ -241,15 +259,17 @@ class PyramidFlowTextEncode: device = mm.get_torch_device() offload_device = mm.unet_offload_device() - autocastcondition = not model.dtype == torch.float32 + text_encoder = model["model"].text_encoder + + autocastcondition = not model["text_encoder_dtype"] == torch.float32 autocast_context = torch.autocast(mm.get_autocast_device(device)) if autocastcondition else nullcontext() - model.text_encoder.to(device) + text_encoder.to(device) with autocast_context: - prompt_embeds, prompt_attention_mask, pooled_prompt_embeds = model.text_encoder(positive_prompt, device) - negative_prompt_embeds, negative_prompt_attention_mask, pooled_negative_prompt_embeds = model.text_encoder(negative_prompt, device) + prompt_embeds, prompt_attention_mask, pooled_prompt_embeds = text_encoder(positive_prompt, device) + negative_prompt_embeds, negative_prompt_attention_mask, pooled_negative_prompt_embeds = text_encoder(negative_prompt, device) if not keep_model_loaded: - model.text_encoder.to(offload_device) + text_encoder.to(offload_device) embeds = { "prompt_embeds": prompt_embeds, @@ -262,6 +282,50 @@ class PyramidFlowTextEncode: return (embeds,) +class PyramidFlowVAEEncode: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "model": ("PYRAMIDFLOWMODEL",), + "image": ("IMAGE",), + }, + } + + RETURN_TYPES = ("LATENT", ) + RETURN_NAMES = ("samples", ) + FUNCTION = "sample" + CATEGORY = "PyramidFlowWrapper" + + def sample(self, model, image): + mm.soft_empty_cache() + + self.vae = model["model"].vae + dtype = model["vae_dtype"] + + device = mm.get_torch_device() + offload_device = mm.unet_offload_device() + self.vae.enable_tiling() + + # For the image latent + self.vae_shift_factor = 0.1490 + self.vae_scale_factor = 1 / 1.8415 + + # For the video latent + self.vae_video_shift_factor = -0.2343 + self.vae_video_scale_factor = 1 / 3.0986 + input_image_tensor = image * 2 - 1 + input_image_tensor = rearrange(input_image_tensor, 'b h w c -> b c h w') + input_image_tensor = input_image_tensor.unsqueeze(2) # Add temporal dimension t=1 + input_image_tensor = input_image_tensor.to(dtype=dtype, device=device) + + self.vae.to(device) + input_image_latent = (self.vae.encode(input_image_tensor).latent_dist.sample() - self.vae_shift_factor) * self.vae_scale_factor # [b c 1 h w] + self.vae.to(offload_device) + + + return (input_image_latent,) + class PyramidFlowVAEDecode: @classmethod def INPUT_TYPES(s): @@ -269,7 +333,7 @@ class PyramidFlowVAEDecode: "required": { "model": ("PYRAMIDFLOWMODEL",), "samples": ("LATENT",), - "tile_sample_min_size": ("INT", {"default": 128, "min": 64, "max": 512, "step": 8}), + "tile_sample_min_size": ("INT", {"default": 256, "min": 64, "max": 512, "step": 8}), "window_size": ("INT", {"default": 2, "min": 1, "max": 4, "step": 1}), }, @@ -284,11 +348,11 @@ class PyramidFlowVAEDecode: mm.soft_empty_cache() latents = samples["samples"] - self.vae = model.vae + self.vae = model["model"].vae device = mm.get_torch_device() offload_device = mm.unet_offload_device() - model.vae.enable_tiling() + self.vae.enable_tiling() # For the image latent self.vae_shift_factor = 0.1490 @@ -324,6 +388,7 @@ NODE_CLASS_MAPPINGS = { "PyramidFlowSampler": PyramidFlowSampler, "PyramidFlowVAEDecode": PyramidFlowVAEDecode, "PyramidFlowTextEncode": PyramidFlowTextEncode, + "PyramidFlowVAEEncode": PyramidFlowVAEEncode, } NODE_DISPLAY_NAME_MAPPINGS = { @@ -331,4 +396,5 @@ NODE_DISPLAY_NAME_MAPPINGS = { "PyramidFlowSampler": "PyramidFlow Sampler", "PyramidFlowVAEDecode" : "PyramidFlow VAE Decode", "PyramidFlowTextEncode": "PyramidFlow Text Encode", + "PyramidFlowVAEEncode": "PyramidFlow VAE Encode", } diff --git a/pyramid_dit/pyramid_dit_for_video_gen_pipeline.py b/pyramid_dit/pyramid_dit_for_video_gen_pipeline.py index e930ddc..e590add 100644 --- a/pyramid_dit/pyramid_dit_for_video_gen_pipeline.py +++ b/pyramid_dit/pyramid_dit_for_video_gen_pipeline.py @@ -58,7 +58,7 @@ class PyramidDiTForVideoGeneration: dit_path = os.path.join(model_path, model_variant) - + self.dit = PyramidDiffusionMMDiT.from_pretrained( dit_path, torch_dtype=torch_dtype, use_gradient_checkpointing=use_gradient_checkpointing, @@ -279,27 +279,25 @@ class PyramidDiTForVideoGeneration: @torch.no_grad() def generate_i2v( self, - #prompt: Union[str, List[str]] = '', prompt_embeds_dict: dict, - input_image: torch.Tensor, + device: torch.device, + input_image_latent: torch.Tensor, temp: int = 1, num_inference_steps: Optional[Union[int, List[int]]] = 28, + height: Optional[int] = None, + width: Optional[int] = None, guidance_scale: float = 7.0, video_guidance_scale: float = 4.0, min_guidance_scale: float = 2.0, use_linear_guidance: bool = False, alpha: float = 0.5, - negative_prompt: Optional[Union[str, List[str]]]="cartoon style, worst quality, low quality, blurry, absolute black, absolute white, low res, extra limbs, extra digits, misplaced objects, mutated anatomy, monochrome, horror", num_images_per_prompt: Optional[int] = 1, generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None, output_type: Optional[str] = "pil", ): - device = self.device + #device = self.device dtype = self.dtype - width = input_image.width - height = input_image.height - assert temp % self.frame_per_unit == 0, "The frames should be divided by frame_per unit" batch_size = 1 # if isinstance(prompt, str): @@ -340,6 +338,11 @@ class PyramidDiTForVideoGeneration: pooled_prompt_embeds = torch.cat([negative_pooled_prompt_embeds, positive_pooled_prompt_embeds], dim=0) prompt_attention_mask = torch.cat([negative_prompt_attention_mask, positive_prompt_attention_mask], dim=0) + prompt_embeds = prompt_embeds.to(dtype) + pooled_prompt_embeds = pooled_prompt_embeds.to(dtype) + prompt_attention_mask = prompt_attention_mask.to(dtype) + + # Create the initial random noise num_channels_latents = self.dit.config.in_channels latents = self.prepare_latents( @@ -366,17 +369,20 @@ class PyramidDiTForVideoGeneration: num_units = temp // self.frame_per_unit stages = self.stages - # encode the image latents - image_transform = transforms.Compose([ - transforms.ToTensor(), - transforms.Normalize(mean=(0.5, 0.5, 0.5), std=(0.5, 0.5, 0.5)), - ]) - input_image_tensor = image_transform(input_image).unsqueeze(0).unsqueeze(2) # [b c 1 h w] - input_image_latent = (self.vae.encode(input_image_tensor.to(device)).latent_dist.sample() - self.vae_shift_factor) * self.vae_scale_factor # [b c 1 h w] - + # # encode the image latents + # image_transform = transforms.Compose([ + # transforms.ToTensor(), + # transforms.Normalize(mean=(0.5, 0.5, 0.5), std=(0.5, 0.5, 0.5)), + # ]) + #input_image_tensor = image_transform(input_image).unsqueeze(0).unsqueeze(2) # [b c 1 h w] + + input_image_latent = input_image_latent.to(dtype).to(device) generated_latents_list = [input_image_latent] # The generated results last_generated_latents = input_image_latent + self.dit.to(device) + comfy_pbar = ProgressBar(num_units) + for unit_index in tqdm(range(1, num_units + 1)): if use_linear_guidance: self._guidance_scale = guidance_scale_list[unit_index] @@ -426,7 +432,7 @@ class PyramidDiTForVideoGeneration: generator, is_first_frame=False, ) - + comfy_pbar.update(1) generated_latents_list.append(intermed_latents[-1]) last_generated_latents = intermed_latents @@ -635,6 +641,10 @@ class PyramidDiTForVideoGeneration: @property def dtype(self): return next(self.dit.parameters()).dtype + + @property + def vae_dtype(self): + return next(self.dit.parameters()).dtype @property def guidance_scale(self):