diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000..b3be918 --- /dev/null +++ b/.gitignore @@ -0,0 +1,7 @@ +__pycache__/ +*.pyc +.ipynb_checkpoints/ +*.ipynb +scripts/batch_condition/train/ +scripts/batch_condition/test/ +scripts/reference/cache/ diff --git a/__init__.py b/__init__.py index 9bce322..5aa4534 100644 --- a/__init__.py +++ b/__init__.py @@ -1,4 +1,5 @@ import importlib +import os NODE_CLASS_MAPPINGS = {} NODE_DISPLAY_NAME_MAPPINGS = {} @@ -19,12 +20,25 @@ scripts = [ "custom_guiders", "custom_noise", "scale_crafter", - "aesthetic_shadow", "for_test", "lora_xy", "reference", ] +def import_from_package(module_name, module_path): + if os.path.isdir(module_path): + module_file = os.path.join(module_path, "__init__.py") + else: + module_file = module_path + + if not os.path.exists(module_file): + raise FileNotFoundError(f"{module_file} not found") + + spec = importlib.util.spec_from_file_location(module_name, module_file) + module = importlib.util.module_from_spec(spec) + spec.loader.exec_module(module) + return module + try: import timm except ImportError: @@ -33,7 +47,8 @@ else: scripts.append("wd-tagger") for script in scripts: - module = importlib.import_module(f"custom_nodes.cgem156-ComfyUI.scripts.{script}") + #module = importlib.import_module(f"custom_nodes.cgem156-ComfyUI.scripts.{script}") + module = import_from_package(f"custom_nodes.cgem156-ComfyUI.scripts.{script}", os.path.join(os.path.dirname(os.path.abspath(__file__)), "scripts", script)) if hasattr(module, 'NODE_CLASS_MAPPINGS'): NODE_CLASS_MAPPINGS.update(getattr(module, 'NODE_CLASS_MAPPINGS')) if hasattr(module, 'NODE_DISPLAY_NAME_MAPPINGS'): diff --git a/js/attention_couple.js b/js/attention_couple.js deleted file mode 100644 index f70171b..0000000 --- a/js/attention_couple.js +++ /dev/null @@ -1,37 +0,0 @@ -import { app } from "/scripts/app.js"; - -app.registerExtension({ - name: "AttentionCouple|cgem156", - async beforeRegisterNodeDef(nodeType, nodeData) { - if (nodeData.name === "AttentionCouple|cgem156") { - const origGetExtraMenuOptions = nodeType.prototype.getExtraMenuOptions; - nodeType.prototype.getExtraMenuOptions = function (_, options) { - const r = origGetExtraMenuOptions?.apply?.(this, arguments); - options.unshift( - { - content: "add input", - callback: () => { - var index = 1; - if (this.inputs != undefined){ - index += this.inputs.length; - } - this.addInput("cond_" + Math.floor(index / 2), "CONDITIONING"); - this.addInput("mask_" + Math.floor(index / 2), "MASK"); - }, - }, - { - content: "remove input", - callback: () => { - if (this.inputs != undefined){ - this.removeInput(this.inputs.length - 1); - this.removeInput(this.inputs.length - 1); - } - }, - }, - ); - return r; - - } - } - }, -}); \ No newline at end of file diff --git a/js/batch_condition.js b/js/batch_condition.js deleted file mode 100644 index 8669a5f..0000000 --- a/js/batch_condition.js +++ /dev/null @@ -1,36 +0,0 @@ -//ref: https://note.com/nyaoki_board/n/na7c54c9ae2a5 - -import { app } from "/scripts/app.js"; - -app.registerExtension({ - name: "BatchString|cgem156", - async beforeRegisterNodeDef(nodeType, nodeData, app) { - if (nodeData.name === "BatchString|cgem156") { - const origGetExtraMenuOptions = nodeType.prototype.getExtraMenuOptions; - nodeType.prototype.getExtraMenuOptions = function (_, options) { - const r = origGetExtraMenuOptions?.apply?.(this, arguments); - options.unshift( - { - content: "add input", - callback: () => { - var index = 1; - if (this.inputs != undefined){ - index += this.inputs.length; - } - this.addInput("text" + index, "STRING", {"multiline": true}); - }, - }, - { - content: "remove input", - callback: () => { - if (this.inputs != undefined){ - this.removeInput(this.inputs.length - 1); - } - }, - }, - ); - return r; - } - } - }, -}); \ No newline at end of file diff --git a/requirements.txt b/requirements.txt new file mode 100644 index 0000000..365f537 --- /dev/null +++ b/requirements.txt @@ -0,0 +1,5 @@ +transformers +timm +pandas +opencv-python +matplotlib diff --git a/scripts/attention_couple/node.py b/scripts/attention_couple/node.py index e8a1739..9846fb6 100644 --- a/scripts/attention_couple/node.py +++ b/scripts/attention_couple/node.py @@ -2,10 +2,15 @@ import torch import torch.nn.functional as F import comfy import math -from ... import ROOT_NAME +from types import SimpleNamespace +from comfy_api.v0_0_2 import io +from ... import ROOT_NAME, NODE_SURFIX, SYMBOL CATEGORY_NAME = ROOT_NAME + "attention_couple" +# Max number of extra cond/mask pairs the UI can grow to via Autogrow. +MAX_PAIRS = 50 + def get_mask(mask, batch_size, num_tokens, original_shape): num_conds = mask.shape[0] @@ -33,49 +38,95 @@ def lcm_for_list(numbers): current_lcm = lcm(current_lcm, number) return current_lcm -class AttentionCouple: +class AttentionCouple(io.ComfyNode): + # NOTE on workflow compatibility: the old V1 node exposed a fixed + # model/base_mask schema and relied on js/attention_couple.js to add + # cond_N (CONDITIONING) / mask_N (MASK) input pairs client-side beyond + # what INPUT_TYPES declared, consumed via an unbounded **kwargs pattern. + # This migrates to the official V3 Autogrow dynamic-input API using two + # parallel Autogrow.TemplateNames templates (one for "cond_N", one for + # "mask_N"), with explicit 1-indexed names so the resolved kwarg names + # match the old JS-generated names exactly (cond_1/mask_1, cond_2/mask_2, + # ...). Old workflows that used pairs within MAX_PAIRS should therefore + # reconnect by name; see the migration report for the caveats (fixed + # upper bound, and cond_N/mask_N no longer forced to be added/removed as + # a strict pair by the UI). + @classmethod + def define_schema(cls) -> io.Schema: + cond_template = io.Autogrow.TemplateNames( + input=io.Conditioning.Input("cond"), + names=[f"cond_{i}" for i in range(1, MAX_PAIRS + 1)], + min=0, + ) + mask_template = io.Autogrow.TemplateNames( + input=io.Mask.Input("mask"), + names=[f"mask_{i}" for i in range(1, MAX_PAIRS + 1)], + min=0, + ) + return io.Schema( + node_id=f"AttentionCouple{NODE_SURFIX}", + display_name=f"Attention Couple {SYMBOL}", + category=CATEGORY_NAME, + inputs=[ + io.Model.Input("model"), + io.Mask.Input("base_mask"), + io.Autogrow.Input("conds", template=cond_template), + io.Autogrow.Input("masks", template=mask_template), + ], + outputs=[ + io.Model.Output(), + ], + ) @classmethod - def INPUT_TYPES(s): - return { - "required": { - "model": ("MODEL", ), - "base_mask": ("MASK",), - } - } - RETURN_TYPES = ("MODEL", ) - FUNCTION = "attention_couple_simple" - CATEGORY = CATEGORY_NAME - - def attention_couple_simple(self, model, base_mask, **kwargs): - + def execute(cls, model, base_mask, conds: io.Autogrow.Type, masks: io.Autogrow.Type) -> io.NodeOutput: new_model = model.clone() - num_conds = len(kwargs) // 2 + 1 - mask = [base_mask] + [kwargs[f"mask_{i}"] for i in range(1, num_conds)] + # Unlike the old JS UI (which always added/removed cond_i/mask_i as + # a pair), the two Autogrow blocks now grow independently, so a + # workflow could connect cond_i without mask_i (or vice versa). + # Fail fast with a clear message instead of silently misaligning + # tensors further down. + cond_indices = {name.split("_", 1)[1] for name in conds} + mask_indices = {name.split("_", 1)[1] for name in masks} + assert cond_indices == mask_indices, ( + f"Mismatched cond_N/mask_N inputs: conds={sorted(conds)}, masks={sorted(masks)}. " + "Every connected cond_N input must have a matching mask_N input, and vice versa." + ) + num_conds = len(conds) + 1 + + mask = [base_mask] + list(masks.values()) mask = torch.stack(mask, dim=0) assert mask.sum(dim=0).min() > 0, "There are areas that are zero in all masks." - self.mask = mask / mask.sum(dim=0, keepdim=True) - self.conds = [kwargs[f"cond_{i}"][0][0] for i in range(1, num_conds)] - num_tokens = [cond.shape[1] for cond in self.conds] + # execute() is a classmethod (no `self`), so the mutable state that + # attn2_patch/attn2_output_patch share across repeated calls (device + # caching, batch_size handoff) lives on this small namespace instead + # of on a node instance. This is a structural translation only; the + # attention-patching math below is unchanged from the V1 node. + state = SimpleNamespace( + mask=mask / mask.sum(dim=0, keepdim=True), + conds=[cond[0][0] for cond in conds.values()], + batch_size=None, + ) + num_tokens = [cond.shape[1] for cond in state.conds] def attn2_patch(q, k, v, extra_options): assert k.mean() == v.mean(), "k and v must be the same." device, dtype = q.device, q.dtype - - if self.conds[0].device != device: - self.conds = [cond.to(device, dtype=dtype) for cond in self.conds] - if self.mask.device != device: - self.mask = self.mask.to(device, dtype=dtype) + + if state.conds[0].device != device: + state.conds = [cond.to(device, dtype=dtype) for cond in state.conds] + if state.mask.device != device: + state.mask = state.mask.to(device, dtype=dtype) cond_or_unconds = extra_options["cond_or_uncond"] num_chunks = len(cond_or_unconds) - self.batch_size = q.shape[0] // num_chunks + state.batch_size = q.shape[0] // num_chunks q_chunks = q.chunk(num_chunks, dim=0) k_chunks = k.chunk(num_chunks, dim=0) lcm_tokens = lcm_for_list(num_tokens + [k.shape[1]]) - conds_tensor = torch.cat([cond.repeat(self.batch_size, lcm_tokens // num_tokens[i], 1) for i, cond in enumerate(self.conds)], dim=0) + conds_tensor = torch.cat([cond.repeat(state.batch_size, lcm_tokens // num_tokens[i], 1) for i, cond in enumerate(state.conds)], dim=0) qs, ks = [], [] for i, cond_or_uncond in enumerate(cond_or_unconds): @@ -95,22 +146,22 @@ class AttentionCouple: def attn2_output_patch(out, extra_options): cond_or_unconds = extra_options["cond_or_uncond"] - mask_downsample = get_mask(self.mask, self.batch_size, out.shape[1], extra_options["original_shape"]) + mask_downsample = get_mask(state.mask, state.batch_size, out.shape[1], extra_options["original_shape"]) outputs = [] pos = 0 for cond_or_uncond in cond_or_unconds: if cond_or_uncond == 1: # uncond - outputs.append(out[pos:pos + self.batch_size]) - pos += self.batch_size + outputs.append(out[pos:pos + state.batch_size]) + pos += state.batch_size else: - masked_output = (out[pos:pos + num_conds * self.batch_size] * mask_downsample).view(num_conds, self.batch_size, out.shape[1], out.shape[2]) + masked_output = (out[pos:pos + num_conds * state.batch_size] * mask_downsample).view(num_conds, state.batch_size, out.shape[1], out.shape[2]) masked_output = masked_output.sum(dim=0) outputs.append(masked_output) - pos += num_conds * self.batch_size + pos += num_conds * state.batch_size return torch.cat(outputs, dim=0) new_model.set_model_attn2_patch(attn2_patch) new_model.set_model_attn2_output_patch(attn2_output_patch) - return (new_model, ) + return io.NodeOutput(new_model) diff --git a/scripts/batch_condition/__init__.py b/scripts/batch_condition/__init__.py index 3159936..de3015e 100644 --- a/scripts/batch_condition/__init__.py +++ b/scripts/batch_condition/__init__.py @@ -1,4 +1,4 @@ -from .node import CLIPTextEncodeBatch, StringInput, BatchString, PrefixString, SaveBatchString, SaveImageBatch, SaveLatentBatch +from .node import CLIPTextEncodeBatch, StringInput, BatchString, PrefixString, SaveBatchString, SaveImageBatch, SaveLatentBatch, RandomColorPrompt from ... import NODE_SURFIX, SYMBOL NODE_CLASS_MAPPINGS = { @@ -8,7 +8,8 @@ NODE_CLASS_MAPPINGS = { f"PrefixString{NODE_SURFIX}": PrefixString, f"SaveBatchString{NODE_SURFIX}": SaveBatchString, f"SaveImageBatch{NODE_SURFIX}": SaveImageBatch, - f"SaveLatentBatch{NODE_SURFIX}": SaveLatentBatch + f"SaveLatentBatch{NODE_SURFIX}": SaveLatentBatch, + f"RandomColorPrompt{NODE_SURFIX}": RandomColorPrompt } NODE_DISPLAY_NAME_MAPPINGS = { @@ -18,7 +19,8 @@ NODE_DISPLAY_NAME_MAPPINGS = { f"PrefixString{NODE_SURFIX}": f"Prefix String {SYMBOL}", f"SaveBatchString{NODE_SURFIX}": f"Save Batch String {SYMBOL}", f"SaveImageBatch{NODE_SURFIX}": f"Save Image Batch {SYMBOL}", - f"SaveLatentBatch{NODE_SURFIX}": f"Save Latent Batch {SYMBOL}" + f"SaveLatentBatch{NODE_SURFIX}": f"Save Latent Batch {SYMBOL}", + f"RandomColorPrompt{NODE_SURFIX}": f"Random Color Prompt {SYMBOL}" } __all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"] \ No newline at end of file diff --git a/scripts/batch_condition/node.py b/scripts/batch_condition/node.py index 273d469..b0020ab 100644 --- a/scripts/batch_condition/node.py +++ b/scripts/batch_condition/node.py @@ -3,7 +3,8 @@ import numpy as np import math import os from PIL import Image -from ... import ROOT_NAME +from comfy_api.v0_0_2 import io +from ... import ROOT_NAME, NODE_SURFIX, SYMBOL CURRENT_DIR = os.path.dirname(os.path.realpath(__file__)) @@ -18,20 +19,24 @@ def lcm_for_list(numbers): current_lcm = lcm(current_lcm, number) return current_lcm -class CLIPTextEncodeBatch: +class CLIPTextEncodeBatch(io.ComfyNode): @classmethod - def INPUT_TYPES(s): - return { - "required": { - "clip": ("CLIP", ), - "texts":("BATCH_STRING", ) - } - } - RETURN_TYPES = ("CONDITIONING",) - FUNCTION = "encode" - CATEGORY = CATEGORY_NAME + def define_schema(cls) -> io.Schema: + return io.Schema( + node_id=f"CLIPTextEncodeBatch{NODE_SURFIX}", + display_name=f"CLIP Text Encode Batch {SYMBOL}", + category=CATEGORY_NAME, + inputs=[ + io.Clip.Input("clip"), + io.Custom("BATCH_STRING").Input("texts"), + ], + outputs=[ + io.Conditioning.Output(), + ], + ) - def encode(self, clip, texts): + @classmethod + def execute(cls, clip, texts) -> io.NodeOutput: conds = [] pooleds = [] num_tokens = [] @@ -41,129 +46,159 @@ class CLIPTextEncodeBatch: conds.append(cond) pooleds.append(pooled) num_tokens.append(cond.shape[1]) - + # Make number of tokens equal # attn(q, k, v) == attn(q, [k]*n, [v]*n) lcm = lcm_for_list(num_tokens) repeats = [lcm//num for num in num_tokens] conds = torch.cat([cond.repeat(1, repeat, 1) for cond, repeat in zip(conds, repeats)]) pooleds = torch.cat(pooleds) - return ([[conds, {"pooled_output": pooleds}]], ) - -class StringInput: + return io.NodeOutput([[conds, {"pooled_output": pooleds}]]) + +class StringInput(io.ComfyNode): @classmethod - def INPUT_TYPES(s): - return { - "required": - { - "text": ("STRING", {"multiline": True}) - } - } - RETURN_TYPES = ("STRING",) - FUNCTION = "encode" + def define_schema(cls) -> io.Schema: + return io.Schema( + node_id=f"StringInput{NODE_SURFIX}", + display_name=f"String Input {SYMBOL}", + category=CATEGORY_NAME, + inputs=[ + io.String.Input("text", multiline=True), + ], + outputs=[ + io.String.Output(), + ], + ) - CATEGORY = CATEGORY_NAME - - def encode(self, text): - return (text, ) - -class BatchString: @classmethod - def INPUT_TYPES(s): - return {"required": {}} - RETURN_TYPES = ("BATCH_STRING",) - FUNCTION = "encode" + def execute(cls, text) -> io.NodeOutput: + return io.NodeOutput(text) - CATEGORY = CATEGORY_NAME +class BatchString(io.ComfyNode): + # NOTE on workflow compatibility: the old V1 node relied on + # js/batch_condition.js to add "text{n}" STRING widget-inputs + # client-side beyond what INPUT_TYPES declared, consumed via an + # unbounded **kwargs pattern (encode() rebuilt the list from + # kwargs["text1"], kwargs["text2"], ...). This migrates to the official + # V3 Autogrow dynamic-input API with explicit names "text1".."textN" so + # the resolved kwarg names match the old JS-generated names exactly. + # Old workflows that used up to MAX_TEXTS inputs should therefore + # reconnect by name. + MAX_TEXTS = 50 - def encode(self, **kwargs): - return ([kwargs[f"text{i+1}"] for i in range(len(kwargs))], ) - -class PrefixString: @classmethod - def INPUT_TYPES(s): - return { - "required": { - "prefix": ("STRING", {"multiline": True}), - "prompts": ("BATCH_STRING", ) - } - } - RETURN_TYPES = ("BATCH_STRING",) - FUNCTION = "encode" + def define_schema(cls) -> io.Schema: + template = io.Autogrow.TemplateNames( + input=io.String.Input("text", multiline=True), + names=[f"text{i}" for i in range(1, cls.MAX_TEXTS + 1)], + min=0, + ) + return io.Schema( + node_id=f"BatchString{NODE_SURFIX}", + display_name=f"Batch String {SYMBOL}", + category=CATEGORY_NAME, + inputs=[ + io.Autogrow.Input("texts", template=template), + ], + outputs=[ + io.Custom("BATCH_STRING").Output(), + ], + ) - CATEGORY = CATEGORY_NAME - - def encode(self, prefix, prompts): - return ([prefix + prompt for prompt in prompts], ) - - -class SaveBatchString: @classmethod - def INPUT_TYPES(s): - return { - "required": { - "prompts": ("BATCH_STRING", ), - "folder": ("STRING", {"default": ""}), - "extension": ("STRING", {"default": "txt"}), - "seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}), - } - } - - RETURN_TYPES = () - FUNCTION = "save" - OUTPUT_NODE = True - CATEGORY = CATEGORY_NAME + def execute(cls, texts: io.Autogrow.Type) -> io.NodeOutput: + return io.NodeOutput(list(texts.values())) - def save(self, prompts, folder, extension, seed): +class PrefixString(io.ComfyNode): + @classmethod + def define_schema(cls) -> io.Schema: + return io.Schema( + node_id=f"PrefixString{NODE_SURFIX}", + display_name=f"Prefix String {SYMBOL}", + category=CATEGORY_NAME, + inputs=[ + io.String.Input("prefix", multiline=True), + io.Custom("BATCH_STRING").Input("prompts"), + ], + outputs=[ + io.Custom("BATCH_STRING").Output(), + ], + ) + + @classmethod + def execute(cls, prefix, prompts) -> io.NodeOutput: + return io.NodeOutput([prefix + prompt for prompt in prompts]) + +class SaveBatchString(io.ComfyNode): + @classmethod + def define_schema(cls) -> io.Schema: + return io.Schema( + node_id=f"SaveBatchString{NODE_SURFIX}", + display_name=f"Save Batch String {SYMBOL}", + category=CATEGORY_NAME, + inputs=[ + io.Custom("BATCH_STRING").Input("prompts"), + io.String.Input("folder", default=""), + io.String.Input("extension", default="txt"), + io.Int.Input("seed", default=0, min=0, max=0xffffffffffffffff), + ], + outputs=[], + is_output_node=True, + ) + + @classmethod + def execute(cls, prompts, folder, extension, seed) -> io.NodeOutput: os.makedirs(os.path.join(CURRENT_DIR, folder), exist_ok=True) for i, prompt in enumerate(prompts): path = os.path.join(CURRENT_DIR, folder, f"{seed:06}_{i:03}.{extension}") with open(path, "w") as f: f.write(prompt) - return {} - -class SaveImageBatch: - @classmethod - def INPUT_TYPES(s): - return { - "required": { - "images": ("IMAGE", ), - "folder": ("STRING", {"default": ""}), - "extension": ("STRING", {"default": "png"}), - "seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}), - } - } - - RETURN_TYPES = () - FUNCTION = "save" - OUTPUT_NODE = True - CATEGORY = CATEGORY_NAME + return io.NodeOutput() - def save(self, images, folder, extension, seed): +class SaveImageBatch(io.ComfyNode): + @classmethod + def define_schema(cls) -> io.Schema: + return io.Schema( + node_id=f"SaveImageBatch{NODE_SURFIX}", + display_name=f"Save Image Batch {SYMBOL}", + category=CATEGORY_NAME, + inputs=[ + io.Image.Input("images"), + io.String.Input("folder", default=""), + io.String.Input("extension", default="png"), + io.Int.Input("seed", default=0, min=0, max=0xffffffffffffffff), + ], + outputs=[], + is_output_node=True, + ) + + @classmethod + def execute(cls, images, folder, extension, seed) -> io.NodeOutput: os.makedirs(os.path.join(CURRENT_DIR, folder), exist_ok=True) for i, image in enumerate(images): path = os.path.join(CURRENT_DIR, folder, f"{seed:06}_{i:03}.{extension}") Image.fromarray((image.float().cpu() * 255).numpy().astype('uint8')).save(path) - return {} - -class SaveLatentBatch: - @classmethod - def INPUT_TYPES(s): - return { - "required": { - "latents": ("LATENT", ), - "folder": ("STRING", {"default": ""}), - "extension": (["npy", "npz"], {"default": "npy"}), - "seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}), - } - } - - RETURN_TYPES = () - FUNCTION = "save" - OUTPUT_NODE = True - CATEGORY = CATEGORY_NAME + return io.NodeOutput() - def save(self, latents, folder, extension, seed): +class SaveLatentBatch(io.ComfyNode): + @classmethod + def define_schema(cls) -> io.Schema: + return io.Schema( + node_id=f"SaveLatentBatch{NODE_SURFIX}", + display_name=f"Save Latent Batch {SYMBOL}", + category=CATEGORY_NAME, + inputs=[ + io.Latent.Input("latents"), + io.String.Input("folder", default=""), + io.Combo.Input("extension", options=["npy", "npz"], default="npy"), + io.Int.Input("seed", default=0, min=0, max=0xffffffffffffffff), + ], + outputs=[], + is_output_node=True, + ) + + @classmethod + def execute(cls, latents, folder, extension, seed) -> io.NodeOutput: os.makedirs(os.path.join(CURRENT_DIR, folder), exist_ok=True) for i, latent in enumerate(latents["samples"]): path = os.path.join(CURRENT_DIR, folder, f"{seed:06}_{i:03}.{extension}") @@ -173,9 +208,47 @@ class SaveLatentBatch: original_size = (latent.shape[1] * 8, latent.shape[2] * 8) crop_ltrb = (0, 0, 0, 0) np.savez( - path, + path, latents=latent.float().cpu().numpy(), original_size=np.array(original_size), crop_ltrb=np.array(crop_ltrb), ) - return {} + return io.NodeOutput() + +class RandomColorPrompt(io.ComfyNode): + MAGIC_WORD = "" + COLORS = [ + "red", "blue", "green", "yellow", "purple", "orange", "pink", "brown", + "black", "white", "gray", "aqua", + ] + + @classmethod + def define_schema(cls) -> io.Schema: + return io.Schema( + node_id=f"RandomColorPrompt{NODE_SURFIX}", + display_name=f"Random Color Prompt {SYMBOL}", + category=CATEGORY_NAME, + inputs=[ + io.String.Input("base_prompt", default="", multiline=True), + io.Int.Input("num_prompts", default=4, min=1, max=100), + io.Int.Input("seed", default=0, min=0, max=0xffffffffffffffff), + ], + outputs=[ + io.Custom("BATCH_STRING").Output(), + io.String.Output(), + ], + ) + + @classmethod + def execute(cls, base_prompt, num_prompts, seed) -> io.NodeOutput: + rng = np.random.RandomState(seed) + prompts = [] + for _ in range(num_prompts): + prompt = base_prompt + while cls.MAGIC_WORD in prompt: + color = rng.choice(cls.COLORS) + prompt = prompt.replace(cls.MAGIC_WORD, color, 1) + prompts.append(prompt) + + return_string = "\n\n".join(prompts) + return io.NodeOutput(prompts, return_string) diff --git a/scripts/cd_tuner/node.py b/scripts/cd_tuner/node.py index 01f8535..09b4e53 100644 --- a/scripts/cd_tuner/node.py +++ b/scripts/cd_tuner/node.py @@ -1,54 +1,31 @@ import torch +from comfy_api.v0_0_2 import io from ... import ROOT_NAME CATEGORY_NAME = ROOT_NAME + "cd-tuner" -class CDTuner: +class CDTuner(io.ComfyNode): @classmethod - def INPUT_TYPES(s): - return { - "required": { - "model": ("MODEL", ), - "detail_1": ("FLOAT", { - "default": 0, - "min": -10, - "max": 10, - "step": 0.1 - }), - "detail_2": ("FLOAT", { - "default": 0, - "min": -10, - "max": 10, - "step": 0.1 - }), - "contrast_1": ("FLOAT", { - "default": 0, - "min": -20, - "max": 20, - "step": 0.1 - }), - "start": ("INT", { - "default": 0, - "min": 0, - "max": 1000, - "step": 1, - "display": "number" - }), - "end": ("INT", { - "default": 1000, - "min": 0, - "max": 1000, - "step": 1, - "display": "number" - }), - }, - } + def define_schema(cls) -> io.Schema: + return io.Schema( + node_id="CD_Tuner|cgem156", + display_name="CD Tuner 🍌", + category=CATEGORY_NAME, + inputs=[ + io.Model.Input("model"), + io.Float.Input("detail_1", default=0, min=-10, max=10, step=0.1), + io.Float.Input("detail_2", default=0, min=-10, max=10, step=0.1), + io.Float.Input("contrast_1", default=0, min=-20, max=20, step=0.1), + io.Int.Input("start", default=0, min=0, max=1000, step=1, display_mode=io.NumberDisplay.number), + io.Int.Input("end", default=1000, min=0, max=1000, step=1, display_mode=io.NumberDisplay.number), + ], + outputs=[ + io.Model.Output(), + ], + ) - RETURN_TYPES = ("MODEL", ) - FUNCTION = "apply" - CATEGORY = CATEGORY_NAME - - def apply(self, model, detail_1, detail_2, contrast_1, start, end): + @classmethod + def execute(cls, model, detail_1, detail_2, contrast_1, start, end) -> io.NodeOutput: ''' detail_1: 最初のConv層のweightを減らしbiasを増やすことで、detailを増やす・・? detail_2: 最後のConv層前のGroupNormの以下略 @@ -56,37 +33,35 @@ class CDTuner: ''' new_model = model.clone() ratios = fineman([detail_1, detail_2, contrast_1]) - self.storedweights = {} - self.start = start - self.end = end + storedweights = {} # unet計算前後のパッチ def apply_cdtuner(model_function, kwargs): t = new_model.model.model_sampling.timestep(kwargs["timestep"]) - if t[0] < (1000 - self.end) or t[0] > (1000 - self.start): + if t[0] < (1000 - end) or t[0] > (1000 - start): return model_function(kwargs["input"], kwargs["timestep"], **kwargs["c"]) for i, name in enumerate(ADJUSTS): # 元の重みをロード - self.storedweights[name] = getset_nested_module_tensor(True, new_model, name).clone() + storedweights[name] = getset_nested_module_tensor(True, new_model, name).clone() if 4 > i: - new_weight = self.storedweights[name] * ratios[i] + new_weight = storedweights[name] * ratios[i] else: - device = self.storedweights[name].device - dtype = self.storedweights[name].dtype - new_weight = self.storedweights[name] + torch.tensor(ratios[i], device=device, dtype=dtype) + device = storedweights[name].device + dtype = storedweights[name].dtype + new_weight = storedweights[name] + torch.tensor(ratios[i], device=device, dtype=dtype) # 重みを書き換え getset_nested_module_tensor(False, new_model, name, new_tensor=new_weight) retval = model_function(kwargs["input"], kwargs["timestep"], **kwargs["c"]) # 重みを元に戻す for name in ADJUSTS: - getset_nested_module_tensor(False, new_model, name, new_tensor=self.storedweights[name]) + getset_nested_module_tensor(False, new_model, name, new_tensor=storedweights[name]) return retval new_model.set_model_unet_function_wrapper(apply_cdtuner) - return (new_model, ) + return io.NodeOutput(new_model) def getset_nested_module_tensor(clone, model, tensor_path, new_tensor=None): @@ -125,4 +100,3 @@ ADJUSTS = [ "model.diffusion_model.out.0.bias", "model.diffusion_model.out.2.bias", ] - diff --git a/scripts/custom_guiders/limited_interval_cfg_guider.py b/scripts/custom_guiders/limited_interval_cfg_guider.py index b0cc23b..8286001 100644 --- a/scripts/custom_guiders/limited_interval_cfg_guider.py +++ b/scripts/custom_guiders/limited_interval_cfg_guider.py @@ -1,4 +1,5 @@ import comfy +from comfy_api.v0_0_2 import io from ... import ROOT_NAME CATEGORY_NAME = ROOT_NAME + "custom_guiders" @@ -7,37 +8,39 @@ class LimitedIntervalCFG(comfy.samplers.CFGGuider): def set_range(self, sigma_low, sigma_high): self.sigma_low = sigma_low self.sigma_high = sigma_high - + def in_range(self, sigma): return self.sigma_low < sigma <= self.sigma_high - + def predict_noise(self, x, timestep, model_options={}, seed=None): cfg = self.cfg if self.in_range(timestep[0].item()) else 1 #print(f"CFG: {cfg} timestep: {timestep} sigma_low: {self.sigma_low} sigma_high: {self.sigma_high}") return comfy.samplers.sampling_function(self.inner_model, x, timestep, self.conds.get("negative", None), self.conds.get("positive", None), cfg, model_options=model_options, seed=seed) -class LimitedIntervalCFGGuider: +class LimitedIntervalCFGGuider(io.ComfyNode): @classmethod - def INPUT_TYPES(s): - return { - "required":{ - "model": ("MODEL",), - "positive": ("CONDITIONING", ), - "negative": ("CONDITIONING", ), - "cfg": ("FLOAT", {"default": 8.0, "min": 0.0, "max": 100.0, "step":0.1, "round": 0.01}), - "start_step": ("FLOAT", {"default": 0, "min": 0, "max": 1, "step": 0.001}), - "end_step": ("FLOAT", {"default": 1, "min": 0, "max": 1, "step": 0.001}), - } - } + def define_schema(cls) -> io.Schema: + return io.Schema( + node_id="LimitedIntervalCFGGuider|cgem156", + display_name="Limited Interval CFG Guider 🍌", + category=CATEGORY_NAME, + inputs=[ + io.Model.Input("model"), + io.Conditioning.Input("positive"), + io.Conditioning.Input("negative"), + io.Float.Input("cfg", default=8.0, min=0.0, max=100.0, step=0.1, round=0.01), + io.Float.Input("start_step", default=0, min=0, max=1, step=0.001), + io.Float.Input("end_step", default=1, min=0, max=1, step=0.001), + ], + outputs=[ + io.Guider.Output(), + ], + ) - RETURN_TYPES = ("GUIDER",) + @classmethod + def execute(cls, model, positive, negative, cfg, start_step, end_step) -> io.NodeOutput: - FUNCTION = "get_guider" - CATEGORY = CATEGORY_NAME - - def get_guider(self, model, positive, negative, cfg, start_step, end_step): - start_sigma = model.model.model_sampling.percent_to_sigma(start_step) end_sigma = model.model.model_sampling.percent_to_sigma(end_step) @@ -45,5 +48,4 @@ class LimitedIntervalCFGGuider: guider.set_conds(positive, negative) guider.set_cfg(cfg) guider.set_range(end_sigma, start_sigma) - return (guider,) - \ No newline at end of file + return io.NodeOutput(guider) diff --git a/scripts/custom_noise/__init__.py b/scripts/custom_noise/__init__.py index 848781e..314555b 100644 --- a/scripts/custom_noise/__init__.py +++ b/scripts/custom_noise/__init__.py @@ -1,16 +1,24 @@ from .variation_noise import VariationNoise, RandomNoiseOffset, RandomNoiseVariationSimple +from .short_distance_noise import ShortDistanceNoise, SameColorNoise +from .tkg_noise import TKGRandomNoise from ... import SYMBOL, NODE_SURFIX NODE_CLASS_MAPPINGS = { f"VariationNoise{NODE_SURFIX}": VariationNoise, f"RandomNoiseOffset{NODE_SURFIX}": RandomNoiseOffset, - f"RandomNoiseVariationSimple{NODE_SURFIX}": RandomNoiseVariationSimple + f"RandomNoiseVariationSimple{NODE_SURFIX}": RandomNoiseVariationSimple, + f"TKGRandomNoise{NODE_SURFIX}": TKGRandomNoise, + f"ShortDistanceNoise{NODE_SURFIX}": ShortDistanceNoise, + f"SameColorNoise{NODE_SURFIX}": SameColorNoise, } NODE_DISPLAY_NAME_MAPPINGS = { f"VariationNoise{NODE_SURFIX}": f"Variation Noise {SYMBOL}", f"RandomNoiseOffset{NODE_SURFIX}": f"Random Noise Offset {SYMBOL}", - f"RandomNoiseVariationSimple{NODE_SURFIX}": f"Random Noise Variation Simple {SYMBOL}" + f"RandomNoiseVariationSimple{NODE_SURFIX}": f"Random Noise Variation Simple {SYMBOL}", + f"TKGRandomNoise{NODE_SURFIX}": f"TKG Random Noise {SYMBOL}", + f"ShortDistanceNoise{NODE_SURFIX}": f"Short Distance Noise {SYMBOL}", + f"SameColorNoise{NODE_SURFIX}": f"Same Color Noise {SYMBOL}", } __all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"] \ No newline at end of file diff --git a/scripts/custom_noise/short_distance_noise.py b/scripts/custom_noise/short_distance_noise.py new file mode 100644 index 0000000..4b62a79 --- /dev/null +++ b/scripts/custom_noise/short_distance_noise.py @@ -0,0 +1,95 @@ +import comfy +from ... import ROOT_NAME, SYMBOL, NODE_SURFIX +import torch +from comfy_api.v0_0_2 import io +CATEGORY_NAME = ROOT_NAME + "custom_noise" + +class Noise_ShortDistance: + def __init__(self, seed, num_samples=32, reference_latents=None): + self.seed = seed + self.num_samples = num_samples + self.reference_latents = reference_latents + + def generate_noise(self, input_latent): + assert self.reference_latents["samples"].shape == input_latent["samples"].shape, "Reference latents and input latents must have the same shape." + latent = self.reference_latents["samples"].to(input_latent["samples"].device, dtype=input_latent["samples"].dtype) + batch_inds = input_latent.get("batch_index", None) + B = latent.shape[0] + K = self.num_samples + + latent_repeat = latent.unsqueeze(1).repeat(1, K, *[1 for _ in latent.shape[1:]]) + noise = comfy.sample.prepare_noise(latent_repeat, self.seed, batch_inds) + + diff = (latent_repeat - noise) ** 2 + dist = diff.flatten(start_dim=2).sum(dim=2) + best_idx = dist.argmin(dim=1) + + # gatherで最短ノイズを選択 + idx_expand = best_idx.view(B, 1, *[1 for _ in latent.shape[1:]]).expand_as(latent_repeat[:, :1]) + best_noise = noise.gather(1, idx_expand).squeeze(1) + + return best_noise + +class ShortDistanceNoise(io.ComfyNode): + @classmethod + def define_schema(cls) -> io.Schema: + return io.Schema( + node_id=f"ShortDistanceNoise{NODE_SURFIX}", + display_name=f"Short Distance Noise {SYMBOL}", + category=CATEGORY_NAME, + inputs=[ + io.Int.Input("seed", default=0, min=0, max=0xffffffffffffffff), + io.Int.Input("num_samples", default=32, min=1, max=4096), + io.Latent.Input("reference_latents"), + ], + outputs=[ + io.Noise.Output(), + ], + ) + + @classmethod + def execute(cls, seed, num_samples, reference_latents) -> io.NodeOutput: + return io.NodeOutput(Noise_ShortDistance(seed, reference_latents)) + +class Noise_SameColor: + def __init__(self, seed, reference_latents, strength, **kwargs): + self.seed = seed + self.reference_latents = reference_latents + self.strength = strength + self.channel_mask = torch.tensor([1.0 if kwargs.get(f"ch_{i:02d}", True) else 0.0 for i in range(16)]) + + def generate_noise(self, input_latent): + assert self.reference_latents["samples"].shape == input_latent["samples"].shape, "Reference latents and input latents must have the same shape." + latent = self.reference_latents["samples"].to(input_latent["samples"].device, dtype=input_latent["samples"].dtype) + batch_inds = input_latent.get("batch_index", None) + noise = comfy.sample.prepare_noise(latent, self.seed, batch_inds) + + latent_mean = latent.mean(dim=1, keepdim=True) + noise_mean = noise.mean(dim=1, keepdim=True) + channel_mask = self.channel_mask.to(latent.device, dtype=latent.dtype).view(1, -1, *[1 for _ in range(len(latent.shape)-2)]) + + noise = noise + (latent_mean - noise_mean) * self.strength * channel_mask + + return noise + +class SameColorNoise(io.ComfyNode): + @classmethod + def define_schema(cls) -> io.Schema: + return io.Schema( + node_id=f"SameColorNoise{NODE_SURFIX}", + display_name=f"Same Color Noise {SYMBOL}", + category=CATEGORY_NAME, + inputs=[ + io.Int.Input("seed", default=0, min=0, max=0xffffffffffffffff), + io.Latent.Input("reference_latents"), + io.Float.Input("strength", default=0.1, min=-1.0, max=1.0, step=0.01), + *[io.Boolean.Input(f"ch_{i:02d}", default=True) for i in range(16)], + ], + outputs=[ + io.Noise.Output(), + ], + ) + + @classmethod + def execute(cls, seed, reference_latents, strength, **kwargs) -> io.NodeOutput: + return io.NodeOutput(Noise_SameColor(seed, reference_latents, strength, **kwargs)) diff --git a/scripts/custom_noise/tkg_noise.py b/scripts/custom_noise/tkg_noise.py new file mode 100644 index 0000000..8a9d55c --- /dev/null +++ b/scripts/custom_noise/tkg_noise.py @@ -0,0 +1,183 @@ +import comfy +from typing import NamedTuple +import torch +import torch.nn.functional as F +from ... import ROOT_NAME, SYMBOL, NODE_SURFIX +from comfy_api.v0_0_2 import io +CATEGORY_NAME = ROOT_NAME + "custom_noise" + +def get_mean_shifted_latents( + latents: torch.Tensor, + shift: float = 0.11, + delta_shift: float = 0.1, + channels: list[float] = [0, 1, 1, 0], # list of {-1, 0, 1} +) -> torch.Tensor: + shifted_latents = latents.clone() + + for idx, sign in enumerate(channels): + if sign == 0: + # skip + continue + + latent_channel = shifted_latents[:, idx, :, :] + + positive_ratio = (latent_channel > 0).float().mean() + target_ratio = positive_ratio + shift * sign + + # gradually shift latent_channel + while True: + latent_channel += delta_shift * sign + new_positive_ratio = (latent_channel > 0).float().mean() + if new_positive_ratio >= target_ratio: + break + + # replace the channel in the original latents + shifted_latents[:, idx, :, :] = latent_channel + + return shifted_latents + + +def get_2d_gaussian( + latent_height: int, + latent_width: int, + std_dev: float, + device: torch.device, + center_x: float = 0.0, + center_y: float = 0.0, + factor: int = 8, # idk why +): + y = torch.linspace(-1, 1, steps=latent_height // factor, device=device) + x = torch.linspace(-1, 1, steps=latent_width // factor, device=device) + + y_grid, x_grid = torch.meshgrid(y, x, indexing="ij") + + x_grid = x_grid - center_x + y_grid = y_grid - center_y + + gauss = torch.exp(-((x_grid**2 + y_grid**2) / (2 * std_dev**2))) + gauss = gauss[None, None, :, :] # add batch and channel dimensions + + return gauss + + +def apply_tkg_noise( + latents: torch.Tensor, + shift: float = 0.11, + delta_shift: float = 0.1, + std_dev: float = 0.5, + factor: int = 8, + channels: list[float] = [0, 1, 1, 0], +): + batch_size, num_channels, latent_height, latent_width = latents.shape + + shifted_latents = get_mean_shifted_latents( + latents, + shift=shift, + delta_shift=delta_shift, + channels=channels, + ) + gauss_mask = get_2d_gaussian( + latent_height=latent_height, + latent_width=latent_width, + std_dev=std_dev, + center_x=0.0, + center_y=0.0, + factor=factor, + device=latents.device, + ) + gauss_mask = F.interpolate( + gauss_mask, + size=(latent_height, latent_width), + mode="bilinear", + align_corners=False, + ) + + gauss_mask = gauss_mask.expand(batch_size, num_channels, -1, -1) + + noised_latents = shifted_latents * (1 - gauss_mask) + latents * gauss_mask + + return noised_latents + + +class ColorSet(NamedTuple): + name: str + channels: list[float] + + +# ref: Figure 28. Additional Result in various color Background with SD +COLOR_SETS: list[ColorSet] = [ + ColorSet("green", [0, 1, 1, 0]), + ColorSet("cyan", [0, 1, 0, 0]), + ColorSet("magenta", [0, -1, -1, -1]), + ColorSet("purple", [0, 0, -1, -1]), + ColorSet("black", [-1, 0, 0, 1]), + ColorSet("orange", [-1, -1, 1, 0]), + ColorSet("white", [0, 0, 0, -1]), + ColorSet("yellow", [0, -1, 1, -1]), +] + +COLOR_SET_MAP: dict[str, ColorSet] = {c.name: c for c in COLOR_SETS} + +class Noise_RandomNoise: + def __init__(self, seed, color="green", shift=0.11, grid_factor=8): + self.seed = seed + self.color = color + self.shift = shift + self.grid_factor = grid_factor + + def generate_noise(self, input_latent): + latent_image = input_latent["samples"] + batch_inds = input_latent["batch_index"] if "batch_index" in input_latent else None + noise = comfy.sample.prepare_noise(latent_image, self.seed, batch_inds) + color_set = COLOR_SET_MAP.get(self.color, COLOR_SET_MAP["green"]) + noise = apply_tkg_noise( + noise, + shift=self.shift, + channels=color_set.channels, + factor=self.grid_factor, + ) + return noise + +class TKGRandomNoise(io.ComfyNode): + @classmethod + def define_schema(cls) -> io.Schema: + return io.Schema( + node_id=f"TKGRandomNoise{NODE_SURFIX}", + display_name=f"TKG Random Noise {SYMBOL}", + category=CATEGORY_NAME, + inputs=[ + io.Int.Input( + "noise_seed", + default=0, + min=0, + max=0xffffffffffffffff, + control_after_generate=True, + ), + io.Combo.Input( + "color", + options=[c.name for c in COLOR_SETS], + default="green", + ), + io.Float.Input( + "shift", + default=0.11, + min=0.0, + max=1.0, + step=0.01, + ), + io.Int.Input( + "grid_factor", + default=8, + min=1, + max=16, + step=1, + ), + ], + outputs=[ + io.Noise.Output(), + ], + ) + + @classmethod + def execute(cls, noise_seed, color, shift, grid_factor) -> io.NodeOutput: + return io.NodeOutput(Noise_RandomNoise(noise_seed, color, shift, grid_factor)) diff --git a/scripts/custom_noise/variation_noise.py b/scripts/custom_noise/variation_noise.py index 1351ad6..4499fb5 100644 --- a/scripts/custom_noise/variation_noise.py +++ b/scripts/custom_noise/variation_noise.py @@ -1,8 +1,9 @@ import comfy -from ... import ROOT_NAME +from ... import ROOT_NAME, SYMBOL, NODE_SURFIX import math import torch import numpy as np +from comfy_api.v0_0_2 import io CATEGORY_NAME = ROOT_NAME + "custom_noise" @@ -21,25 +22,28 @@ class VariationNoiseGenarator: noise = base_noise[self.batch_index].unsqueeze(0) * self.similarity + variation_noise * math.sqrt(1 - self.similarity ** 2) return noise -class VariationNoise: +class VariationNoise(io.ComfyNode): @classmethod - def INPUT_TYPES(s): - return { - "required":{ - "base_seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}), - "seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}), - "similarity": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}), - "batch_index": ("INT", {"default": 1, "min": 1, "max": 4096}), - } - } + def define_schema(cls) -> io.Schema: + return io.Schema( + node_id=f"VariationNoise{NODE_SURFIX}", + display_name=f"Variation Noise {SYMBOL}", + category=CATEGORY_NAME, + inputs=[ + io.Int.Input("base_seed", default=0, min=0, max=0xffffffffffffffff), + io.Int.Input("seed", default=0, min=0, max=0xffffffffffffffff), + io.Float.Input("similarity", default=0.0, min=0.0, max=1.0, step=0.001), + io.Int.Input("batch_index", default=1, min=1, max=4096), + ], + outputs=[ + io.Noise.Output(), + ], + ) - RETURN_TYPES = ("NOISE",) - FUNCTION = "get_noise" - CATEGORY = CATEGORY_NAME + @classmethod + def execute(cls, base_seed, seed, similarity, batch_index) -> io.NodeOutput: + return io.NodeOutput(VariationNoiseGenarator(base_seed, seed, similarity, batch_index-1)) - def get_noise(self, base_seed, seed, similarity, batch_index): - return (VariationNoiseGenarator(base_seed, seed, similarity, batch_index-1),) - def prepare_noise(latent_image, seed, noise_inds=None, offset=0): """ creates random noise given a latent image and a seed. @@ -63,7 +67,7 @@ def prepare_noise(latent_image, seed, noise_inds=None, offset=0): noise_offset = torch.randn([1] + list(latent_image.size())[1:2] + [1,1], dtype=latent_image.dtype, generator=generator, device="cpu") if i in unique_inds: noise_offsets.append(noise_offset) - + noises = [noises[i] for i in inverse] noises = torch.cat(noises, axis=0) @@ -81,22 +85,26 @@ class Noise_RandomNoiseOffset: batch_inds = input_latent["batch_index"] if "batch_index" in input_latent else None return prepare_noise(latent_image, self.seed, batch_inds, self.offset) -class RandomNoiseOffset: +class RandomNoiseOffset(io.ComfyNode): @classmethod - def INPUT_TYPES(s): - return {"required":{ - "noise_seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}), - "offset": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 10.0, "step": 0.01}), - } - } - - RETURN_TYPES = ("NOISE",) - FUNCTION = "get_noise" - CATEGORY = CATEGORY_NAME + def define_schema(cls) -> io.Schema: + return io.Schema( + node_id=f"RandomNoiseOffset{NODE_SURFIX}", + display_name=f"Random Noise Offset {SYMBOL}", + category=CATEGORY_NAME, + inputs=[ + io.Int.Input("noise_seed", default=0, min=0, max=0xffffffffffffffff), + io.Float.Input("offset", default=0.0, min=0.0, max=10.0, step=0.01), + ], + outputs=[ + io.Noise.Output(), + ], + ) + + @classmethod + def execute(cls, noise_seed, offset) -> io.NodeOutput: + return io.NodeOutput(Noise_RandomNoiseOffset(noise_seed, offset)) - def get_noise(self, noise_seed, offset): - return (Noise_RandomNoiseOffset(noise_seed, offset),) - class Noise_RandomNoiseVariationSimple: def __init__(self, seed, similarity): self.seed = seed @@ -109,20 +117,23 @@ class Noise_RandomNoiseVariationSimple: noise = torch.cat([noise[:1], noise[:1] * self.similarity + noise[1:] * math.sqrt(1 - self.similarity ** 2)]) return noise - -class RandomNoiseVariationSimple: - @classmethod - def INPUT_TYPES(s): - return { - "required":{ - "seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}), - "similarity": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}), - } - } - - RETURN_TYPES = ("NOISE",) - FUNCTION = "get_noise" - CATEGORY = CATEGORY_NAME - def get_noise(self, seed, similarity): - return (Noise_RandomNoiseVariationSimple(seed, similarity),) \ No newline at end of file +class RandomNoiseVariationSimple(io.ComfyNode): + @classmethod + def define_schema(cls) -> io.Schema: + return io.Schema( + node_id=f"RandomNoiseVariationSimple{NODE_SURFIX}", + display_name=f"Random Noise Variation Simple {SYMBOL}", + category=CATEGORY_NAME, + inputs=[ + io.Int.Input("seed", default=0, min=0, max=0xffffffffffffffff), + io.Float.Input("similarity", default=0.0, min=0.0, max=1.0, step=0.001), + ], + outputs=[ + io.Noise.Output(), + ], + ) + + @classmethod + def execute(cls, seed, similarity) -> io.NodeOutput: + return io.NodeOutput(Noise_RandomNoiseVariationSimple(seed, similarity)) diff --git a/scripts/custom_samplers/euler_ancestral_fixed_noise.py b/scripts/custom_samplers/euler_ancestral_fixed_noise.py index 93610c7..138b726 100644 --- a/scripts/custom_samplers/euler_ancestral_fixed_noise.py +++ b/scripts/custom_samplers/euler_ancestral_fixed_noise.py @@ -1,7 +1,8 @@ -from comfy.samplers import KSAMPLER +from comfy.samplers import KSAMPLER from comfy.k_diffusion.sampling import sample_euler_ancestral import torch -from ... import ROOT_NAME +from comfy_api.v0_0_2 import io +from ... import ROOT_NAME, SYMBOL, NODE_SURFIX def fixed_noise_sampler(x, seed=None): if seed is not None: @@ -20,21 +21,24 @@ def sample_euler_ancestral_fixed_noise(model, x, sigmas, extra_args=None, callba noise_sampler = fixed_noise_sampler(x, seed=seed) if noise_sampler is None else noise_sampler return sample_euler_ancestral(model, x, sigmas, extra_args=extra_args, callback=callback, disable=disable, eta=eta, s_noise=s_noise, noise_sampler=noise_sampler) -class SamplerEulerAncestralFixedNoise: +class SamplerEulerAncestralFixedNoise(io.ComfyNode): @classmethod - def INPUT_TYPES(s): - return { - "required":{ - "noise": (["fixed", "random"], {"default": "fixed"}), - "eta": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step":0.01, "round": False}), - "s_noise": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step":0.01, "round": False}), - }, - } - RETURN_TYPES = ("SAMPLER",) - CATEGORY = ROOT_NAME + "custom_samplers" + def define_schema(cls) -> io.Schema: + return io.Schema( + node_id=f"SamplerEulerAncestralFixedNoise{NODE_SURFIX}", + display_name=f"Sampler Euler Ancestral Fixed Noise {SYMBOL}", + category=ROOT_NAME + "custom_samplers", + inputs=[ + io.Combo.Input("noise", options=["fixed", "random"], default="fixed"), + io.Float.Input("eta", default=1.0, min=0.0, max=100.0, step=0.01, round=False), + io.Float.Input("s_noise", default=1.0, min=0.0, max=100.0, step=0.01, round=False), + ], + outputs=[ + io.Sampler.Output(), + ], + ) - FUNCTION = "get_sampler" - - def get_sampler(self, noise, eta, s_noise): + @classmethod + def execute(cls, noise, eta, s_noise) -> io.NodeOutput: sampler = KSAMPLER(sample_euler_ancestral_fixed_noise if noise=="fixed" else sample_euler_ancestral, {"eta": eta, "s_noise": s_noise}) - return (sampler, ) \ No newline at end of file + return io.NodeOutput(sampler) \ No newline at end of file diff --git a/scripts/custom_samplers/gradual_latent.py b/scripts/custom_samplers/gradual_latent.py index 951248b..c3e4564 100644 --- a/scripts/custom_samplers/gradual_latent.py +++ b/scripts/custom_samplers/gradual_latent.py @@ -3,8 +3,9 @@ import torch from torchvision.transforms.functional import gaussian_blur from comfy.k_diffusion.sampling import default_noise_sampler, get_ancestral_step, to_d, BrownianTreeNoiseSampler from tqdm.auto import trange +from comfy_api.v0_0_2 import io -from ... import ROOT_NAME +from ... import ROOT_NAME, SYMBOL, NODE_SURFIX def interpolate(x, size, unsharp_strength=0.0, unsharp_kernel_size=3, unsharp_sigma=0.5, unsharp=False, mode="bicubic", align_corners=False): x = torch.nn.functional.interpolate(x, size=size, mode=mode, align_corners=align_corners) @@ -61,7 +62,7 @@ def sample_euler_ancestral( callback({"x": x, "i": i, "sigma": sigmas[i], "sigma_hat": sigmas[i], "denoised": denoised}) # Euler method - d = to_d(x, sigmas[i], denoised) + d = to_d(x, sigmas[i], denoised) if i not in upscale_info: x = denoised + d * sigma_down elif unsharp_target == "x": @@ -115,7 +116,7 @@ def sample_dpmpp_2s_ancestral( callback({"x": x, "i": i, "sigma": sigmas[i], "sigma_hat": sigmas[i], "denoised": denoised}) if sigma_down == 0: # Euler method - d = to_d(x, sigmas[i], denoised) + d = to_d(x, sigmas[i], denoised) if i not in upscale_info: x = denoised + d * sigma_down elif unsharp_target == "x": @@ -220,7 +221,7 @@ def sample_dpmpp_2m_sde( if eta: noise_sampler = BrownianTreeNoiseSampler(x, sigma_min, sigma_max, seed=seed, cpu=True) x = x + noise_sampler(sigmas[i], sigmas[i + 1]) * sigmas[i + 1] * (-2 * eta_h).expm1().neg().sqrt() * s_noise - + h_last = h return x @@ -270,33 +271,35 @@ def sample_lcm( return x -class GradualLatentSampler: +class GradualLatentSampler(io.ComfyNode): # kernel_sizeのstepを2にすると、2,4,6,8... となるので、stepを1にしておく @classmethod - def INPUT_TYPES(s): - return { - "required": { - "sampler_name": (["euler_ancestral", "dpmpp_2s_ancestral", "dpmpp_2m_sde", "lcm"],), - "eta": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01, "round": False}), - "s_noise": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01, "round": False}), - "upscale_ratio": ("FLOAT", {"default": 2.0, "min": 0.0, "max": 16.0, "step": 0.01, "round": False}), - "start_step": ("INT", {"default": 5, "min": 0, "max": 1000, "step": 1}), - "end_step": ("INT", {"default": 15, "min": 0, "max": 1000, "step": 1}), - "upscale_n_step": ("INT", {"default": 3, "min": 0, "max": 1000, "step": 1}), - "unsharp_kernel_size": ("INT", {"default": 3, "min": 1, "max": 21, "step": 1}), - "unsharp_sigma": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 10.0, "step": 0.01, "round": False}), - "unsharp_strength": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 10.0, "step": 0.01, "round": False}), - "unsharp_target": (["x", "denoised"],), - } - } + def define_schema(cls) -> io.Schema: + return io.Schema( + node_id=f"GradualLatentSampler{NODE_SURFIX}", + display_name=f"Gradual Latent Sampler {SYMBOL}", + category=ROOT_NAME + "custom_samplers", + inputs=[ + io.Combo.Input("sampler_name", options=["euler_ancestral", "dpmpp_2s_ancestral", "dpmpp_2m_sde", "lcm"]), + io.Float.Input("eta", default=1.0, min=0.0, max=10.0, step=0.01, round=False), + io.Float.Input("s_noise", default=1.0, min=0.0, max=10.0, step=0.01, round=False), + io.Float.Input("upscale_ratio", default=2.0, min=0.0, max=16.0, step=0.01, round=False), + io.Int.Input("start_step", default=5, min=0, max=1000, step=1), + io.Int.Input("end_step", default=15, min=0, max=1000, step=1), + io.Int.Input("upscale_n_step", default=3, min=0, max=1000, step=1), + io.Int.Input("unsharp_kernel_size", default=3, min=1, max=21, step=1), + io.Float.Input("unsharp_sigma", default=0.5, min=0.0, max=10.0, step=0.01, round=False), + io.Float.Input("unsharp_strength", default=0.0, min=0.0, max=10.0, step=0.01, round=False), + io.Combo.Input("unsharp_target", options=["x", "denoised"]), + ], + outputs=[ + io.Sampler.Output(), + ], + ) - RETURN_TYPES = ("SAMPLER",) - CATEGORY = ROOT_NAME + "custom_samplers" - - FUNCTION = "get_sampler" - - def get_sampler( - self, + @classmethod + def execute( + cls, sampler_name, eta, s_noise, @@ -308,7 +311,7 @@ class GradualLatentSampler: unsharp_sigma, unsharp_strength, unsharp_target, - ): + ) -> io.NodeOutput: if sampler_name == "euler_ancestral": sample_function = sample_euler_ancestral elif sampler_name == "dpmpp_2s_ancestral": @@ -319,7 +322,7 @@ class GradualLatentSampler: sample_function = sample_lcm else: raise ValueError("Unknown sampler name") - + unsharp_target = unsharp_target if unsharp_strength > 0 else "x" # interpの位置が違うので調整 unsharp_kernel_size = unsharp_kernel_size if unsharp_kernel_size % 2 == 1 else unsharp_kernel_size + 1 @@ -339,14 +342,4 @@ class GradualLatentSampler: "unsharp_target": unsharp_target, }, ) - return (sampler,) - - -NODE_CLASS_MAPPINGS = { - "GradualLatentSampler": GradualLatentSampler, -} - -NODE_DISPLAY_NAME_MAPPINGS = { - # Sampling - "GradualLatentSampler": "GradualLatentSampler", -} \ No newline at end of file + return io.NodeOutput(sampler) diff --git a/scripts/custom_samplers/lcm_sampler_rcfg.py b/scripts/custom_samplers/lcm_sampler_rcfg.py index cef21fb..58bcf93 100644 --- a/scripts/custom_samplers/lcm_sampler_rcfg.py +++ b/scripts/custom_samplers/lcm_sampler_rcfg.py @@ -12,8 +12,9 @@ import torch from comfy.k_diffusion.sampling import default_noise_sampler from tqdm.auto import trange import copy +from comfy_api.v0_0_2 import io -from ... import ROOT_NAME +from ... import ROOT_NAME, SYMBOL, NODE_SURFIX @torch.no_grad() def sampler_lcm_rcfg(model, x, sigmas, extra_args=None, callback=None, disable=None, noise_sampler=None, enable=True, delta=1.0, cfg=1.0, original_latent=None, **kwargs): @@ -55,30 +56,27 @@ def sampler_lcm_rcfg(model, x, sigmas, extra_args=None, callback=None, disable=N return x -class LCMSamplerRCFG: +class LCMSamplerRCFG(io.ComfyNode): @classmethod - def INPUT_TYPES(s): - return { - "required":{ - "enable": ("BOOLEAN", {"default": True}), - "delta": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 5.0, "step":0.01, "round": False}), - "cfg": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 5.0, "step":0.01, "round": False}), - }, - "optional":{ - "original_latent": ("LATENT",), - } - } - RETURN_TYPES = ("SAMPLER",) - CATEGORY = ROOT_NAME + "custom_samplers" + def define_schema(cls) -> io.Schema: + return io.Schema( + node_id=f"LCMSamplerRCFG{NODE_SURFIX}", + display_name=f"LCM Sampler RCFG {SYMBOL}", + category=ROOT_NAME + "custom_samplers", + inputs=[ + io.Boolean.Input("enable", default=True), + io.Float.Input("delta", default=1.0, min=0.0, max=5.0, step=0.01, round=False), + io.Float.Input("cfg", default=1.0, min=0.0, max=5.0, step=0.01, round=False), + io.Latent.Input("original_latent", optional=True), + ], + outputs=[ + io.Sampler.Output(), + ], + ) - FUNCTION = "get_sampler" - - def get_sampler(self, enable, delta, cfg, original_latent=None): + @classmethod + def execute(cls, enable, delta, cfg, original_latent=None) -> io.NodeOutput: original_latent = original_latent["samples"] if original_latent is not None else None sampler = KSAMPLER(sampler_lcm_rcfg, {"enable": enable, "delta":delta, "cfg":cfg, "original_latent":original_latent}) - return (sampler, ) - -NODE_CLASS_MAPPINGS = { - "LCMSamplerRCFG": LCMSamplerRCFG, -} \ No newline at end of file + return io.NodeOutput(sampler) \ No newline at end of file diff --git a/scripts/custom_samplers/sampler_custom_preview.py b/scripts/custom_samplers/sampler_custom_preview.py index 0d541ff..f1f1978 100644 --- a/scripts/custom_samplers/sampler_custom_preview.py +++ b/scripts/custom_samplers/sampler_custom_preview.py @@ -2,7 +2,8 @@ import comfy from latent_preview import get_previewer import numpy as np import torch -from ... import ROOT_NAME +from comfy_api.v0_0_2 import io +from ... import ROOT_NAME, SYMBOL, NODE_SURFIX def image_to_tensor(image): return torch.tensor(np.array(image).astype(np.float32)) / 255.0 @@ -29,26 +30,30 @@ def prepare_callback(model, steps, x0_output_dict=None, previews=None): pbar.update_absolute(step + 1, total_steps, preview_bytes) return callback -class SamplerCustomAdvancedPreview: +class SamplerCustomAdvancedPreview(io.ComfyNode): @classmethod - def INPUT_TYPES(s): - return {"required": - {"noise": ("NOISE", ), - "guider": ("GUIDER", ), - "sampler": ("SAMPLER", ), - "sigmas": ("SIGMAS", ), - "latent_image": ("LATENT", ), - } - } + def define_schema(cls) -> io.Schema: + return io.Schema( + node_id=f"SamplerCustomAdvancedPreview{NODE_SURFIX}", + display_name=f"Sampler Custom Advanced Preview {SYMBOL}", + category=ROOT_NAME + "custom_samplers", + inputs=[ + io.Noise.Input("noise"), + io.Guider.Input("guider"), + io.Sampler.Input("sampler"), + io.Sigmas.Input("sigmas"), + io.Latent.Input("latent_image"), + ], + outputs=[ + io.Latent.Output(display_name="output"), + io.Latent.Output(display_name="denoised_output"), + io.Image.Output(display_name="previews"), + ], + ) - RETURN_TYPES = ("LATENT", "LATENT", "IMAGE") - RETURN_NAMES = ("output", "denoised_output", "previews") - - FUNCTION = "sample" - CATEGORY = ROOT_NAME + "custom_samplers" - - def sample(self, noise, guider, sampler, sigmas, latent_image): + @classmethod + def execute(cls, noise, guider, sampler, sigmas, latent_image) -> io.NodeOutput: latent = latent_image latent_image = latent["samples"] latent = latent.copy() @@ -76,4 +81,4 @@ class SamplerCustomAdvancedPreview: out_denoised = out previews = torch.stack(previews) - return (out, out_denoised, previews) \ No newline at end of file + return io.NodeOutput(out, out_denoised, previews) \ No newline at end of file diff --git a/scripts/custom_samplers/tcd_sampler.py b/scripts/custom_samplers/tcd_sampler.py index d0f8a03..9657a71 100644 --- a/scripts/custom_samplers/tcd_sampler.py +++ b/scripts/custom_samplers/tcd_sampler.py @@ -2,8 +2,9 @@ from comfy.samplers import KSAMPLER import torch from comfy.k_diffusion.sampling import default_noise_sampler, to_d from tqdm.auto import trange +from comfy_api.v0_0_2 import io -from ... import ROOT_NAME +from ... import ROOT_NAME, SYMBOL, NODE_SURFIX @torch.no_grad() def sampler_tcd(model, x, sigmas, extra_args=None, callback=None, disable=None, noise_sampler=None, gamma=None): @@ -37,23 +38,22 @@ def sampler_tcd(model, x, sigmas, extra_args=None, callback=None, disable=None, x = x + noise_sampler(sigma_from, sigma_to) * sigma_up return x -class TCDSampler: +class TCDSampler(io.ComfyNode): @classmethod - def INPUT_TYPES(s): - return { - "required":{ - "gamma": ("FLOAT", {"default": 0.3, "min": 0.0, "max": 1.0, "step":0.01}), - }, - } - RETURN_TYPES = ("SAMPLER",) - CATEGORY = ROOT_NAME + "custom_samplers" + def define_schema(cls) -> io.Schema: + return io.Schema( + node_id=f"TCDSampler{NODE_SURFIX}", + display_name=f"TCD Sampler {SYMBOL}", + category=ROOT_NAME + "custom_samplers", + inputs=[ + io.Float.Input("gamma", default=0.3, min=0.0, max=1.0, step=0.01), + ], + outputs=[ + io.Sampler.Output(), + ], + ) - FUNCTION = "get_sampler" - - def get_sampler(self, gamma): + @classmethod + def execute(cls, gamma) -> io.NodeOutput: sampler = KSAMPLER(sampler_tcd, {"gamma": gamma}) - return (sampler, ) - -NODE_CLASS_MAPPINGS = { - "TCDSampler": TCDSampler, -} \ No newline at end of file + return io.NodeOutput(sampler) \ No newline at end of file diff --git a/scripts/custom_schedulers/text_scheduler.py b/scripts/custom_schedulers/text_scheduler.py index e61ee5f..e954a45 100644 --- a/scripts/custom_schedulers/text_scheduler.py +++ b/scripts/custom_schedulers/text_scheduler.py @@ -5,27 +5,37 @@ connect to SamplerCustom ''' import torch +from comfy_api.v0_0_2 import io from ... import ROOT_NAME CATEGORY_NAME = ROOT_NAME + "custom_schedulers" -class TextScheduler: +class TextScheduler(io.ComfyNode): @classmethod - def INPUT_TYPES(s): - return {"required":{"model": ("MODEL",), "timesteps": ("STRING", {"multiline": True}), "verbose": ("BOOLEAN", )}} - RETURN_TYPES = ("SIGMAS",) - CATEGORY = CATEGORY_NAME + def define_schema(cls) -> io.Schema: + return io.Schema( + node_id="TextScheduler|cgem156", + display_name="Text Scheduler 🍌", + category=CATEGORY_NAME, + inputs=[ + io.Model.Input("model"), + io.String.Input("timesteps", multiline=True), + io.Boolean.Input("verbose"), + ], + outputs=[ + io.Sigmas.Output(), + ], + ) - FUNCTION = "get_sigmas" - - def get_sigmas(self, model, timesteps, verbose): + @classmethod + def execute(cls, model, timesteps, verbose) -> io.NodeOutput: timesteps = [float(timestep) for timestep in timesteps.replace(" ", "").split(",")] sigmas = model.model.model_sampling.sigma(torch.tensor(timesteps)) sigmas = torch.cat([sigmas, torch.tensor([0])]) if verbose: print("sigmas:", sigmas.tolist()) - return (sigmas, ) + return io.NodeOutput(sigmas) NODE_CLASS_MAPPINGS = { "TextScheduler": TextScheduler, diff --git a/scripts/dart/node.py b/scripts/dart/node.py index effcd19..f5b0e07 100644 --- a/scripts/dart/node.py +++ b/scripts/dart/node.py @@ -3,47 +3,55 @@ from transformers.generation.logits_process import UnbatchedClassifierFreeGuidan import comfy import torch import re -from ... import ROOT_NAME +from comfy_api.v0_0_2 import io +from ... import ROOT_NAME, NODE_SURFIX, SYMBOL CATEGORY_NAME = ROOT_NAME + "dart" -class LoadDart: +class LoadDart(io.ComfyNode): @classmethod - def INPUT_TYPES(s): - return { - "required": { - "tokenizer": ("STRING", {"default": "p1atdev/dart-v1-sft"}), - "model": ("STRING", {"default": "p1atdev/dart-v1-sft"}), - } - } - RETURN_TYPES = ("DART_TOKENIZER", "DART_MODEL", ) - FUNCTION = "load" + def define_schema(cls) -> io.Schema: + return io.Schema( + node_id=f"LoadDart{NODE_SURFIX}", + display_name=f"Load Dart {SYMBOL}", + category=CATEGORY_NAME, + inputs=[ + io.String.Input("tokenizer", default="p1atdev/dart-v1-sft"), + io.String.Input("model", default="p1atdev/dart-v1-sft"), + ], + outputs=[ + io.Custom("DART_TOKENIZER").Output(), + io.Custom("DART_MODEL").Output(), + ], + ) - CATEGORY = CATEGORY_NAME - - def load(self, tokenizer, model): + @classmethod + def execute(cls, tokenizer, model) -> io.NodeOutput: tokenizer = AutoTokenizer.from_pretrained(tokenizer, trust_remote_code=True) model = AutoModelForCausalLM.from_pretrained(model, trust_remote_code=True) - return (tokenizer, model, ) - -class DartPrompt: + return io.NodeOutput(tokenizer, model) + +class DartPrompt(io.ComfyNode): @classmethod - def INPUT_TYPES(s): - return { - "required": { - "rating": (["general", "sensitive", "questionable", "explicit", "sfw", "nsfw"], ), - "copyright": ("STRING", {"default": "original"}), - "character": ("STRING", {"default": ""}), - "general": ("STRING", {"multiline": True}), - "long": (["very_short", "short", "long", "very_long"], {"default": "long"}), - } - } - RETURN_TYPES = ("STRING", ) - FUNCTION = "load" + def define_schema(cls) -> io.Schema: + return io.Schema( + node_id=f"DartPrompt{NODE_SURFIX}", + display_name=f"Dart Prompt {SYMBOL}", + category=CATEGORY_NAME, + inputs=[ + io.Combo.Input("rating", options=["general", "sensitive", "questionable", "explicit", "sfw", "nsfw"]), + io.String.Input("copyright", default="original"), + io.String.Input("character", default=""), + io.String.Input("general", multiline=True), + io.Combo.Input("long", options=["very_short", "short", "long", "very_long"], default="long"), + ], + outputs=[ + io.String.Output(), + ], + ) - CATEGORY = CATEGORY_NAME - - def load(self, rating, copyright, character, general, long): + @classmethod + def execute(cls, rating, copyright, character, general, long) -> io.NodeOutput: prompt = "<|bos|>" prompt += f"rating:{rating}" prompt += f"{copyright}" @@ -52,128 +60,125 @@ class DartPrompt: prompt += f"{general}" prompt += "<|input_end|>" - return (prompt, ) + return io.NodeOutput(prompt) -class DartPromptV2: +class DartPromptV2(io.ComfyNode): @classmethod - def INPUT_TYPES(s): - return { - "required": { - "rating": (["general", "sensitive", "questionable", "explicit", "sfw", "nsfw"], ), - "copyright": ("STRING", {"default": "original"}), - "character": ("STRING", {"default": ""}), - "general": ("STRING", {"multiline": True}), - "aspect_ratio": (["ultra_wide", "wide", "square", "tall", "ultra_tall"], {"default": "tall"}), - "length": (["very_short", "short", "medium", "long", "very_long"], {"default": "medium"}), - "identity": (["none", "lax", "strict"], {"default": "none"}), - } - } - RETURN_TYPES = ("STRING", ) - FUNCTION = "load" + def define_schema(cls) -> io.Schema: + return io.Schema( + node_id=f"DartPromptV2{NODE_SURFIX}", + display_name=f"Dart Prompt V2 {SYMBOL}", + category=CATEGORY_NAME, + inputs=[ + io.Combo.Input("rating", options=["general", "sensitive", "questionable", "explicit", "sfw", "nsfw"]), + io.String.Input("copyright", default="original"), + io.String.Input("character", default=""), + io.String.Input("general", multiline=True), + io.Combo.Input("aspect_ratio", options=["ultra_wide", "wide", "square", "tall", "ultra_tall"], default="tall"), + io.Combo.Input("length", options=["very_short", "short", "medium", "long", "very_long"], default="medium"), + io.Combo.Input("identity", options=["none", "lax", "strict"], default="none"), + ], + outputs=[ + io.String.Output(), + ], + ) - CATEGORY = CATEGORY_NAME - - def load(self, rating, copyright, character, general, aspect_ratio, length, identity): - prompt = "<|bos|>" + @classmethod + def execute(cls, rating, copyright, character, general, aspect_ratio, length, identity) -> io.NodeOutput: + prompt = "<|bos|>" prompt += f"{copyright}" prompt += f"{character}" prompt += f"<|rating:{rating}|>" + f"<|aspect_ratio:{aspect_ratio}|>" + f"<|length:{length}|>" + f"<|identity:{identity}|>" prompt += f"{general}<|identity:{identity}|><|input_end|>" - return (prompt, ) - -class DartConfig: + return io.NodeOutput(prompt) + +class DartConfig(io.ComfyNode): @classmethod - def INPUT_TYPES(s): - input_types = { - "required": { - "max_new_tokens": ( - "INT", - {"default": 128, "min": 1, "max": 256, "step": 1}, - ), - "min_new_tokens": ( - "INT", - {"default": 0, "min": 0, "max": 255, "step": 1}, - ), - "temperature": ( - "FLOAT", - {"default": 1.0, "min": 0.0, "max": 5.0, "step": 0.01}, - ), - "top_p": ( - "FLOAT", - {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}, - ), - "top_k": ( - "INT", - {"default": 20, "min": 1, "max": 500, "step": 1}, - ), - "num_beams": ( - "INT", - {"default": 1, "min": 1, "max": 10, "step": 1}, - ), - "cfg_scale": ( - "FLOAT", - {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01}, - ), - }, + def define_schema(cls) -> io.Schema: + return io.Schema( + node_id=f"DartConfig{NODE_SURFIX}", + display_name=f"Dart Config {SYMBOL}", + category=CATEGORY_NAME, + inputs=[ + io.Int.Input("max_new_tokens", default=128, min=1, max=256, step=1), + io.Int.Input("min_new_tokens", default=0, min=0, max=255, step=1), + io.Float.Input("temperature", default=1.0, min=0.0, max=5.0, step=0.01), + io.Float.Input("top_p", default=1.0, min=0.0, max=1.0, step=0.01), + io.Int.Input("top_k", default=20, min=1, max=500, step=1), + io.Int.Input("num_beams", default=1, min=1, max=10, step=1), + io.Float.Input("cfg_scale", default=1.0, min=0.0, max=10.0, step=0.01), + ], + outputs=[ + io.Custom("DART_CONFIG").Output(), + ], + ) + + @classmethod + def execute(cls, max_new_tokens, min_new_tokens, temperature, top_p, top_k, num_beams, cfg_scale) -> io.NodeOutput: + kwargs = { + "max_new_tokens": max_new_tokens, + "min_new_tokens": min_new_tokens, + "temperature": temperature, + "top_p": top_p, + "top_k": top_k, + "num_beams": num_beams, + "cfg_scale": cfg_scale, } - - return input_types - - RETURN_TYPES = ("DART_CONFIG",) - FUNCTION = "compose" - CATEGORY = CATEGORY_NAME - - def compose(self, **kwargs): kwargs["temperature"] = float(kwargs["temperature"]) # avoid error - return (kwargs,) - -class BanTags: + return io.NodeOutput(kwargs) + +class BanTags(io.ComfyNode): @classmethod - def INPUT_TYPES(s): - return { - "required":{ - "tokenizer": ("DART_TOKENIZER", ), - "ban_tags": ("STRING", {"multiline": True}), - } - } - RETURN_TYPES = ("STRING", ) - FUNCTION = "generate" + def define_schema(cls) -> io.Schema: + return io.Schema( + node_id=f"BanTags{NODE_SURFIX}", + display_name=f"Ban Tags {SYMBOL}", + category=CATEGORY_NAME, + inputs=[ + io.Custom("DART_TOKENIZER").Input("tokenizer"), + io.String.Input("ban_tags", multiline=True), + ], + outputs=[ + io.String.Output(), + ], + ) - CATEGORY = CATEGORY_NAME - - def generate(self, tokenizer, ban_tags): + @classmethod + def execute(cls, tokenizer, ban_tags) -> io.NodeOutput: ban_tags_result = set() patterns = [re.compile(ban_tag) for ban_tag in ban_tags.splitlines()] for pattern in patterns: for tag in tokenizer.vocab: if pattern.match(tag): ban_tags_result.add(tag) - return (", ".join(ban_tags_result), ) - -class DartGenerate: + return io.NodeOutput(", ".join(ban_tags_result)) + +class DartGenerate(io.ComfyNode): @classmethod - def INPUT_TYPES(s): - return { - "required": { - "tokenizer": ("DART_TOKENIZER", ), - "model": ("DART_MODEL", ), - "prompt": ("STRING", {"default": ""}), - "batch_size": ("INT", {"default": 1, "min": 1, "max": 4096}), - "seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}), - }, - "optional":{ - "config": ("DART_CONFIG", ), - "negative": ("STRING", {"default": ""}), - "ban_tags": ("STRING", {"default": ""}), - } - } - RETURN_TYPES = ("BATCH_STRING", "STRING") - FUNCTION = "generate" + def define_schema(cls) -> io.Schema: + return io.Schema( + node_id=f"DartGenerate{NODE_SURFIX}", + display_name=f"Dart Generate {SYMBOL}", + category=CATEGORY_NAME, + inputs=[ + io.Custom("DART_TOKENIZER").Input("tokenizer"), + io.Custom("DART_MODEL").Input("model"), + io.String.Input("prompt", default=""), + io.Int.Input("batch_size", default=1, min=1, max=4096), + io.Int.Input("seed", default=0, min=0, max=0xffffffffffffffff), + io.Custom("DART_CONFIG").Input("config", optional=True), + io.String.Input("negative", default="", optional=True), + io.String.Input("ban_tags", default="", optional=True), + ], + outputs=[ + io.Custom("BATCH_STRING").Output(), + io.String.Output(), + ], + ) - CATEGORY = CATEGORY_NAME - - def generate(self, tokenizer, model, prompt, batch_size, seed, config=None, negative=None, ban_tags=None): + @classmethod + def execute(cls, tokenizer, model, prompt, batch_size, seed, config=None, negative=None, ban_tags=None) -> io.NodeOutput: if config: config = config else: @@ -185,14 +190,14 @@ class DartGenerate: "top_k": 100, "num_beams": 1, } - + rng_state = torch.get_rng_state() cuda_rng_state = torch.cuda.get_rng_state() - + if seed is not None: torch.manual_seed(seed) torch.cuda.manual_seed(seed) - + generation_config = GenerationConfig.from_pretrained("p1atdev/dart-v1-sft", **config) # こんなんでいいの? model.to(comfy.model_management.get_torch_device(), dtype=torch.float16).eval() inputs = tokenizer([prompt], return_tensors="pt").input_ids.to(comfy.model_management.get_torch_device()).repeat(batch_size, 1) @@ -230,5 +235,4 @@ class DartGenerate: torch.set_rng_state(rng_state) torch.cuda.set_rng_state(cuda_rng_state) - return (prompts, strings) - + return io.NodeOutput(prompts, strings) diff --git a/scripts/for_test/__init__.py b/scripts/for_test/__init__.py index c932632..5ec9338 100644 --- a/scripts/for_test/__init__.py +++ b/scripts/for_test/__init__.py @@ -1,12 +1,23 @@ from .attention_scale import AttentionScale +from .kv_token_multiplier import CLIPTextEncodeBatchKVMultiply +from .kmeans_quant import KmeansQuantize +from .mse_heatmap import MSEHeatmap, MSEHeatmapTagger from ... import SYMBOL, NODE_SURFIX NODE_CLASS_MAPPINGS = { - f"AttentionScale{NODE_SURFIX}": AttentionScale + f"AttentionScale{NODE_SURFIX}": AttentionScale, + f"CLIPTextEncodeBatchKVMultiply{NODE_SURFIX}": CLIPTextEncodeBatchKVMultiply, + f"KmeansQuantize{NODE_SURFIX}": KmeansQuantize, + f"MSEHeatmap{NODE_SURFIX}": MSEHeatmap, + f"MSEHeatmapTagger{NODE_SURFIX}": MSEHeatmapTagger } NODE_DISPLAY_NAME_MAPPINGS = { - f"AttentionScale{NODE_SURFIX}": f"Attention Scale {SYMBOL}" + f"AttentionScale{NODE_SURFIX}": f"Attention Scale {SYMBOL}", + f"CLIPTextEncodeBatchKVMultiply{NODE_SURFIX}": f"CLIP Text Encode Batch KV Multiply {SYMBOL}", + f"KmeansQuantize{NODE_SURFIX}": f"Kmeans Quantize {SYMBOL}", + f"MSEHeatmap{NODE_SURFIX}": f"MSE Heatmap {SYMBOL}", + f"MSEHeatmapTagger{NODE_SURFIX}": f"MSE Heatmap Tagger {SYMBOL}" } __all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"] \ No newline at end of file diff --git a/scripts/for_test/attention_scale.py b/scripts/for_test/attention_scale.py index 18c2463..caa7dff 100644 --- a/scripts/for_test/attention_scale.py +++ b/scripts/for_test/attention_scale.py @@ -1,6 +1,7 @@ import torch from comfy.ldm.modules.attention import optimized_attention -from ... import ROOT_NAME +from ... import ROOT_NAME, SYMBOL, NODE_SURFIX +from comfy_api.v0_0_2 import io def attention_pytorch(q, k, v, heads, temperature=1.0, mask=None): b, _, dim_head = q.shape @@ -18,50 +19,53 @@ def attention_pytorch(q, k, v, heads, temperature=1.0, mask=None): ) return out -class AttentionScale: +class AttentionScale(io.ComfyNode): @classmethod - def INPUT_TYPES(s): - return { - "required": { - "model": ("MODEL", ), - "temperature": ("FLOAT", {"default": 1.0, "min": -1000.0, "max": 1000.0, "step": 0.01}), - "start_step": ("FLOAT", {"default": 0, "min": 0, "max": 1, "step": 0.001}), - "end_step": ("FLOAT", {"default": 1, "min": 0, "max": 1, "step": 0.001}), - "attn1": ("BOOLEAN", {"default": True}), - "attn2": ("BOOLEAN", {"default": True}), - } - } + def define_schema(cls) -> io.Schema: + return io.Schema( + node_id=f"AttentionScale{NODE_SURFIX}", + display_name=f"Attention Scale {SYMBOL}", + category=ROOT_NAME + "for_test", + inputs=[ + io.Model.Input("model"), + io.Float.Input("temperature", default=1.0, min=-1000.0, max=1000.0, step=0.01), + io.Float.Input("start_step", default=0, min=0, max=1, step=0.001), + io.Float.Input("end_step", default=1, min=0, max=1, step=0.001), + io.Boolean.Input("attn1", default=True), + io.Boolean.Input("attn2", default=True), + ], + outputs=[ + io.Model.Output(), + ], + ) - RETURN_TYPES = ("MODEL", ) - FUNCTION = "apply" - CATEGORY = ROOT_NAME + "for_test" - - def apply(self, model, temperature, start_step, end_step, attn1, attn2): + @classmethod + def execute(cls, model, temperature, start_step, end_step, attn1, attn2) -> io.NodeOutput: new_model = model.clone() - self.temperature = temperature - self.start_sigma = new_model.model.model_sampling.percent_to_sigma(start_step) - self.end_sigma = new_model.model.model_sampling.percent_to_sigma(end_step) + temperature_ = temperature + start_sigma = new_model.model.model_sampling.percent_to_sigma(start_step) + end_sigma = new_model.model.model_sampling.percent_to_sigma(end_step) def attn_patch(q, k, v, extra_options): sigma = extra_options["sigmas"][0].item() - if self.end_sigma <= sigma <= self.start_sigma: - output = attention_pytorch(q, k, v, extra_options["n_heads"], temperature = self.temperature) + if end_sigma <= sigma <= start_sigma: + output = attention_pytorch(q, k, v, extra_options["n_heads"], temperature = temperature_) else: output = attention_pytorch(q, k, v, extra_options["n_heads"], temperature = 1.0) return output - + def dummy_attn_path(q, k, v, extra_options): return optimized_attention(q, k, v, extra_options["n_heads"]) - self.sdxl = hasattr(new_model.model.diffusion_model, "label_emb") + sdxl = hasattr(new_model.model.diffusion_model, "label_emb") attn1_patch = attn_patch if attn1 else dummy_attn_path attn2_patch = attn_patch if attn2 else dummy_attn_path - if not self.sdxl: + if not sdxl: for id in [1,2,4,5,7,8]: # id of input_blocks that have cross attention new_model.set_model_attn1_replace(attn1_patch, "input", id) new_model.set_model_attn2_replace(attn2_patch, "input", id) @@ -85,5 +89,4 @@ class AttentionScale: new_model.set_model_attn1_replace(attn1_patch, "output", id, index) new_model.set_model_attn2_replace(attn2_patch, "output", id, index) - return (new_model, ) - + return io.NodeOutput(new_model) diff --git a/scripts/for_test/kmeans_quant.py b/scripts/for_test/kmeans_quant.py new file mode 100644 index 0000000..1cb7d3c --- /dev/null +++ b/scripts/for_test/kmeans_quant.py @@ -0,0 +1,101 @@ +from ... import ROOT_NAME, SYMBOL, NODE_SURFIX +import torch +import numpy as np +import cv2 +from comfy_api.v0_0_2 import io + +# ref:https://qiita.com/fdsafdfadsa/items/4e8046998be9627ca85d +def kmeans_quant(img, K, kmeans_pp): + + flags = cv2.KMEANS_RANDOM_CENTERS if not kmeans_pp else cv2.KMEANS_PP_CENTERS + criteria = (cv2.TERM_CRITERIA_EPS + cv2.TERM_CRITERIA_MAX_ITER, 10, 1e-4) + _, label, center = cv2.kmeans(img, K, None, criteria, 10, flags) + res = center[label.flatten()] + + return res + +class KMeansManhattan: + def __init__(self, n_clusters, max_iters=10, tol=1e-4): + self.n_clusters = n_clusters + self.max_iters = max_iters + self.tol = tol + + def fit(self, X): + # データセットのサイズ + n_samples, n_features = X.shape + + # クラスタ中心をデータポイントの中からランダムに初期化 + rng = np.random.default_rng() + self.centroids = X[rng.choice(n_samples, self.n_clusters, replace=False)] + + for i in range(self.max_iters): + # 各データポイントを最も近いクラスタに割り当てる + self.labels = self._assign_clusters(X) + + # 新しいクラスタ中心を計算 (マンハッタン距離のためには中央値を使用) + new_centroids = np.array([np.median(X[self.labels == j], axis=0) for j in range(self.n_clusters)]) + + # クラスタ中心の変化が許容範囲内であれば終了 + if np.all(np.abs(self.centroids - new_centroids).sum(axis=1) < self.tol): + break + + self.centroids = new_centroids + + def _assign_clusters(self, X): + # 各データポイントとクラスタ中心とのマンハッタン距離を計算 + distances = np.sum(np.abs(X[:, np.newaxis] - self.centroids), axis=2) + # 最も近いクラスタにラベルを割り当てる + return np.argmin(distances, axis=1) + + def predict(self, X): + # 新しいデータに対してクラスタを予測 + return self._assign_clusters(X) + +def kmeans(img, K, kmeans_pp, manhattan, seed): + orogin_state = np.random.get_state() + np.random.seed(seed) + + if manhattan: + kmeans = KMeansManhattan(n_clusters=K) + kmeans.fit(img) + retval = kmeans.centroids[kmeans.predict(img)] + else: + retval = kmeans_quant(img, K, kmeans_pp) + + np.random.set_state(orogin_state) + return retval + +class KmeansQuantize(io.ComfyNode): + @classmethod + def define_schema(cls) -> io.Schema: + return io.Schema( + node_id=f"KmeansQuantize{NODE_SURFIX}", + display_name=f"Kmeans Quantize {SYMBOL}", + category=ROOT_NAME + "for_test", + inputs=[ + io.Image.Input("image"), + io.Int.Input("colors", default=256, min=1, max=256, step=1), + io.Boolean.Input("individual"), + io.Boolean.Input("kmeans_pp"), + io.Boolean.Input("manhattan"), + io.Int.Input("seed", default=0, min=0, max=0xffffffffffffffff), + ], + outputs=[ + io.Image.Output(), + ], + ) + + @classmethod + def execute(cls, image: torch.Tensor, colors: int, individual: bool, kmeans_pp:bool, manhattan: bool, seed: int) -> io.NodeOutput: + batch_size, height, width, channels = image.shape + image = image.reshape(batch_size, height * width, channels).float().cpu().numpy() + + if individual: + result = np.zeros_like(image) + for i in range(batch_size): + result[i] = kmeans(image[i], colors, kmeans_pp, manhattan, seed) + else: + result = kmeans(image.reshape(-1, channels), colors, kmeans_pp, manhattan, seed).reshape(batch_size, height * width, channels) + + result = torch.from_numpy(result).float().reshape(batch_size, height, width, channels) + return io.NodeOutput(result) diff --git a/scripts/for_test/kv_token_multiplier.py b/scripts/for_test/kv_token_multiplier.py new file mode 100644 index 0000000..6ec3699 --- /dev/null +++ b/scripts/for_test/kv_token_multiplier.py @@ -0,0 +1,70 @@ +import comfy +import torch +from ... import ROOT_NAME, SYMBOL, NODE_SURFIX +from comfy_api.v0_0_2 import io + +CATEGORY_NAME = ROOT_NAME + "for_test" + +def reset_weight(tokens): + ret_dic = {} + for key in tokens: + ret_dic[key] = [[(token, 1) for token, weight in tokens[key][0]]] + weights = [weight for token, weight in tokens[key][0]] + return ret_dic, weights + +class CLIPTextEncodeBatchKVMultiply(io.ComfyNode): + @classmethod + def define_schema(cls) -> io.Schema: + return io.Schema( + node_id=f"CLIPTextEncodeBatchKVMultiply{NODE_SURFIX}", + display_name=f"CLIP Text Encode Batch KV Multiply {SYMBOL}", + category=CATEGORY_NAME, + inputs=[ + io.Model.Input("model"), + io.Clip.Input("clip"), + io.String.Input("text_k", multiline=True), + io.String.Input("text_v", multiline=True), + ], + outputs=[ + io.Model.Output(), + io.Conditioning.Output(), + ], + ) + + @classmethod + def execute(cls, model, clip, text_k, text_v) -> io.NodeOutput: + + tokens_k = clip.tokenize(text_k) + tokens_v = clip.tokenize(text_v) + + tokens_no_weight_k, k_weights = reset_weight(tokens_k) + tokens_no_weight_v, v_weights = reset_weight(tokens_v) + + assert tokens_no_weight_k == tokens_no_weight_v, "tokens_k and tokens_v must be the same." + cond, pooled = clip.encode_from_tokens(tokens_no_weight_k, return_pooled=True) + + state = { + "k_weights": torch.tensor(k_weights).view(1, -1, 1), + "v_weights": torch.tensor(v_weights).view(1, -1, 1), + } + + new_model = model.clone() + def attn2_patch(q, k, v, extra_options): + + assert k.mean() == v.mean(), "k and v must be the same." + if k.shape[1] != state["k_weights"].shape[1]: + state["k_weights"].repeat(1, k.shape[1] // state["k_weights"].shape[1], 1) + state["v_weights"].repeat(1, v.shape[1] // state["v_weights"].shape[1], 1) + + if state["k_weights"].device != k.device: + state["k_weights"] = state["k_weights"].to(k) + state["v_weights"] = state["v_weights"].to(v) + + ks = k * state["k_weights"] + vs = v * state["v_weights"] + + return q, ks, vs + + new_model.set_model_attn2_patch(attn2_patch) + + return io.NodeOutput(new_model, [[cond, {"pooled_output": pooled}]]) diff --git a/scripts/for_test/mse_heatmap.py b/scripts/for_test/mse_heatmap.py new file mode 100644 index 0000000..ec00d0d --- /dev/null +++ b/scripts/for_test/mse_heatmap.py @@ -0,0 +1,110 @@ +import torch +import numpy as np +import matplotlib.pyplot as plt +from matplotlib.colors import Normalize +from ... import ROOT_NAME, SYMBOL, NODE_SURFIX +from comfy_api.v0_0_2 import io + +WDTaggerFeatures = io.Custom("WD-TAGGER-FEATURES") + +def heatmap_to_numpy(heatmap, cmap="jet"): + norm = Normalize(vmin=np.min(heatmap), vmax=np.max(heatmap)) # 正規化 + colormap = plt.get_cmap(cmap) + heatmap_rgb = colormap(norm(heatmap))[:, :, :3] + return heatmap_rgb + +class MSEHeatmap(io.ComfyNode): + @classmethod + def define_schema(cls) -> io.Schema: + return io.Schema( + node_id=f"MSEHeatmap{NODE_SURFIX}", + display_name=f"MSE Heatmap {SYMBOL}", + category=ROOT_NAME + "for_test", + inputs=[ + io.Latent.Input("latent1"), + io.Latent.Input("latent2"), + io.Image.Input("image"), + io.Float.Input("alpha", default=0.3, min=0, max=1, step=0.01), + ], + outputs=[ + io.Image.Output(), + ], + ) + + @classmethod + def execute(cls, latent1, latent2, image, alpha) -> io.NodeOutput: + latent1 = latent1["samples"] + latent2 = latent2["samples"] + print(latent1.size(), latent2.size(), image.size()) + error = torch.norm(latent1 - latent2, dim=1, keepdim=False) + heatmaps = [heatmap_to_numpy(error[i].cpu().numpy()) for i in range(error.size(0))] + heatmaps = torch.from_numpy(np.array(heatmaps)) + h, w = image.size(1), image.size(2) + print(heatmaps.size()) + heatmaps = heatmaps.permute(0, 3, 1, 2) + heatmaps = torch.nn.functional.interpolate(heatmaps, size=(h, w), mode="bilinear") + heatmaps = heatmaps.permute(0, 2, 3, 1) + print(heatmaps.size()) + heatmaps = heatmaps * alpha + image * (1 - alpha) + return io.NodeOutput(heatmaps) + +class MSEHeatmapTagger(io.ComfyNode): + @classmethod + def define_schema(cls) -> io.Schema: + return io.Schema( + node_id=f"MSEHeatmapTagger{NODE_SURFIX}", + display_name=f"MSE Heatmap Tagger {SYMBOL}", + category=ROOT_NAME + "for_test", + inputs=[ + WDTaggerFeatures.Input("features"), + io.Image.Input("image"), + io.Float.Input("alpha", default=0.3, min=0, max=1, step=0.01), + ], + outputs=[ + io.Image.Output(), + ], + ) + + @classmethod + def execute(cls, features, image, alpha) -> io.NodeOutput: + features = features["feature"].detach().clone().cpu() + bsz = features.shape[0] + if features.shape[1] == 1025: # eva02-large + feature_size = 32 + channel_dim = 2 + hw_dim = 1 + features = features[:,1:] + features = features.view(bsz, feature_size, feature_size, -1).permute(0, 3, 1, 2) + elif features.shape[1] == 1024 and len(features.shape) == 3: # vit-large + feature_size = 32 + channel_dim = 2 + hw_dim = 1 + features = features.view(bsz, feature_size, feature_size, -1).permute(0, 3, 1, 2) + elif features.shape[1] == 1024: # convnext + feature_size = 14 + channel_dim = 1 + hw_dim = (2, 3) + features = features.view(bsz, -1, feature_size, feature_size) + elif features.shape[2] == 768: # vit + feature_size = 28 + channel_dim = 2 + hw_dim = 1 + features = features.view(bsz, feature_size, feature_size, -1).permute(0, 3, 1, 2) + elif features.shape[3] == 1024: # swin + feature_size = 14 + channel_dim = 3 + hw_dim = (1, 2) + features = features.permute(0, 3, 1, 2) + + print(features.size(), image.size()) + error = torch.norm(features[:1] - features[1:], dim=1, keepdim=False) + heatmaps = [heatmap_to_numpy(error[i].cpu().numpy()) for i in range(error.size(0))] + heatmaps = torch.from_numpy(np.array(heatmaps)) + h, w = image.size(1), image.size(2) + print(heatmaps.size()) + heatmaps = heatmaps.permute(0, 3, 1, 2) + heatmaps = torch.nn.functional.interpolate(heatmaps, size=(h, w), mode="bilinear") + heatmaps = heatmaps.permute(0, 2, 3, 1) + print(heatmaps.size()) + heatmaps = heatmaps * alpha + image[1:] * (1 - alpha) + return io.NodeOutput(heatmaps) diff --git a/scripts/lora_merger/load.py b/scripts/lora_merger/load.py index 35a1c50..bd6aad4 100644 --- a/scripts/lora_merger/load.py +++ b/scripts/lora_merger/load.py @@ -2,7 +2,8 @@ import comfy import folder_paths import os import re -from ... import ROOT_NAME +from comfy_api.v0_0_2 import io +from ... import ROOT_NAME, NODE_SURFIX, SYMBOL CATEGORY_NAME = ROOT_NAME + "lora_merger" @@ -63,73 +64,72 @@ LBW12TO20 = [1, 2, 3, 4, 7, 17, 18, 19] MID_ID = {26:13, 20:10} -class LoraLoaderFromWeight: - def __init__(self): - self.loaded_lora = None +class LoraLoaderFromWeight(io.ComfyNode): + @classmethod + def define_schema(cls) -> io.Schema: + return io.Schema( + node_id=f"LoraLoaderFromWeight{NODE_SURFIX}", + display_name=f"LoRA Loader From Weight {SYMBOL}", + category=CATEGORY_NAME, + inputs=[ + io.Custom("LoRA").Input("lora"), + io.Model.Input("model"), + io.Clip.Input("clip_optional", optional=True), + ], + outputs=[ + io.Model.Output(), + io.Clip.Output(), + ], + ) @classmethod - def INPUT_TYPES(s): - return { - "required": { - "lora": ("LoRA", ), - "model": ("MODEL",), - }, - "optional": { - "clip_optional": ("CLIP", ), - } - } - RETURN_TYPES = ("MODEL", "CLIP") - FUNCTION = "load_lora_from_weight" - - CATEGORY = CATEGORY_NAME - - def load_lora_from_weight(self, lora, model, clip_optional=None): + def execute(cls, lora, model, clip_optional=None) -> io.NodeOutput: lora_weight = lora["lora"] strength_model = lora["strength_model"] strength_clip = lora["strength_clip"] if strength_model == 0 and strength_clip == 0: - return (model, clip_optional) + return io.NodeOutput(model, clip_optional) model_lora, clip_lora = comfy.sd.load_lora_for_models(model, clip_optional, lora_weight, strength_model, strength_clip) - return (model_lora, clip_lora) + return io.NodeOutput(model_lora, clip_lora) -class LoraLoaderWeightOnly: - def __init__(self): - self.loaded_lora = None - self.lbw = None +# module-level cache replacing the old per-instance `self.loaded_lora` / +# `self.lbw` state (execute() is a classmethod, no `self` to cache on). +_weight_only_cache = {"loaded_lora": None, "lbw": None} + +class LoraLoaderWeightOnly(io.ComfyNode): + @classmethod + def define_schema(cls) -> io.Schema: + return io.Schema( + node_id=f"LoraLoaderWeightOnly{NODE_SURFIX}", + display_name=f"LoRA Loader Weight Only {SYMBOL}", + category=CATEGORY_NAME, + inputs=[ + io.Combo.Input("lora_name", options=folder_paths.get_filename_list("loras")), + io.Float.Input("strength_model", default=1.0, min=-20.0, max=20.0, step=0.01), + io.Float.Input("strength_clip", default=1.0, min=-20.0, max=20.0, step=0.01), + io.String.Input("lbw", multiline=False, default=""), + ], + outputs=[ + io.Custom("LoRA").Output(), + ], + ) @classmethod - def INPUT_TYPES(s): - return { - "required": { - "lora_name": (folder_paths.get_filename_list("loras"), ), - "strength_model": ("FLOAT", {"default": 1.0, "min": -20.0, "max": 20.0, "step": 0.01}), - "strength_clip": ("FLOAT", {"default": 1.0, "min": -20.0, "max": 20.0, "step": 0.01}), - "lbw": ("STRING", { - "multiline": False, - "default": "" - }), - } - } - RETURN_TYPES = ("LoRA", ) - FUNCTION = "load_lora_weight_only" - - CATEGORY = CATEGORY_NAME - - def load_lora_weight_only(self, lora_name, strength_model, strength_clip, lbw): + def execute(cls, lora_name, strength_model, strength_clip, lbw) -> io.NodeOutput: lora_path = folder_paths.get_full_path("loras", lora_name) lora = None - if self.loaded_lora is not None: - if self.loaded_lora[0] == lora_path: - lora = self.loaded_lora[1] + if _weight_only_cache["loaded_lora"] is not None: + if _weight_only_cache["loaded_lora"][0] == lora_path: + lora = _weight_only_cache["loaded_lora"][1] else: - temp = self.loaded_lora - self.loaded_lora = None + temp = _weight_only_cache["loaded_lora"] + _weight_only_cache["loaded_lora"] = None del temp - if lora is None or self.lbw != lbw: + if lora is None or _weight_only_cache["lbw"] != lbw: lora = comfy.utils.load_torch_file(lora_path, safe_load=True) if lbw != "": weight_list = parse_weight_list(lbw) @@ -139,7 +139,7 @@ class LoraLoaderWeightOnly: strength_clip = strength_clip * weight_list[0] - up_keys = [key for key in lora.keys() if "lora_up" in key and not "lora_te" in key] + up_keys = [key for key in lora.keys() if ("lora_up" in key or "lora_B" in key) and not "lora_te" in key] for key in up_keys: ids = extract_numbers(key) @@ -166,11 +166,18 @@ class LoraLoaderWeightOnly: if weight != 0.0: lora[key] = lora[key] * weight else: + if "lora_up" in key: + down_key = key.replace("lora_up", "lora_down") + alpha_key = key.replace("lora_up.weight", "alpha") + else: + down_key = key.replace("lora_B", "lora_A") + alpha_key = key.replace("lora_B.weight", "alpha") del lora[key] - del lora[key.replace("lora_up", "lora_down")] - del lora[key.replace("lora_up.weight", "alpha")] + del lora[down_key] + if alpha_key in lora: + del lora[alpha_key] - self.loaded_lora = (lora_path, lora) - self.lbw = lbw + _weight_only_cache["loaded_lora"] = (lora_path, lora) + _weight_only_cache["lbw"] = lbw - return ({"lora": lora, "strength_model": strength_model, "strength_clip": strength_clip}, ) + return io.NodeOutput({"lora": lora, "strength_model": strength_model, "strength_clip": strength_clip}) diff --git a/scripts/lora_merger/merge.py b/scripts/lora_merger/merge.py index e7bf103..3b855d5 100644 --- a/scripts/lora_merger/merge.py +++ b/scripts/lora_merger/merge.py @@ -1,54 +1,58 @@ import comfy import math import torch -from ... import ROOT_NAME +from comfy_api.v0_0_2 import io +from ... import ROOT_NAME, NODE_SURFIX, SYMBOL CATEGORY_NAME = ROOT_NAME + "lora_merger" CLAMP_QUANTILE = 0.99 +REGULAR_LORA = "regular" +DIFFUSERS_LORA = "diffusers" -class LoraMerge: - def __init__(self): - self.loaded_lora = None +class LoraMerge(io.ComfyNode): + @classmethod + def define_schema(cls) -> io.Schema: + return io.Schema( + node_id=f"LoraMerger{NODE_SURFIX}", + display_name=f"LoRA Merge {SYMBOL}", + category=CATEGORY_NAME, + inputs=[ + io.Custom("LoRA").Input("lora_1"), + io.Combo.Input("mode", options=["add", "concat", "svd", "svd_fast"]), + io.Int.Input( + "rank", + default=16, # Minimum value + min=1, + max=320, # Maximum value + step=1, # Slider's step + display_mode=io.NumberDisplay.number, # Cosmetic only: display as "number" or "slider" + ), + io.Float.Input( + "threshold", + default=1.0, + min=0, + max=1, + step=0.01, + ), + io.Combo.Input("device", options=["cuda", "cpu"]), + io.Combo.Input("dtype", options=["float32", "float16", "bfloat16"]), + io.Custom("LoRA").Input("lora_2", optional=True), + ], + outputs=[ + io.Custom("LoRA").Output(), + ], + ) @classmethod - def INPUT_TYPES(s): - return { - "required": { - "lora_1": ("LoRA",), - "mode": (["add", "concat", "svd"], ), - "rank": ("INT", { - "default": 16, - "min": 1, #Minimum value - "max": 320, #Maximum value - "step": 1, #Slider's step - "display": "number" # Cosmetic only: display as "number" or "slider" - }), - "threshold": ("FLOAT", { - "default": 1.0, - "min": 0, - "max": 1, - "step": 0.01, - }), - "device": (["cuda", "cpu"], ), - "dtype": (["float32", "float16", "bfloat16"], ), - }, - "optional": { - "lora_2": ("LoRA",), - } - } - RETURN_TYPES = ("LoRA", ) - FUNCTION = "lora_merge" + def execute(cls, lora_1, lora_2=None, mode=None, rank=None, threshold=None, device=None, dtype=None) -> io.NodeOutput: - CATEGORY = CATEGORY_NAME + lora = cls.merge(lora_1, lora_2, mode, rank, threshold, device, dtype) - def lora_merge(self, lora_1, lora_2=None, mode=None, rank=None, threshold=None, device=None, dtype=None): - - lora = self.merge(lora_1, lora_2, mode, rank, threshold, device, dtype) + return io.NodeOutput(lora) - return (lora, ) - + @staticmethod @torch.no_grad() - def merge(self, lora_1, lora_2, mode, rank, threshold, device, dtype): + def merge(lora_1, lora_2, mode, rank, threshold, device, dtype): # lora = up @ down * alpha / rank weight = {} @@ -57,32 +61,36 @@ class LoraMerge: if lora_2 is None: lora_2 = {"lora":{}, "strength_model":0, "strength_clip":0} - keys_1 = [key[: key.rfind(".lora_down")] for key in lora_1["lora"].keys() if ".lora_down" in key] - keys_2 = [key[: key.rfind(".lora_down")] for key in lora_2["lora"].keys() if ".lora_down" in key] + keys_1 = lora_module_keys(lora_1) + keys_2 = lora_module_keys(lora_2) keys = list(set(keys_1 + keys_2)) print(f"Merging {len(keys)} modules") - print(f"{len(keys)-len(keys_1)} modules only in lora_1") - print(f"{len(keys)-len(keys_2)} modules only in lora_2") + print(f"{len(keys)-len(keys_2)} modules only in lora_1") + print(f"{len(keys)-len(keys_1)} modules only in lora_2") pber = comfy.utils.ProgressBar(len(keys)) for key in keys: + output_format = lora_key_format(key, lora_1) or lora_key_format(key, lora_2) or REGULAR_LORA + if key not in keys_1: up, down, alpha = calc_up_down_alpha(key, lora_2) - if mode == "svd": - up, down = svd_merge(up, down, None, None, rank, threshold, device) + if mode in ("svd", "svd_fast"): + up, down = svd_merge(up, down, None, None, rank, threshold, device, fast=mode=="svd_fast") elif key not in keys_2: up, down, alpha = calc_up_down_alpha(key, lora_1) - if mode == "svd": - up, down = svd_merge(up, down, None, None, rank, threshold, device) + if mode in ("svd", "svd_fast"): + up, down = svd_merge(up, down, None, None, rank, threshold, device, fast=mode=="svd_fast") else: up_1, down_1, alpha_1 = calc_up_down_alpha(key, lora_1, add=mode!="add") up_2, down_2, alpha_2 = calc_up_down_alpha(key, lora_2, add=mode!="add") alpha = alpha_1 + alpha_1_value = alpha_to_float(alpha_1) + alpha_2_value = alpha_to_float(alpha_2) # Scale to match alpha_1 - up_2 = up_2 * math.sqrt(alpha_2/alpha) - down_2 = down_2 * math.sqrt(alpha_2/alpha) + up_2 = up_2 * math.sqrt(alpha_2_value/alpha_1_value) + down_2 = down_2 * math.sqrt(alpha_2_value/alpha_1_value) up_1 = up_1.to(dtype=dtype) down_1 = down_1.to(dtype=dtype) @@ -104,12 +112,10 @@ class LoraMerge: scale_2 = math.sqrt((r_1+r_2)/r_2) up = torch.cat([up_1*scale_1, up_2*scale_2], dim=1) down = torch.cat([down_1*scale_1, down_2*scale_2], dim=0) - elif mode == "svd": - up, down = svd_merge(up_1, down_1, up_2, down_2, rank, threshold, device) + elif mode in ("svd", "svd_fast"): + up, down = svd_merge(up_1, down_1, up_2, down_2, rank, threshold, device, fast=mode=="svd_fast") - weight[key + ".lora_up.weight"] = up - weight[key + ".lora_down.weight"] = down - weight[key + ".alpha"] = alpha + set_up_down_alpha(weight, key, up, down, alpha, output_format) pber.update(1) @@ -119,33 +125,34 @@ class LoraMerge: return {"lora":weight, "strength_model":1, "strength_clip":1} -class LoraSVDRank: - def __init__(self): - self.loaded_lora = None +class LoraSVDRank(io.ComfyNode): + @classmethod + def define_schema(cls) -> io.Schema: + return io.Schema( + node_id=f"LoraSVDRank{NODE_SURFIX}", + display_name=f"LoRA SVD Rank {SYMBOL}", + category=CATEGORY_NAME, + inputs=[ + io.Custom("LoRA").Input("lora"), + io.Float.Input( + "threshold", + default=1.0, + min=0, + max=1, + step=0.001, + ), + io.Combo.Input("device", options=["cuda", "cpu"]), + ], + outputs=[ + io.String.Output(), + ], + ) @classmethod - def INPUT_TYPES(s): - return { - "required": { - "lora": ("LoRA",), - "threshold": ("FLOAT", { - "default": 1.0, - "min": 0, - "max": 1, - "step": 0.01, - }), - "device": (["cuda", "cpu"], ), - }, - } - RETURN_TYPES = ("STRING", ) - FUNCTION = "show" - - CATEGORY = CATEGORY_NAME - @torch.no_grad() - def show(self, lora, threshold, device): + def execute(cls, lora, threshold, device) -> io.NodeOutput: - keys = [key[: key.rfind(".lora_down")] for key in lora["lora"].keys() if ".lora_down" in key] + keys = lora_module_keys(lora) pber = comfy.utils.ProgressBar(len(keys)) content = "" @@ -154,13 +161,18 @@ class LoraSVDRank: index = svd_show(up, down, threshold, device) content += f"{key}: {index}\n" pber.update(1) - - return (content, ) + + return io.NodeOutput(content) @torch.no_grad() def calc_up_down_alpha(key, lora, add=True): - up_key = key + ".lora_up.weight" - down_key = key + ".lora_down.weight" + lora_format = lora_key_format(key, lora) + if lora_format == DIFFUSERS_LORA: + up_key = key + ".lora_B.weight" + down_key = key + ".lora_A.weight" + else: + up_key = key + ".lora_up.weight" + down_key = key + ".lora_down.weight" alpha_key = key + ".alpha" is_te = "lora_te" in key @@ -171,14 +183,47 @@ def calc_up_down_alpha(key, lora, add=True): up = lora["lora"][up_key] * sqrt_scale * sign_scale down = lora["lora"][down_key] * sqrt_scale - alpha = lora["lora"][alpha_key] + alpha = lora["lora"].get(alpha_key) + if alpha is None: + alpha = torch.tensor(down.shape[0], dtype=torch.float32, device=down.device) return up, down, alpha +def lora_module_keys(lora): + keys = set() + for key in lora["lora"].keys(): + if key.endswith(".lora_down.weight"): + keys.add(key[: key.rfind(".lora_down.weight")]) + elif key.endswith(".lora_A.weight"): + keys.add(key[: key.rfind(".lora_A.weight")]) + return list(keys) + +def lora_key_format(key, lora): + state_dict = lora["lora"] + if key + ".lora_up.weight" in state_dict and key + ".lora_down.weight" in state_dict: + return REGULAR_LORA + if key + ".lora_B.weight" in state_dict and key + ".lora_A.weight" in state_dict: + return DIFFUSERS_LORA + return None + +def set_up_down_alpha(weight, key, up, down, alpha, lora_format): + if lora_format == DIFFUSERS_LORA: + weight[key + ".lora_B.weight"] = up + weight[key + ".lora_A.weight"] = down + else: + weight[key + ".lora_up.weight"] = up + weight[key + ".lora_down.weight"] = down + weight[key + ".alpha"] = alpha + +def alpha_to_float(alpha): + if torch.is_tensor(alpha): + return float(alpha.detach().cpu()) + return float(alpha) + # frovenius normによるrankの計算 -def index_sv_fro(S, target): +def index_sv_fro(S, target, total_sq=None): S_squared = S.pow(2) - s_fro_sq = float(torch.sum(S_squared)) + s_fro_sq = float(torch.sum(S_squared) if total_sq is None else total_sq) sum_S_squared = torch.cumsum(S_squared, dim=0)/s_fro_sq index = int(torch.searchsorted(sum_S_squared, target**2)) + 1 index = max(1, min(index, len(S)-1)) @@ -186,10 +231,13 @@ def index_sv_fro(S, target): return index @torch.no_grad() -def svd_merge(up_1, down_1, up_2, down_2, rank, threshold, device=None): +def svd_merge(up_1, down_1, up_2, down_2, rank, threshold, device=None, fast=False): org_device = up_1.device org_dtype = up_1.dtype + if up_2 is None and threshold >= 1 and rank == up_1.shape[1]: + return up_1.contiguous(), down_1.contiguous() + up_1 = up_1.to(device) down_1 = down_1.to(device) r_1 = up_1.shape[1] @@ -206,11 +254,17 @@ def svd_merge(up_1, down_1, up_2, down_2, rank, threshold, device=None): weight = weight.to(dtype=torch.float32) # SVD only supports float32 - U, S, Vh = torch.linalg.svd(weight) + total_sq = torch.sum(weight.pow(2)) if fast and threshold < 1 else None + + if fast: + U, S, Vh = svd_lowrank(weight, rank, threshold, total_sq=total_sq) + else: + U, S, Vh = torch.linalg.svd(weight, full_matrices=False) if threshold < 1: - rank = index_sv_fro(S, threshold) + 1 + rank = index_sv_fro(S, threshold, total_sq=total_sq) + 1 + rank = min(rank, len(S)) U = U[:, :rank] S = S[:rank] U = U @ torch.diag(S) @@ -228,11 +282,40 @@ def svd_merge(up_1, down_1, up_2, down_2, rank, threshold, device=None): U = U.reshape(up_1.shape[0], rank, 1, 1) Vh = Vh.reshape(rank, down_1.shape[1], down_1.shape[2], down_1.shape[3]) - up = U.to(org_device, dtype=org_dtype) * math.sqrt(rank) - down = Vh.to(org_device, dtype=org_dtype) * math.sqrt(rank) + up = (U.to(org_device, dtype=org_dtype) * math.sqrt(rank)).contiguous() + down = (Vh.to(org_device, dtype=org_dtype) * math.sqrt(rank)).contiguous() return up, down +@torch.no_grad() +def svd_lowrank(weight, rank, threshold, total_sq=None, oversample=8, niter=2): + max_rank = min(weight.shape) + q = min(max_rank, max(1, rank + oversample)) + + if threshold < 1: + q = min(max_rank, max(q, 32)) + + while True: + U, S, V = torch.svd_lowrank(weight, q=q, niter=niter) + order = torch.argsort(S, descending=True) + U = U[:, order] + S = S[order] + V = V[:, order] + + if threshold >= 1 or q >= max_rank: + break + + estimated_rank = index_sv_fro(S, threshold, total_sq=total_sq) + 1 + if estimated_rank < len(S) - 1: + break + + next_q = min(max_rank, q * 2) + if next_q == q: + break + q = next_q + + return U, S, V.T + @torch.no_grad() def svd_show(up, down, threshold, device): up = up.to(device) @@ -241,8 +324,10 @@ def svd_show(up, down, threshold, device): weight = up.view(-1, rank) @ down.view(rank, -1) weight = weight.to(dtype=torch.float32) # SVD only supports float32 - U, S, Vh = torch.linalg.svd(weight) + U, S, Vh = torch.linalg.svd(weight, full_matrices=False) if threshold < 1: index = index_sv_fro(S, threshold) + else: + index = rank - return index \ No newline at end of file + return index diff --git a/scripts/lora_merger/save.py b/scripts/lora_merger/save.py index ee07175..59acb22 100644 --- a/scripts/lora_merger/save.py +++ b/scripts/lora_merger/save.py @@ -2,45 +2,50 @@ import comfy import folder_paths import math import os -from ... import ROOT_NAME +from comfy_api.v0_0_2 import io +from ... import ROOT_NAME, NODE_SURFIX, SYMBOL CATEGORY_NAME = ROOT_NAME + "lora_merger" -class LoraSave: - def __init__(self): - self.loaded_lora = None +class LoraSave(io.ComfyNode): + @classmethod + def define_schema(cls) -> io.Schema: + return io.Schema( + node_id=f"LoraSave{NODE_SURFIX}", + display_name=f"LoRA Save {SYMBOL}", + category=CATEGORY_NAME, + inputs=[ + io.Custom("LoRA").Input("lora"), + io.String.Input("file_name", multiline=False, default="merged"), + io.Combo.Input("extension", options=["safetensors"]), + ], + outputs=[], + is_output_node=True, + ) @classmethod - def INPUT_TYPES(s): - return {"required": { "lora": ("LoRA",), - "file_name": ("STRING", {"multiline": False, "default": "merged"}), - "extension": (["safetensors"], ), - }} - RETURN_TYPES = () - FUNCTION = "lora_save" - - CATEGORY = CATEGORY_NAME - - OUTPUT_NODE = True - - def lora_save(self, lora, file_name, extension): + def execute(cls, lora, file_name, extension) -> io.NodeOutput: save_path = os.path.join(folder_paths.folder_names_and_paths["loras"][0][0], file_name + "." + extension) if lora["strength_model"] == 1 and lora["strength_clip"] == 1: - new_state_dict = lora["lora"] + new_state_dict = make_contiguous(lora["lora"]) else: new_state_dict = {} for key in lora["lora"].keys(): scale = lora["strength_clip"] if "lora_te" in key else lora["strength_model"] sqrt_scale = math.sqrt(abs(scale)) sign_scale = 1 if scale >= 0 else -1 - if "lora_up" in key: + if "lora_up" in key or "lora_B" in key: new_state_dict[key] = lora["lora"][key] * sqrt_scale * sign_scale - elif "lora_down" in key: + elif "lora_down" in key or "lora_A" in key: new_state_dict[key] = lora["lora"][key] * sqrt_scale else: new_state_dict[key] = lora["lora"][key] + new_state_dict = make_contiguous(new_state_dict) print(f"Saving LoRA to {save_path}") comfy.utils.save_torch_file(new_state_dict, save_path) - return {} \ No newline at end of file + return io.NodeOutput() + +def make_contiguous(state_dict): + return {key: value.contiguous() if hasattr(value, "contiguous") else value for key, value in state_dict.items()} diff --git a/scripts/lora_xy/node.py b/scripts/lora_xy/node.py index 44bd9bd..2905be3 100644 --- a/scripts/lora_xy/node.py +++ b/scripts/lora_xy/node.py @@ -1,8 +1,12 @@ import comfy +import comfy.samplers +import comfy.sd +import comfy.utils from comfy_extras.nodes_custom_sampler import SamplerCustom +import nodes import folder_paths -from nodes import LoraLoader, PreviewImage, KSampler, KSamplerAdvanced -from ... import ROOT_NAME +from ... import ROOT_NAME, NODE_SURFIX, SYMBOL +from comfy_api.v0_0_2 import io, ui import torch from PIL import Image, ImageFont, ImageDraw import numpy as np @@ -15,10 +19,10 @@ def generate_image_matrix(images, xy_list): num_images = len(images) cols = len(xy_list) # 列数 rows = num_images // cols - + fig, axes = plt.subplots(rows, cols, figsize=(cols * 2, rows * 2)) axes = axes.flatten() # 1次元配列化 - + for i in range(len(axes)): if i < num_images: axes[i].imshow(images[i]) @@ -26,143 +30,218 @@ def generate_image_matrix(images, xy_list): axes[i].axis("off") else: axes[i].axis("off") # 余ったスペースを空白にする - + plt.tight_layout() - + # Figure をバイナリデータとして保存し、PIL画像に変換 buf = BytesIO() plt.savefig(buf, format='png', bbox_inches='tight', pad_inches=0) plt.close(fig) buf.seek(0) - - return Image.open(buf) + + return Image.open(buf) CATEGORY_NAME = ROOT_NAME + "lora_xy" -class LoraLoaderModelOnlyXY(LoraLoader): +# module-level cache replacing the old per-instance `self.loaded_lora` state from +# nodes.py's LoraLoader (execute() is a classmethod, no `self` to cache on). +_lora_xy_cache = {"loaded_lora": None} + +def _load_lora_model_only(model, lora_name, strength_model): + # Mirrors nodes.py LoraLoader.load_lora(model, clip=None, lora_name, strength_model, strength_clip=0). + if strength_model == 0: + return model + + lora_path = folder_paths.get_full_path_or_raise("loras", lora_name) + lora = None + lora_metadata = None + loaded_lora = _lora_xy_cache["loaded_lora"] + if loaded_lora is not None: + if loaded_lora[0] == lora_path: + lora = loaded_lora[1] + lora_metadata = loaded_lora[2] if len(loaded_lora) > 2 else None + else: + _lora_xy_cache["loaded_lora"] = None + + if lora is None: + lora, lora_metadata = comfy.utils.load_torch_file(lora_path, safe_load=True, return_metadata=True) + _lora_xy_cache["loaded_lora"] = (lora_path, lora, lora_metadata) + + model_lora, _ = comfy.sd.load_lora_for_models(model, None, lora, strength_model, 0, lora_metadata=lora_metadata) + return model_lora + +class LoraLoaderModelOnlyXY(io.ComfyNode): @classmethod - def INPUT_TYPES(s): - return { - "required": { - "model": ("MODEL",), - "lora_name": (folder_paths.get_filename_list("loras"), ), - "strength_list": ("STRING", {"multiline": True}), - } - } - RETURN_TYPES = ("XY_MODEL","XY_LIST", ) - FUNCTION = "load_lora_model_only_xy" - CATEGORY = CATEGORY_NAME + def define_schema(cls) -> io.Schema: + return io.Schema( + node_id=f"LoraLoaderModelOnlyXY{NODE_SURFIX}", + display_name=f"Lora Loader Model Only XY {SYMBOL}", + category=CATEGORY_NAME, + inputs=[ + io.Model.Input("model"), + io.Combo.Input("lora_name", options=folder_paths.get_filename_list("loras")), + io.String.Input("strength_list", multiline=True), + ], + outputs=[ + io.Custom("XY_MODEL").Output(), + io.Custom("XY_LIST").Output(), + ], + ) - - def load_lora_model_only_xy(self, model, lora_name, strength_list): + @classmethod + def execute(cls, model, lora_name, strength_list) -> io.NodeOutput: models = [] xy_list = [] weights = [float(x.strip()) for x in strength_list.strip().strip(",").split(",")] for value in weights: - models.append(self.load_lora(model, None, lora_name, value, 0)[0]) + models.append(_load_lora_model_only(model, lora_name, value)) xy_list.append(f"{lora_name.split('.')[0]}:{value}") - return (models, xy_list) + return io.NodeOutput(models, xy_list) -class SamplerCustomXY(SamplerCustom): +class SamplerCustomXY(io.ComfyNode): @classmethod - def INPUT_TYPES(s): - return {"required": - {"model_xy": ("XY_MODEL",), - "add_noise": ("BOOLEAN", {"default": True}), - "noise_seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}), - "cfg": ("FLOAT", {"default": 8.0, "min": 0.0, "max": 100.0, "step":0.1, "round": 0.01}), - "positive": ("CONDITIONING", ), - "negative": ("CONDITIONING", ), - "sampler": ("SAMPLER", ), - "sigmas": ("SIGMAS", ), - "latent_image": ("LATENT", ), - } - } + def define_schema(cls) -> io.Schema: + return io.Schema( + node_id=f"SamplerCustomXY{NODE_SURFIX}", + display_name=f"Sampler Custom XY {SYMBOL}", + category=CATEGORY_NAME, + inputs=[ + io.Custom("XY_MODEL").Input("model_xy"), + io.Boolean.Input("add_noise", default=True), + io.Int.Input("noise_seed", default=0, min=0, max=0xffffffffffffffff), + io.Float.Input("cfg", default=8.0, min=0.0, max=100.0, step=0.1, round=0.01), + io.Conditioning.Input("positive"), + io.Conditioning.Input("negative"), + io.Sampler.Input("sampler"), + io.Sigmas.Input("sigmas"), + io.Latent.Input("latent_image"), + ], + outputs=[ + io.Latent.Output(display_name="output"), + io.Latent.Output(display_name="denoised_output"), + ], + ) - FUNCTION = "sample_xy" - CATEGORY = CATEGORY_NAME - - def sample_xy(self, model_xy, **kwargs): + @classmethod + def execute(cls, model_xy, add_noise, noise_seed, cfg, positive, negative, sampler, sigmas, latent_image) -> io.NodeOutput: outputs = [] denoised_outputs = [] + # Composition, not inheritance: SamplerCustom is itself a V3 io.ComfyNode now, so we + # call its public `execute` classmethod per model instead of subclassing it. This keeps + # us in sync with upstream's noise/x0-output/nested-tensor handling without duplicating it. for model in model_xy: - output, denoised_output = self.sample(model, **kwargs) + result = SamplerCustom.execute( + model=model, + add_noise=add_noise, + noise_seed=noise_seed, + cfg=cfg, + positive=positive, + negative=negative, + sampler=sampler, + sigmas=sigmas, + latent_image=latent_image, + ) + output, denoised_output = result.result outputs.append(output["samples"]) denoised_outputs.append(denoised_output["samples"]) - return ({"samples":torch.cat(outputs)}, {"samples":torch.cat(denoised_outputs)}) - -class KSamplerXY(KSampler): - @classmethod - def INPUT_TYPES(s): - return {"required": - {"model_xy": ("XY_MODEL",), - "seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}), - "steps": ("INT", {"default": 20, "min": 1, "max": 10000}), - "cfg": ("FLOAT", {"default": 8.0, "min": 0.0, "max": 100.0, "step":0.1, "round": 0.01}), - "sampler_name": (comfy.samplers.KSampler.SAMPLERS, ), - "scheduler": (comfy.samplers.KSampler.SCHEDULERS, ), - "positive": ("CONDITIONING", ), - "negative": ("CONDITIONING", ), - "latent_image": ("LATENT", ), - "denoise": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}), - } - } - - FUNCTION = "sample_xy" - CATEGORY = CATEGORY_NAME + return io.NodeOutput({"samples": torch.cat(outputs)}, {"samples": torch.cat(denoised_outputs)}) - def sample_xy(self, model_xy, **kwargs): +class KSamplerXY(io.ComfyNode): + @classmethod + def define_schema(cls) -> io.Schema: + return io.Schema( + node_id=f"KSamplerXY{NODE_SURFIX}", + display_name=f"KSampler XY {SYMBOL}", + category=CATEGORY_NAME, + inputs=[ + io.Custom("XY_MODEL").Input("model_xy"), + io.Int.Input("seed", default=0, min=0, max=0xffffffffffffffff), + io.Int.Input("steps", default=20, min=1, max=10000), + io.Float.Input("cfg", default=8.0, min=0.0, max=100.0, step=0.1, round=0.01), + io.Combo.Input("sampler_name", options=comfy.samplers.KSampler.SAMPLERS), + io.Combo.Input("scheduler", options=comfy.samplers.KSampler.SCHEDULERS), + io.Conditioning.Input("positive"), + io.Conditioning.Input("negative"), + io.Latent.Input("latent_image"), + io.Float.Input("denoise", default=1.0, min=0.0, max=1.0, step=0.01), + ], + outputs=[ + io.Latent.Output(), + ], + ) + + @classmethod + def execute(cls, model_xy, seed, steps, cfg, sampler_name, scheduler, positive, negative, latent_image, denoise) -> io.NodeOutput: outputs = [] + # Composition: nodes.common_ksampler is the stable module-level function that both + # KSampler and KSamplerAdvanced wrap; calling it directly avoids depending on the + # KSampler node class itself. for model in model_xy: - output = self.sample(model, **kwargs)[0] + output = nodes.common_ksampler(model, seed, steps, cfg, sampler_name, scheduler, positive, negative, latent_image, denoise=denoise)[0] outputs.append(output["samples"]) - return ({"samples":torch.cat(outputs)},) - -class KSamplerAdvancedXY(KSamplerAdvanced): - @classmethod - def INPUT_TYPES(s): - return {"required": - {"model_xy": ("XY_MODEL",), - "add_noise": (["enable", "disable"], ), - "noise_seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}), - "steps": ("INT", {"default": 20, "min": 1, "max": 10000}), - "cfg": ("FLOAT", {"default": 8.0, "min": 0.0, "max": 100.0, "step":0.1, "round": 0.01}), - "sampler_name": (comfy.samplers.KSampler.SAMPLERS, ), - "scheduler": (comfy.samplers.KSampler.SCHEDULERS, ), - "positive": ("CONDITIONING", ), - "negative": ("CONDITIONING", ), - "latent_image": ("LATENT", ), - "start_at_step": ("INT", {"default": 0, "min": 0, "max": 10000}), - "end_at_step": ("INT", {"default": 10000, "min": 0, "max": 10000}), - "return_with_leftover_noise": (["disable", "enable"], ), - } - } - - FUNCTION = "sample_xy" - CATEGORY = CATEGORY_NAME + return io.NodeOutput({"samples": torch.cat(outputs)}) - def sample_xy(self, model_xy, **kwargs): +class KSamplerAdvancedXY(io.ComfyNode): + @classmethod + def define_schema(cls) -> io.Schema: + return io.Schema( + node_id=f"KSamplerAdvancedXY{NODE_SURFIX}", + display_name=f"KSampler Advanced XY {SYMBOL}", + category=CATEGORY_NAME, + inputs=[ + io.Custom("XY_MODEL").Input("model_xy"), + io.Combo.Input("add_noise", options=["enable", "disable"]), + io.Int.Input("noise_seed", default=0, min=0, max=0xffffffffffffffff), + io.Int.Input("steps", default=20, min=1, max=10000), + io.Float.Input("cfg", default=8.0, min=0.0, max=100.0, step=0.1, round=0.01), + io.Combo.Input("sampler_name", options=comfy.samplers.KSampler.SAMPLERS), + io.Combo.Input("scheduler", options=comfy.samplers.KSampler.SCHEDULERS), + io.Conditioning.Input("positive"), + io.Conditioning.Input("negative"), + io.Latent.Input("latent_image"), + io.Int.Input("start_at_step", default=0, min=0, max=10000), + io.Int.Input("end_at_step", default=10000, min=0, max=10000), + io.Combo.Input("return_with_leftover_noise", options=["disable", "enable"]), + ], + outputs=[ + io.Latent.Output(), + ], + ) + + @classmethod + def execute(cls, model_xy, add_noise, noise_seed, steps, cfg, sampler_name, scheduler, positive, negative, latent_image, + start_at_step, end_at_step, return_with_leftover_noise) -> io.NodeOutput: outputs = [] + force_full_denoise = True + if return_with_leftover_noise == "enable": + force_full_denoise = False + disable_noise = False + if add_noise == "disable": + disable_noise = True + + # Composition: same nodes.common_ksampler function that KSamplerAdvanced.sample wraps. for model in model_xy: - output = self.sample(model, **kwargs)[0] + output = nodes.common_ksampler(model, noise_seed, steps, cfg, sampler_name, scheduler, positive, negative, latent_image, + denoise=1.0, disable_noise=disable_noise, start_step=start_at_step, last_step=end_at_step, + force_full_denoise=force_full_denoise)[0] outputs.append(output["samples"]) - return ({"samples":torch.cat(outputs)},) - + return io.NodeOutput({"samples": torch.cat(outputs)}) + class XYImage: @classmethod def INPUT_TYPES(s): return { "required":{"images": ("IMAGE", ), "xy_list": ("XY_LIST", )}, } - + FUNCTION = "xy_images" CATEGORY_NAME = ROOT_NAME @@ -172,7 +251,7 @@ class XYImage: i = 255. * image.cpu().numpy() img = Image.fromarray(np.clip(i, 0, 255).astype(np.uint8)) pil_images.append(img) - + imgs = generate_image_matrix(pil_images, xy_list) img = np.array(imgs).astype(np.float32) / 255. img = img * 2. - 1. @@ -180,79 +259,70 @@ class XYImage: return {"images": img} - +def _xy_text_to_image(text): + font = ImageFont.load_default() + img = Image.new('RGB', (256, 20), 'white') + draw = ImageDraw.Draw(img) + text_width, text_height = draw.textbbox((0,0), text, font=font)[2:] + text_x = (256 - text_width) / 2 + text_y = (20 - text_height) / 2 + draw.text((text_x, text_y), text, font=font, fill='black') + return img +def _xy_plot(images, xy_list, text_height=100): + n = len(xy_list) + m = len(images) // n + image_width, image_height = images[0].width, images[0].height -class PreviewXY(PreviewImage): + # キャンバスのサイズを再計算(全画像が同じサイズの場合) + canvas_width = image_width * n + canvas_height = (image_height * m) + text_height # 文字列の高さ分を追加 + + # キャンバスを再作成 + canvas = Image.new('RGB', (canvas_width, canvas_height), 'white') + + # 画像と文字列の画像をキャンバスに配置(全画像が同じサイズの場合の最適化) + for i, img in enumerate(images): + # 画像を配置する位置を計算 + x_offset = (i // m) * image_width + y_offset = (i % m) * (image_height) + text_height # 文字列の高さ分をオフセットして再計算 + canvas.paste(img, (x_offset, y_offset)) + + text_images = [_xy_text_to_image(title).resize((image_width, text_height)) for title in xy_list] + + # 文字列の画像をキャンバスに配置(各列の上部に) + for i, text_img in enumerate(text_images): + canvas.paste(text_img, (i * image_width, 0)) + + return canvas + +class PreviewXY(io.ComfyNode): @classmethod - def INPUT_TYPES(s): - return { - "required":{"images": ("IMAGE", ), "xy_list": ("XY_LIST", )}, - "hidden": {"prompt": "PROMPT", "extra_pnginfo": "EXTRA_PNGINFO"}, - } - - CATEGORY_NAME = ROOT_NAME - - def save_images(self, images, xy_list, filename_prefix="ComfyUI", prompt=None, extra_pnginfo=None): - filename_prefix += self.prefix_append - full_output_folder, filename, counter, subfolder, filename_prefix = folder_paths.get_save_image_path(filename_prefix, self.output_dir, images[0].shape[1], images[0].shape[0]) - results = list() - + def define_schema(cls) -> io.Schema: + return io.Schema( + node_id=f"PreviewXY{NODE_SURFIX}", + display_name=f"Preview XY {SYMBOL}", + category=CATEGORY_NAME, + inputs=[ + io.Image.Input("images"), + io.Custom("XY_LIST").Input("xy_list"), + ], + outputs=[], + is_output_node=True, + ) + + @classmethod + def execute(cls, images, xy_list) -> io.NodeOutput: pil_images = [] - for (batch_number, image) in enumerate(images): + for image in images: i = 255. * image.cpu().numpy() img = Image.fromarray(np.clip(i, 0, 255).astype(np.uint8)) pil_images.append(img) - img = self.xy_plot(pil_images, xy_list) - + canvas = _xy_plot(pil_images, xy_list) - file = "lora_xy_.png" - img.save(os.path.join(full_output_folder, file), compress_level=self.compress_level) - results.append({ - "filename": file, - "subfolder": subfolder, - "type": self.type - }) - counter += 1 + canvas_np = np.array(canvas).astype(np.float32) / 255. + canvas_tensor = torch.from_numpy(canvas_np).unsqueeze(0) - return { "ui": { "images": results } } - - def xy_plot(self, images, xy_list, text_height=100): - n = len(xy_list) - m = len(images) // n - - image_width, image_height = images[0].width, images[0].height - - # キャンバスのサイズを再計算(全画像が同じサイズの場合) - canvas_width = image_width * n - canvas_height = (image_height * m) + text_height # 文字列の高さ分を追加 - - # キャンバスを再作成 - canvas = Image.new('RGB', (canvas_width, canvas_height), 'white') - - # 画像と文字列の画像をキャンバスに配置(全画像が同じサイズの場合の最適化) - for i, img in enumerate(images): - # 画像を配置する位置を計算 - x_offset = (i // m) * image_width - y_offset = (i % m) * (image_height) + text_height # 文字列の高さ分をオフセットして再計算 - canvas.paste(img, (x_offset, y_offset)) - - text_images = [self.text_to_image(title).resize((image_width, text_height)) for title in xy_list] - - # 文字列の画像をキャンバスに配置(各列の上部に) - for i, text_img in enumerate(text_images): - canvas.paste(text_img, (i * image_width, 0)) - - return canvas - - def text_to_image(self, text): - font = ImageFont.load_default() - img = Image.new('RGB', (256, 20), 'white') - draw = ImageDraw.Draw(img) - text_width, text_height = draw.textbbox((0,0), text, font=font)[2:] - text_x = (256 - text_width) / 2 - text_y = (20 - text_height) / 2 - draw.text((text_x, text_y), text, font=font, fill='black') - return img + return io.NodeOutput(ui=ui.PreviewImage(canvas_tensor, cls=cls)) diff --git a/scripts/lortnoc/node.py b/scripts/lortnoc/node.py index fba06e5..4a6b45f 100644 --- a/scripts/lortnoc/node.py +++ b/scripts/lortnoc/node.py @@ -2,64 +2,72 @@ import comfy import folder_paths from .input_hint import ControlNetConditioningEmbedding import torch.nn.functional as F +from comfy_api.v0_0_2 import io from ... import ROOT_NAME CATEGORY_NAME = ROOT_NAME + "lortnoc" -class LortnocLoader: - def __init__(self): - self.loaded_lora = None +# module-level cache replacing the old per-instance `self.loaded_lora` / +# `self.input_hint` state (execute() is a classmethod, no `self` to cache on). +_lortnoc_cache = {"loaded_lora": None, "input_hint": None} + +class LortnocLoader(io.ComfyNode): + @classmethod + def define_schema(cls) -> io.Schema: + return io.Schema( + node_id="LortnocLoader|cgem156", + display_name="Lortnoc Loader 🍌", + category=CATEGORY_NAME, + inputs=[ + io.Model.Input("model"), + io.Image.Input("image"), + io.Combo.Input("file_name", options=folder_paths.get_filename_list("controlnet")), + io.Float.Input("strength_lora", default=1.0, min=-20.0, max=20.0, step=0.01), + io.Float.Input("strength_hint", default=1.0, min=-20.0, max=20.0, step=0.01), + ], + outputs=[ + io.Model.Output(), + ], + ) @classmethod - def INPUT_TYPES(s): - return {"required": { "model": ("MODEL",), - "image": ("IMAGE", ), - "file_name": (folder_paths.get_filename_list("controlnet"), ), - "strength_lora": ("FLOAT", {"default": 1.0, "min": -20.0, "max": 20.0, "step": 0.01}), - "strength_hint": ("FLOAT", {"default": 1.0, "min": -20.0, "max": 20.0, "step": 0.01}), - }} - RETURN_TYPES = ("MODEL", ) - FUNCTION = "load_lortnoc" - - CATEGORY = CATEGORY_NAME - - def load_lortnoc(self, model, image, file_name, strength_lora, strength_hint): + def execute(cls, model, image, file_name, strength_lora, strength_hint) -> io.NodeOutput: if strength_lora == 0 and strength_hint == 0: - return (model, ) + return io.NodeOutput(model) lora_path = folder_paths.get_full_path("controlnet", file_name) lora = None - if self.loaded_lora is not None: - if self.loaded_lora[0] == lora_path: - lora = self.loaded_lora[1] + loaded_lora = _lortnoc_cache["loaded_lora"] + if loaded_lora is not None: + if loaded_lora[0] == lora_path: + lora = loaded_lora[1] else: - temp = self.loaded_lora - self.loaded_lora = None - del temp + _lortnoc_cache["loaded_lora"] = None if lora is None: state_dict = comfy.utils.load_torch_file(lora_path, safe_load=True) lora = {k:v for k, v in state_dict.items() if "lora" in k} - self.input_hint_sd = {".".join(k.split(".")[1:]):v for k, v in state_dict.items() if "lora" not in k} - self.loaded_lora = (lora_path, lora) - - self.input_hint = ControlNetConditioningEmbedding(320, 3) - self.input_hint.load_state_dict(self.input_hint_sd) + input_hint_sd = {".".join(k.split(".")[1:]):v for k, v in state_dict.items() if "lora" not in k} + _lortnoc_cache["loaded_lora"] = (lora_path, lora) - self.hint = self.input_hint(image.permute(0, 3, 1, 2)) + input_hint = ControlNetConditioningEmbedding(320, 3) + input_hint.load_state_dict(input_hint_sd) + _lortnoc_cache["input_hint"] = input_hint + + hint = _lortnoc_cache["input_hint"](image.permute(0, 3, 1, 2)) model_lora, _ = comfy.sd.load_lora_for_models(model, None, lora, strength_lora, None) def input_block_patch(h, transformer_options): if transformer_options["block"][1] == 0: size = h.shape[2:] - if size != self.hint.shape[2:]: - hint = F.interpolate(self.hint, size, mode="bilinear", align_corners=False).to(h) + if size != hint.shape[2:]: + hint_resized = F.interpolate(hint, size, mode="bilinear", align_corners=False).to(h) else: - hint = self.hint.to(h) - h = h + hint * strength_hint - + hint_resized = hint.to(h) + h = h + hint_resized * strength_hint + return h - + model_lora.set_model_input_block_patch(input_block_patch) - return (model_lora, ) \ No newline at end of file + return io.NodeOutput(model_lora) diff --git a/scripts/multiple_lora_loader/__init__.py b/scripts/multiple_lora_loader/__init__.py index 96a185f..0d83785 100644 --- a/scripts/multiple_lora_loader/__init__.py +++ b/scripts/multiple_lora_loader/__init__.py @@ -11,7 +11,6 @@ num_loras = [int(i) for i in config.replace(" ", "").split(",")] NODE_CLASS_MAPPINGS = { f"MultipleLoraLoader{i}{NODE_SURFIX}": create_class(i) for i in num_loras } - NODE_DISPLAY_NAME_MAPPINGS = { f"MultipleLoraLoader{i}{NODE_SURFIX}": f"MultipleLoraLoader{i} {SYMBOL}" for i in num_loras } diff --git a/scripts/multiple_lora_loader/node.py b/scripts/multiple_lora_loader/node.py index 064c28d..efb252c 100644 --- a/scripts/multiple_lora_loader/node.py +++ b/scripts/multiple_lora_loader/node.py @@ -1,96 +1,134 @@ import comfy import folder_paths -from ... import ROOT_NAME +from comfy_api.v0_0_2 import io +from ... import ROOT_NAME, NODE_SURFIX, SYMBOL from .flux_map import FLUX_MAP CATEGORY_NAME = ROOT_NAME + "multiple_lora_loader" + +# Module-level cache replacing the old per-instance `self.loaded_lora` dict. +# execute() is now a classmethod (no `self` to hold state), so the cache is +# keyed by (unique_id, slot_key): unique_id identifies the node instance in the +# graph (via the hidden UNIQUE_ID input) and slot_key identifies the lora slot +# within that node (an int index for the fixed loaders, a slot name for the +# dynamic loader). This reproduces the exact old granularity -- one cache entry +# per lora slot per node instance -- just relocated out of `self`. +_lora_cache = {} + + +def _load_lora(unique_id, slot_key, model, clip, lora_name, strength_model, strength_clip): + """Load (with caching + flux key remapping) and apply a single LoRA slot.""" + if strength_model == 0 and strength_clip == 0: + return model, clip + + lora_path = folder_paths.get_full_path("loras", lora_name) + cache_key = (unique_id, slot_key) + cached = _lora_cache.get(cache_key) + + if cached is not None and cached[0] == lora_path: + new_lora = cached[1] + else: + if cached is not None: + del _lora_cache[cache_key] + state_dict = comfy.utils.load_torch_file(lora_path, safe_load=True) + new_lora = {} + for key, value in state_dict.items(): + new_lora[FLUX_MAP.get(key, key)] = value + del state_dict + _lora_cache[cache_key] = (lora_path, new_lora) + + model_lora, clip_lora = comfy.sd.load_lora_for_models(model, clip, new_lora, strength_model, strength_clip) + return model_lora, clip_lora + + +def _multiple_lora_loader(unique_id, model, clip, normalize, normalize_sum, slots): + """Shared merge logic used by both the fixed-slot and dynamic loaders. + + `slots` is an ordered list of (slot_key, lora_name, strength_model, apply). + Behavior (including the normalize division) is byte-for-byte the same math + as the original per-instance implementation; only the cache storage moved. + """ + lora_names = [s[1] for s in slots] + strength_models = [s[2] for s in slots] + applys = [s[3] for s in slots] + + for i, lora_name in enumerate(lora_names): + if lora_name == "None": + applys[i] = False + + strength_sum = 0 + for i in range(len(slots)): + if applys[i]: + strength_sum += strength_models[i] + + if normalize: + scale = normalize_sum / strength_sum + else: + scale = 1.0 + + for i, (slot_key, lora_name, strength_model, apply) in enumerate(slots): + if not applys[i]: + continue + scaled_strength = strength_model * scale + model, clip = _load_lora(unique_id, slot_key, model, clip, lora_name, scaled_strength, scaled_strength) + + return model, clip + + def create_class(num_loras): - class MultipleLoraLoader: - def __init__(self): - self.loaded_lora = {k: None for k in range(num_loras)} + """Build a V3 (io.ComfyNode) class exposing `num_loras` fixed LoRA slots. - @classmethod - def INPUT_TYPES(s): - required = {"model": ("MODEL", )} + Kept for backward compatibility with existing workflows (config.txt still + drives how many fixed-size variants get registered). Node ids, input + names/order and defaults are unchanged from the pre-V3 implementation. + """ - required["normalize"] = ("BOOLEAN", {"default": False}) - required["normalize_sum"] = ("FLOAT", {"default": 1.0, "min": -50.0, "max": 50.0, "step": 0.01}) + @classmethod + def define_schema(cls) -> io.Schema: + inputs = [ + io.Model.Input("model"), + io.Boolean.Input("normalize", default=False), + io.Float.Input("normalize_sum", default=1.0, min=-50.0, max=50.0, step=0.01, round=0.001), + ] + lora_options = ["None"] + folder_paths.get_filename_list("loras") + for i in range(num_loras): + inputs.append(io.Combo.Input(f"lora_name_{i}", options=lora_options)) + inputs.append(io.Float.Input(f"strength_model_{i}", default=1.0, min=-20.0, max=20.0, step=0.01, round=0.001)) + inputs.append(io.Boolean.Input(f"apply_{i}", default=True)) + inputs.append(io.Clip.Input("clip_optional", optional=True)) - for i in range(num_loras): - required[f"lora_name_{i}"] = (["None"] + folder_paths.get_filename_list("loras"), ) - required[f"strength_model_{i}"] = ("FLOAT", {"default": 1.0, "min": -20.0, "max": 20.0, "step": 0.01}) - required[f"apply_{i}"] = ("BOOLEAN", {"default": True}) + return io.Schema( + node_id=f"MultipleLoraLoader{num_loras}{NODE_SURFIX}", + display_name=f"MultipleLoraLoader{num_loras} {SYMBOL}", + category=CATEGORY_NAME, + description=f"Fixed {num_loras}-slot multi-LoRA loader. Slot counts are configured in config.txt.", + inputs=inputs, + outputs=[ + io.Model.Output(), + io.Clip.Output(), + ], + hidden=[io.Hidden.unique_id], + ) - return {"required": required, "optional": {"clip_optional": ("CLIP", )}} - - RETURN_TYPES = ("MODEL", "CLIP") - FUNCTION = "multiple_lora_loader" - CATEGORY = CATEGORY_NAME + @classmethod + def execute(cls, model, normalize, normalize_sum, clip_optional=None, **kwargs) -> io.NodeOutput: + clip = clip_optional - def multiple_lora_loader(self, **kwargs): + slots = [ + (i, kwargs[f"lora_name_{i}"], kwargs[f"strength_model_{i}"], kwargs[f"apply_{i}"]) + for i in range(num_loras) + ] - model = kwargs.get("model") - clip = kwargs.get("clip_optional", None) + model, clip = _multiple_lora_loader(cls.hidden.unique_id, model, clip, normalize, normalize_sum, slots) + return io.NodeOutput(model, clip) - normalize = kwargs.get("normalize") - normalize_sum = kwargs.get("normalize_sum") + return type( + f"MultipleLoraLoader{num_loras}", + (io.ComfyNode,), + { + "define_schema": define_schema, + "execute": execute, + }, + ) - lora_names = [kwargs.get(f"lora_name_{i}") for i in range(num_loras)] - strength_models = [kwargs.get(f"strength_model_{i}") for i in range(num_loras)] - applys = [kwargs.get(f"apply_{i}") for i in range(num_loras)] - - strength_sum = 0 - for i in range(num_loras): - if lora_names[i] == "None": - applys[i] = False - - if applys[i]: - strength_sum += strength_models[i] - - if normalize: - scale = normalize_sum / strength_sum - else: - scale = 1.0 - - for i in range(num_loras): - lora_name = lora_names[i] - strength_model = strength_models[i] * scale - apply = applys[i] - - #print(lora_name, strength_model, apply) - - if apply: - model, clip = self.load_lora(model, clip, lora_name, strength_model, strength_model, i) - - return (model, clip) - - def load_lora(self, model, clip, lora_name, strength_model, strength_clip, index): - if strength_model == 0 and strength_clip == 0: - return (model, clip) - - lora_path = folder_paths.get_full_path("loras", lora_name) - lora = None - if self.loaded_lora[index] is not None: - if self.loaded_lora[index][0] == lora_path: - lora = self.loaded_lora[index][1] - else: - temp = self.loaded_lora[index] - self.loaded_lora[index] = None - del temp - - if lora is None: - lora = comfy.utils.load_torch_file(lora_path, safe_load=True) - new_lora = {} - for key, value in lora.items(): - new_lora[FLUX_MAP.get(key, key)] = value - del lora - - self.loaded_lora[index] = (lora_path, new_lora) - else: - new_lora = lora - - model_lora, clip_lora = comfy.sd.load_lora_for_models(model, clip, new_lora, strength_model, strength_clip) - return (model_lora, clip_lora) - - return MultipleLoraLoader diff --git a/scripts/reference/__init__.py b/scripts/reference/__init__.py index 897d35d..ad721c0 100644 --- a/scripts/reference/__init__.py +++ b/scripts/reference/__init__.py @@ -1,14 +1,18 @@ -from .reference import ReferenceApply, ReferenceLatent +from .reference import ReferenceApply, ReferenceLatent, MultipleReferenceApply, MultipleReferenceLatent from ... import SYMBOL, NODE_SURFIX NODE_CLASS_MAPPINGS = { f"ReferenceApply{NODE_SURFIX}": ReferenceApply, f"ReferenceLatent{NODE_SURFIX}": ReferenceLatent, + f"MultipleReferenceApply{NODE_SURFIX}": MultipleReferenceApply, + f"MultipleReferenceLatent{NODE_SURFIX}": MultipleReferenceLatent, } NODE_DISPLAY_NAME_MAPPINGS = { f"ReferenceApply{NODE_SURFIX}": f"Reference Apply {SYMBOL}", f"ReferenceLatent{NODE_SURFIX}": f"Reference Latent {SYMBOL}", + f"MultipleReferenceApply{NODE_SURFIX}": f"Multiple Reference Apply {SYMBOL}", + f"MultipleReferenceLatent{NODE_SURFIX}": f"Multiple Reference Latent {SYMBOL}", } __all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"] \ No newline at end of file diff --git a/scripts/reference/reference.py b/scripts/reference/reference.py index a052d08..47e13c5 100644 --- a/scripts/reference/reference.py +++ b/scripts/reference/reference.py @@ -1,39 +1,40 @@ import torch -from ... import ROOT_NAME +from comfy_api.v0_0_2 import io +from ... import ROOT_NAME, SYMBOL, NODE_SURFIX CATEGORY_NAME = ROOT_NAME + "reference" -class ReferenceApply: +class ReferenceApply(io.ComfyNode): @classmethod - def INPUT_TYPES(s): - return { - "required": { - "model": ("MODEL",), - "index": ("INT", {"default": 0, "min": 0, "max": 256}), - "mode": (["concat", "replace"], {"default": "concat"}), - "depth": ("INT", {"default": 12, "min": -1, "max": 12}), - "start_step": ("FLOAT", {"default": 0,"min": 0, "max": 1, "step": 0.01}), - "end_step": ("FLOAT", {"default": 1, "min": 0, "max": 1, "step": 0.01}), - "apply_input": ("BOOLEAN", {"default": True}), - "apply_middle": ("BOOLEAN", {"default": True}), - "apply_output": ("BOOLEAN", {"default": True}), - } - } + def define_schema(cls) -> io.Schema: + return io.Schema( + node_id=f"ReferenceApply{NODE_SURFIX}", + display_name=f"Reference Apply {SYMBOL}", + category=CATEGORY_NAME, + inputs=[ + io.Model.Input("model"), + io.Int.Input("index", default=0, min=0, max=256), + io.Combo.Input("mode", options=["concat", "replace"], default="concat"), + io.Int.Input("depth", default=12, min=-1, max=12), + io.Float.Input("start_step", default=0, min=0, max=1, step=0.01), + io.Float.Input("end_step", default=1, min=0, max=1, step=0.01), + io.Boolean.Input("apply_input", default=True), + io.Boolean.Input("apply_middle", default=True), + io.Boolean.Input("apply_output", default=True), + ], + outputs=[ + io.Model.Output(), + ], + ) - RETURN_TYPES = ("MODEL", ) - FUNCTION = "reference_only" - - CATEGORY = CATEGORY_NAME - - def reference_only(self, model, index, mode, depth, start_step, end_step, apply_input, apply_middle, apply_output): + @classmethod + def execute(cls, model, index, mode, depth, start_step, end_step, apply_input, apply_middle, apply_output) -> io.NodeOutput: model_reference = model.clone() start_sigma = model_reference.model.model_sampling.percent_to_sigma(start_step) end_sigma = model_reference.model.model_sampling.percent_to_sigma(end_step) - self.depth = depth - - self.sdxl = hasattr(model_reference.model.diffusion_model, "label_emb") - self.num_blocks = 8 if self.sdxl else 11 + sdxl = hasattr(model_reference.model.diffusion_model, "label_emb") + num_blocks = 8 if sdxl else 11 def reference_apply(q, k, v, extra_options): block_name, block_id = extra_options["block"] @@ -46,9 +47,9 @@ class ReferenceApply: return q, k, v if block_name == "output" and not apply_output: return q, k, v - + if block_name == "output": - block_number = self.num_blocks - block_id + block_number = num_blocks - block_id else: block_number = block_id @@ -58,36 +59,38 @@ class ReferenceApply: sigma = extra_options["sigmas"][0].item() - - if end_sigma <= sigma <= start_sigma and block_number <= self.depth: + if end_sigma <= sigma <= start_sigma and block_number <= depth: k_ref = k_out[index::batch_size].repeat_interleave(batch_size, dim=0).clone() v_ref = v_out[index::batch_size].repeat_interleave(batch_size, dim=0).clone() k_out = torch.cat([k_out, k_ref], dim=1) if mode == "concat" else k_ref v_out = torch.cat([v_out, v_ref], dim=1) if mode == "concat" else v_ref - + return q_out, k_out, v_out model_reference.set_model_attn1_patch(reference_apply) - return (model_reference, ) - -class ReferenceLatent: - @classmethod - def INPUT_TYPES(s): - return { - "required": { - "latent": ("LATENT",), - "index": ("INT", {"default": 0, "min": 0, "max": 256}), - "batch_size": ("INT", {"default": 1, "min": 1, "max": 256}), - } - } - - RETURN_TYPES = ("LATENT", ) - FUNCTION = "reference_latent" - CATEGORY = CATEGORY_NAME + return io.NodeOutput(model_reference) - def reference_latent(self, latent, index, batch_size): +class ReferenceLatent(io.ComfyNode): + @classmethod + def define_schema(cls) -> io.Schema: + return io.Schema( + node_id=f"ReferenceLatent{NODE_SURFIX}", + display_name=f"Reference Latent {SYMBOL}", + category=CATEGORY_NAME, + inputs=[ + io.Latent.Input("latent"), + io.Int.Input("index", default=0, min=0, max=256), + io.Int.Input("batch_size", default=1, min=1, max=256), + ], + outputs=[ + io.Latent.Output(), + ], + ) + + @classmethod + def execute(cls, latent, index, batch_size) -> io.NodeOutput: latent_new = latent.copy() sample = latent_new["samples"] @@ -101,5 +104,120 @@ class ReferenceLatent: latent_new["samples"] = empty_latent latent_new["noise_mask"] = noise_mask - return (latent_new, ) + return io.NodeOutput(latent_new) +class MultipleReferenceApply(io.ComfyNode): + @classmethod + def define_schema(cls) -> io.Schema: + return io.Schema( + node_id=f"MultipleReferenceApply{NODE_SURFIX}", + display_name=f"Multiple Reference Apply {SYMBOL}", + category=CATEGORY_NAME, + inputs=[ + io.Model.Input("model"), + io.String.Input("indices", default="0"), + io.Int.Input("depth", default=12, min=-1, max=12), + io.Float.Input("start_step", default=0, min=0, max=1, step=0.01), + io.Float.Input("end_step", default=1, min=0, max=1, step=0.01), + io.Boolean.Input("apply_input", default=True), + io.Boolean.Input("apply_middle", default=True), + io.Boolean.Input("apply_output", default=True), + io.String.Input("weights", default=""), + ], + outputs=[ + io.Model.Output(), + ], + ) + + @classmethod + def execute(cls, model, indices, depth, start_step, end_step, apply_input, apply_middle, apply_output, weights) -> io.NodeOutput: + model_reference = model.clone() + start_sigma = model_reference.model.model_sampling.percent_to_sigma(start_step) + end_sigma = model_reference.model.model_sampling.percent_to_sigma(end_step) + + sdxl = hasattr(model_reference.model.diffusion_model, "label_emb") + num_blocks = 8 if sdxl else 11 + + indices = [int(i) for i in indices.split(",") if i.strip().isdigit()] + weights = [float(i) for i in weights.split(",") if i.strip()] if weights else [1.0] * len(indices) + + def reference_apply(q, k, v, extra_options): + block_name, block_id = extra_options["block"] + + + if block_name == "input" and not apply_input: + return q, k, v + if block_name == "middle" and not apply_middle: + return q, k, v + if block_name == "output" and not apply_output: + return q, k, v + + if block_name == "output": + block_number = num_blocks - block_id + else: + block_number = block_id + + q_out = q.clone() + k_out = k.clone() + v_out = v.clone() + + sigma = extra_options["sigmas"][0].item() + + + if end_sigma <= sigma <= start_sigma and block_number <= depth: + chunks = len(extra_options["cond_or_uncond"]) + batch_size = q.shape[0] // chunks + num_tokens = q.shape[1] + + k_refs = torch.cat([k_out[i::batch_size] for i in indices], dim=1) + v_refs = torch.cat([v_out[i::batch_size] * weight for i, weight in zip(indices, weights)], dim=1) + + k_out = k_out.repeat(1, len(indices)+1, 1).clone() + v_out = v_out.repeat(1, len(indices)+1, 1).clone() + for i in range(batch_size): + if i not in indices: + k_out[i::batch_size, num_tokens:] = k_refs.clone() + v_out[i::batch_size, num_tokens:] = v_refs.clone() + + return q_out, k_out, v_out + + model_reference.set_model_attn1_patch(reference_apply) + + return io.NodeOutput(model_reference) + +class MultipleReferenceLatent(io.ComfyNode): + @classmethod + def define_schema(cls) -> io.Schema: + return io.Schema( + node_id=f"MultipleReferenceLatent{NODE_SURFIX}", + display_name=f"Multiple Reference Latent {SYMBOL}", + category=CATEGORY_NAME, + inputs=[ + io.Latent.Input("latent"), + io.String.Input("indices", default="0"), + io.Int.Input("batch_size", default=1, min=1, max=256), + ], + outputs=[ + io.Latent.Output(), + ], + ) + + @classmethod + def execute(cls, latent, indices, batch_size) -> io.NodeOutput: + latent_new = latent.copy() + indices = [int(i) for i in indices.split(",") if i.strip().isdigit()] + + sample = latent_new["samples"] + b, _, height, width = sample.shape + + assert len(indices) == b + + empty_latent = torch.zeros_like(latent["samples"][:1]).repeat(batch_size , 1, 1, 1) + empty_latent[torch.tensor(indices)] = sample + noise_mask = torch.ones(batch_size, 1, height * 8, width * 8).to(sample) + noise_mask[torch.tensor(indices)] = 0.0 + + latent_new["samples"] = empty_latent + latent_new["noise_mask"] = noise_mask + + return io.NodeOutput(latent_new) diff --git a/scripts/scale_crafter/node.py b/scripts/scale_crafter/node.py index 596588b..97c3303 100644 --- a/scripts/scale_crafter/node.py +++ b/scripts/scale_crafter/node.py @@ -3,40 +3,78 @@ import math import comfy.ops import torch.nn.functional as F +from comfy_api.v0_0_2 import io ops = comfy.ops.disable_weight_init from ... import ROOT_NAME CATEGORY_NAME = ROOT_NAME + "scale-crafter" -class ScaleCrafter: +class ScaleCrafter(io.ComfyNode): @classmethod - def INPUT_TYPES(s): - return { - "required": { - "model": ("MODEL", ), - "dilation_rate": ("FLOAT", {"default": 1, "min": 0.01, "max": 10, "step": 0.01 }), - "depth": ("INT", {"default": 0, "min": 0, "max": 12, "step": 1, "display": "number"}), - "start": ("INT", {"default": 0, "min": 0, "max": 1000, "step": 1, "display": "number"}), - "end": ("INT", {"default": 500, "min": 0, "max": 1000, "step": 1, "display": "number"}), - }, - } + def define_schema(cls) -> io.Schema: + return io.Schema( + node_id="ScaleCrafter|cgem156", + display_name="Scale Crafter 🍌", + category=CATEGORY_NAME, + inputs=[ + io.Model.Input("model"), + io.Float.Input("dilation_rate", default=1, min=0.01, max=10, step=0.01), + io.Int.Input("depth", default=0, min=0, max=12, step=1, display_mode=io.NumberDisplay.number), + io.Int.Input("start", default=0, min=0, max=1000, step=1, display_mode=io.NumberDisplay.number), + io.Int.Input("end", default=500, min=0, max=1000, step=1, display_mode=io.NumberDisplay.number), + ], + outputs=[ + io.Model.Output(), + ], + ) - RETURN_TYPES = ("MODEL", ) - FUNCTION = "apply" - CATEGORY = CATEGORY_NAME - - def apply(self, model, dilation_rate, depth, start, end): + @classmethod + def execute(cls, model, dilation_rate, depth, start, end) -> io.NodeOutput: new_model = model.clone() - self.org_forwards = {} - self.start = start - self.end = end - self.dilation_rate = dilation_rate - self.depth = depth + org_forwards = {} - self.target_dilation = (math.ceil(self.dilation_rate), math.ceil(self.dilation_rate)) - self.target_padding = self.target_dilation - self.interp_rate = self.target_dilation[0] / self.dilation_rate + target_dilation = (math.ceil(dilation_rate), math.ceil(dilation_rate)) + target_padding = target_dilation + interp_rate = target_dilation[0] / dilation_rate + + def forward_hooker(module, forward): + def forward_hook(x): + org_size = x.shape[2:] + module.dilation = target_dilation + module.padding = target_padding + if interp_rate != 1.0: + x = F.interpolate(x, scale_factor=interp_rate, mode='bicubic', align_corners=False) + x = forward(x) + if interp_rate != 1.0: + x = F.interpolate(x, size=org_size, mode='bicubic', align_corners=False) + module.dilation = (1, 1) + module.padding = (1, 1) + return x + return forward_hook + + def replace_conv2d(model): + for name, module in model.model.diffusion_model.named_modules(): + if isinstance(module, ops.Conv2d) and module.kernel_size == (3, 3) and module.stride == (1, 1) and module.padding == (1, 1): + if name.split(".")[0] == "input_blocks": + cur_depth = int(name.split(".")[1]) + max_depth = cur_depth + elif name.split(".")[0] == "middle_block": + cur_depth = max_depth + 1 + elif name.split(".")[0] == "output_blocks": + cur_depth = max_depth - int(name.split(".")[1]) + else: + cur_depth = 0 + + if cur_depth >= depth: + org_forwards[name] = module.forward + module.forward = forward_hooker(module, org_forwards[name]) + + def restore_conv2d(model): + for name, module in model.model.diffusion_model.named_modules(): + if name in org_forwards: + module.forward = org_forwards[name] + org_forwards.clear() # unet計算前後のパッチ def apply_dilate(model_function, kwargs): @@ -44,51 +82,12 @@ class ScaleCrafter: t = new_model.model.model_sampling.timestep(sigmas) if t[0] < (1000 - end) or t[0] > (1000 - start): return model_function(kwargs["input"], kwargs["timestep"], **kwargs["c"]) - - self.replace_conv2d(new_model) + + replace_conv2d(new_model) retval = model_function(kwargs["input"], kwargs["timestep"], **kwargs["c"]) - self.restore_conv2d(new_model) + restore_conv2d(new_model) return retval new_model.set_model_unet_function_wrapper(apply_dilate) - return (new_model, ) - - def replace_conv2d(self, model): - for name, module in model.model.diffusion_model.named_modules(): - if isinstance(module, ops.Conv2d) and module.kernel_size == (3, 3) and module.stride == (1, 1) and module.padding == (1, 1): - if name.split(".")[0] == "input_blocks": - depth = int(name.split(".")[1]) - max_depth = depth - elif name.split(".")[0] == "middle_block": - depth = max_depth + 1 - elif name.split(".")[0] == "output_blocks": - depth = max_depth - int(name.split(".")[1]) - else: - depth = 0 - - if depth >= self.depth: - self.org_forwards[name] = module.forward - module.forward = self.forward_hooker(module, self.org_forwards[name]) - - def restore_conv2d(self, model): - for name, module in model.model.diffusion_model.named_modules(): - if name in self.org_forwards: - module.forward = self.org_forwards[name] - self.org_forwards = {} - - def forward_hooker(self, module, forward): - def forward_hook(x): - org_size = x.shape[2:] - module.dilation = self.target_dilation - module.padding = self.target_padding - if self.interp_rate != 1.0: - x = F.interpolate(x, scale_factor=self.interp_rate, mode='bicubic', align_corners=False) - x = forward(x) - if self.interp_rate != 1.0: - x = F.interpolate(x, size=org_size, mode='bicubic', align_corners=False) - module.dilation = (1, 1) - module.padding = (1, 1) - return x - return forward_hook - + return io.NodeOutput(new_model) diff --git a/scripts/wd-tagger/__init__.py b/scripts/wd-tagger/__init__.py index 627c2db..72ebe8c 100644 --- a/scripts/wd-tagger/__init__.py +++ b/scripts/wd-tagger/__init__.py @@ -1,4 +1,4 @@ -from .node import LoadTagger, PredictTag, GradCam, GradCamAuto, GradPair +from .node import LoadTagger, PredictTag, GradCam, GradCamAuto, GradPair, WDTaggerSimilarity from ... import SYMBOL, NODE_SURFIX NODE_CLASS_MAPPINGS = { @@ -7,6 +7,7 @@ NODE_CLASS_MAPPINGS = { f"GradCam{NODE_SURFIX}": GradCam, f"GradCamAuto{NODE_SURFIX}": GradCamAuto, f"GradPair{NODE_SURFIX}": GradPair, + f"WDTaggerSimilarity{NODE_SURFIX}": WDTaggerSimilarity, } NODE_DISPLAY_NAME_MAPPINGS = { @@ -15,6 +16,7 @@ NODE_DISPLAY_NAME_MAPPINGS = { f"GradCam{NODE_SURFIX}": f"Grad Cam {SYMBOL}", f"GradCamAuto{NODE_SURFIX}": f"Grad Cam Auto {SYMBOL}", f"GradPair{NODE_SURFIX}": f"Grad Pair {SYMBOL}", + f"WDTaggerSimilarity{NODE_SURFIX}": f"WD Tagger Similarity {SYMBOL}", } __all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"] diff --git a/scripts/wd-tagger/node.py b/scripts/wd-tagger/node.py index f8427c4..d72ff50 100644 --- a/scripts/wd-tagger/node.py +++ b/scripts/wd-tagger/node.py @@ -5,11 +5,17 @@ import cv2 import pandas as pd import torch import matplotlib.pyplot as plt +from comfy_api.v0_0_2 import io -from ... import ROOT_NAME +from ... import ROOT_NAME, SYMBOL, NODE_SURFIX CATEGORY_NAME = ROOT_NAME + "wd-tagger" +WDTagger = io.Custom("WD_TAGGER") +WDTaggerLabels = io.Custom("WD_TAGGER_LABELS") +WDTaggerFeatures = io.Custom("WD-TAGGER-FEATURES") +BatchString = io.Custom("BATCH_STRING") + MODEL_REPO_MAP = [ "SmilingWolf/wd-vit-tagger-v3", "SmilingWolf/wd-swinv2-tagger-v3", @@ -18,57 +24,68 @@ MODEL_REPO_MAP = [ "SmilingWolf/wd-eva02-large-tagger-v3", ] -class LoadTagger: - def __init__(self): - self.loaded_model = None - self.loaded_df = None - self.loaded_model_name = None +# module-level cache (V3 nodes execute as classmethods, so instance attributes are not available) +_TAGGER_CACHE = { + "loaded_model": None, + "loaded_df": None, + "loaded_model_name": None, +} + +class LoadTagger(io.ComfyNode): + @classmethod + def define_schema(cls) -> io.Schema: + return io.Schema( + node_id=f"LoadTagger{NODE_SURFIX}", + display_name=f"Load Tagger {SYMBOL}", + category=CATEGORY_NAME, + inputs=[ + io.Combo.Input("tagger", options=MODEL_REPO_MAP), + io.Combo.Input("dtype", options=["fp16", "fp32", "bf16"]), + ], + outputs=[ + WDTagger.Output(), + WDTaggerLabels.Output(), + ], + ) @classmethod - def INPUT_TYPES(s): - return { - "required": { - "tagger": (MODEL_REPO_MAP,), - "dtype": (["fp16", "fp32", "bf16"], ), - } - } - RETURN_TYPES = ("WD_TAGGER", "WD_TAGGER_LABELS") - FUNCTION = "load_tagger" - - CATEGORY = CATEGORY_NAME - @torch.inference_mode(False) - def load_tagger(self, tagger, dtype): - - if self.loaded_model_name != tagger: - self.loaded_model_name = tagger - self.loaded_model = timm.create_model(f"hf_hub:{tagger}", pretrained=True) - self.loaded_df = pd.read_csv(f"https://huggingface.co/{tagger}/resolve/main/selected_tags.csv") - self.dtype = torch.float16 if dtype == "fp16" else torch.float32 if dtype == "fp32" else torch.bfloat16 - self.loaded_model = self.loaded_model.to("cuda", dtype=self.dtype).eval() + def execute(cls, tagger, dtype) -> io.NodeOutput: - return (self.loaded_model, self.loaded_df) - -class PredictTag: + if _TAGGER_CACHE["loaded_model_name"] != tagger: + _TAGGER_CACHE["loaded_model_name"] = tagger + _TAGGER_CACHE["loaded_model"] = timm.create_model(f"hf_hub:{tagger}", pretrained=True) + _TAGGER_CACHE["loaded_df"] = pd.read_csv(f"https://huggingface.co/{tagger}/resolve/main/selected_tags.csv") + torch_dtype = torch.float16 if dtype == "fp16" else torch.float32 if dtype == "fp32" else torch.bfloat16 + _TAGGER_CACHE["loaded_model"] = _TAGGER_CACHE["loaded_model"].to("cuda", dtype=torch_dtype).eval() + + return io.NodeOutput(_TAGGER_CACHE["loaded_model"], _TAGGER_CACHE["loaded_df"]) + +class PredictTag(io.ComfyNode): @classmethod - def INPUT_TYPES(s): - return { - "required": { - "tagger": ("WD_TAGGER",), - "labels": ("WD_TAGGER_LABELS",), - "image": ("IMAGE",), - "rating": ("BOOLEAN", {"default": False}), - "character_thereshold": ("FLOAT", {"default": 0.85, "min": 0.0, "max": 1.001, "step": 0.001}), - "general_thereshold": ("FLOAT", {"default": 0.35, "min": 0.0, "max": 1.001, "step": 0.001}), - } - } - - RETURN_TYPES = ("BATCH_STRING", "STRING", "WD-TAGGER-FEATURES") - FUNCTION = "predict_tag" - CATEGORY = CATEGORY_NAME + def define_schema(cls) -> io.Schema: + return io.Schema( + node_id=f"PredictTag{NODE_SURFIX}", + display_name=f"Predict Tag {SYMBOL}", + category=CATEGORY_NAME, + inputs=[ + WDTagger.Input("tagger"), + WDTaggerLabels.Input("labels"), + io.Image.Input("image"), + io.Boolean.Input("rating", default=False), + io.Float.Input("character_thereshold", default=0.85, min=0.0, max=1.001, step=0.001), + io.Float.Input("general_thereshold", default=0.35, min=0.0, max=1.001, step=0.001), + ], + outputs=[ + BatchString.Output(), + io.String.Output(), + WDTaggerFeatures.Output(), + ], + ) + @classmethod @torch.inference_mode(False) - def predict_tag(self, tagger, labels, image, rating, character_thereshold, general_thereshold): + def execute(cls, tagger, labels, image, rating, character_thereshold, general_thereshold) -> io.NodeOutput: dtype = tagger.parameters().__next__().dtype preprocessed_image = preprocess(image).to("cuda", dtype=dtype) with torch.no_grad(): @@ -86,7 +103,7 @@ class PredictTag: tags.append(sorted_labels[sorted_labels["category"] == 9]["name"].to_list()[0]) character_tags = sorted_labels[(sorted_labels["prob"] > character_thereshold) & (sorted_labels["category"] == 4)]["name"].to_list() general_tags = sorted_labels[(sorted_labels["prob"] > general_thereshold) & (sorted_labels["category"] == 0)]["name"].to_list() - + tags += character_tags + general_tags prompt = ", ".join([tag.replace("_", " ") for tag in tags]) prompts.append(prompt) @@ -94,44 +111,47 @@ class PredictTag: string = "\n".join([f"prompt:{i}\n{prompt}" for i, prompt in enumerate(prompts)]) id_to_tag = labels['name'].to_dict() tag_to_id = {v:k for k,v in id_to_tag.items()} - + features = { "feature": feature, "image": ((preprocessed_image + 1) / 2).flip(1).permute(0, 2, 3, 1).float().cpu(), # なにこれは・・・ "tag_to_id": tag_to_id, "prob": probs } - - return (prompts, string, features) - -class GradCam: - @classmethod - def INPUT_TYPES(s): - return { - "required": { - "tagger": ("WD_TAGGER",), - "features": ("WD-TAGGER-FEATURES",), - "target_tag": ("STRING",{"default": "", "multiline": True}), - "heat_map_alpha": ("FLOAT", {"default": 0.3, "min": 0.0, "max": 1.0, "step": 0.01}), - "intepolate": (["nearest", "linear", "bilinear", "bicubic", "trilinear", "area", "nearest-exact"], {"default": "bilinear"}), - "negative": ("BOOLEAN", ), - } - } - - RETURN_TYPES = ("IMAGE", ) - FUNCTION = "grad_cam" - CATEGORY = CATEGORY_NAME + return io.NodeOutput(prompts, string, features) + +class GradCam(io.ComfyNode): + @classmethod + def define_schema(cls) -> io.Schema: + return io.Schema( + node_id=f"GradCam{NODE_SURFIX}", + display_name=f"Grad Cam {SYMBOL}", + category=CATEGORY_NAME, + inputs=[ + WDTagger.Input("tagger"), + WDTaggerFeatures.Input("features"), + io.String.Input("target_tag", default="", multiline=True), + io.Float.Input("heat_map_alpha", default=0.3, min=0.0, max=1.0, step=0.01), + io.Combo.Input("intepolate", options=["nearest", "linear", "bilinear", "bicubic", "trilinear", "area", "nearest-exact"], default="bilinear"), + io.Boolean.Input("negative"), + ], + outputs=[ + io.Image.Output(), + ], + ) + + @classmethod @torch.inference_mode(False) - def grad_cam(self, tagger, features, target_tag, heat_map_alpha, intepolate, negative): - + def execute(cls, tagger, features, target_tag, heat_map_alpha, intepolate, negative) -> io.NodeOutput: + image = features["image"] - + size = (image.shape[1], image.shape[2]) target_ids = [features["tag_to_id"][tag.strip().replace(" ", "_")] for tag in target_tag.strip().strip(",").split(",")] features = features["feature"].detach().clone().requires_grad_(True) - + gradients = [] if features.shape[1] == 1025: # eva02-large feature_size = 32 @@ -158,7 +178,7 @@ class GradCam: for i in range(len(features)): feature = features[i].unsqueeze(0) outputs = tagger.forward_head(feature).sigmoid() - + output = outputs[0, torch.tensor(target_ids)].sum(dim=-1) gradients.append(torch.autograd.grad(output, feature, retain_graph=True)[0]) @@ -183,36 +203,39 @@ class GradCam: heat_map = torch.nn.functional.interpolate(heat_map, size=size, mode=intepolate) heat_map = heat_map.permute(0, 2, 3, 1) - return (image * (1 - heat_map_alpha) + heat_map * heat_map_alpha, ) + return io.NodeOutput(image * (1 - heat_map_alpha) + heat_map * heat_map_alpha) -class GradCamAuto: +class GradCamAuto(io.ComfyNode): @classmethod - def INPUT_TYPES(s): - return { - "required": { - "tagger": ("WD_TAGGER",), - "features": ("WD-TAGGER-FEATURES",), - "threshold": ("FLOAT", {"default": 0.3, "min": 0.0, "max": 1.0, "step": 0.01}), - "heat_map_alpha": ("FLOAT", {"default": 0.3, "min": 0.0, "max": 1.0, "step": 0.01}), - "intepolate": (["nearest", "linear", "bilinear", "bicubic", "trilinear", "area", "nearest-exact"], {"default": "bilinear"}), - } - } - - RETURN_TYPES = ("IMAGE", ) - FUNCTION = "grad_cam" - CATEGORY = CATEGORY_NAME + def define_schema(cls) -> io.Schema: + return io.Schema( + node_id=f"GradCamAuto{NODE_SURFIX}", + display_name=f"Grad Cam Auto {SYMBOL}", + category=CATEGORY_NAME, + inputs=[ + WDTagger.Input("tagger"), + WDTaggerFeatures.Input("features"), + io.Float.Input("threshold", default=0.3, min=0.0, max=1.0, step=0.01), + io.Float.Input("heat_map_alpha", default=0.3, min=0.0, max=1.0, step=0.01), + io.Combo.Input("intepolate", options=["nearest", "linear", "bilinear", "bicubic", "trilinear", "area", "nearest-exact"], default="bilinear"), + ], + outputs=[ + io.Image.Output(), + ], + ) + @classmethod @torch.inference_mode(False) - def grad_cam(self, tagger, features, threshold, heat_map_alpha, intepolate): - + def execute(cls, tagger, features, threshold, heat_map_alpha, intepolate) -> io.NodeOutput: + image = features["image"].detach().clone() if image.shape[0] > 1: raise ValueError("Batch size must be 1") - + size = (image.shape[1], image.shape[2]) id_to_tag = {v:k for k,v in features["tag_to_id"].items()} features = features["feature"].detach().clone().requires_grad_(True) - + gradients = [] if features.shape[1] == 1025: # eva02-large feature_size = 32 @@ -245,7 +268,7 @@ class GradCamAuto: gradients.append(torch.autograd.grad(output, features, retain_graph=True)[0]) tagger.zero_grad() features.grad = None - + gradients = torch.cat(gradients) weight = torch.mean(gradients, dim=hw_dim, keepdim=True) @@ -271,7 +294,7 @@ class GradCamAuto: score = outputs[0, target_id] cv2.putText(image, f"{target_tag}:", (10, 30), cv2.FONT_HERSHEY_SIMPLEX, 1, (255, 255, 255), 2) cv2.putText(image, f"{score:.2f}", (10, 60), cv2.FONT_HERSHEY_SIMPLEX, 1, (255, 255, 255), 2) - + # sort by score image_score = [(image, output.item()) for image, output in zip(images, outputs_filtered)] image_score.sort(key=lambda x: x[1], reverse=True) @@ -279,28 +302,32 @@ class GradCamAuto: output_image = torch.from_numpy(np.array(images)) output_image = output_image.float() / 255 - return (output_image, ) + return io.NodeOutput(output_image) -class GradPair: +class GradPair(io.ComfyNode): @classmethod - def INPUT_TYPES(s): - return { - "required": { - "tagger": ("WD_TAGGER",), - "features": ("WD-TAGGER-FEATURES",), - "heat_map_alpha": ("FLOAT", {"default": 0.3, "min": 0.0, "max": 1.0, "step": 0.01}), - "intepolate": (["nearest", "linear", "bilinear", "bicubic", "trilinear", "area", "nearest-exact"], {"default": "bilinear"}), - "negative": ("BOOLEAN", ), - } - } - - RETURN_TYPES = ("IMAGE", "STRING") - FUNCTION = "grad_cam" - CATEGORY = CATEGORY_NAME + def define_schema(cls) -> io.Schema: + return io.Schema( + node_id=f"GradPair{NODE_SURFIX}", + display_name=f"Grad Pair {SYMBOL}", + category=CATEGORY_NAME, + inputs=[ + WDTagger.Input("tagger"), + WDTaggerFeatures.Input("features"), + io.Float.Input("heat_map_alpha", default=0.3, min=0.0, max=1.0, step=0.01), + io.Combo.Input("intepolate", options=["nearest", "linear", "bilinear", "bicubic", "trilinear", "area", "nearest-exact"], default="bilinear"), + io.Boolean.Input("negative"), + ], + outputs=[ + io.Image.Output(), + io.String.Output(), + ], + ) + @classmethod @torch.inference_mode(False) - def grad_cam(self, tagger, features, heat_map_alpha, intepolate, negative): - + def execute(cls, tagger, features, heat_map_alpha, intepolate, negative) -> io.NodeOutput: + prob_diff = (features["prob"][0] - features["prob"][1]) prob_diff_data = pd.DataFrame({"label": features["tag_to_id"].keys(), "prob_diff": prob_diff}) prob_diff_data = prob_diff_data.sort_values(by="prob_diff", ascending=False) @@ -309,15 +336,15 @@ class GradPair: bottom_20 = prob_diff_data.tail(20).sort_values(by="prob_diff") output_string = f"Top 20 difference:\n{top_20.to_string(index=False)}\n ... \n:\n{bottom_20.to_string(index=False)}" - + image = features["image"] if image.shape[0] != 2: raise ValueError("Batch size must be 2") - + size = (image.shape[1], image.shape[2]) features = features["feature"].detach().clone().requires_grad_(True) - + gradients = [] if features.shape[1] == 1025: # eva02-large feature_size = 32 @@ -371,4 +398,52 @@ class GradPair: heat_map = torch.nn.functional.interpolate(heat_map, size=size, mode=intepolate) heat_map = heat_map.permute(0, 2, 3, 1) - return (image * (1 - heat_map_alpha) + heat_map * heat_map_alpha, output_string) \ No newline at end of file + return io.NodeOutput(image * (1 - heat_map_alpha) + heat_map * heat_map_alpha, output_string) + +class WDTaggerSimilarity(io.ComfyNode): + @classmethod + def define_schema(cls) -> io.Schema: + return io.Schema( + node_id=f"WDTaggerSimilarity{NODE_SURFIX}", + display_name=f"WD Tagger Similarity {SYMBOL}", + category=CATEGORY_NAME, + inputs=[ + WDTagger.Input("tagger"), + WDTaggerLabels.Input("labels"), + io.String.Input("tag", multiline=True), + io.Combo.Input("category", options=["all", "general", "character"]), + io.Boolean.Input("ascending", default=False), + ], + outputs=[ + io.String.Output(), + ], + ) + + @classmethod + def execute(cls, tagger, labels, tag, category, ascending) -> io.NodeOutput: + dtype = tagger.parameters().__next__().dtype + tag_list = [t.strip().replace(" ", "_") for t in tag.strip().strip(",").split(",")] + tag_ids = [labels[labels["name"] == t].index[0] for t in tag_list if t in labels["name"].values] + if len(tag_ids) == 0: + return io.NodeOutput(f"No valid tags found in input: {tag}") + + with torch.no_grad(): + tag_embeddings = tagger.get_classifier().weight[tag_ids].to("cpu", dtype=dtype) + all_embeddings = tagger.get_classifier().weight.to("cpu", dtype=dtype) + + tag_embeddings = tag_embeddings / tag_embeddings.norm(dim=1, keepdim=True) + all_embeddings = all_embeddings / all_embeddings.norm(dim=1, keepdim=True) + + similarity = torch.matmul(all_embeddings, tag_embeddings.T).min(dim=1).values.cpu().numpy() + + labels["similarity"] = similarity + if category == "general": + labels = labels[labels["category"] == 0] + elif category == "character": + labels = labels[labels["category"] == 4] + + labels = labels.sort_values(by="similarity", ascending=ascending) + output_string = f"Similarity result for tags: {', '.join(tag_list)}\n" + output_string += labels[["name", "similarity"]].head(50).to_string(index=False) + + return io.NodeOutput(output_string)