Latent add and interpolate transform

This commit is contained in:
RomanKuschanow
2024-03-08 14:07:15 +02:00
parent d98c6d4541
commit 68f7fd42bf
14 changed files with 243 additions and 17 deletions
+8
View File
@@ -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",
}
@@ -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)
@@ -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)
@@ -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)
@@ -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)
@@ -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)
@@ -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)
@@ -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)
@@ -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)
@@ -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)
@@ -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)
@@ -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
@@ -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
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
+4
View File
@@ -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