improve the math nodes

This commit is contained in:
cubiq
2024-08-18 16:49:20 +02:00
parent 1526e2c18a
commit 99a1423f7b
3 changed files with 234 additions and 25 deletions
+185 -2
View File
@@ -6,6 +6,70 @@ from nodes import MAX_RESOLUTION
any = AnyType("*")
class SimpleMathFloat:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"value": ("FLOAT", { "default": 0.0, "min": -0xffffffffffffffff, "max": 0xffffffffffffffff, "step": 0.05 }),
},
}
RETURN_TYPES = ("FLOAT", )
FUNCTION = "execute"
CATEGORY = "essentials/utilities"
def execute(self, value):
return (float(value), )
class SimpleMathPercent:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"value": ("FLOAT", { "default": 0.0, "min": 0, "max": 1, "step": 0.05 }),
},
}
RETURN_TYPES = ("FLOAT", )
FUNCTION = "execute"
CATEGORY = "essentials/utilities"
def execute(self, value):
return (float(value), )
class SimpleMathInt:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"value": ("INT", { "default": 0, "min": -0xffffffffffffffff, "max": 0xffffffffffffffff, "step": 1 }),
},
}
RETURN_TYPES = ("INT",)
FUNCTION = "execute"
CATEGORY = "essentials/utilities"
def execute(self, value):
return (int(value), )
class SimpleMathSlider:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"value": ("FLOAT", { "display": "slider", "default": 0.5, "min": 0.0, "max": 1.0, "step": 0.001 }),
},
}
RETURN_TYPES = ("FLOAT",)
FUNCTION = "execute"
CATEGORY = "essentials/utilities"
def execute(self, value):
return (value, )
class SimpleMath:
@classmethod
def INPUT_TYPES(s):
@@ -13,6 +77,7 @@ class SimpleMath:
"optional": {
"a": (any, { "default": 0.0 }),
"b": (any, { "default": 0.0 }),
"c": (any, { "default": 0.0 }),
},
"required": {
"value": ("STRING", { "multiline": False, "default": "" }),
@@ -23,12 +88,13 @@ class SimpleMath:
FUNCTION = "execute"
CATEGORY = "essentials/utilities"
def execute(self, value, a = 0.0, b = 0.0):
def execute(self, value, a = 0.0, b = 0.0, c = 0.0):
import ast
import operator as op
a = float(a)
b = float(b)
c = float(c)
operators = {
ast.Add: op.add,
@@ -40,6 +106,15 @@ class SimpleMath:
ast.BitXor: op.xor,
ast.USub: op.neg,
ast.Mod: op.mod,
ast.Eq: op.eq,
ast.NotEq: op.ne,
ast.Lt: op.lt,
ast.LtE: op.le,
ast.Gt: op.gt,
ast.GtE: op.ge,
#ast.And: op.and_,
#ast.Or: op.or_,
ast.Not: op.not_
}
op_functions = {
@@ -58,10 +133,23 @@ class SimpleMath:
return a
if node.id == "b":
return b
if node.id == "c":
return c
elif isinstance(node, ast.BinOp): # <left> <operator> <right>
return operators[type(node.op)](eval_(node.left), eval_(node.right))
elif isinstance(node, ast.UnaryOp): # <operator> <operand> e.g., -1
return operators[type(node.op)](eval_(node.operand))
elif isinstance(node, ast.Compare): # comparison operators
left = eval_(node.left)
for op, comparator in zip(node.ops, node.comparators):
if not operators[type(op)](left, eval_(comparator)):
return 0
return 1
elif isinstance(node, ast.BoolOp): # boolean operators (And, Or)
if isinstance(node.op, ast.And):
return all(eval_(value) for value in node.values)
elif isinstance(node.op, ast.Or):
return any(eval_(value) for value in node.values)
elif isinstance(node, ast.Call): # custom function
if node.func.id in op_functions:
args =[eval_(arg) for arg in node.args]
@@ -82,6 +170,87 @@ class SimpleMath:
return (round(result), result, )
class SimpleMathCondition:
@classmethod
def INPUT_TYPES(s):
return {
"optional": {
"a": (any, { "default": 0.0 }),
"b": (any, { "default": 0.0 }),
"c": (any, { "default": 0.0 }),
},
"required": {
"evaluate": (any, {"default": 0}),
"on_true": ("STRING", { "multiline": False, "default": "" }),
"on_false": ("STRING", { "multiline": False, "default": "" }),
},
}
RETURN_TYPES = ("INT", "FLOAT", )
FUNCTION = "execute"
CATEGORY = "essentials/utilities"
def execute(self, evaluate, on_true, on_false, a = 0.0, b = 0.0, c = 0.0):
return SimpleMath().execute(on_true if evaluate else on_false, a, b, c)
class SimpleCondition:
def __init__(self):
pass
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"evaluate": (any, {"default": 0}),
"on_true": (any, {"default": 0}),
},
"optional": {
"on_false": (any, {"default": 0}),
},
}
RETURN_TYPES = (any,)
RETURN_NAMES = ("value",)
FUNCTION = "execute"
CATEGORY = "essentials/utilities"
def execute(self, evaluate, on_true, on_false=0):
return (on_true if evaluate else on_false,)
class SimpleComparison:
def __init__(self):
pass
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"a": (any, {"default": 0}),
"b": (any, {"default": 0}),
"comparison": (["==", "!=", "<", "<=", ">", ">="],),
},
}
RETURN_TYPES = ("BOOLEAN",)
FUNCTION = "execute"
CATEGORY = "essentials/utilities"
def execute(self, a, b, comparison):
if comparison == "==":
return (a == b,)
elif comparison == "!=":
return (a != b,)
elif comparison == "<":
return (a < b,)
elif comparison == "<=":
return (a <= b,)
elif comparison == ">":
return (a > b,)
elif comparison == ">=":
return (a >= b,)
class ConsoleDebug:
@classmethod
def INPUT_TYPES(s):
@@ -232,7 +401,14 @@ MISC_CLASS_MAPPINGS = {
"ModelCompile+": ModelCompile,
"RemoveLatentMask+": RemoveLatentMask,
"SDXLEmptyLatentSizePicker+": SDXLEmptyLatentSizePicker,
"SimpleComparison+": SimpleComparison,
"SimpleCondition+": SimpleCondition,
"SimpleMath+": SimpleMath,
"SimpleMathCondition+": SimpleMathCondition,
"SimpleMathFloat+": SimpleMathFloat,
"SimpleMathInt+": SimpleMathInt,
"SimpleMathPercent+": SimpleMathPercent,
"SimpleMathSlider+": SimpleMathSlider,
}
MISC_NAME_MAPPINGS = {
@@ -241,6 +417,13 @@ MISC_NAME_MAPPINGS = {
"DebugTensorShape+": "🔧 Debug Tensor Shape",
"ModelCompile+": "🔧 Model Compile",
"RemoveLatentMask+": "🔧 Remove Latent Mask",
"SDXLEmptyLatentSizePicker+": "🔧 SDXL Empty Latent Size Picker",
"SDXLEmptyLatentSizePicker+": "🔧 Empty Latent Size Picker",
"SimpleComparison+": "🔧 Simple Comparison",
"SimpleCondition+": "🔧 Simple Condition",
"SimpleMath+": "🔧 Simple Math",
"SimpleMathCondition+": "🔧 Simple Math Condition",
"SimpleMathFloat+": "🔧 Simple Math Float",
"SimpleMathInt+": "🔧 Simple Math Int",
"SimpleMathPercent+": "🔧 Simple Math Percent",
"SimpleMathSlider+": "🔧 Simple Math Slider",
}
+45 -22
View File
@@ -9,6 +9,25 @@ import torchvision.transforms.v2 as T
import torch.nn.functional as F
import logging
# From https://github.com/BlenderNeko/ComfyUI_Noise/
def slerp(val, low, high):
dims = low.shape
low = low.reshape(dims[0], -1)
high = high.reshape(dims[0], -1)
low_norm = low/torch.norm(low, dim=1, keepdim=True)
high_norm = high/torch.norm(high, dim=1, keepdim=True)
low_norm[low_norm != low_norm] = 0.0
high_norm[high_norm != high_norm] = 0.0
omega = torch.acos((low_norm*high_norm).sum(1))
so = torch.sin(omega)
res = (torch.sin((1.0-val)*omega)/so).unsqueeze(1)*low + (torch.sin(val*omega)/so).unsqueeze(1) * high
return res.reshape(dims)
class KSamplerVariationsWithNoise:
@classmethod
def INPUT_TYPES(s):
@@ -34,25 +53,6 @@ class KSamplerVariationsWithNoise:
FUNCTION = "execute"
CATEGORY = "essentials/sampling"
# From https://github.com/BlenderNeko/ComfyUI_Noise/
def slerp(self, val, low, high):
dims = low.shape
low = low.reshape(dims[0], -1)
high = high.reshape(dims[0], -1)
low_norm = low/torch.norm(low, dim=1, keepdim=True)
high_norm = high/torch.norm(high, dim=1, keepdim=True)
low_norm[low_norm != low_norm] = 0.0
high_norm[high_norm != high_norm] = 0.0
omega = torch.acos((low_norm*high_norm).sum(1))
so = torch.sin(omega)
res = (torch.sin((1.0-val)*omega)/so).unsqueeze(1)*low + (torch.sin(val*omega)/so).unsqueeze(1) * high
return res.reshape(dims)
def prepare_mask(self, mask, shape):
mask = torch.nn.functional.interpolate(mask.reshape((-1, 1, mask.shape[-2], mask.shape[-1])), size=(shape[2], shape[3]), mode="bilinear")
mask = mask.expand((-1,shape[1],-1,-1))
@@ -81,7 +81,7 @@ class KSamplerVariationsWithNoise:
generator = torch.manual_seed(variation_seed)
variation_noise = torch.randn((batch_size, 4, height, width), dtype=torch.float32, device="cpu", generator=generator).cpu()
slerp_noise = self.slerp(variation_strength, base_noise, variation_noise)
slerp_noise = slerp(variation_strength, base_noise, variation_noise)
# Calculate sigma
comfy.model_management.load_model_gpu(model)
@@ -159,16 +159,39 @@ class InjectLatentNoise:
"latent": ("LATENT", ),
"noise_seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
"noise_strength": ("FLOAT", {"default": 1.0, "min": -20.0, "max": 20.0, "step":0.01, "round": 0.01}),
"normalize": (["false", "true"], {"default": "false"}),
},
"optional": {
"mask": ("MASK", ),
}}
RETURN_TYPES = ("LATENT",)
FUNCTION = "execute"
CATEGORY = "essentials/sampling"
def execute(self, latent, noise_seed, noise_strength):
def execute(self, latent, noise_seed, noise_strength, normalize="false", mask=None):
torch.manual_seed(noise_seed)
noise_latent = latent.copy()
noise_latent["samples"] = noise_latent["samples"].clone() + torch.randn_like(noise_latent["samples"]) * noise_strength
original_samples = noise_latent["samples"].clone()
random_noise = torch.randn_like(original_samples)
if normalize == "true":
mean = original_samples.mean()
std = original_samples.std()
random_noise = random_noise * std + mean
random_noise = original_samples + random_noise * noise_strength
if mask is not None:
mask = F.interpolate(mask.reshape((-1, 1, mask.shape[-2], mask.shape[-1])), size=(random_noise.shape[2], random_noise.shape[3]), mode="bilinear")
mask = mask.expand((-1,random_noise.shape[1],-1,-1)).clamp(0.0, 1.0)
if mask.shape[0] < random_noise.shape[0]:
mask = mask.repeat((random_noise.shape[0] -1) // mask.shape[0] + 1, 1, 1, 1)[:random_noise.shape[0]]
elif mask.shape[0] > random_noise.shape[0]:
mask = mask[:random_noise.shape[0]]
random_noise = mask * random_noise + (1-mask) * original_samples
noise_latent["samples"] = random_noise
return (noise_latent, )
+4 -1
View File
@@ -21,6 +21,7 @@ class DrawText:
"vertical_align": (["top", "center", "bottom"],),
"offset_x": ("INT", { "default": 0, "min": -MAX_RESOLUTION, "max": MAX_RESOLUTION, "step": 1 }),
"offset_y": ("INT", { "default": 0, "min": -MAX_RESOLUTION, "max": MAX_RESOLUTION, "step": 1 }),
"direction": (["ltr", "rtl"],),
},
"optional": {
"img_composite": ("IMAGE",),
@@ -31,12 +32,14 @@ class DrawText:
FUNCTION = "execute"
CATEGORY = "essentials/text"
def execute(self, text, font, size, color, background_color, shadow_distance, shadow_blur, shadow_color, horizontal_align, vertical_align, offset_x, offset_y, img_composite=None):
def execute(self, text, font, size, color, background_color, shadow_distance, shadow_blur, shadow_color, horizontal_align, vertical_align, offset_x, offset_y, direction, img_composite=None):
from PIL import Image, ImageDraw, ImageFont, ImageColor, ImageFilter
font = ImageFont.truetype(os.path.join(FONTS_DIR, font), size)
lines = text.split("\n")
if direction == "rtl":
lines = [line[::-1] for line in lines]
# Calculate the width and height of the text
text_width = max(font.getbbox(line)[2] for line in lines)