From b33eee0c8974e067b532fe8183fe2e458e04fc96 Mon Sep 17 00:00:00 2001 From: Austin Mroz Date: Wed, 6 Dec 2023 11:54:04 -0600 Subject: [PATCH] Initial proof of concept Novel results have been obtained with the current nodes, but results are finicky --- .gitignore | 1 + README.md | 2 ++ __init__.py | 4 +++ nodes.py | 94 +++++++++++++++++++++++++++++++++++++++++++++++++++++ 4 files changed, 101 insertions(+) create mode 100644 .gitignore create mode 100644 README.md create mode 100644 __init__.py create mode 100644 nodes.py diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000..c18dd8d --- /dev/null +++ b/.gitignore @@ -0,0 +1 @@ +__pycache__/ diff --git a/README.md b/README.md new file mode 100644 index 0000000..36a6e48 --- /dev/null +++ b/README.md @@ -0,0 +1,2 @@ +# SpliceTools +Experimental utility nodes with a focus on manipulation of noised latents diff --git a/__init__.py b/__init__.py new file mode 100644 index 0000000..d9738a3 --- /dev/null +++ b/__init__.py @@ -0,0 +1,4 @@ +from .nodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS + +WEB_DIRECTORY = "./web" +__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS", "WEB_DIRECTORY"] diff --git a/nodes.py b/nodes.py new file mode 100644 index 0000000..d93addc --- /dev/null +++ b/nodes.py @@ -0,0 +1,94 @@ +import torch + +import torch.nn.functional as F +from comfy_extras.nodes_post_processing import gaussian_kernel + +class LogSigmas: + """For testing, simply prints the input sigmas""" + @classmethod + def INPUT_TYPES(s): + return {"required": { "sigmas": ("SIGMAS",), + }} + + FUNCTION = "log_sigmas" + OUTPUT_NODE = True + RETURN_TYPES = () + CATEGORY = "_for_testing" + + def log_sigmas(self, sigmas): + print(sigmas) + return () + +#Blur functions shameless stolen borrowed comfy_extras/nodes_post_processing +#with slight modifications for latent dimensions + +def gaussian_blur(latents, kernel, radius=20): + padded_latents = F.pad(latents, [radius]*4, 'reflect') + blurred = F.conv2d(padded_latents, kernel, padding=(radius*2+1) // 2, groups=4) + return blurred[:, :, radius:-radius, radius:-radius] + +class SpliceLatents: + """Performs a fast approximate splice of 2 latents by bluring. + Intended to eventually automatically calculate blur strength from sigmas""" + @classmethod + def INPUT_TYPES(s): + #These numbers are likely flawed + return {"required": {"mult": ("FLOAT", {"default": 1.0, "precision": 3, + "step": 0.1, "round": .001}), + "size": ("INT", {"default": 4, "min": 1, "step": 1}), + "wetness": ("FLOAT", {"default": 1.0, "max": 1, + "min": 0, "precision": 3, + "step": 0.1, "round": .01})}, + "optional": {"lower": ("LATENT",), + "upper": ("LATENT",)}} + + FUNCTION = "splice_latents" + RETURN_TYPES = ("LATENT",) + CATEGORY = "latent/advanced" + + def splice_latents(s, mult, size, wetness=1.0, lower=None, upper=None): + if lower is None and upper is None: + raise "lower and upper can't both be none" + if lower is None: + lower = torch.zeros_like(upper['samples']) + else: + lower = lower['samples'] + if upper is None: + upper = torch.zeros_like(lower) + else: + upper = upper['samples'] + radius = size + kernel = gaussian_kernel(radius * 2 + 1, mult, device=lower.device).repeat(4,1,1).unsqueeze(1) + + lower_b = gaussian_blur(lower, kernel, radius) + upper_b = gaussian_blur(upper, kernel, radius) + upper_e = upper - upper_b + lower_out = lower_b * wetness + lower * (1 - wetness) + upper_out = upper_e * wetness + upper * (1 - wetness) + + return ({"samples": lower_out + upper_out},) + +class SpliceDenoised: + """A convenience node to splice latents when both noised and denoised outputs exist""" + @classmethod + def INPUT_TYPES(s): + return {"required": { + "noised_latent" : ("LATENT",), + "denoised_latent" : ("LATENT",), + "donor_latent" : ("LATENT",), + }} + + RETURN_TYPES = ("LATENT",) + FUNCTION = "splice_denoised" + CATEGORY = "_for_testing" + + def splice_denoised(self, noised_latent, denoised_latent, donor_latent): + samples = noised_latent['samples'] - denoised_latent['samples'] + donor_latent['samples'] + return ({"samples": samples},) + +NODE_CLASS_MAPPINGS = { + "LogSigmas": LogSigmas, + "SpliceLatents": SpliceLatents, + "SpliceDenoised": SpliceDenoised +} +NODE_DISPLAY_NAME_MAPPINGS = {}