Merge pull request #14 from dmarx/dev

convenience nodes
This commit is contained in:
David Marx
2023-12-10 21:47:00 -08:00
committed by GitHub
5 changed files with 194 additions and 8 deletions
+9 -1
View File
@@ -132,15 +132,23 @@ Generates a batch of `n` conditionings multiplying each conditioning by th value
### Add Conditions
![Apply Curve To Conditioning](assets/node_add-conditions.png)
![Add Conditions](assets/node_add-conditions.png)
![Add Conditions (x10)](node_add-conditions-x10.png)
If you're using the `x10` node, at least `curve_0` must be non-empty. The other cond positions are all optionally populated.
### Curve Arithmetic Operators
![Curve Arithmetic](assets/nodes_curve-arithmetic.png)
Arithmetic is performed at the union of keyframes of the provided curves.
NB: the division operator is unreliable at the time of this writing (2023-12-09).
![Curve Arithmetic - batch pooling](assets/node_curve-arithmetic-x10.png)
If you have lots of curve objects to multiply together or add together, here are some convenience nodes.
## Scheduling
Binary file not shown.

After

Width:  |  Height:  |  Size: 42 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 40 KiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 1.6 MiB

After

Width:  |  Height:  |  Size: 1.4 MiB

+185 -7
View File
@@ -9,6 +9,8 @@ import numpy as np
import io
from PIL import Image
import torchvision.transforms as TT
logging.basicConfig(level=logging.INFO,
format='%(asctime)s - %(name)s - %(levelname)s - %(message)s')
@@ -193,6 +195,40 @@ class KfConditioningAdd:
return (outv, )
class KfConditioningAddx10:
CATEGORY = CATEGORY
FUNCTION = "main"
RETURN_TYPES = ("CONDITIONING",)
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"cond_0": ("CONDITIONING",{"forceInput": True,}),
},
"optional": {
"cond_1": ("CONDITIONING",{"forceInput": True, "default": 0}),
"cond_2": ("CONDITIONING",{"forceInput": True, "default": 0}),
"cond_3": ("CONDITIONING",{"forceInput": True, "default": 0}),
"cond_4": ("CONDITIONING",{"forceInput": True, "default": 0}),
"cond_5": ("CONDITIONING",{"forceInput": True, "default": 0}),
"cond_6": ("CONDITIONING",{"forceInput": True, "default": 0}),
"cond_7": ("CONDITIONING",{"forceInput": True, "default": 0}),
"cond_8": ("CONDITIONING",{"forceInput": True, "default": 0}),
"cond_9": ("CONDITIONING",{"forceInput": True, "default": 0}),
},
}
def main(self, cond_0, **kwargs):
((cond_t_out, cond_d_out),) = deepcopy(cond_0)
for ((cond_t,cond_d),) in kwargs.values():
cond_t, cond_d = deepcopy(cond_t), deepcopy(cond_d)
cond_t_out = cond_t_out + cond_t
cond_d_out["pooled_output"] = cond_d_out["pooled_output"] + cond_d["pooled_output"]
return [((cond_t_out, cond_d_out),)] #((cond_t_out, cond_d_out),)
# class KfCurveInverse:
# CATEGORY = CATEGORY
# FUNCTION = "main"
@@ -224,11 +260,12 @@ class KfCurveDraw:
def INPUT_TYPES(cls):
return {
"required": {
"curve": ("KEYFRAMED_CURVE",)
"curve": ("KEYFRAMED_CURVE", {"forceInput": True,}),
"n": ("INT", {"default": 64}),
}
}
def main(self, curve):
def main(self, curve, n):
"""
"""
@@ -238,25 +275,79 @@ class KfCurveDraw:
# Build the plot using the provided function
#build_plot(ax)
#curve.plot(ax=ax)
curve.plot()
width, height = 10, 5 #inches
#curve.plot(n=n)
eps:float=1e-9
# value to be subtracted from keyframe to produce additional points for plotting.
# Plotting these additional values is important for e.g. visualizing step function behavior.
m=3
if n < m:
n = self.duration + 1
n = max(m, n)
xs_base = list(range(int(n))) + list(curve.keyframes)
logger.debug(f"xs_base:{xs_base}")
xs = set()
for x in xs_base:
xs.add(x)
xs.add(x-eps)
xs = [x for x in list(set(xs)) if (x >= 0)]
xs.sort()
ys = [curve[x] for x in xs]
width, height = 12,8 #inches
plt.figure(figsize=(width, height))
#line = plt.plot(xs, ys, *args, **kargs)
line = plt.plot(xs, ys)
kfx = curve.keyframes
kfy = [curve[x] for x in kfx]
plt.scatter(kfx, kfy, color=line[0].get_color())
#width, height = 10, 5 #inches
#plt.figure(figsize=(width, height))
# Save the plot to a BytesIO object
buf = io.BytesIO()
plt.savefig(buf, format='png', bbox_inches='tight')
plt.close() # no idea if this makes a difference
buf.seek(0)
# Read the image into a numpy array, converting it to RGB mode
pil_image = Image.open(buf).convert('RGB')
plot_array = np.array(pil_image) #.astype(np.uint8)
#plot_array = np.array(pil_image) #.astype(np.uint8)
# Convert the array to the desired shape [batch, channels, width, height]
#plot_array = np.transpose(plot_array, (2, 0, 1)) # Reorder to [channels, width, height]
#plot_array = np.expand_dims(plot_array, axis=0) # Add the batch dimension
#plot_array = torch.tensor(plot_array) #.float()
plot_array = torch.from_numpy(plot_array)
return (plot_array,)
#plot_array = torch.from_numpy(plot_array)
img_tensor = TT.ToTensor()(pil_image)
img_tensor = img_tensor.unsqueeze(0)
img_tensor = img_tensor.permute([0, 2, 3, 1])
return (img_tensor,)
#return (plot_array,)
# buffer_io = BytesIO()
# plt.savefig(buffer_io, format='png', bbox_inches='tight')
# plt.close()
# buffer_io.seek(0)
# img = Image.open(buffer_io)
# img_tensor = TT.ToTensor()(img)
# img_tensor = img_tensor.unsqueeze(0)
# img_tensor = img_tensor.permute([0, 2, 3, 1])
# return (img_tensor,)
###########################################
@@ -283,6 +374,48 @@ class KfCurvesAdd:
return (curve_1 + curve_2, )
class KfCurvesAddx10:
CATEGORY = CATEGORY
FUNCTION = "main"
RETURN_TYPES = ("KEYFRAMED_CURVE",)
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"curve_0": ("KEYFRAMED_CURVE",{"forceInput": True,}),
},
"optional": {
"curve_1": ("KEYFRAMED_CURVE",{"forceInput": True, "default": 0}),
"curve_2": ("KEYFRAMED_CURVE",{"forceInput": True, "default": 0}),
"curve_3": ("KEYFRAMED_CURVE",{"forceInput": True, "default": 0}),
"curve_4": ("KEYFRAMED_CURVE",{"forceInput": True, "default": 0}),
"curve_5": ("KEYFRAMED_CURVE",{"forceInput": True, "default": 0}),
"curve_6": ("KEYFRAMED_CURVE",{"forceInput": True, "default": 0}),
"curve_7": ("KEYFRAMED_CURVE",{"forceInput": True, "default": 0}),
"curve_8": ("KEYFRAMED_CURVE",{"forceInput": True, "default": 0}),
"curve_9": ("KEYFRAMED_CURVE",{"forceInput": True, "default": 0}),
},
}
def main(self, curve_0, curve_1, curve_2, curve_3, curve_4, curve_5, curve_6, curve_7, curve_8, curve_9):
#curve_1 = deepcopy(curve_1)
#curve_2 = deepcopy(curve_2)
#return (curve_1 + curve_2, )
curve_out = (
curve_0 +
curve_1 +
curve_2 +
curve_3 +
curve_4 +
curve_5 +
curve_6 +
curve_7 +
curve_8 +
curve_9)
return (curve_out,)
class KfCurvesSubtract:
CATEGORY = CATEGORY
FUNCTION = "main"
@@ -323,6 +456,48 @@ class KfCurvesMultiply:
return (curve_1 * curve_2, )
class KfCurvesMultiplyx10:
CATEGORY = CATEGORY
FUNCTION = "main"
RETURN_TYPES = ("KEYFRAMED_CURVE",)
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"curve_0": ("KEYFRAMED_CURVE",{"forceInput": True,}),
},
"optional": {
"curve_1": ("KEYFRAMED_CURVE",{"forceInput": True, "default": 0}),
"curve_2": ("KEYFRAMED_CURVE",{"forceInput": True, "default": 0}),
"curve_3": ("KEYFRAMED_CURVE",{"forceInput": True, "default": 0}),
"curve_4": ("KEYFRAMED_CURVE",{"forceInput": True, "default": 0}),
"curve_5": ("KEYFRAMED_CURVE",{"forceInput": True, "default": 0}),
"curve_6": ("KEYFRAMED_CURVE",{"forceInput": True, "default": 0}),
"curve_7": ("KEYFRAMED_CURVE",{"forceInput": True, "default": 0}),
"curve_8": ("KEYFRAMED_CURVE",{"forceInput": True, "default": 0}),
"curve_9": ("KEYFRAMED_CURVE",{"forceInput": True, "default": 0}),
},
}
def main(self, curve_0, curve_1, curve_2, curve_3, curve_4, curve_5, curve_6, curve_7, curve_8, curve_9):
#curve_1 = deepcopy(curve_1)
#curve_2 = deepcopy(curve_2)
#return (curve_1 + curve_2, )
curve_out = (
curve_0 *
curve_1 *
curve_2 *
curve_3 *
curve_4 *
curve_5 *
curve_6 *
curve_7 *
curve_8 *
curve_9)
return (curve_out,)
## This seems to not be working properly. I think the issue is upstream in Keyframed
# TODO: set as experimental?
class KfCurvesDivide:
@@ -424,6 +599,9 @@ NODE_CLASS_MAPPINGS = {
"KfCurvesDivide": KfCurvesDivide,
"KfCurveConstant": KfCurveConstant,
#########################
"KfConditioningAddx10":KfConditioningAddx10,
"KfCurvesAddx10":KfCurvesAddx10,
"KfCurvesMultiplyx10":KfCurvesMultiplyx10,
}