first results

This commit is contained in:
kijai
2024-10-10 15:29:00 +03:00
parent 3e57ab03c1
commit a7033bdd26
2 changed files with 208 additions and 75 deletions
+124 -16
View File
@@ -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",
}
@@ -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: