add ModelIter, CLIPIter, VAEIter
This commit is contained in:
@@ -26,6 +26,9 @@
|
|||||||
|model|Dict2Model|`DICT`, (config_file)|`MODEL`|instantiate a model from given state_dict|
|
|model|Dict2Model|`DICT`, (config_file)|`MODEL`|instantiate a model from given state_dict|
|
||||||
|model|StateDictMerger|`DICT`, `DICT`, `FLOAT`|`MODEL`, `CLIP`, `VAE`|merge two or three models|
|
|model|StateDictMerger|`DICT`, `DICT`, `FLOAT`|`MODEL`, `CLIP`, `VAE`|merge two or three models|
|
||||||
|model|StateDictMergerBlockWeighted|`DICT`, `DICT`|`DICT`|merge two models with per-block weights|
|
|model|StateDictMergerBlockWeighted|`DICT`, `DICT`|`DICT`|merge two models with per-block weights|
|
||||||
|
|model|ModelIter|`MODEL`, `MODEL`|`MODEL`|iterate models|
|
||||||
|
|model|CLIPlIter|`CLIP`, `CLIP`|`CLIP`|iterate CLIPs|
|
||||||
|
|model|VAElIter|`VAE`, `VAE`|`VAE`|iterate VAEs|
|
||||||
|
|
||||||
## Output nodes
|
## Output nodes
|
||||||
|
|
||||||
|
|||||||
+12
-2
@@ -1,8 +1,9 @@
|
|||||||
from .randomlatent import RandomLatentImage
|
from .randomlatent import RandomLatentImage
|
||||||
from .vae import VAEDecodeBatched, VAEEncodeBatched
|
from .vae import VAEDecodeBatched, VAEEncodeBatched
|
||||||
from .sample import KSamplerSetting, KSamplerOverrided, KSamplerXYZ
|
from .sample import KSamplerSetting, KSamplerOverrided, KSamplerXYZ
|
||||||
from .model import StateDictLoader, Dict2Model
|
from .model.loader import StateDictLoader, Dict2Model
|
||||||
from .model_merge import StateDictMerger, StateDictMergerBlockWeighted
|
from .model.iter import ModelIter, CLIPIter, VAEIter
|
||||||
|
from .model.merge import StateDictMerger, StateDictMergerBlockWeighted
|
||||||
from .image import GridImage
|
from .image import GridImage
|
||||||
|
|
||||||
NODE_CLASS_MAPPINGS = {
|
NODE_CLASS_MAPPINGS = {
|
||||||
@@ -35,6 +36,15 @@ NODE_CLASS_MAPPINGS = {
|
|||||||
## creates model from state_dict loaded by `StateDictLoader`
|
## creates model from state_dict loaded by `StateDictLoader`
|
||||||
'Dict2Model': Dict2Model,
|
'Dict2Model': Dict2Model,
|
||||||
|
|
||||||
|
## iterate two models for KSamplerXYZ
|
||||||
|
'ModelIter': ModelIter,
|
||||||
|
|
||||||
|
## iterate two CLIPs for KSamplerXYZ
|
||||||
|
'CLIPIter': CLIPIter,
|
||||||
|
|
||||||
|
## iterate two VAEs for KSamplerXYZ
|
||||||
|
'VAEIter': VAEIter,
|
||||||
|
|
||||||
## merge two (weighted sum) or three (add difference) state_dict
|
## merge two (weighted sum) or three (add difference) state_dict
|
||||||
'StateDictMerger': StateDictMerger,
|
'StateDictMerger': StateDictMerger,
|
||||||
|
|
||||||
|
|||||||
+143
@@ -0,0 +1,143 @@
|
|||||||
|
from typing import List, Callable
|
||||||
|
import torch
|
||||||
|
import tqdm
|
||||||
|
from comfy.sd import ModelPatcher, CLIP, VAE
|
||||||
|
|
||||||
|
class CondForModels(torch.Tensor):
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def __new__(cls, x, ex, *args, **kwargs):
|
||||||
|
return super().__new__(cls, x, *args, **kwargs) # type: ignore
|
||||||
|
|
||||||
|
def __init__(self, x, ex: List[torch.Tensor], *args, **kwargs):
|
||||||
|
super().__init__()
|
||||||
|
self.ex = ex
|
||||||
|
|
||||||
|
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)
|
||||||
|
|
||||||
|
setattr(clip, ATTR_NAME, [lambda: clip])
|
||||||
|
|
||||||
|
old_encode = CLIP.encode
|
||||||
|
|
||||||
|
def new_encode(*args, **kwargs):
|
||||||
|
xs = []
|
||||||
|
clips = getattr(clip, ATTR_NAME)
|
||||||
|
for fn in tqdm.tqdm(clips):
|
||||||
|
clip_: CLIP = fn()
|
||||||
|
if clip_ == clip:
|
||||||
|
x = old_encode(clip_, *args, **kwargs)
|
||||||
|
else:
|
||||||
|
x = clip_.encode(*args, **kwargs)
|
||||||
|
if x.dim() == 2:
|
||||||
|
x = x.unsqueeze(0)
|
||||||
|
xs.append(x)
|
||||||
|
return CondForModels(xs[0], xs)
|
||||||
|
|
||||||
|
clip.encode = new_encode
|
||||||
|
|
||||||
|
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)
|
||||||
|
|
||||||
|
setattr(vae, ATTR_NAME, [lambda: vae])
|
||||||
|
|
||||||
|
old_decode = VAE.decode
|
||||||
|
|
||||||
|
def new_decode(*args, **kwargs):
|
||||||
|
xs = []
|
||||||
|
vaes = getattr(vae, ATTR_NAME)
|
||||||
|
for fn in tqdm.tqdm(vaes):
|
||||||
|
vae_: VAE = fn()
|
||||||
|
if vae_ == vae:
|
||||||
|
x = old_decode(vae_, *args, **kwargs)
|
||||||
|
else:
|
||||||
|
x = vae_.decode(*args, **kwargs)
|
||||||
|
if x.dim() == 3:
|
||||||
|
x = x.unsqueeze(0)
|
||||||
|
xs.append(x)
|
||||||
|
return torch.cat(xs)
|
||||||
|
|
||||||
|
vae.decode = new_decode
|
||||||
|
|
||||||
|
return getattr(vae, ATTR_NAME)
|
||||||
|
|
||||||
|
|
||||||
|
class ModelIter:
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(cls):
|
||||||
|
return {
|
||||||
|
'required': {
|
||||||
|
'model1': ('MODEL', ),
|
||||||
|
'model2': ('MODEL', )
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
RETURN_TYPES = ('MODEL',)
|
||||||
|
|
||||||
|
FUNCTION = 'execute'
|
||||||
|
|
||||||
|
CATEGORY = 'model'
|
||||||
|
|
||||||
|
def execute(self, model1, model2):
|
||||||
|
fns = iterize_model(model1)
|
||||||
|
fns.append(lambda: model2)
|
||||||
|
return (model1,)
|
||||||
|
|
||||||
|
|
||||||
|
class CLIPIter:
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(cls):
|
||||||
|
return {
|
||||||
|
'required': {
|
||||||
|
'clip1': ('CLIP', ),
|
||||||
|
'clip2': ('CLIP', )
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
RETURN_TYPES = ('CLIP',)
|
||||||
|
|
||||||
|
FUNCTION = 'execute'
|
||||||
|
|
||||||
|
CATEGORY = 'model'
|
||||||
|
|
||||||
|
def execute(self, clip1, clip2):
|
||||||
|
fns = iterize_clip(clip1)
|
||||||
|
fns.append(lambda: clip2)
|
||||||
|
return (clip1,)
|
||||||
|
|
||||||
|
|
||||||
|
class VAEIter:
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(cls):
|
||||||
|
return {
|
||||||
|
'required': {
|
||||||
|
'vae1': ('VAE', ),
|
||||||
|
'vae2': ('VAE', )
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
RETURN_TYPES = ('VAE',)
|
||||||
|
|
||||||
|
FUNCTION = 'execute'
|
||||||
|
|
||||||
|
CATEGORY = 'model'
|
||||||
|
|
||||||
|
def execute(self, vae1, vae2):
|
||||||
|
fns = iterize_vae(vae1)
|
||||||
|
fns.append(lambda: vae2)
|
||||||
|
return (vae1,)
|
||||||
@@ -1,11 +1,12 @@
|
|||||||
import re
|
import re
|
||||||
from itertools import product
|
from itertools import product
|
||||||
from typing import Callable, List, Dict, Any, Union, Iterable
|
from typing import Callable, List, Dict, Any, Union, Tuple, cast
|
||||||
import torch
|
import torch
|
||||||
import model_management # type: ignore
|
import model_management # type: ignore
|
||||||
import comfy.samplers
|
import comfy.samplers
|
||||||
from nodes import common_ksampler
|
from nodes import common_ksampler
|
||||||
from comfy.sd import ModelPatcher
|
from comfy.sd import ModelPatcher
|
||||||
|
from .model.iter import iterize_model, CondForModels
|
||||||
|
|
||||||
re_int = re.compile(r"\s*([+-]?\s*\d+)\s*")
|
re_int = re.compile(r"\s*([+-]?\s*\d+)\s*")
|
||||||
re_float = re.compile(r"\s*([+-]?\s*\d+(?:.\d*)?)\s*")
|
re_float = re.compile(r"\s*([+-]?\s*\d+(?:.\d*)?)\s*")
|
||||||
@@ -46,8 +47,83 @@ def get_cfg(noises: torch.Tensor, latent_image: torch.Tensor, cfgs: List[float])
|
|||||||
cf = torch.FloatTensor(cfgs * noises.shape[0])
|
cf = torch.FloatTensor(cfgs * noises.shape[0])
|
||||||
return torch.cat(ns), torch.cat(lat), cf[...,None,None,None]
|
return torch.cat(ns), torch.cat(lat), cf[...,None,None,None]
|
||||||
|
|
||||||
|
def process_cond(
|
||||||
|
conds: List[List[Union[torch.Tensor,dict]]],
|
||||||
|
control_nets: Union[list,None],
|
||||||
|
noise_shape0: int,
|
||||||
|
device
|
||||||
|
):
|
||||||
|
conds_copy = []
|
||||||
|
for p in conds:
|
||||||
|
t: torch.Tensor = p[0] # type: ignore
|
||||||
|
if t.shape[0] < noise_shape0:
|
||||||
|
t = torch.cat([t] * noise_shape0)
|
||||||
|
t = t.to(device)
|
||||||
|
if control_nets is not None and 'control' in p[1]:
|
||||||
|
control_nets += [p[1]['control']] # type: ignore
|
||||||
|
conds_copy += [[t] + p[1:]]
|
||||||
|
return conds_copy
|
||||||
|
|
||||||
|
def process_cond_for_models(
|
||||||
|
conds: List[List[Union[torch.Tensor,CondForModels,dict]]],
|
||||||
|
control_nets: list,
|
||||||
|
noise_shape0: int,
|
||||||
|
device
|
||||||
|
):
|
||||||
|
assert (
|
||||||
|
all(isinstance(p[0], CondForModels) for p in conds)
|
||||||
|
or not any(isinstance(p[0], CondForModels) for p in conds)
|
||||||
|
)
|
||||||
|
|
||||||
|
if isinstance(conds[0][0], CondForModels):
|
||||||
|
sizes = set( len(cast(CondForModels, p[0]).ex) for p in conds )
|
||||||
|
assert len(sizes) == 1, f'number of conditions: {sizes}'
|
||||||
|
size = sizes.pop()
|
||||||
|
|
||||||
|
#
|
||||||
|
# conds
|
||||||
|
# + [ CondForModels, dictA ]
|
||||||
|
# | .ex + condA for model1
|
||||||
|
# | + condA for model2
|
||||||
|
# | ...
|
||||||
|
# | L condA for model{size}
|
||||||
|
# + [ CondForModels, dictB ]
|
||||||
|
# | .ex + condB for model1
|
||||||
|
# | + condB for model2
|
||||||
|
# | ...
|
||||||
|
# | L condB for model{size}
|
||||||
|
# ...
|
||||||
|
#
|
||||||
|
# vvv
|
||||||
|
#
|
||||||
|
# conds
|
||||||
|
# + [ [ condA_for_model1, dictA ], [ condB_for_model1, dictB ], ... ]
|
||||||
|
# + [ [ condA_for_model2, dictA ], [ condB_for_model2, dictB ], ... ]
|
||||||
|
# ...
|
||||||
|
#
|
||||||
|
|
||||||
|
result = []
|
||||||
|
for model_index in range(size):
|
||||||
|
cs = []
|
||||||
|
for cond_for_models in conds:
|
||||||
|
c: CondForModels = cond_for_models[0] # type: ignore
|
||||||
|
rest = cond_for_models[1:]
|
||||||
|
cond = c.ex[model_index]
|
||||||
|
cs.append([cond, *rest])
|
||||||
|
result.append(process_cond(
|
||||||
|
cs,
|
||||||
|
control_nets if model_index == 0 else None,
|
||||||
|
noise_shape0,
|
||||||
|
device
|
||||||
|
))
|
||||||
|
return result
|
||||||
|
|
||||||
|
else:
|
||||||
|
return [ process_cond(conds, control_nets, noise_shape0, device) ]
|
||||||
|
|
||||||
|
|
||||||
def common_ksampler_xyz(
|
def common_ksampler_xyz(
|
||||||
model: Union[ModelPatcher,Iterable[ModelPatcher]],
|
model: ModelPatcher,
|
||||||
seed: Union[int,List[int]],
|
seed: Union[int,List[int]],
|
||||||
steps: Union[int,List[int]],
|
steps: Union[int,List[int]],
|
||||||
cfg: Union[float,List[float]],
|
cfg: Union[float,List[float]],
|
||||||
@@ -66,9 +142,6 @@ def common_ksampler_xyz(
|
|||||||
noise_mask = None
|
noise_mask = None
|
||||||
device = model_management.get_torch_device()
|
device = model_management.get_torch_device()
|
||||||
|
|
||||||
if not isinstance(model, Iterable):
|
|
||||||
model = (model,)
|
|
||||||
|
|
||||||
if not isinstance(seed, list):
|
if not isinstance(seed, list):
|
||||||
seed = [seed]
|
seed = [seed]
|
||||||
|
|
||||||
@@ -99,41 +172,24 @@ def common_ksampler_xyz(
|
|||||||
latent_image = latent_image.to(device)
|
latent_image = latent_image.to(device)
|
||||||
cfg_ = cfg_.to(device)
|
cfg_ = cfg_.to(device)
|
||||||
|
|
||||||
positive_copy = []
|
|
||||||
negative_copy = []
|
|
||||||
|
|
||||||
control_nets = []
|
control_nets = []
|
||||||
for p in positive:
|
positive_copies = process_cond_for_models(positive, control_nets, noise.shape[0], device)
|
||||||
t = p[0]
|
negative_copies = process_cond_for_models(negative, control_nets, noise.shape[0], device)
|
||||||
if t.shape[0] < noise.shape[0]:
|
|
||||||
t = torch.cat([t] * noise.shape[0])
|
|
||||||
t = t.to(device)
|
|
||||||
if 'control' in p[1]:
|
|
||||||
control_nets += [p[1]['control']]
|
|
||||||
positive_copy += [[t] + p[1:]]
|
|
||||||
for n in negative:
|
|
||||||
t = n[0]
|
|
||||||
if t.shape[0] < noise.shape[0]:
|
|
||||||
t = torch.cat([t] * noise.shape[0])
|
|
||||||
t = t.to(device)
|
|
||||||
if 'control' in n[1]:
|
|
||||||
control_nets += [n[1]['control']]
|
|
||||||
negative_copy += [[t] + n[1:]]
|
|
||||||
|
|
||||||
control_net_models = []
|
control_net_models = []
|
||||||
for x in control_nets:
|
for x in control_nets:
|
||||||
control_net_models += x.get_control_models()
|
control_net_models += x.get_control_models()
|
||||||
model_management.load_controlnet_gpu(control_net_models)
|
model_management.load_controlnet_gpu(control_net_models)
|
||||||
|
|
||||||
#samplers: List[comfy.samplers.KSampler] = []
|
|
||||||
samplers: List[Dict[str,Any]] = []
|
samplers: List[Dict[str,Any]] = []
|
||||||
for model_, sampler_name_, scheduler_, steps_ in product(model, sampler_name, scheduler, steps):
|
for (model_index, model_fn), sampler_name_, scheduler_, steps_ in product(enumerate(iterize_model(model)), sampler_name, scheduler, steps):
|
||||||
if sampler_name_ not in comfy.samplers.KSampler.SAMPLERS:
|
if sampler_name_ not in comfy.samplers.KSampler.SAMPLERS:
|
||||||
raise ValueError(f'unknown sampler name: {sampler_name_}')
|
raise ValueError(f'unknown sampler name: {sampler_name_}')
|
||||||
if scheduler_ not in comfy.samplers.KSampler.SCHEDULERS:
|
if scheduler_ not in comfy.samplers.KSampler.SCHEDULERS:
|
||||||
raise ValueError(f'unknown scheduler name: {scheduler_}')
|
raise ValueError(f'unknown scheduler name: {scheduler_}')
|
||||||
samplers.append(dict(
|
samplers.append(dict(
|
||||||
model=model_,
|
model_index=model_index,
|
||||||
|
model=model_fn,
|
||||||
steps=steps_,
|
steps=steps_,
|
||||||
device=device,
|
device=device,
|
||||||
sampler=sampler_name_,
|
sampler=sampler_name_,
|
||||||
@@ -143,12 +199,16 @@ def common_ksampler_xyz(
|
|||||||
|
|
||||||
all_samples: List[torch.Tensor] = []
|
all_samples: List[torch.Tensor] = []
|
||||||
for sampler_args in samplers:
|
for sampler_args in samplers:
|
||||||
model_ = sampler_args['model']
|
model_ = sampler_args['model']()
|
||||||
model_management.load_model_gpu(model_)
|
model_management.load_model_gpu(model_)
|
||||||
sampler_args['model'] = model_.model
|
sampler_args['model'] = model_.model
|
||||||
|
|
||||||
|
model_index = sampler_args.pop('model_index')
|
||||||
|
positive_copy = positive_copies[model_index]
|
||||||
|
negative_copy = negative_copies[model_index]
|
||||||
|
|
||||||
sampler = comfy.samplers.KSampler(**sampler_args)
|
sampler = comfy.samplers.KSampler(**sampler_args)
|
||||||
print(f'XYZ sampler={sampler.sampler}/{sampler.scheduler} {sampler.steps}steps')
|
print(f'XYZ sampler=model@{model_index}/{sampler.sampler}/{sampler.scheduler} {sampler.steps}steps')
|
||||||
|
|
||||||
samples = sampler.sample(noise, positive_copy, negative_copy, cfg=cfg_, latent_image=latent_image, start_step=start_step, last_step=last_step, force_full_denoise=force_full_denoise, denoise_mask=noise_mask)
|
samples = sampler.sample(noise, positive_copy, negative_copy, cfg=cfg_, latent_image=latent_image, start_step=start_step, last_step=last_step, force_full_denoise=force_full_denoise, denoise_mask=noise_mask)
|
||||||
samples = samples.cpu()
|
samples = samples.cpu()
|
||||||
|
|||||||
Reference in New Issue
Block a user