replaced modelMerge with multiModelMerge node

This commit is contained in:
TinyTerra
2023-09-25 22:38:56 +10:00
parent eea645a8aa
commit d10afc025c
3 changed files with 165 additions and 41 deletions
+1 -1
View File
@@ -1,3 +1,3 @@
config.ini
nsp_pantry.json
__pycache__
+12 -31
View File
@@ -156,38 +156,18 @@ function widgetLogic(node, widget) {
toggleWidget(node, findWidgetByName(node, 'control_after_generate'), true)
}
}
if (widget.name === 'model_merge_type') {
if (widget.value === "Simple") {
toggleWidget(node, findWidgetByName(node, 'simple_ratio'), true)
toggleWidget(node, findWidgetByName(node, 'block_input'))
toggleWidget(node, findWidgetByName(node, 'block_middle'))
toggleWidget(node, findWidgetByName(node, 'block_out'))
toggleWidget(node, findWidgetByName(node, 'subtract_multiplier'))
} else if (widget.value === "Blocks") {
toggleWidget(node, findWidgetByName(node, 'simple_ratio'))
toggleWidget(node, findWidgetByName(node, 'block_input'), true)
toggleWidget(node, findWidgetByName(node, 'block_middle'), true)
toggleWidget(node, findWidgetByName(node, 'block_out'), true)
toggleWidget(node, findWidgetByName(node, 'subtract_multiplier'))
} else if (widget.value === "Subtract") {
toggleWidget(node, findWidgetByName(node, 'simple_ratio'))
toggleWidget(node, findWidgetByName(node, 'block_input'))
toggleWidget(node, findWidgetByName(node, 'block_middle'))
toggleWidget(node, findWidgetByName(node, 'block_out'))
toggleWidget(node, findWidgetByName(node, 'subtract_multiplier'), true)
} else if (["Add", "model1", "model2"].includes(widget.value)) {
toggleWidget(node, findWidgetByName(node, 'simple_ratio'))
toggleWidget(node, findWidgetByName(node, 'block_input'))
toggleWidget(node, findWidgetByName(node, 'block_middle'))
toggleWidget(node, findWidgetByName(node, 'block_out'))
toggleWidget(node, findWidgetByName(node, 'subtract_multiplier'))
if (widget.name === 'ckpt_B_name') {
if (widget.value === "None") {
toggleWidget(node, findWidgetByName(node, 'config_B_name'))
} else {
toggleWidget(node, findWidgetByName(node, 'config_B_name'), true)
}
}
if (widget.name === 'clip_merge_type') {
if (widget.value === "Simple") {
toggleWidget(node, findWidgetByName(node, 'clip_simple_ration'), true)
if (widget.name === 'ckpt_C_name') {
if (widget.value === "None") {
toggleWidget(node, findWidgetByName(node, 'config_C_name'))
} else {
toggleWidget(node, findWidgetByName(node, 'clip_simple_ration'))
toggleWidget(node, findWidgetByName(node, 'config_C_name'), true)
}
}
if (widget.name === 'save_model') {
@@ -203,7 +183,8 @@ function widgetLogic(node, widget) {
const getSetWidgets = ['rescale_after_model', 'rescale', 'image_output',
'lora_name', 'lora1_name', 'lora2_name', 'lora3_name',
'refiner_lora1_name', 'refiner_lora2_name', 'upscale_method',
'image_output', 'add_noise', 'model_merge_type', 'clip_merge_type', 'save_model']
'image_output', 'add_noise',
'ckpt_B_name', 'ckpt_C_name', 'save_model']
function getSetters(node) {
if (node.widgets)
@@ -240,7 +221,7 @@ app.registerExtension({
node.getTitle() == "pipeKSamplerSDXL" ||
node.getTitle() == "imageRemBG" ||
node.getTitle() == "imageOutput"||
node.getTitle() == "modelMerge") {
node.getTitle() == "multiModelMerge") {
getSetters(node)
}
}
+152 -9
View File
@@ -18,9 +18,11 @@ import numpy as np
import folder_paths
import comfy.samplers
import latent_preview
import comfy.model_base
from pathlib import Path
import comfy.model_management
from comfy.sd import CLIP, VAE
from comfy.cli_args import args
from urllib.request import urlopen
from collections import defaultdict
from PIL.PngImagePlugin import PngInfo
@@ -1042,8 +1044,6 @@ class ttN_TSC_pipeLoader:
ttNcache.update_loaded_objects(prompt)
# Load models
if config_name == "Default":
config_name = None
model, clip, vae = ttNcache.load_checkpoint(ckpt_name, config_name)
@@ -2237,8 +2237,6 @@ class ttN_modelMerge:
if clip_override is not None:
clip = clip_override
print("VAE:",vae)
if vae == "bvae1":
vae = vae1
elif vae == "bvae2":
@@ -2247,12 +2245,154 @@ class ttN_modelMerge:
if vae_override is not None:
vae = vae_override
print(model, clip, vae)
if save_model == "True":
Nmm.CheckpointSave.save(Nmm.CheckpointSave(), model, clip, vae, save_prefix, prompt, extra_pnginfo)
return (model, vae, clip)
class ttN_multiModelMerge:
version = '1.0.0'
def __init__(self):
pass
@classmethod
def INPUT_TYPES(s):
return {"required": {
"ckpt_A_name": (folder_paths.get_filename_list("checkpoints"), ),
"config_A_name": (["Default",] + folder_paths.get_filename_list("configs"), {"default": "Default"} ),
"ckpt_B_name": (["None",] + folder_paths.get_filename_list("checkpoints"), ),
"config_B_name": (["Default",] + folder_paths.get_filename_list("configs"), {"default": "Default"} ),
"ckpt_C_name": (["None",] + folder_paths.get_filename_list("checkpoints"), ),
"config_C_name": (["Default",] + folder_paths.get_filename_list("configs"), {"default": "Default"} ),
"interpolation": (["None (A only)", "Weighted sum = ( A*(1-M) + B*M )", "Add difference = ( A + (B-C)*M )"],),
"multiplier": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}),
"save_model": (["True", "False"],),
"save_prefix": ("STRING", {"default": "checkpoints/ComfyUI"}),
},
"optional": {
"model_A_override": ("MODEL",),
"model_B_override": ("MODEL",),
"model_C_override": ("MODEL",),
"clip_A_override": ("CLIP",),
"clip_B_override": ("CLIP",),
"clip_C_override": ("CLIP",),
"optional_vae": ("VAE",),
},
"hidden": {"prompt": "PROMPT", "extra_pnginfo": "EXTRA_PNGINFO", "ttNnodeVersion": ttN_modelMerge.version},
}
RETURN_TYPES = ("MODEL", "CLIP", "VAE",)
RETURN_NAMES = ("model", "clip", "vae",)
FUNCTION = "mergificate"
CATEGORY = "ttN"
def mergificate(self, ckpt_A_name, config_A_name, ckpt_B_name, config_B_name, ckpt_C_name, config_C_name,
interpolation, multiplier, save_model, save_prefix,
model_A_override=None, model_B_override=None, model_C_override=None,
clip_A_override=None, clip_B_override=None, clip_C_override=None,
optional_vae=None, prompt=None, extra_pnginfo=None):
def compile_letter(model_o, clip_o, ckpt_name, config_name):
if model_o and clip_o:
return model_o, clip_o
else:
model, clip, _ = ttNcache.load_checkpoint(ckpt_name, config_name)
model = model_o if model_o is not None else model
clip = clip_o if clip_o is not None else clip
return model, clip
def merge(base_model, base_strength, patch_model, patch_strength):
m = base_model.clone()
kp = patch_model.get_key_patches("diffusion_model.")
for k in kp:
m.add_patches({k: kp[k]}, patch_strength, base_strength)
return m
def clip_merge(base_clip, base_strength, patch_clip, patch_strength):
m = base_clip.clone()
kp = patch_clip.get_key_patches()
for k in kp:
if k.endswith(".position_ids") or k.endswith(".logit_scale"):
continue
m.add_patches({k: kp[k]}, patch_strength, base_strength)
return m
if (interpolation == "None (A only)") or ((ckpt_B_name == "None") and (model_B_override == None)):
model, clip = compile_letter(model_A_override, clip_A_override, ckpt_A_name, config_A_name)
elif (interpolation == "Weighted sum = ( A*(1-M) + B*M )") or ((ckpt_C_name == "None") and (model_C_override == None)):
model_A, clip_A = compile_letter(model_A_override, clip_A_override, ckpt_A_name, config_A_name)
model_B, clip_B = compile_letter(model_B_override, clip_B_override, ckpt_B_name, config_B_name)
model = merge(model_A, (1.0 - multiplier), model_B, multiplier)
clip = clip_merge(clip_A, (1.0 - multiplier), clip_B, multiplier)
elif (interpolation == "Add difference = ( A + (B-C)*M )"):
model_A, clip_A = compile_letter(model_A_override, clip_A_override, ckpt_A_name, config_A_name)
model_B, clip_B = compile_letter(model_B_override, clip_B_override, ckpt_B_name, config_B_name)
model_C, clip_C = compile_letter(model_C_override, clip_C_override, ckpt_C_name, config_C_name)
B_minus_C = merge(model_B, 1.0, model_C, -1.0)
model = merge(model_A, 1.0, B_minus_C, multiplier)
B_minus_C = clip_merge(clip_B, 1.0, clip_C, -1.0)
clip = clip_merge(clip_A, 1.0, B_minus_C, multiplier)
if optional_vae not in ["None", None]:
vae_sd = optional_vae.get_sd()
vae = optional_vae
else:
vae_sd = {}
vae = None
if save_model == "True":
full_output_folder, filename, counter, subfolder, save_prefix = folder_paths.get_save_image_path(save_prefix, folder_paths.get_output_directory())
prompt_info = ""
if prompt is not None:
prompt_info = json.dumps(prompt)
metadata = {}
enable_modelspec = True
if isinstance(model.model, comfy.model_base.SDXL):
metadata["modelspec.architecture"] = "stable-diffusion-xl-v1-base"
elif isinstance(model.model, comfy.model_base.SDXLRefiner):
metadata["modelspec.architecture"] = "stable-diffusion-xl-v1-refiner"
else:
enable_modelspec = False
if enable_modelspec:
metadata["modelspec.sai_model_spec"] = "1.0.0"
metadata["modelspec.implementation"] = "sgm"
metadata["modelspec.title"] = "{} {}".format(filename, counter)
if model.model.model_type == comfy.model_base.ModelType.EPS:
metadata["modelspec.predict_key"] = "epsilon"
elif model.model.model_type == comfy.model_base.ModelType.V_PREDICTION:
metadata["modelspec.predict_key"] = "v"
if not args.disable_metadata:
metadata["prompt"] = prompt_info
if extra_pnginfo is not None:
for x in extra_pnginfo:
metadata[x] = json.dumps(extra_pnginfo[x])
output_checkpoint = f"{filename}_{counter:05}_.safetensors"
output_checkpoint = os.path.join(full_output_folder, output_checkpoint)
comfy.model_management.load_models_gpu([model, clip.load_model()])
sd = model.model.state_dict_for_saving(clip.get_sd(), vae_sd)
comfy.utils.save_torch_file(sd, output_checkpoint, metadata=metadata)
return (model, clip, vae)
#---------------------------------------------------------------ttN/text START----------------------------------------------------------------------#
class ttN_text:
version = '1.0.0'
@@ -2750,7 +2890,8 @@ TTN_VERSIONS = {
"pipe2DETAILER": ttN_pipe_2DETAILER.version,
"xyPlot": ttN_XYPlot.version,
"pipeEncodeConcat": ttN_pipeEncodeConcat.version,
"modelMerge": ttN_modelMerge.version,
#"modelMerge": ttN_modelMerge.version,
"multiModelMerge": ttN_multiModelMerge.version,
"text": ttN_text.version,
"textDebug": ttN_textDebug.version,
"concat": ttN_concat.version,
@@ -2779,7 +2920,8 @@ NODE_CLASS_MAPPINGS = {
#ttN/encode
"ttN pipeEncodeConcat": ttN_pipeEncodeConcat,
"ttN modelMerge": ttN_modelMerge,
#"ttN modelMerge": ttN_modelMerge,
"ttN multiModelMerge": ttN_multiModelMerge,
#ttN/text
"ttN text": ttN_text,
@@ -2814,7 +2956,8 @@ NODE_DISPLAY_NAME_MAPPINGS = {
#ttN/encode
"ttN pipeEncodeConcat": "pipeEncodeConcat",
"ttN modelMerge": "modelMerge",
#"ttN modelMerge": "modelMerge",
"ttN multiModelMerge": "multiModelMerge",
#ttN/text
"ttN text": "text",