From d9d37bbf32e191942e9bda9699efdc878d627b8b Mon Sep 17 00:00:00 2001 From: spacepxl Date: Mon, 27 May 2024 23:43:40 -0400 Subject: [PATCH] optimized batch processing in prep for inversion --- README.md | 2 +- nodes.py | 20 ++++++++++++++++---- 2 files changed, 17 insertions(+), 5 deletions(-) diff --git a/README.md b/README.md index 6af1a74..f4e58f3 100644 --- a/README.md +++ b/README.md @@ -39,7 +39,7 @@ to `ComfyUI_windows_portable/python_embeded/Include/*` And from `C:/Users/username/AppData/Local/Programs/Python/Python310/libs/*` to `ComfyUI_windows_portable/python_embeded/libs/*` -If all of that is set up correctly, when you run a StyleGAN workflow, it will first build the necessary PyTorch plugins (should take 30-60s), then generate an image. There will be a message in the console, and then subsequent images will be much faster to generate, about 0.1-0.3s on a 3090. +If all of that is set up correctly, when you run a StyleGAN workflow, it will first build the necessary PyTorch plugins (should take 30-60s), then generate an image. There will be a message in the console, and then subsequent images will be much faster to generate (measured at 64 images/sec on a 3090 with a large batch, although ComfyUI's tensor to PIL for previews will bottleneck realtime generation to more like 8 fps) StyleGAN2: ``` diff --git a/nodes.py b/nodes.py index 4958baf..7edb3f4 100644 --- a/nodes.py +++ b/nodes.py @@ -3,6 +3,7 @@ import sys import numpy as np import pickle import torch +from tqdm import trange from .slerp import slerp @@ -12,6 +13,7 @@ sys.modules["dnnlib"] = dnnlib sys.modules["torch_utils"] = torch_utils import folder_paths +from comfy.utils import PROGRESS_BAR_ENABLED, ProgressBar # set the models directory if "stylegan" not in folder_paths.folder_names_and_paths: @@ -77,11 +79,21 @@ class StyleGANSampler: if class_label < 0: class_label = None - img = stylegan_model(stylegan_latent, class_label) - img = torch.permute(img, (0, 2, 3, 1)) # BCHW -> BHWC - img = torch.clip(img / 2 + 0.5, 0, 1) # [-1, 1] -> [0, 1] + imgs = [] + batch_size = stylegan_latent.size(0) + pbar = None + if PROGRESS_BAR_ENABLED and batch_size > 1: + pbar = ProgressBar(batch_size) + for i in trange(batch_size): + img = stylegan_model(stylegan_latent[i].unsqueeze(0), class_label) + img = torch.permute(img, (0, 2, 3, 1)) # BCHW -> BHWC + img = torch.clip(img / 2 + 0.5, 0, 1) # [-1, 1] -> [0, 1] + imgs.append(img) + if pbar is not None: + pbar.update(1) - return (img, ) + imgs = torch.cat(imgs, dim=0) + return (imgs, ) class BlendStyleGANLatents: @classmethod