diff --git a/__init__.py b/__init__.py index b766454..b45203b 100644 --- a/__init__.py +++ b/__init__.py @@ -8,9 +8,13 @@ NODE_CLASS_MAPPINGS = { "MirrorTransform": MirrorTransform, "ShiftTransform": ShiftTransform, "MultiplyTransform": MultiplyTransform, + "LatentInterpolateTransform": LatentInterpolateTransform, + "LatentAddTransform": LatentAddTransform, "OneTimeMirrorTransform": OneTimeMirrorTransform, "OneTimeMultiplyTransform": OneTimeMultiplyTransform, "OneTimeShiftTransform": OneTimeShiftTransform, + "OneTimeLatentInterpolateTransform": OneTimeLatentInterpolateTransform, + "OneTimeLatentAddTransform": OneTimeLatentAddTransform, "TransformsCombine": TransformsCombine, "TransformOffset": TransformOffset, } @@ -23,9 +27,13 @@ NODE_DISPLAY_NAME_MAPPINGS = { "MirrorTransform": "Mirror transform", "ShiftTransform": "Shift transform", "MultiplyTransform": "Multiply transform", + "LatentInterpolateTransform": "Latent interpolate transform", + "LatentAddTransform": "Latent add transform", "OneTimeMirrorTransform": "Mirror transform (one time)", "OneTimeMultiplyTransform": "Shift transform (one time)", "OneTimeShiftTransform": "Multiply transform (one time)", + "OneTimeLatentInterpolateTransform": "Latent interpolate transform (one time)", + "OneTimeLatentAddTransform": "Latent add transform (one time)", "TransformsCombine": "Combine transforms", "TransformOffset": "Transform offset", } diff --git a/nodes/KSamplerNodes/Transforms/LatentAddTransform.py b/nodes/KSamplerNodes/Transforms/LatentAddTransform.py new file mode 100644 index 0000000..f050923 --- /dev/null +++ b/nodes/KSamplerNodes/Transforms/LatentAddTransform.py @@ -0,0 +1,47 @@ +from .transform_functions import latent_add_transform + + +class LatentAddTransform: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "latent": ("LATENT",), + "offset": ("OFFSET",), + "start_at": ("FLOAT", {"default": 0, "min": 0.0, "max": 1.0, "step": 0.01}), + "stop_at": ("FLOAT", {"default": 0, "min": 0.0, "max": 1.0, "step": 0.01}), + "multiplier": ("FLOAT", {"default": 1, "min": -10, "max": 10, "step": 0.01}), + } + } + + RETURN_TYPES = ("TRANSFORM",) + FUNCTION = "process" + + CATEGORY = "sampling/transforms" + + def process(self, + offset, + latent, + start_at=0, + stop_at=0, + multiplier=1): + return ([{ + "params": { + "latent": latent["samples"][0], + "start_at": start_at, + "stop_at": stop_at, + "multiplier": multiplier, + "offset": offset, + "offset_status": offset["process_every"] - offset["offset"] - 1, + }, + "function": self.func + }],) + + def func(self, step, x0, total_steps, params): + if (total_steps * params["start_at"] <= step <= total_steps * params["stop_at"] and + (params["offset_status"] == 0 if params["offset"]["mode"] == "process_every" else params["offset_status"] != 0)): + if params["offset_status"] == 0: + params["offset_status"] = params["offset"]["process_every"] - 1 + else: + params["offset_status"] -= 1 + return latent_add_transform(x0, params) diff --git a/nodes/KSamplerNodes/Transforms/LatentIntrpolateTransform.py b/nodes/KSamplerNodes/Transforms/LatentIntrpolateTransform.py new file mode 100644 index 0000000..32019ab --- /dev/null +++ b/nodes/KSamplerNodes/Transforms/LatentIntrpolateTransform.py @@ -0,0 +1,50 @@ +from .transform_functions import latent_interpolate_transform + + +class LatentInterpolateTransform: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "latent": ("LATENT",), + "offset": ("OFFSET",), + "start_at": ("FLOAT", {"default": 0, "min": 0.0, "max": 1.0, "step": 0.01}), + "stop_at": ("FLOAT", {"default": 0, "min": 0.0, "max": 1.0, "step": 0.01}), + "factor": ("FLOAT", {"default": 0.5, "min": 0, "max": 1, "step": 0.01}), + "multiplier": ("FLOAT", {"default": 1, "min": -10, "max": 10, "step": 0.01}), + } + } + + RETURN_TYPES = ("TRANSFORM",) + FUNCTION = "process" + + CATEGORY = "sampling/transforms" + + def process(self, + offset, + latent, + start_at=0, + stop_at=0, + factor=0.5, + multiplier=1): + return ([{ + "params": { + "latent": latent["samples"][0], + "start_at": start_at, + "stop_at": stop_at, + "factor": factor, + "multiplier": multiplier, + "offset": offset, + "offset_status": offset["process_every"] - offset["offset"] - 1, + }, + "function": self.func + }],) + + def func(self, step, x0, total_steps, params): + if (total_steps * params["start_at"] <= step <= total_steps * params["stop_at"] and + (params["offset_status"] == 0 if params["offset"]["mode"] == "process_every" else params["offset_status"] != 0)): + if params["offset_status"] == 0: + params["offset_status"] = params["offset"]["process_every"] - 1 + else: + params["offset_status"] -= 1 + return latent_interpolate_transform(x0, params) diff --git a/nodes/KSamplerNodes/Transforms/MirrorTransform.py b/nodes/KSamplerNodes/Transforms/MirrorTransform.py index c5899ef..0dc6cd9 100644 --- a/nodes/KSamplerNodes/Transforms/MirrorTransform.py +++ b/nodes/KSamplerNodes/Transforms/MirrorTransform.py @@ -33,13 +33,17 @@ class MirrorTransform: "stop_at": stop_at, "mode": mode, "direction": direction, + "offset": offset, + "offset_status": offset["process_every"] - offset["offset"] - 1, }, - "offset": offset, - "offset_status": offset["process_every"] - offset["offset"] - 1, "function": self.func }],) - def func(self, step, x0, total_steps, params) -> list: + def func(self, step, x0, total_steps, params): if (total_steps * params["start_at"] <= step <= total_steps * params["stop_at"] and (params["offset_status"] == 0 if params["offset"]["mode"] == "process_every" else params["offset_status"] != 0)): + if params["offset_status"] == 0: + params["offset_status"] = params["offset"]["process_every"] - 1 + else: + params["offset_status"] -= 1 return mirror_transform(x0, params) diff --git a/nodes/KSamplerNodes/Transforms/MultiplyTransform.py b/nodes/KSamplerNodes/Transforms/MultiplyTransform.py index a3bd7d7..de30f78 100644 --- a/nodes/KSamplerNodes/Transforms/MultiplyTransform.py +++ b/nodes/KSamplerNodes/Transforms/MultiplyTransform.py @@ -31,13 +31,17 @@ class MultiplyTransform: "stop_at": stop_at, "mode": mode, "multiplier": multiplier, + "offset": offset, + "offset_status": offset["process_every"] - offset["offset"] - 1, }, - "offset": offset, - "offset_status": offset["process_every"] - offset["offset"] - 1, "function": self.func }],) - def func(self, step, x0, total_steps, params) -> list: + def func(self, step, x0, total_steps, params): if (total_steps * params["start_at"] <= step <= total_steps * params["stop_at"] and (params["offset_status"] == 0 if params["offset"]["mode"] == "process_every" else params["offset_status"] != 0)): + if params["offset_status"] == 0: + params["offset_status"] = params["offset"]["process_every"] - 1 + else: + params["offset_status"] -= 1 return multiply_transform(x0, params) diff --git a/nodes/KSamplerNodes/Transforms/OneTimeLatentAddTransform.py b/nodes/KSamplerNodes/Transforms/OneTimeLatentAddTransform.py new file mode 100644 index 0000000..ed4895c --- /dev/null +++ b/nodes/KSamplerNodes/Transforms/OneTimeLatentAddTransform.py @@ -0,0 +1,35 @@ +from .transform_functions import latent_add_transform + + +class OneTimeLatentAddTransform: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "latent": ("LATENT",), + "step": ("INT", {"default": 1, "min": 1, "max": 10000}), + "multiplier": ("FLOAT", {"default": 1, "min": -10, "max": 10, "step": 0.01}), + } + } + + RETURN_TYPES = ("TRANSFORM",) + FUNCTION = "process" + + CATEGORY = "sampling/transforms" + + def process(self, + latent, + step=1, + multiplier=1): + return ([{ + "params": { + "latent": latent["samples"][0], + "step": step, + "multiplier": multiplier, + }, + "function": self.func + }],) + + def func(self, step, x0, total_steps, params): + if step == params["step"]: + return latent_add_transform(x0, params) diff --git a/nodes/KSamplerNodes/Transforms/OneTimeLatentIntrpolateTransform.py b/nodes/KSamplerNodes/Transforms/OneTimeLatentIntrpolateTransform.py new file mode 100644 index 0000000..c205431 --- /dev/null +++ b/nodes/KSamplerNodes/Transforms/OneTimeLatentIntrpolateTransform.py @@ -0,0 +1,38 @@ +from .transform_functions import latent_interpolate_transform + + +class OneTimeLatentInterpolateTransform: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "latent": ("LATENT",), + "step": ("INT", {"default": 1, "min": 1, "max": 10000}), + "factor": ("FLOAT", {"default": 0.5, "min": 0, "max": 1, "step": 0.01}), + "multiplier": ("FLOAT", {"default": 1, "min": -10, "max": 10, "step": 0.01}), + } + } + + RETURN_TYPES = ("TRANSFORM",) + FUNCTION = "process" + + CATEGORY = "sampling/transforms" + + def process(self, + latent, + step=1, + factor=0.5, + multiplier=1): + return ([{ + "params": { + "latent": latent["samples"][0], + "step": step, + "factor": factor, + "multiplier": multiplier, + }, + "function": self.func + }],) + + def func(self, step, x0, total_steps, params): + if step == params["step"]: + return latent_interpolate_transform(x0, params) diff --git a/nodes/KSamplerNodes/Transforms/OneTimeMirrorTransform.py b/nodes/KSamplerNodes/Transforms/OneTimeMirrorTransform.py index 867d144..9376078 100644 --- a/nodes/KSamplerNodes/Transforms/OneTimeMirrorTransform.py +++ b/nodes/KSamplerNodes/Transforms/OneTimeMirrorTransform.py @@ -32,6 +32,6 @@ class OneTimeMirrorTransform: "function": self.func }],) - def func(self, step, x0, total_steps, params) -> list: + def func(self, step, x0, total_steps, params): if step == params["step"]: return mirror_transform(x0, params) diff --git a/nodes/KSamplerNodes/Transforms/OneTimeMultiplyTransform.py b/nodes/KSamplerNodes/Transforms/OneTimeMultiplyTransform.py index a3df778..084e428 100644 --- a/nodes/KSamplerNodes/Transforms/OneTimeMultiplyTransform.py +++ b/nodes/KSamplerNodes/Transforms/OneTimeMultiplyTransform.py @@ -30,6 +30,6 @@ class OneTimeMultiplyTransform: "function": self.func }],) - def func(self, step, x0, total_steps, params) -> list: + def func(self, step, x0, total_steps, params): if step == params["step"]: return multiply_transform(x0, params) diff --git a/nodes/KSamplerNodes/Transforms/OneTimeShiftTransform.py b/nodes/KSamplerNodes/Transforms/OneTimeShiftTransform.py index bb68600..9b0ba2d 100644 --- a/nodes/KSamplerNodes/Transforms/OneTimeShiftTransform.py +++ b/nodes/KSamplerNodes/Transforms/OneTimeShiftTransform.py @@ -33,6 +33,6 @@ class OneTimeShiftTransform: "function": self.func }],) - def func(self, step, x0, total_steps, params) -> list: + def func(self, step, x0, total_steps, params): if step == params["step"]: return shift_transform(x0, params) diff --git a/nodes/KSamplerNodes/Transforms/ShiftTransform.py b/nodes/KSamplerNodes/Transforms/ShiftTransform.py index 7dcf2cd..c688567 100644 --- a/nodes/KSamplerNodes/Transforms/ShiftTransform.py +++ b/nodes/KSamplerNodes/Transforms/ShiftTransform.py @@ -34,13 +34,17 @@ class ShiftTransform: "mode": mode, "x_shift": x_shift, "y_shift": y_shift, + "offset": offset, + "offset_status": offset["process_every"] - offset["offset"] - 1, }, - "offset": offset, - "offset_status": offset["process_every"] - offset["offset"] - 1, "function": self.func }],) - def func(self, step, x0, total_steps, params) -> list: + def func(self, step, x0, total_steps, params): if (total_steps * params["start_at"] <= step <= total_steps * params["stop_at"] and (params["offset_status"] == 0 if params["offset"]["mode"] == "process_every" else params["offset_status"] != 0)): + if params["offset_status"] == 0: + params["offset_status"] = params["offset"]["process_every"] - 1 + else: + params["offset_status"] -= 1 return shift_transform(x0, params) diff --git a/nodes/KSamplerNodes/Transforms/__init__.py b/nodes/KSamplerNodes/Transforms/__init__.py index 05cfb9b..85dbcdf 100644 --- a/nodes/KSamplerNodes/Transforms/__init__.py +++ b/nodes/KSamplerNodes/Transforms/__init__.py @@ -2,7 +2,11 @@ from .TransformsCombine import TransformsCombine from .MirrorTransform import MirrorTransform from .ShiftTransform import ShiftTransform from .MultiplyTransform import MultiplyTransform +from .LatentIntrpolateTransform import LatentInterpolateTransform +from .LatentAddTransform import LatentAddTransform from .OneTimeMultiplyTransform import OneTimeMultiplyTransform from .OneTimeShiftTransform import OneTimeShiftTransform from .OneTimeMirrorTransform import OneTimeMirrorTransform +from .OneTimeLatentIntrpolateTransform import OneTimeLatentInterpolateTransform +from .OneTimeLatentAddTransform import OneTimeLatentAddTransform from .TransformOffset import TransformOffset diff --git a/nodes/KSamplerNodes/Transforms/transform_functions.py b/nodes/KSamplerNodes/Transforms/transform_functions.py index 6dd9611..72bec10 100644 --- a/nodes/KSamplerNodes/Transforms/transform_functions.py +++ b/nodes/KSamplerNodes/Transforms/transform_functions.py @@ -1,7 +1,7 @@ import torch +import comfy - -def multiply_transform(x0, params) -> list: +def multiply_transform(x0, params): x = x0 if params["mode"] == "replace": @@ -12,7 +12,7 @@ def multiply_transform(x0, params) -> list: return x -def shift_transform(x0, params) -> list: +def shift_transform(x0, params): x = x0 if params["mode"] == "replace": @@ -29,7 +29,7 @@ def shift_transform(x0, params) -> list: return x -def mirror_transform(x0, params) -> list: +def mirror_transform(x0, params): x = x0 if params["mode"] == "replace": @@ -55,4 +55,32 @@ def mirror_transform(x0, params) -> list: elif params["direction"] == "180 degree rotation": x = (torch.rot90(torch.rot90(x, dims=[1, 2]), dims=[1, 2]) + x) / 2 - return x \ No newline at end of file + return x + + +def latent_interpolate_transform(x0, params): + latent = params["latent"] + + if x0.shape != latent.shape: + latent.permute(0, 3, 1, 2) + latent = comfy.utils.common_upscale(latent, x0.shape[3], x0.shape[2], 'bicubic') + latent.permute(0, 2, 3, 1) + + x = x0 * params["factor"] + latent * (1 - params["factor"]) + x *= params["multiplier"] + + return x + + +def latent_add_transform(x0, params): + latent = params["latent"] + + if x0.shape != latent.shape: + latent.permute(0, 3, 1, 2) + latent = comfy.utils.common_upscale(latent, x0.shape[3], x0.shape[2], 'bicubic') + latent.permute(0, 2, 3, 1) + + x = x0 + latent + x *= params["multiplier"] + + return x diff --git a/nodes/__init__.py b/nodes/__init__.py index 8691b0d..492509f 100644 --- a/nodes/__init__.py +++ b/nodes/__init__.py @@ -5,8 +5,12 @@ from .KSamplerNodes.KSamplerMirroringApart import KSamplerMirroringApart from .KSamplerNodes.Transforms import MirrorTransform from .KSamplerNodes.Transforms import MultiplyTransform from .KSamplerNodes.Transforms import ShiftTransform +from .KSamplerNodes.Transforms import LatentInterpolateTransform +from .KSamplerNodes.Transforms import LatentAddTransform from .KSamplerNodes.Transforms import OneTimeMirrorTransform from .KSamplerNodes.Transforms import OneTimeMultiplyTransform from .KSamplerNodes.Transforms import OneTimeShiftTransform +from .KSamplerNodes.Transforms import OneTimeLatentInterpolateTransform +from .KSamplerNodes.Transforms import OneTimeLatentAddTransform from .KSamplerNodes.Transforms import TransformsCombine from .KSamplerNodes.Transforms import TransformOffset