Latent add and interpolate transform
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user