Merge pull request #15 from dmarx/dev
parameter groups and drawing improvements
This commit is contained in:
@@ -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 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
@@ -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",
|
||||
}
|
||||
@@ -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",)
|
||||
|
||||
|
||||
# ###########################
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user