diff --git a/model/iter.py b/model/iter.py index 1d0933c..351a406 100644 --- a/model/iter.py +++ b/model/iter.py @@ -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 +ATTR_NAME = 'iter_fn' + def iterize_model(model: ModelPatcher) -> List[Callable[[],ModelPatcher]]: - ATTR_NAME = 'iter_fn' 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: diff --git a/vae.py b/vae.py index e7ce6f7..e18eedd 100644 --- a/vae.py +++ b/vae.py @@ -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) + + vaes = torch.chunk(v, vae_num) + + for vn, vv in enumerate(vaes): + vae_results[vn].append(vv) - vs = torch.cat(results) + results = [] + for k in sorted(vae_results.keys()): + v = vae_results[k] + results.extend(v) + + vs = torch.cat(results).contiguous() return (vs,)