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 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:
|
||||
|
||||
Reference in New Issue
Block a user