From a7033bdd262a3e3c8ad2d3a8a3e84db06affa8fd Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Thu, 10 Oct 2024 15:29:00 +0300 Subject: [PATCH] first results --- nodes.py | 140 +++++++++++++++-- .../pyramid_dit_for_video_gen_pipeline.py | 143 ++++++++++-------- 2 files changed, 208 insertions(+), 75 deletions(-) diff --git a/nodes.py b/nodes.py index fd943c9..e9baa7e 100644 --- a/nodes.py +++ b/nodes.py @@ -5,7 +5,7 @@ import comfy.model_management as mm from comfy.utils import ProgressBar, load_torch_file from contextlib import nullcontext - +from einops import rearrange from .pyramid_dit import PyramidDiTForVideoGeneration import logging @@ -153,15 +153,15 @@ class PyramidFlowSampler: return { "required": { "model": ("PYRAMIDFLOWMODEL",), + "prompt_embeds": ("PYRAMIDFLOWPROMPT",), + "width": ("INT", {"default": 640, "min": 128, "max": 2048, "step": 8}), "height": ("INT", {"default": 384, "min": 128, "max": 2048, "step": 8}), - "width": ("INT", {"default": 656, "min": 128, "max": 2048, "step": 8}), "steps": ("INT", {"default": 20, "min": 1, "max": 200, "step": 1}), "video_steps": ("INT", {"default": 10, "min": 5, "max": 2048, "step": 4}), - "temp": ("INT", {"default": 8, "min": 1}), + "temp": ("INT", {"default": 8, "min": 1, "tooltip": "temp=16: 5s, temp=31: 10s"}), "guidance_scale": ("FLOAT", {"default": 9.0, "min": 0.0, "max": 30.0, "step": 0.01, "tooltip": "The guidance for the first frame"}), "video_guidance_scale": ("FLOAT", {"default": 5.0, "min": 0.0, "max": 30.0, "step": 0.01, "tooltip": "The guidance for the other video latent"}), "seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}), - "prompt": ("STRING", {"default": "", "multiline": True}), "keep_model_loaded": ("BOOLEAN", {"default": False}), }, @@ -170,12 +170,12 @@ class PyramidFlowSampler: # } } - RETURN_TYPES = ("IMAGE", ) - RETURN_NAMES = ("images", ) + RETURN_TYPES = ("PYRAMIDFLOWMODEL", "LATENT", ) + RETURN_NAMES = ("model","samples", ) FUNCTION = "sample" CATEGORY = "PyramidFlowWrapper" - def sample(self, model, steps, prompt, 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): mm.soft_empty_cache() device = mm.get_torch_device() @@ -188,35 +188,143 @@ class PyramidFlowSampler: autocastcondition = not model.dtype == torch.float32 autocast_context = torch.autocast(mm.get_autocast_device(device)) if autocastcondition else nullcontext() - model.dit.to(device) + #model.dit.to(device) #model.vae.to(device) #model.text_encoder.to(device) with autocast_context: - frames = model.generate( - prompt=prompt, - num_inference_steps=[steps, steps, steps], - video_num_inference_steps=[video_steps, video_steps, video_steps], + 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="pt", + output_type="latent", ) - print(frames.shape) if not keep_model_loaded: - model.to(offload_device) + model.dit.to(offload_device) - return (frames,) + return (model, {"samples": latents},) + +class PyramidFlowTextEncode: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "model": ("PYRAMIDFLOWMODEL",), + "positive_prompt": ("STRING", {"default": "hyper quality, Ultra HD, 8K", "multiline": True} ), + "negative_prompt": ("STRING", {"default": "", "multiline": True} ), + "keep_model_loaded": ("BOOLEAN", {"default": False}), + + }, + # "optional": { + # "samples": ("LATENT", ), + # } + } + + RETURN_TYPES = ("PYRAMIDFLOWPROMPT", ) + RETURN_NAMES = ("prompt_embeds", ) + FUNCTION = "sample" + CATEGORY = "PyramidFlowWrapper" + + def sample(self, model, positive_prompt, negative_prompt, keep_model_loaded): + mm.soft_empty_cache() + + device = mm.get_torch_device() + offload_device = mm.unet_offload_device() + model.vae.enable_tiling() + + autocastcondition = not model.dtype == torch.float32 + autocast_context = torch.autocast(mm.get_autocast_device(device)) if autocastcondition else nullcontext() + + model.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) + if not keep_model_loaded: + model.text_encoder.to(offload_device) + + embeds = { + "prompt_embeds": prompt_embeds, + "attention_mask": prompt_attention_mask, + "pooled_embeds": pooled_prompt_embeds, + "negative_prompt_embeds": negative_prompt_embeds, + "negative_attention_mask": negative_prompt_attention_mask, + "negative_pooled_embeds": pooled_negative_prompt_embeds + } + + return (embeds,) + +class PyramidFlowVAEDecode: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "model": ("PYRAMIDFLOWMODEL",), + "samples": ("LATENT",), + "tile_sample_min_size": ("INT", {"default": 128, "min": 64, "max": 512, "step": 8}), + "window_size": ("INT", {"default": 2, "min": 1, "max": 4, "step": 1}), + + }, + } + + RETURN_TYPES = ("IMAGE", ) + RETURN_NAMES = ("images", ) + FUNCTION = "sample" + CATEGORY = "PyramidFlowWrapper" + + def sample(self, model, samples, tile_sample_min_size, window_size): + mm.soft_empty_cache() + + latents = samples["samples"] + self.vae = model.vae + + device = mm.get_torch_device() + offload_device = mm.unet_offload_device() + model.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 + + self.vae.to(device) + if latents.shape[2] == 1: + latents = (latents / self.vae_scale_factor) + self.vae_shift_factor + else: + latents[:, :, :1] = (latents[:, :, :1] / self.vae_scale_factor) + self.vae_shift_factor + latents[:, :, 1:] = (latents[:, :, 1:] / self.vae_video_scale_factor) + self.vae_video_shift_factor + + image = self.vae.decode(latents, temporal_chunk=True, window_size=window_size, tile_sample_min_size=tile_sample_min_size).sample + + self.vae.to(offload_device) + + image = image.float() + image = (image / 2 + 0.5).clamp(0, 1) + image = rearrange(image, "B C T H W -> (B T) H W C") + image = image.cpu().float() + + + return (image,) NODE_CLASS_MAPPINGS = { "DownloadAndLoadPyramidFlowModel": DownloadAndLoadPyramidFlowModel, "PyramidFlowSampler": PyramidFlowSampler, + "PyramidFlowVAEDecode": PyramidFlowVAEDecode, + "PyramidFlowTextEncode": PyramidFlowTextEncode, } NODE_DISPLAY_NAME_MAPPINGS = { "DownloadAndLoadPyramidFlowModel": "(Down)load PyramidFlow Model", "PyramidFlowSampler": "PyramidFlow Sampler", + "PyramidFlowVAEDecode" : "PyramidFlow VAE Decode", + "PyramidFlowTextEncode": "PyramidFlow Text Encode", } diff --git a/pyramid_dit/pyramid_dit_for_video_gen_pipeline.py b/pyramid_dit/pyramid_dit_for_video_gen_pipeline.py index 82b5beb..b2c05a2 100644 --- a/pyramid_dit/pyramid_dit_for_video_gen_pipeline.py +++ b/pyramid_dit/pyramid_dit_for_video_gen_pipeline.py @@ -1,38 +1,26 @@ import torch import os -import sys -import torch.nn as nn + import torch.nn.functional as F from collections import OrderedDict from einops import rearrange from diffusers.utils.torch_utils import randn_tensor -import numpy as np + import math -import random import PIL from PIL import Image from tqdm import tqdm from torchvision import transforms from copy import deepcopy from typing import Any, Callable, Dict, List, Optional, Union -from accelerate import Accelerator from ..diffusion_schedulers import PyramidFlowMatchEulerDiscreteScheduler from ..video_vae.modeling_causal_vae import CausalVideoVAE -# from ..trainer_misc import ( -# all_to_all, -# is_sequence_parallel_initialized, -# get_sequence_parallel_group, -# get_sequence_parallel_group_rank, -# get_sequence_parallel_rank, -# get_sequence_parallel_world_size, -# get_rank, -# ) - from .modeling_pyramid_mmdit import PyramidDiffusionMMDiT from .modeling_text_encoder import SD3TextEncoderWithMask +from comfy.utils import ProgressBar def compute_density_for_timestep_sampling( weighting_scheme: str, batch_size: int, logit_mean: float = None, logit_std: float = None, mode_scale: float = None @@ -205,12 +193,22 @@ class PyramidDiTForVideoGeneration: latents = randn_tensor(shape, generator=generator, device=device, dtype=dtype) return latents + # def sample_block_noise(self, bs, ch, temp, height, width): + # gamma = self.scheduler.config.gamma + # dist = torch.distributions.multivariate_normal.MultivariateNormal(torch.zeros(4), torch.eye(4) * (1 + gamma) - torch.ones(4, 4) * gamma) + # block_number = bs * ch * temp * (height // 2) * (width // 2) + # noise = torch.stack([dist.sample() for _ in range(block_number)]) # [block number, 4] + # noise = rearrange(noise, '(b c t h w) (p q) -> b c t (h p) (w q)',b=bs,c=ch,t=temp,h=height//2,w=width//2,p=2,q=2) + # return noise + def sample_block_noise(self, bs, ch, temp, height, width): gamma = self.scheduler.config.gamma - dist = torch.distributions.multivariate_normal.MultivariateNormal(torch.zeros(4), torch.eye(4) * (1 + gamma) - torch.ones(4, 4) * gamma) + epsilon = 1e-5 # Small value to ensure positive definiteness + covariance_matrix = torch.eye(4) * (1 + gamma) - torch.ones(4, 4) * gamma + torch.eye(4) * epsilon + dist = torch.distributions.multivariate_normal.MultivariateNormal(torch.zeros(4), covariance_matrix) block_number = bs * ch * temp * (height // 2) * (width // 2) - noise = torch.stack([dist.sample() for _ in range(block_number)]) # [block number, 4] - noise = rearrange(noise, '(b c t h w) (p q) -> b c t (h p) (w q)',b=bs,c=ch,t=temp,h=height//2,w=width//2,p=2,q=2) + noise = torch.stack([dist.sample() for _ in range(block_number)]) # [block number, 4] + noise = rearrange(noise, '(b c t h w) (p q) -> b c t (h p) (w q)', b=bs, c=ch, t=temp, h=height//2, w=width//2, p=2, q=2) return noise @torch.no_grad() @@ -232,6 +230,7 @@ class PyramidDiTForVideoGeneration: ): stages = self.stages intermed_latents = [] + #print(f"Start generating one unit, the latents shape is {latents.shape}") for i_s in range(len(stages)): self.scheduler.set_timesteps(num_inference_steps[i_s], i_s, device=device) @@ -295,8 +294,9 @@ class PyramidDiTForVideoGeneration: @torch.no_grad() def generate_i2v( self, - prompt: Union[str, List[str]] = '', - input_image: PIL.Image = None, + #prompt: Union[str, List[str]] = '', + prompt_embeds_dict: dict, + input_image: torch.Tensor, temp: int = 1, num_inference_steps: Optional[Union[int, List[int]]] = 28, guidance_scale: float = 7.0, @@ -316,23 +316,23 @@ class PyramidDiTForVideoGeneration: 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): + # batch_size = 1 + # prompt = prompt + ", hyper quality, Ultra HD, 8K" # adding this prompt to improve aesthetics + # else: + # assert isinstance(prompt, list) + # batch_size = len(prompt) + # prompt = [_ + ", hyper quality, Ultra HD, 8K" for _ in prompt] - if isinstance(prompt, str): - batch_size = 1 - prompt = prompt + ", hyper quality, Ultra HD, 8K" # adding this prompt to improve aesthetics - else: - assert isinstance(prompt, list) - batch_size = len(prompt) - prompt = [_ + ", hyper quality, Ultra HD, 8K" for _ in prompt] - - if isinstance(num_inference_steps, int): - num_inference_steps = [num_inference_steps] * len(self.stages) + # if isinstance(num_inference_steps, int): + # num_inference_steps = [num_inference_steps] * len(self.stages) - negative_prompt = negative_prompt or "" + # negative_prompt = negative_prompt or "" - # Get the text embeddings - prompt_embeds, prompt_attention_mask, pooled_prompt_embeds = self.text_encoder(prompt, device) - negative_prompt_embeds, negative_prompt_attention_mask, negative_pooled_prompt_embeds = self.text_encoder(negative_prompt, device) + # # Get the text embeddings + # prompt_embeds, prompt_attention_mask, pooled_prompt_embeds = self.text_encoder(prompt, device) + # negative_prompt_embeds, negative_prompt_attention_mask, negative_pooled_prompt_embeds = self.text_encoder(negative_prompt, device) if use_linear_guidance: max_guidance_scale = guidance_scale @@ -342,10 +342,18 @@ class PyramidDiTForVideoGeneration: self._guidance_scale = guidance_scale self._video_guidance_scale = video_guidance_scale + positive_prompt_embeds = prompt_embeds_dict['prompt_embeds'] + positive_pooled_prompt_embeds = prompt_embeds_dict['pooled_embeds'] + positive_prompt_attention_mask = prompt_embeds_dict['attention_mask'] + + negative_prompt_embeds = prompt_embeds_dict['negative_prompt_embeds'] + negative_pooled_prompt_embeds = prompt_embeds_dict['negative_pooled_embeds'] + negative_prompt_attention_mask = prompt_embeds_dict['negative_attention_mask'] + if self.do_classifier_free_guidance: - prompt_embeds = torch.cat([negative_prompt_embeds, prompt_embeds], dim=0) - pooled_prompt_embeds = torch.cat([negative_pooled_prompt_embeds, pooled_prompt_embeds], dim=0) - prompt_attention_mask = torch.cat([negative_prompt_attention_mask, prompt_attention_mask], dim=0) + prompt_embeds = torch.cat([negative_prompt_embeds, positive_prompt_embeds], dim=0) + 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) # Create the initial random noise num_channels_latents = self.dit.config.in_channels @@ -449,7 +457,7 @@ class PyramidDiTForVideoGeneration: @torch.no_grad() def generate( self, - prompt: Union[str, List[str]] = None, + prompt_embeds_dict: dict, height: Optional[int] = None, width: Optional[int] = None, temp: int = 1, @@ -464,19 +472,20 @@ class PyramidDiTForVideoGeneration: num_images_per_prompt: Optional[int] = 1, generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None, output_type: Optional[str] = "pil", + device: Optional[torch.device] = None, ): - device = self.device + #device = self.device dtype = self.dtype assert (temp - 1) % self.frame_per_unit == 0, "The frames should be divided by frame_per unit" - if isinstance(prompt, str): - batch_size = 1 - prompt = prompt + ", hyper quality, Ultra HD, 8K" # adding this prompt to improve aesthetics - else: - assert isinstance(prompt, list) - batch_size = len(prompt) - prompt = [_ + ", hyper quality, Ultra HD, 8K" for _ in prompt] + # if isinstance(prompt, str): + # batch_size = 1 + # prompt = prompt + ", hyper quality, Ultra HD, 8K" # adding this prompt to improve aesthetics + # else: + # assert isinstance(prompt, list) + # batch_size = len(prompt) + # prompt = [_ + ", hyper quality, Ultra HD, 8K" for _ in prompt] if isinstance(num_inference_steps, int): num_inference_steps = [num_inference_steps] * len(self.stages) @@ -484,13 +493,15 @@ class PyramidDiTForVideoGeneration: if isinstance(video_num_inference_steps, int): video_num_inference_steps = [video_num_inference_steps] * len(self.stages) - negative_prompt = negative_prompt or "" + #negative_prompt = negative_prompt or "" - # Get the text embeddings - self.text_encoder.to(device) - prompt_embeds, prompt_attention_mask, pooled_prompt_embeds = self.text_encoder(prompt, device) - negative_prompt_embeds, negative_prompt_attention_mask, negative_pooled_prompt_embeds = self.text_encoder(negative_prompt, device) - self.text_encoder.to('cpu') + # # Get the text embeddings + # self.text_encoder.to(device) + # prompt_embeds, prompt_attention_mask, pooled_prompt_embeds = self.text_encoder(prompt, device) + # negative_prompt_embeds, negative_prompt_attention_mask, negative_pooled_prompt_embeds = self.text_encoder(negative_prompt, device) + # self.text_encoder.to('cpu') + + batch_size=1 if use_linear_guidance: max_guidance_scale = guidance_scale @@ -501,10 +512,19 @@ class PyramidDiTForVideoGeneration: self._guidance_scale = guidance_scale self._video_guidance_scale = video_guidance_scale + positive_prompt_embeds = prompt_embeds_dict['prompt_embeds'] + positive_pooled_prompt_embeds = prompt_embeds_dict['pooled_embeds'] + positive_prompt_attention_mask = prompt_embeds_dict['attention_mask'] + + negative_prompt_embeds = prompt_embeds_dict['negative_prompt_embeds'] + negative_pooled_prompt_embeds = prompt_embeds_dict['negative_pooled_embeds'] + negative_prompt_attention_mask = prompt_embeds_dict['negative_attention_mask'] + + if self.do_classifier_free_guidance: - prompt_embeds = torch.cat([negative_prompt_embeds, prompt_embeds], dim=0) - pooled_prompt_embeds = torch.cat([negative_pooled_prompt_embeds, pooled_prompt_embeds], dim=0) - prompt_attention_mask = torch.cat([negative_prompt_attention_mask, prompt_attention_mask], dim=0) + prompt_embeds = torch.cat([negative_prompt_embeds, positive_prompt_embeds], dim=0) + 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) # Create the initial random noise num_channels_latents = self.dit.config.in_channels @@ -535,6 +555,10 @@ class PyramidDiTForVideoGeneration: generated_latents_list = [] # The generated results last_generated_latents = None + #self.dit.to(torch.float8_e4m3fn) + self.dit.to(device) + comfy_pbar = ProgressBar(num_units) + for unit_index in tqdm(range(num_units)): if use_linear_guidance: self._guidance_scale = guidance_scale_list[unit_index] @@ -602,21 +626,22 @@ class PyramidDiTForVideoGeneration: generator, is_first_frame=False, ) - + comfy_pbar.update(1) generated_latents_list.append(intermed_latents[-1]) last_generated_latents = intermed_latents + self.dit.to('cpu') generated_latents = torch.cat(generated_latents_list, dim=2) if output_type == "latent": image = generated_latents else: - image = self.decode_latent(generated_latents) + image = self.decode_latent(generated_latents, device) return image - def decode_latent(self, latents): - self.vae.to(self.device) + def decode_latent(self, latents, device): + self.vae.to(device) if latents.shape[2] == 1: latents = (latents / self.vae_scale_factor) + self.vae_shift_factor else: