Merge pull request #15 from dmarx/dev

parameter groups and drawing improvements
This commit is contained in:
David Marx
2023-12-10 23:32:34 -08:00
committed by GitHub
4 changed files with 149 additions and 43 deletions
+10
View File
@@ -15,6 +15,7 @@ Related project: https://github.com/FizzleDorf/ComfyUI_FizzNodes
* [AnimateDiff Prompt Superposition - Complex Workflow](#animatediff-prompt-superposition---complex-workflow)
* [Simple Curved Parameter](#simple-curved-parameter)
* [Multi-Prompt Transition With Manually Specified Curves](#multi-prompt-transition-with-manually-specified-curves)
* [Parameter Groups and Curve Drawing Utilities](#parameter-groups-and-curve-drawing-utilities)
* [Nodes](#nodes)
* [Curve Constructors](#curve-constructors)
* [Curve From String](#curve-from-string)
@@ -86,6 +87,15 @@ Workflow output: https://twitter.com/DigThatData/status/1733416414864957484
If you're feeling adventurous, this workflow demonstrates how you would use the curve objects directly to acheive the same thing as a schedule. Each prompt gets its own seaparate curve indicating the weight of the prompt at that time (you probably want the various conditionings weights to sum to 1 when combined. If you need to "fill" missing conditioning weight, try using an empty prompt).
## Parameter Groups and Curve Drawing Utilities
![Parameter Groups and Curve Drawing Utilities](examples/pgroups-and-curve-drawing-utilities.png)
Parameter Groups can be used to carry around multiple curves around. If you want to attach a curve to a parameter group, it will need a label.
We provide utilities for drawing curves. You can draw multiple curves simultaneously by collecting them in a parameter group. The labels you give the curves will also be used as labels for plotting. You can also plot labels you haven't curves (a label was randomly generated when the curve was initialized).
# Nodes
## Curve Constructors
Binary file not shown.

After

Width:  |  Height:  |  Size: 137 KiB

+126 -43
View File
@@ -70,6 +70,25 @@ label: foo"""
return (curve,)
class KfSetCurveLabel:
CATEGORY=CATEGORY
FUNCTION = 'main'
RETURN_TYPES = ("KEYFRAMED_CURVE",)
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"curve": ("KEYFRAMED_CURVE",{"forceInput": True,}),
"label": ("STRING", {
"multiline": False,
"default": "~curve~"})}}
def main(self, curve, label):
curve = deepcopy(curve)
curve.label = label
return (curve,)
class KfEvaluateCurveAtT:
CATEGORY=CATEGORY # TODO: create a "utils" group
FUNCTION = 'main'
@@ -251,21 +270,7 @@ class KfConditioningAddx10:
# return (curve,)
class KfCurveDraw:
CATEGORY = f"{CATEGORY}/experimental"
FUNCTION = "main"
RETURN_TYPES = ("IMAGE",)
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"curve": ("KEYFRAMED_CURVE", {"forceInput": True,}),
"n": ("INT", {"default": 64}),
}
}
def main(self, curve, n):
def plot_curve(curve, n):
"""
"""
@@ -293,20 +298,27 @@ class KfCurveDraw:
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())
plt.figure(figsize=(width, height))
xs = [x for x in list(set(xs)) if (x >= 0)]
xs.sort()
def draw_curve(curve):
ys = [curve[x] for x in xs]
#line = plt.plot(xs, ys, *args, **kargs)
line = plt.plot(xs, ys, label=curve.label)
kfx = curve.keyframes
kfy = [curve[x] for x in kfx]
plt.scatter(kfx, kfy, color=line[0].get_color())
if isinstance(curve, kf.ParameterGroup):
for c in curve.parameters.values():
draw_curve(c)
else:
draw_curve(curve)
plt.legend()
#width, height = 10, 5 #inches
@@ -331,24 +343,43 @@ class KfCurveDraw:
img_tensor = TT.ToTensor()(pil_image)
img_tensor = img_tensor.unsqueeze(0)
img_tensor = img_tensor.permute([0, 2, 3, 1])
return img_tensor
class KfCurveDraw:
CATEGORY = f"{CATEGORY}/experimental"
FUNCTION = "main"
RETURN_TYPES = ("IMAGE",)
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"curve": ("KEYFRAMED_CURVE", {"forceInput": True,}),
"n": ("INT", {"default": 64}),
}
}
def main(self, curve, n):
img_tensor = plot_curve(curve, n)
return (img_tensor,)
#return (plot_array,)
# buffer_io = BytesIO()
# plt.savefig(buffer_io, format='png', bbox_inches='tight')
# plt.close()
class KfPGroupDraw:
CATEGORY = f"{CATEGORY}/experimental"
FUNCTION = "main"
RETURN_TYPES = ("IMAGE",)
# 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,)
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"parameter_group": ("PARAMETER_GROUP", {"forceInput": True,}),
"n": ("INT", {"default": 64}),
}
}
def main(self, parameter_group, n):
img_tensor = plot_curve(parameter_group, n)
return (img_tensor,)
###########################################
# curve arithmetic
@@ -541,15 +572,58 @@ class KfCurveConstant:
### TODO: Working with parameter groups
# Label curve
## inputs: curve, label (text widget)
# add curve(s) to parameter group
## inputs: pgroup, curve
## returns pgroup
## if pgroup not provided, new one created
class KfAddCurveToPGroup:
CATEGORY = CATEGORY
FUNCTION = "main"
#RETURN_TYPES = ("KEYFRAMED_CURVE",)
RETURN_TYPES = ("PARAMETER_GROUP",)
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"curve": ("KEYFRAMED_CURVE",{"forceInput": True,}),
},
"optional": {
"parameter_group": ("PARAMETER_GROUP",{"forceInput": True,}),
},
}
def main(self, curve, parameter_group=None):
curve = deepcopy(curve)
if parameter_group is None:
#parameter_group = kf.ParameterGroup({curve.label:curve})
parameter_group = kf.ParameterGroup([curve])
else:
parameter_group = deepcopy(parameter_group)
parameter_group.parameters[curve.label] = curve
return (parameter_group,)
class KfGetCurveFromPGroup:
CATEGORY = CATEGORY
FUNCTION = "main"
RETURN_TYPES = ("KEYFRAMED_CURVE",)
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"curve_label": ("STRING",{"default": "My Curve",}),
"parameter_group": ("PARAMETER_GROUP",{"forceInput": True,}),
},
}
def main(self, curve_label, parameter_group):
curve = parameter_group.parameters[curve_label]
return (deepcopy(curve),)
# get curve from parameter group
## inputs: pgroup, label
## returns curve
@@ -586,13 +660,19 @@ class KfCurveConstant:
NODE_CLASS_MAPPINGS = {
"KfCurveFromString": KfCurveFromString,
"KfCurveFromYAML": KfCurveFromYAML,
"KfSetCurveLabel": KfSetCurveLabel,
"KfEvaluateCurveAtT": KfEvaluateCurveAtT,
"KfApplyCurveToCond": KfApplyCurveToCond,
"KfConditioningAdd": KfConditioningAdd,
#######################################
"KfAddCurveToPGroup": KfAddCurveToPGroup,
"KfGetCurveFromPGroup": KfGetCurveFromPGroup,
#######################################
#"KfCurveToAcnLatentKeyframe": KfCurveToAcnLatentKeyframe,
#######################################
#"KfCurveInverse": KfCurveInverse,
"KfCurveDraw": KfCurveDraw,
"KfPGroupDraw": KfPGroupDraw,
"KfCurvesAdd": KfCurvesAdd,
"KfCurvesSubtract": KfCurvesSubtract,
"KfCurvesMultiply": KfCurvesMultiply,
@@ -618,4 +698,7 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"KfCurvesMultiply": "Curve_1 * Curve_2",
"KfCurvesDivide": "Curve_1 / Curve_2",
"KfCurveConstant": "Constant-Valued Curve",
"KfSetCurveLabel": "Set Curve Label",
"KfAddCurveToPGroup": "Add Curve To Parameter Group",
"KfGetCurveFromPGroup": "Get Curve From Parameter Group",
}
+13
View File
@@ -3,6 +3,8 @@ import torch
import numpy as np
from PIL.Image import Image
import keyframed as kf
logging.basicConfig(level=logging.INFO,
format='%(asctime)s - %(name)s - %(levelname)s - %(message)s')
logger = logging.getLogger(__name__)
@@ -36,6 +38,10 @@ def _inspect(item, depth=0):
if isinstance(item, Image):
logger.info(f"{pad}item.mode: {item.mode}")
if isinstance(item, kf.ParameterGroup):
logger.info(f"{pad}item.parameters: {item.parameters}")
# to do: be fancy and change to a match statement
@@ -147,6 +153,13 @@ class KfDebug_Curve(KfDebug_Passthrough):
RETURN_TYPES = ("KEYFRAMED_CURVE",)
class KfDebug_PGroup(KfDebug_Passthrough):
RETURN_TYPES = ("PARAMETER_GROUP",)
class KfDebug_Schedule(KfDebug_Passthrough):
RETURN_TYPES = ("SCHEDULE",)
# ###########################