add ModelIter, CLIPIter, VAEIter

This commit is contained in:
hnmr293
2023-04-02 19:18:28 +09:00
parent 280341941f
commit c34b344de8
6 changed files with 247 additions and 31 deletions
+3
View File
@@ -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
View File
@@ -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
View File
@@ -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,)
View File
View File
+89 -29
View File
@@ -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()