From 7284af0580d6b90baf62396acfa40a05510719fa Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Jukka=20Sepp=C3=A4nen?= <40791699+kijai@users.noreply.github.com> Date: Fri, 11 Oct 2024 02:12:36 +0300 Subject: [PATCH] Update nodes.py --- nodes.py | 12 +++++------- 1 file changed, 5 insertions(+), 7 deletions(-) diff --git a/nodes.py b/nodes.py index 61e8446..adcad45 100644 --- a/nodes.py +++ b/nodes.py @@ -7,7 +7,7 @@ from comfy.utils import ProgressBar, load_torch_file from contextlib import nullcontext from einops import rearrange from .pyramid_dit import PyramidDiTForVideoGeneration - +import torchvision.transforms as transforms import logging logging.basicConfig(level=logging.INFO, format='%(asctime)s - %(levelname)s - %(message)s') log = logging.getLogger(__name__) @@ -302,17 +302,15 @@ class PyramidFlowVAEEncode: device = mm.get_torch_device() offload_device = mm.unet_offload_device() - self.vae.enable_tiling() + self.vae.disable_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') + normalize = transforms.Normalize(mean=(0.5, 0.5, 0.5), std=(0.5, 0.5, 0.5)) + input_image_tensor = rearrange(image, 'b h w c -> b c h w') + input_image_tensor = normalize(input_image_tensor) input_image_tensor = input_image_tensor.unsqueeze(2) # Add temporal dimension t=1 input_image_tensor = input_image_tensor.to(dtype=dtype, device=device)