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 tqdm
from comfy.sd import ModelPatcher, CLIP, VAE
@@ -13,14 +13,14 @@ class CondForModels(torch.Tensor):
super().__init__()
self.ex = ex
def iterize_model(model: ModelPatcher) -> List[Callable[[],ModelPatcher]]:
ATTR_NAME = 'iter_fn'
def iterize_model(model: ModelPatcher) -> List[Callable[[],ModelPatcher]]:
if not hasattr(model, ATTR_NAME):
setattr(model, ATTR_NAME, [lambda: model])
return getattr(model, ATTR_NAME)
def iterize_clip(clip: CLIP) -> List[Callable[[],CLIP]]:
ATTR_NAME = 'iter_fn'
if hasattr(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)
def iterize_vae(vae: VAE) -> List[Callable[[],VAE]]:
ATTR_NAME = 'iter_fn'
if hasattr(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)
def try_get_iter(obj) -> Optional[List[Callable[[],Any]]]:
return getattr(obj, ATTR_NAME, None)
class ModelIter:
+22 -3
View File
@@ -1,5 +1,8 @@
from collections import defaultdict
from typing import Dict, List
import torch
from tqdm import trange
from .model.iter import try_get_iter
class VAEDecodeBatched:
def __init__(self, device="cpu"):
@@ -29,14 +32,30 @@ class VAEDecodeBatched:
s = samples['samples']
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):
e = min([i+batch_size, n])
t = s[i:e, ...]
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,)