fix VAEDecodeBatched for VAEIter

This commit is contained in:
hnmr293
2023-04-02 23:50:19 +09:00
parent d0dec53bce
commit 886466a95a
2 changed files with 27 additions and 7 deletions
+5 -4
View File
@@ -1,4 +1,4 @@
from typing import List, Callable from typing import List, Callable, Any, Optional
import torch import torch
import tqdm import tqdm
from comfy.sd import ModelPatcher, CLIP, VAE from comfy.sd import ModelPatcher, CLIP, VAE
@@ -13,14 +13,14 @@ class CondForModels(torch.Tensor):
super().__init__() super().__init__()
self.ex = ex self.ex = ex
ATTR_NAME = 'iter_fn'
def iterize_model(model: ModelPatcher) -> List[Callable[[],ModelPatcher]]: def iterize_model(model: ModelPatcher) -> List[Callable[[],ModelPatcher]]:
ATTR_NAME = 'iter_fn'
if not hasattr(model, ATTR_NAME): if not hasattr(model, ATTR_NAME):
setattr(model, ATTR_NAME, [lambda: model]) setattr(model, ATTR_NAME, [lambda: model])
return getattr(model, ATTR_NAME) return getattr(model, ATTR_NAME)
def iterize_clip(clip: CLIP) -> List[Callable[[],CLIP]]: def iterize_clip(clip: CLIP) -> List[Callable[[],CLIP]]:
ATTR_NAME = 'iter_fn'
if hasattr(clip, ATTR_NAME): if hasattr(clip, ATTR_NAME):
return getattr(clip, ATTR_NAME) return getattr(clip, ATTR_NAME)
@@ -47,7 +47,6 @@ def iterize_clip(clip: CLIP) -> List[Callable[[],CLIP]]:
return getattr(clip, ATTR_NAME) return getattr(clip, ATTR_NAME)
def iterize_vae(vae: VAE) -> List[Callable[[],VAE]]: def iterize_vae(vae: VAE) -> List[Callable[[],VAE]]:
ATTR_NAME = 'iter_fn'
if hasattr(vae, ATTR_NAME): if hasattr(vae, ATTR_NAME):
return getattr(vae, ATTR_NAME) return getattr(vae, ATTR_NAME)
@@ -73,6 +72,8 @@ def iterize_vae(vae: VAE) -> List[Callable[[],VAE]]:
return getattr(vae, ATTR_NAME) return getattr(vae, ATTR_NAME)
def try_get_iter(obj) -> Optional[List[Callable[[],Any]]]:
return getattr(obj, ATTR_NAME, None)
class ModelIter: class ModelIter:
+22 -3
View File
@@ -1,5 +1,8 @@
from collections import defaultdict
from typing import Dict, List
import torch import torch
from tqdm import trange from tqdm import trange
from .model.iter import try_get_iter
class VAEDecodeBatched: class VAEDecodeBatched:
def __init__(self, device="cpu"): def __init__(self, device="cpu"):
@@ -29,14 +32,30 @@ class VAEDecodeBatched:
s = samples['samples'] s = samples['samples']
n = s.shape[0] n = s.shape[0]
results = [] iters = try_get_iter(vae)
if iters is None:
vae_num = 1
else:
vae_num = len(iters)
vae_results: Dict[int,List[torch.Tensor]] = defaultdict(lambda: [])
for i in trange(0, n, batch_size): for i in trange(0, n, batch_size):
e = min([i+batch_size, n]) e = min([i+batch_size, n])
t = s[i:e, ...] t = s[i:e, ...]
v = vae.decode(t) v = vae.decode(t)
results.append(v)
vs = torch.cat(results) vaes = torch.chunk(v, vae_num)
for vn, vv in enumerate(vaes):
vae_results[vn].append(vv)
results = []
for k in sorted(vae_results.keys()):
v = vae_results[k]
results.extend(v)
vs = torch.cat(results).contiguous()
return (vs,) return (vs,)