Merge pull request #25 from laksjdjf/node-v3-migration
Node V3スキーマへの全面移行(全60ノード)
This commit is contained in:
@@ -0,0 +1,7 @@
|
||||
__pycache__/
|
||||
*.pyc
|
||||
.ipynb_checkpoints/
|
||||
*.ipynb
|
||||
scripts/batch_condition/train/
|
||||
scripts/batch_condition/test/
|
||||
scripts/reference/cache/
|
||||
+17
-2
@@ -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'):
|
||||
|
||||
@@ -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;
|
||||
|
||||
}
|
||||
}
|
||||
},
|
||||
});
|
||||
@@ -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;
|
||||
}
|
||||
}
|
||||
},
|
||||
});
|
||||
@@ -0,0 +1,5 @@
|
||||
transformers
|
||||
timm
|
||||
pandas
|
||||
opencv-python
|
||||
matplotlib
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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"]
|
||||
+185
-112
@@ -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 = "<color>"
|
||||
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)
|
||||
|
||||
+30
-56
@@ -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",
|
||||
]
|
||||
|
||||
|
||||
@@ -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,)
|
||||
|
||||
return io.NodeOutput(guider)
|
||||
|
||||
@@ -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"]
|
||||
@@ -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))
|
||||
@@ -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))
|
||||
@@ -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),)
|
||||
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))
|
||||
|
||||
@@ -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, )
|
||||
return io.NodeOutput(sampler)
|
||||
@@ -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",
|
||||
}
|
||||
return io.NodeOutput(sampler)
|
||||
|
||||
@@ -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,
|
||||
}
|
||||
return io.NodeOutput(sampler)
|
||||
@@ -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)
|
||||
return io.NodeOutput(out, out_denoised, previews)
|
||||
@@ -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,
|
||||
}
|
||||
return io.NodeOutput(sampler)
|
||||
@@ -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,
|
||||
|
||||
+142
-138
@@ -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:{rating}</rating>"
|
||||
prompt += f"<copylight>{copyright}</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"<copylight>{copyright}</copyright>"
|
||||
prompt += f"<character>{character}</character>"
|
||||
prompt += f"<|rating:{rating}|>" + f"<|aspect_ratio:{aspect_ratio}|>" + f"<|length:{length}|>" + f"<|identity:{identity}|>"
|
||||
prompt += f"<general>{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)
|
||||
|
||||
@@ -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"]
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
@@ -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}]])
|
||||
@@ -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)
|
||||
+63
-56
@@ -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})
|
||||
|
||||
+175
-90
@@ -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
|
||||
return index
|
||||
|
||||
+26
-21
@@ -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 {}
|
||||
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()}
|
||||
|
||||
+231
-161
@@ -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))
|
||||
|
||||
+44
-36
@@ -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, )
|
||||
return io.NodeOutput(model_lora)
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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"]
|
||||
+167
-49
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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"]
|
||||
|
||||
+191
-116
@@ -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)
|
||||
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)
|
||||
|
||||
Reference in New Issue
Block a user