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 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:
|
||||||
|
|
||||||
|
|||||||
@@ -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,)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user