first results
This commit is contained in:
@@ -5,7 +5,7 @@ import comfy.model_management as mm
|
|||||||
from comfy.utils import ProgressBar, load_torch_file
|
from comfy.utils import ProgressBar, load_torch_file
|
||||||
|
|
||||||
from contextlib import nullcontext
|
from contextlib import nullcontext
|
||||||
|
from einops import rearrange
|
||||||
from .pyramid_dit import PyramidDiTForVideoGeneration
|
from .pyramid_dit import PyramidDiTForVideoGeneration
|
||||||
|
|
||||||
import logging
|
import logging
|
||||||
@@ -153,15 +153,15 @@ class PyramidFlowSampler:
|
|||||||
return {
|
return {
|
||||||
"required": {
|
"required": {
|
||||||
"model": ("PYRAMIDFLOWMODEL",),
|
"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}),
|
"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}),
|
"steps": ("INT", {"default": 20, "min": 1, "max": 200, "step": 1}),
|
||||||
"video_steps": ("INT", {"default": 10, "min": 5, "max": 2048, "step": 4}),
|
"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"}),
|
"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"}),
|
"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}),
|
"seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
|
||||||
"prompt": ("STRING", {"default": "", "multiline": True}),
|
|
||||||
"keep_model_loaded": ("BOOLEAN", {"default": False}),
|
"keep_model_loaded": ("BOOLEAN", {"default": False}),
|
||||||
|
|
||||||
},
|
},
|
||||||
@@ -170,12 +170,12 @@ class PyramidFlowSampler:
|
|||||||
# }
|
# }
|
||||||
}
|
}
|
||||||
|
|
||||||
RETURN_TYPES = ("IMAGE", )
|
RETURN_TYPES = ("PYRAMIDFLOWMODEL", "LATENT", )
|
||||||
RETURN_NAMES = ("images", )
|
RETURN_NAMES = ("model","samples", )
|
||||||
FUNCTION = "sample"
|
FUNCTION = "sample"
|
||||||
CATEGORY = "PyramidFlowWrapper"
|
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()
|
mm.soft_empty_cache()
|
||||||
|
|
||||||
device = mm.get_torch_device()
|
device = mm.get_torch_device()
|
||||||
@@ -188,35 +188,143 @@ class PyramidFlowSampler:
|
|||||||
autocastcondition = not model.dtype == torch.float32
|
autocastcondition = not model.dtype == torch.float32
|
||||||
autocast_context = torch.autocast(mm.get_autocast_device(device)) if autocastcondition else nullcontext()
|
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.vae.to(device)
|
||||||
#model.text_encoder.to(device)
|
#model.text_encoder.to(device)
|
||||||
with autocast_context:
|
with autocast_context:
|
||||||
frames = model.generate(
|
latents = model.generate(
|
||||||
prompt=prompt,
|
prompt_embeds_dict = prompt_embeds,
|
||||||
num_inference_steps=[steps, steps, steps],
|
device=device,
|
||||||
video_num_inference_steps=[video_steps, video_steps, video_steps],
|
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,
|
height=height,
|
||||||
width=width,
|
width=width,
|
||||||
temp=temp,
|
temp=temp,
|
||||||
guidance_scale=guidance_scale, # The guidance for the first frame
|
guidance_scale=guidance_scale, # The guidance for the first frame
|
||||||
video_guidance_scale=video_guidance_scale, # The guidance for the other video latent
|
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:
|
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 = {
|
NODE_CLASS_MAPPINGS = {
|
||||||
"DownloadAndLoadPyramidFlowModel": DownloadAndLoadPyramidFlowModel,
|
"DownloadAndLoadPyramidFlowModel": DownloadAndLoadPyramidFlowModel,
|
||||||
"PyramidFlowSampler": PyramidFlowSampler,
|
"PyramidFlowSampler": PyramidFlowSampler,
|
||||||
|
"PyramidFlowVAEDecode": PyramidFlowVAEDecode,
|
||||||
|
"PyramidFlowTextEncode": PyramidFlowTextEncode,
|
||||||
|
|
||||||
}
|
}
|
||||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||||
"DownloadAndLoadPyramidFlowModel": "(Down)load PyramidFlow Model",
|
"DownloadAndLoadPyramidFlowModel": "(Down)load PyramidFlow Model",
|
||||||
"PyramidFlowSampler": "PyramidFlow Sampler",
|
"PyramidFlowSampler": "PyramidFlow Sampler",
|
||||||
|
"PyramidFlowVAEDecode" : "PyramidFlow VAE Decode",
|
||||||
|
"PyramidFlowTextEncode": "PyramidFlow Text Encode",
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,38 +1,26 @@
|
|||||||
import torch
|
import torch
|
||||||
import os
|
import os
|
||||||
import sys
|
|
||||||
import torch.nn as nn
|
|
||||||
import torch.nn.functional as F
|
import torch.nn.functional as F
|
||||||
|
|
||||||
from collections import OrderedDict
|
from collections import OrderedDict
|
||||||
from einops import rearrange
|
from einops import rearrange
|
||||||
from diffusers.utils.torch_utils import randn_tensor
|
from diffusers.utils.torch_utils import randn_tensor
|
||||||
import numpy as np
|
|
||||||
import math
|
import math
|
||||||
import random
|
|
||||||
import PIL
|
import PIL
|
||||||
from PIL import Image
|
from PIL import Image
|
||||||
from tqdm import tqdm
|
from tqdm import tqdm
|
||||||
from torchvision import transforms
|
from torchvision import transforms
|
||||||
from copy import deepcopy
|
from copy import deepcopy
|
||||||
from typing import Any, Callable, Dict, List, Optional, Union
|
from typing import Any, Callable, Dict, List, Optional, Union
|
||||||
from accelerate import Accelerator
|
|
||||||
from ..diffusion_schedulers import PyramidFlowMatchEulerDiscreteScheduler
|
from ..diffusion_schedulers import PyramidFlowMatchEulerDiscreteScheduler
|
||||||
from ..video_vae.modeling_causal_vae import CausalVideoVAE
|
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_pyramid_mmdit import PyramidDiffusionMMDiT
|
||||||
from .modeling_text_encoder import SD3TextEncoderWithMask
|
from .modeling_text_encoder import SD3TextEncoderWithMask
|
||||||
|
|
||||||
|
from comfy.utils import ProgressBar
|
||||||
|
|
||||||
def compute_density_for_timestep_sampling(
|
def compute_density_for_timestep_sampling(
|
||||||
weighting_scheme: str, batch_size: int, logit_mean: float = None, logit_std: float = None, mode_scale: float = None
|
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)
|
latents = randn_tensor(shape, generator=generator, device=device, dtype=dtype)
|
||||||
return latents
|
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):
|
def sample_block_noise(self, bs, ch, temp, height, width):
|
||||||
gamma = self.scheduler.config.gamma
|
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)
|
block_number = bs * ch * temp * (height // 2) * (width // 2)
|
||||||
noise = torch.stack([dist.sample() for _ in range(block_number)]) # [block number, 4]
|
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 = 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
|
return noise
|
||||||
|
|
||||||
@torch.no_grad()
|
@torch.no_grad()
|
||||||
@@ -232,6 +230,7 @@ class PyramidDiTForVideoGeneration:
|
|||||||
):
|
):
|
||||||
stages = self.stages
|
stages = self.stages
|
||||||
intermed_latents = []
|
intermed_latents = []
|
||||||
|
#print(f"Start generating one unit, the latents shape is {latents.shape}")
|
||||||
|
|
||||||
for i_s in range(len(stages)):
|
for i_s in range(len(stages)):
|
||||||
self.scheduler.set_timesteps(num_inference_steps[i_s], i_s, device=device)
|
self.scheduler.set_timesteps(num_inference_steps[i_s], i_s, device=device)
|
||||||
@@ -295,8 +294,9 @@ class PyramidDiTForVideoGeneration:
|
|||||||
@torch.no_grad()
|
@torch.no_grad()
|
||||||
def generate_i2v(
|
def generate_i2v(
|
||||||
self,
|
self,
|
||||||
prompt: Union[str, List[str]] = '',
|
#prompt: Union[str, List[str]] = '',
|
||||||
input_image: PIL.Image = None,
|
prompt_embeds_dict: dict,
|
||||||
|
input_image: torch.Tensor,
|
||||||
temp: int = 1,
|
temp: int = 1,
|
||||||
num_inference_steps: Optional[Union[int, List[int]]] = 28,
|
num_inference_steps: Optional[Union[int, List[int]]] = 28,
|
||||||
guidance_scale: float = 7.0,
|
guidance_scale: float = 7.0,
|
||||||
@@ -316,23 +316,23 @@ class PyramidDiTForVideoGeneration:
|
|||||||
height = input_image.height
|
height = input_image.height
|
||||||
|
|
||||||
assert temp % self.frame_per_unit == 0, "The frames should be divided by frame_per unit"
|
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):
|
# if isinstance(num_inference_steps, int):
|
||||||
batch_size = 1
|
# num_inference_steps = [num_inference_steps] * len(self.stages)
|
||||||
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)
|
|
||||||
|
|
||||||
negative_prompt = negative_prompt or ""
|
# negative_prompt = negative_prompt or ""
|
||||||
|
|
||||||
# Get the text embeddings
|
# # Get the text embeddings
|
||||||
prompt_embeds, prompt_attention_mask, pooled_prompt_embeds = self.text_encoder(prompt, 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)
|
# negative_prompt_embeds, negative_prompt_attention_mask, negative_pooled_prompt_embeds = self.text_encoder(negative_prompt, device)
|
||||||
|
|
||||||
if use_linear_guidance:
|
if use_linear_guidance:
|
||||||
max_guidance_scale = guidance_scale
|
max_guidance_scale = guidance_scale
|
||||||
@@ -342,10 +342,18 @@ class PyramidDiTForVideoGeneration:
|
|||||||
self._guidance_scale = guidance_scale
|
self._guidance_scale = guidance_scale
|
||||||
self._video_guidance_scale = video_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:
|
if self.do_classifier_free_guidance:
|
||||||
prompt_embeds = torch.cat([negative_prompt_embeds, prompt_embeds], dim=0)
|
prompt_embeds = torch.cat([negative_prompt_embeds, positive_prompt_embeds], dim=0)
|
||||||
pooled_prompt_embeds = torch.cat([negative_pooled_prompt_embeds, pooled_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, prompt_attention_mask], dim=0)
|
prompt_attention_mask = torch.cat([negative_prompt_attention_mask, positive_prompt_attention_mask], dim=0)
|
||||||
|
|
||||||
# Create the initial random noise
|
# Create the initial random noise
|
||||||
num_channels_latents = self.dit.config.in_channels
|
num_channels_latents = self.dit.config.in_channels
|
||||||
@@ -449,7 +457,7 @@ class PyramidDiTForVideoGeneration:
|
|||||||
@torch.no_grad()
|
@torch.no_grad()
|
||||||
def generate(
|
def generate(
|
||||||
self,
|
self,
|
||||||
prompt: Union[str, List[str]] = None,
|
prompt_embeds_dict: dict,
|
||||||
height: Optional[int] = None,
|
height: Optional[int] = None,
|
||||||
width: Optional[int] = None,
|
width: Optional[int] = None,
|
||||||
temp: int = 1,
|
temp: int = 1,
|
||||||
@@ -464,19 +472,20 @@ class PyramidDiTForVideoGeneration:
|
|||||||
num_images_per_prompt: Optional[int] = 1,
|
num_images_per_prompt: Optional[int] = 1,
|
||||||
generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,
|
generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,
|
||||||
output_type: Optional[str] = "pil",
|
output_type: Optional[str] = "pil",
|
||||||
|
device: Optional[torch.device] = None,
|
||||||
):
|
):
|
||||||
device = self.device
|
#device = self.device
|
||||||
dtype = self.dtype
|
dtype = self.dtype
|
||||||
|
|
||||||
assert (temp - 1) % self.frame_per_unit == 0, "The frames should be divided by frame_per unit"
|
assert (temp - 1) % self.frame_per_unit == 0, "The frames should be divided by frame_per unit"
|
||||||
|
|
||||||
if isinstance(prompt, str):
|
# if isinstance(prompt, str):
|
||||||
batch_size = 1
|
# batch_size = 1
|
||||||
prompt = prompt + ", hyper quality, Ultra HD, 8K" # adding this prompt to improve aesthetics
|
# prompt = prompt + ", hyper quality, Ultra HD, 8K" # adding this prompt to improve aesthetics
|
||||||
else:
|
# else:
|
||||||
assert isinstance(prompt, list)
|
# assert isinstance(prompt, list)
|
||||||
batch_size = len(prompt)
|
# batch_size = len(prompt)
|
||||||
prompt = [_ + ", hyper quality, Ultra HD, 8K" for _ in prompt]
|
# prompt = [_ + ", hyper quality, Ultra HD, 8K" for _ in prompt]
|
||||||
|
|
||||||
if isinstance(num_inference_steps, int):
|
if isinstance(num_inference_steps, int):
|
||||||
num_inference_steps = [num_inference_steps] * len(self.stages)
|
num_inference_steps = [num_inference_steps] * len(self.stages)
|
||||||
@@ -484,13 +493,15 @@ class PyramidDiTForVideoGeneration:
|
|||||||
if isinstance(video_num_inference_steps, int):
|
if isinstance(video_num_inference_steps, int):
|
||||||
video_num_inference_steps = [video_num_inference_steps] * len(self.stages)
|
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
|
# # Get the text embeddings
|
||||||
self.text_encoder.to(device)
|
# self.text_encoder.to(device)
|
||||||
prompt_embeds, prompt_attention_mask, pooled_prompt_embeds = self.text_encoder(prompt, 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)
|
# negative_prompt_embeds, negative_prompt_attention_mask, negative_pooled_prompt_embeds = self.text_encoder(negative_prompt, device)
|
||||||
self.text_encoder.to('cpu')
|
# self.text_encoder.to('cpu')
|
||||||
|
|
||||||
|
batch_size=1
|
||||||
|
|
||||||
if use_linear_guidance:
|
if use_linear_guidance:
|
||||||
max_guidance_scale = guidance_scale
|
max_guidance_scale = guidance_scale
|
||||||
@@ -501,10 +512,19 @@ class PyramidDiTForVideoGeneration:
|
|||||||
self._guidance_scale = guidance_scale
|
self._guidance_scale = guidance_scale
|
||||||
self._video_guidance_scale = video_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:
|
if self.do_classifier_free_guidance:
|
||||||
prompt_embeds = torch.cat([negative_prompt_embeds, prompt_embeds], dim=0)
|
prompt_embeds = torch.cat([negative_prompt_embeds, positive_prompt_embeds], dim=0)
|
||||||
pooled_prompt_embeds = torch.cat([negative_pooled_prompt_embeds, pooled_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, prompt_attention_mask], dim=0)
|
prompt_attention_mask = torch.cat([negative_prompt_attention_mask, positive_prompt_attention_mask], dim=0)
|
||||||
|
|
||||||
# Create the initial random noise
|
# Create the initial random noise
|
||||||
num_channels_latents = self.dit.config.in_channels
|
num_channels_latents = self.dit.config.in_channels
|
||||||
@@ -535,6 +555,10 @@ class PyramidDiTForVideoGeneration:
|
|||||||
generated_latents_list = [] # The generated results
|
generated_latents_list = [] # The generated results
|
||||||
last_generated_latents = None
|
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)):
|
for unit_index in tqdm(range(num_units)):
|
||||||
if use_linear_guidance:
|
if use_linear_guidance:
|
||||||
self._guidance_scale = guidance_scale_list[unit_index]
|
self._guidance_scale = guidance_scale_list[unit_index]
|
||||||
@@ -602,21 +626,22 @@ class PyramidDiTForVideoGeneration:
|
|||||||
generator,
|
generator,
|
||||||
is_first_frame=False,
|
is_first_frame=False,
|
||||||
)
|
)
|
||||||
|
comfy_pbar.update(1)
|
||||||
generated_latents_list.append(intermed_latents[-1])
|
generated_latents_list.append(intermed_latents[-1])
|
||||||
last_generated_latents = intermed_latents
|
last_generated_latents = intermed_latents
|
||||||
|
self.dit.to('cpu')
|
||||||
|
|
||||||
generated_latents = torch.cat(generated_latents_list, dim=2)
|
generated_latents = torch.cat(generated_latents_list, dim=2)
|
||||||
|
|
||||||
if output_type == "latent":
|
if output_type == "latent":
|
||||||
image = generated_latents
|
image = generated_latents
|
||||||
else:
|
else:
|
||||||
image = self.decode_latent(generated_latents)
|
image = self.decode_latent(generated_latents, device)
|
||||||
|
|
||||||
return image
|
return image
|
||||||
|
|
||||||
def decode_latent(self, latents):
|
def decode_latent(self, latents, device):
|
||||||
self.vae.to(self.device)
|
self.vae.to(device)
|
||||||
if latents.shape[2] == 1:
|
if latents.shape[2] == 1:
|
||||||
latents = (latents / self.vae_scale_factor) + self.vae_shift_factor
|
latents = (latents / self.vae_scale_factor) + self.vae_shift_factor
|
||||||
else:
|
else:
|
||||||
|
|||||||
Reference in New Issue
Block a user