fix VAEDecodeBatched for VAEIter
This commit is contained in:
+5
-4
@@ -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:
|
||||
|
||||
|
||||
@@ -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,)
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user