optimized batch processing in prep for inversion

This commit is contained in:
spacepxl
2024-05-27 23:43:40 -04:00
parent d57ac7012e
commit d9d37bbf32
2 changed files with 17 additions and 5 deletions
+1 -1
View File
@@ -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:
```
+16 -4
View File
@@ -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