diff --git a/__init__.py b/__init__.py index 8bc0b8d..ccea139 100644 --- a/__init__.py +++ b/__init__.py @@ -4,6 +4,7 @@ from .sample import KSamplerSetting, KSamplerOverrided, KSamplerXYZ from .model.loader import StateDictLoader, Dict2Model from .model.iter import ModelIter, CLIPIter, VAEIter from .model.merge import StateDictMerger, StateDictMergerBlockWeighted +from .model.merge2 import StateDictMergerBlockWeightedMulti from .image import GridImage NODE_CLASS_MAPPINGS = { @@ -52,6 +53,10 @@ NODE_CLASS_MAPPINGS = { ## weights should be specified by Text 'StateDictMergerBlockWeighted': StateDictMergerBlockWeighted, + ## merge block weighted + ## weights should be specified by Text + 'StateDictMergerBlockWeightedMulti': StateDictMergerBlockWeightedMulti, + # image ## rearrange images to single image with specified columns and gap diff --git a/model/merge.py b/model/merge.py index e244d37..bd032ef 100644 --- a/model/merge.py +++ b/model/merge.py @@ -110,6 +110,21 @@ re_inp = re.compile(r'\.input_blocks\.(\d+)\.') re_mid = re.compile(r'\.middle_block\.(\d+)\.') re_out = re.compile(r'\.output_blocks\.(\d+)\.') +def block_index(key: str): + if not key.startswith('model.diffusion_model.'): + return None + if 'time_embed' in key: + return 0 + if '.out.' in key: + return 24 + m = re_inp.search(key) + if m: return int(m.group(1)) + m = re_mid.search(key) + if m: return 12 + int(m.group(1)) + m = re_out.search(key) + if m: return 13 + int(m.group(1)) + return None + def weighted_sum_block( model_A: Dict[str,torch.Tensor], model_B: Dict[str,torch.Tensor], @@ -121,23 +136,8 @@ def weighted_sum_block( print('merging ...') print('mode: Block Weighted') - def index(key: str): - if not key.startswith('model.diffusion_model.'): - return None - if 'time_embed' in key: - return 0 - if '.out.' in key: - return 24 - m = re_inp.search(key) - if m: return int(m.group(1)) - m = re_mid.search(key) - if m: return 12 + int(m.group(1)) - m = re_out.search(key) - if m: return 13 + int(m.group(1)) - return None - def merge_fn(key, t1, t2): - weight_index = index(key) + weight_index = block_index(key) if weight_index is None: alpha = base_alpha elif 25 <= weight_index: diff --git a/model/merge2.py b/model/merge2.py new file mode 100644 index 0000000..7d0aed8 --- /dev/null +++ b/model/merge2.py @@ -0,0 +1,170 @@ +from typing import Dict, Union, List, Callable, Optional +import torch +import folder_paths +from .loader import Dict2Model +from .merge import block_index, weighted_sum_block, StateDictMergerBlockWeighted +from .iter import iterize_model + +from comfy.ldm.models.diffusion.ddpm import LatentDiffusion + +class MergedModule(torch.nn.Module): + + def __init__(self, name: str, a: torch.nn.Module, b: torch.nn.Module, alpha: Callable[[str],float]): + super().__init__() + + assert hasattr(a, 'weight') + assert hasattr(b, 'weight') + + self._name = name + self.a = a + self.b = b + self.alpha = alpha + # + #self.a._apply = self._apply_a + #self.b._apply = self._apply_b + + def forward(self, *args, **kwargs): + va = self.a(*args, **kwargs) + vb = self.b(*args, **kwargs) + a = self.alpha(self._name) + return (1-a)*va + a*vb + + #def _apply_a(self, *args, **kwargs): + # torch.nn.Module._apply(self.b, *args, **kwargs) + # return torch.nn.Module._apply(self.a, *args, **kwargs) + # + #def _apply_b(self, *args, **kwargs): + # torch.nn.Module._apply(self.a, *args, **kwargs) + # return torch.nn.Module._apply(self.b, *args, **kwargs) + + +ATTR_ALPHAS = 'mbw_alphas' +ATTR_INDEX = 'mbw_index' + +def get_current_alpha(model: LatentDiffusion) -> Optional[List[float]]: + if hasattr(model, ATTR_ALPHAS): + return getattr(model, ATTR_ALPHAS)[getattr(model, ATTR_INDEX)] + else: + return None + +def mbw_on_the_fly( + model_A: LatentDiffusion, + model_B: LatentDiffusion, + alphas_list: List[List[float]], + base_alpha: float, +): + setattr(model_A, ATTR_ALPHAS, alphas_list) + setattr(model_A, ATTR_INDEX, 0) + + def alpha_fn(name: str): + block = block_index(name) + if block is not None and 25 <= block: + raise ValueError('must not happen') + + if block is None: + return base_alpha + else: + index: int = getattr(model_A, ATTR_INDEX) + return alphas_list[index][block] + + def replace(parent_name: str, mod_A: torch.nn.Module, mod_B: torch.nn.Module, alpha: Callable[[str],float]): + for name, a in list(mod_A.named_children()): + b = getattr(mod_B, name, None) + if b is None: + continue + + long_name = f'{parent_name}.{name}' if len(parent_name) != 0 else name + if not hasattr(a, 'weight') and not hasattr(b, 'weight'): + replace(long_name, a, b, alpha) + + elif hasattr(a, 'weight') and hasattr(b, 'weight'): + setattr(mod_A, name, MergedModule(long_name, a, b, alpha)) + + else: + a_with = 'with' if hasattr(a, 'weight') else 'without' + b_with = 'with' if hasattr(b, 'weight') else 'without' + print(f'mismatch: model_A has key {long_name} {a_with} weights, and model_B {b_with} weights.') + + replace('', model_A, model_B, alpha_fn) + + +class StateDictMergerBlockWeightedMulti: + + @classmethod + def INPUT_TYPES(cls): + d = StateDictMergerBlockWeighted.INPUT_TYPES() + d['required']['config_name'] = (folder_paths.get_filename_list('configs'), ) + return d + + RETURN_TYPES = ('MODEL','CLIP','VAE') + + FUNCTION = 'execute' + + CATEGORY = 'model' + + def execute( + self, + model_A: Dict[str,torch.Tensor], + model_B: Dict[str,torch.Tensor], + position_ids: str, + half: str, + base_alpha: float, + alphas: str, + config_name: str, + ): + alphas_list = self.get_alphas(alphas) + + clip_vae = self.merge_clip_vae(model_A, model_B, base_alpha, position_ids, half) + + modelA, clipA, vaeA = self.get_model(model_A, config_name) + modelB, clipB, vaeB = self.get_model(model_B, config_name) + + class WeightLoader(torch.nn.Module): + pass + + w = WeightLoader() + w.cond_stage_model = clipA.cond_stage_model + w.first_stage_model = vaeA.first_stage_model + w.load_state_dict(clip_vae, strict=False) + + mbw_on_the_fly(modelA.model, modelB.model, alphas_list, base_alpha) + + model_fn = iterize_model(modelA) + model_fn.clear() + for index in range(len(alphas_list)): + def fn(index=index): + setattr(modelA.model, ATTR_INDEX, index) + return modelA + model_fn.append(fn) + + return (modelA, clipA, vaeA) + + def get_alphas(self, alphas: str): + alphas_line = [ [ float(x.strip()) for x in line.strip().split(',') if 0 < len(x.strip()) ] for line in alphas.split('\n') ] + alphas_line = list(filter(lambda vs: len(vs) != 0, alphas_line)) # ignore empty line + + for row, line in enumerate(alphas_line, 1): + if len(line) != 25: + raise ValueError(f'line {row}: given {len(line)} values, expected 25.') + + return alphas_line + + def merge_clip_vae( + self, + model_A: Dict[str,torch.Tensor], + model_B: Dict[str,torch.Tensor], + base_alpha: float, + position_ids: str, + half: str + ): + def filter_(dic, ss): + return { k: v for k, v in dic.items() if any(k.startswith(s) for s in ss) } + + clip_vae_A = filter_(model_A, ['cond_stage_model', 'first_stage_model']) + clip_vae_B = filter_(model_B, ['cond_stage_model', 'first_stage_model']) + + clip_vae = weighted_sum_block(clip_vae_A, clip_vae_B, base_alpha, [0]*25, position_ids, half) + return clip_vae + + def get_model(self, model: Dict[str,torch.Tensor], config_name: str): + return Dict2Model().execute(model, config_name) diff --git a/sample.py b/sample.py index d9ed7bd..0013dc3 100644 --- a/sample.py +++ b/sample.py @@ -7,6 +7,7 @@ import comfy.samplers from nodes import common_ksampler from comfy.sd import ModelPatcher from .model.iter import iterize_model, CondForModels +from .model import merge2 re_int = re.compile(r"\s*([+-]?\s*\d+)\s*") re_float = re.compile(r"\s*([+-]?\s*\d+(?:.\d*)?)\s*") @@ -210,6 +211,10 @@ def common_ksampler_xyz( sampler = comfy.samplers.KSampler(**sampler_args) print(f'XYZ sampler=model@{model_index}/{sampler.sampler}/{sampler.scheduler} {sampler.steps}steps') + alphas = merge2.get_current_alpha(model_.model) + if alphas is not None: + print(f'alpha = {alphas}') + 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() all_samples.append(samples)