diff --git a/.gitignore b/.gitignore index 68bc17f..e8e5aa5 100644 --- a/.gitignore +++ b/.gitignore @@ -1,3 +1,7 @@ +_venv/ + +############### + # Byte-compiled / optimized / DLL files __pycache__/ *.py[cod] diff --git a/README.md b/README.md index 672ab67..c04c515 100644 --- a/README.md +++ b/README.md @@ -1,2 +1,13 @@ # ComfyUI-Keyframed -[Work In Progress] ComfyUI nodes to facilitate value keyframing + +🚧 Work In Progress 🚧 - ComfyUI nodes to facilitate value keyframing by providing an interface for using [keyframed](https://github.com/dmarx/keyframed) in ComfyUI workflows. + + +...Open question: if I make this, what will differentiate it from https://github.com/FizzleDorf/ComfyUI_FizzNodes ? + +* easier curve composition +* easier to change interpolators/easing functions + +# Philosophy + +Curves, interpolators, and keyframes are objects that can be passed around, plugged and unplugged, and interchanged. \ No newline at end of file diff --git a/__init__.py b/__init__.py new file mode 100644 index 0000000..af13844 --- /dev/null +++ b/__init__.py @@ -0,0 +1,39 @@ +import os +import subprocess +import importlib.util +import sys + +# scavenged install sequence from https://github.com/FizzleDorf/ComfyUI_FizzNodes/blob/main/__init__.py +def is_installed(package, package_overwrite=None): + try: + spec = importlib.util.find_spec(package) + except ModuleNotFoundError: + pass + + package = package_overwrite or package + + python = sys.executable + if spec is None: + print(f"Installing {package}...") + command = f'"{python}" -m pip install {package}' + + result = subprocess.run(command, stdout=subprocess.PIPE, stderr=subprocess.PIPE, shell=True, env=os.environ) + + if result.returncode != 0: + print(f"Couldn't install\nCommand: {command}\nError code: {result.returncode}") + +# to do: read from requirements.txt +is_installed("keyframed") + +from .nodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS + +print(os.environ.get('COMFYUI_DEBUG_MODE')) + +from .debug import NODE_CLASS_MAPPINGS as ncm0, NODE_DISPLAY_NAME_MAPPINGS as ndnm0 + +# there's probably a cleaner, more-dummy-proof way to do this. +# feels like an accident waiting to happen. low risk though. +NODE_CLASS_MAPPINGS.update(ncm0) +NODE_DISPLAY_NAME_MAPPINGS.update(ndnm0) + +__all__ =["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"] diff --git a/debug.py b/debug.py new file mode 100644 index 0000000..8082700 --- /dev/null +++ b/debug.py @@ -0,0 +1,169 @@ +import logging +import torch +import numpy as np +from PIL.Image import Image + +logging.basicConfig(level=logging.INFO, + format='%(asctime)s - %(name)s - %(levelname)s - %(message)s') +logger = logging.getLogger(__name__) + +CATEGORY="keyframed/debug" + +# maybe use icecream here instead? +# https://github.com/gruns/icecream +def _inspect(item, depth=0): + pad="\t"*depth + if depth > 0: + pad +="- " + # NB: Linter says using f-strings in log statements can hinder performance + logger.info(f"{pad}type: {type(item)}") + log_item=True + + if hasattr(item, "dtype"): + logger.info(f"{pad}item.dtype: {item.dtype}") + log_item=False + + # maybe a bit overengineered. whatever. + if hasattr(item, "shape"): + logger.info(f"{pad}item.shape: {item.shape}") + log_item=False + elif hasattr(item, "size"): + try: + logger.info(f"{pad}item.shape: {item.size()}") + except TypeError: + logger.info(f"{pad}item.shape: {item.size}") + log_item=False + + if isinstance(item, Image): + logger.info(f"{pad}item.mode: {item.mode}") + + + # to do: be fancy and change to a match statement + #if isinstance(item, dict): + if hasattr(item, 'keys'): + logger.info(f"{pad}item.keys(): {item.keys()}") + if hasattr(item, 'items'): + for k,v in item.items(): + logger.info(f"{pad}key: {k}") + #logger.info(f"{pad}value: {_inspect(v, depth=depth+1)}") + _inspect(v, depth=depth+1) + log_item=False + + if isinstance(item, list) or isinstance(item, tuple): + logger.info(f"{pad}len(item): {len(item)}") + for entry in item: + _inspect(entry, depth=depth+1) + log_item=False + + if log_item: + logger.info(f"{pad}item: {item}") + + +class KfDebug_Passthrough: + CATEGORY=CATEGORY + FUNCTION = 'main' + OUTPUT_NODE=True + + _FORCED_INPUT = {"label": ("STRING", { + "multiline": True, #True if you want the field to look like the one on the ClipTextEncode node + "default": "debugging passthrough"})} + + @classmethod + def INPUT_TYPES(cls): + outv = { + "required": { + "item": (cls.RETURN_TYPES[0],{"forceInput": True,}), + }, + } + outv["required"].update(cls._FORCED_INPUT) + return outv + + def main(self, item, label): + #if label: + logger.info(f"label: {label}") + _inspect(item) + return (item,) # pretty sure it's gotta be a tuple? + + +# class KfDebug_DummyOutput(KfDebug_Passthrough): +# OUTPUT_NODE=True +# _FORCED_INPUT = {"label": ("STRING", { +# "multiline": True, #True if you want the field to look like the one on the ClipTextEncode node +# "default": "dummy output"})} + +########################### + +### Built-in Types + + +# there should be a way to create a type-agnostic passthrough node + +class KfDebug_Clip(KfDebug_Passthrough): + RETURN_TYPES = ("CLIP",) + + +class KfDebug_Cond(KfDebug_Passthrough): + RETURN_TYPES = ("CONDITIONING",) + + +class KfDebug_Float(KfDebug_Passthrough): + RETURN_TYPES = ("FLOAT",) + + +class KfDebug_Image(KfDebug_Passthrough): + RETURN_TYPES = ("IMAGE",) + + +class KfDebug_Int(KfDebug_Passthrough): + RETURN_TYPES = ("INT",) + + +class KfDebug_Latent(KfDebug_Passthrough): + RETURN_TYPES = ("LATENT",) + + +class KfDebug_Model(KfDebug_Passthrough): + RETURN_TYPES = ("MODEL",) + + +class KfDebug_String(KfDebug_Passthrough): + RETURN_TYPES = ("STRING",) + + +class KfDebug_Vae(KfDebug_Passthrough): + RETURN_TYPES = ("VAE",) + + +############################################## + +### Custom Node Types + + +class KfDebug_Segs(KfDebug_Passthrough): + RETURN_TYPES = ("SEGS",) + + +class KfDebug_Curve(KfDebug_Passthrough): + RETURN_TYPES = ("KEYFRAMED_CURVE",) + + +# ########################### + + +NODE_CLASS_MAPPINGS = { + #"KfDebug_Passthrough": KfDebug_Passthrough, + "KfDebug_Clip": KfDebug_Clip, + "KfDebug_Cond": KfDebug_Cond, + "KfDebug_Curve": KfDebug_Curve, + "KfDebug_Float": KfDebug_Float, + "KfDebug_Image": KfDebug_Image, + "KfDebug_Int": KfDebug_Int, + "KfDebug_Latent": KfDebug_Latent, + "KfDebug_Model": KfDebug_Model, + "KfDebug_Segs": KfDebug_Segs, + "KfDebug_String": KfDebug_String, + "KfDebug_Vae": KfDebug_Vae, +} + +# A dictionary that contains the friendly/humanly readable titles for the nodes +NODE_DISPLAY_NAME_MAPPINGS = {k:k for k in NODE_CLASS_MAPPINGS} \ No newline at end of file diff --git a/example.py b/example.py new file mode 100644 index 0000000..0b61415 --- /dev/null +++ b/example.py @@ -0,0 +1,102 @@ +class Example: + """ + A example node + + Class methods + ------------- + INPUT_TYPES (dict): + Tell the main program input parameters of nodes. + + Attributes + ---------- + RETURN_TYPES (`tuple`): + The type of each element in the output tulple. + RETURN_NAMES (`tuple`): + Optional: The name of each output in the output tulple. + FUNCTION (`str`): + The name of the entry-point method. For example, if `FUNCTION = "execute"` then it will run Example().execute() + OUTPUT_NODE ([`bool`]): + If this node is an output node that outputs a result/image from the graph. The SaveImage node is an example. + The backend iterates on these output nodes and tries to execute all their parents if their parent graph is properly connected. + Assumed to be False if not present. + CATEGORY (`str`): + The category the node should appear in the UI. + execute(s) -> tuple || None: + The entry point method. The name of this method must be the same as the value of property `FUNCTION`. + For example, if `FUNCTION = "execute"` then this method's name must be `execute`, if `FUNCTION = "foo"` then it must be `foo`. + """ + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + """ + Return a dictionary which contains config for all input fields. + Some types (string): "MODEL", "VAE", "CLIP", "CONDITIONING", "LATENT", "IMAGE", "INT", "STRING", "FLOAT". + Input types "INT", "STRING" or "FLOAT" are special values for fields on the node. + The type can be a list for selection. + + Returns: `dict`: + - Key input_fields_group (`string`): Can be either required, hidden or optional. A node class must have property `required` + - Value input_fields (`dict`): Contains input fields config: + * Key field_name (`string`): Name of a entry-point method's argument + * Value field_config (`tuple`): + + First value is a string indicate the type of field or a list for selection. + + Secound value is a config for type "INT", "STRING" or "FLOAT". + """ + return { + "required": { + "image": ("IMAGE",), + "int_field": ("INT", { + "default": 0, + "min": 0, #Minimum value + "max": 4096, #Maximum value + "step": 64, #Slider's step + "display": "number" # Cosmetic only: display as "number" or "slider" + }), + "float_field": ("FLOAT", { + "default": 1.0, + "min": 0.0, + "max": 10.0, + "step": 0.01, + "round": 0.001, #The value represeting the precision to round to, will be set to the step value by default. Can be set to False to disable rounding. + "display": "number"}), + "print_to_screen": (["enable", "disable"],), + "string_field": ("STRING", { + "multiline": False, #True if you want the field to look like the one on the ClipTextEncode node + "default": "Hello World!" + }), + }, + } + + RETURN_TYPES = ("IMAGE",) + #RETURN_NAMES = ("image_output_name",) + + FUNCTION = "test" + + #OUTPUT_NODE = False + + CATEGORY = "Example" + + def test(self, image, string_field, int_field, float_field, print_to_screen): + if print_to_screen == "enable": + print(f"""Your input contains: + string_field aka input text: {string_field} + int_field: {int_field} + float_field: {float_field} + """) + #do some processing on the image, in this example I just invert it + image = 1.0 - image + return (image,) + + +# A dictionary that contains all nodes you want to export with their names +# NOTE: names should be globally unique +NODE_CLASS_MAPPINGS = { + "Example": Example +} + +# A dictionary that contains the friendly/humanly readable titles for the nodes +NODE_DISPLAY_NAME_MAPPINGS = { + "Example": "Example Node" +} \ No newline at end of file diff --git a/nodes.py b/nodes.py new file mode 100644 index 0000000..42bc7fb --- /dev/null +++ b/nodes.py @@ -0,0 +1,383 @@ +import keyframed as kf +from keyframed.dsl import curve_from_cn_string +import logging +import torch +from copy import deepcopy +#import warnings +import matplotlib.pyplot as plt +import numpy as np +import io +from PIL import Image + + +logging.basicConfig(level=logging.INFO, + format='%(asctime)s - %(name)s - %(levelname)s - %(message)s') +logger = logging.getLogger(__name__) + + +CATEGORY = "keyframed" + + +class KfCurveFromString: + CATEGORY=CATEGORY + FUNCTION = 'main' + RETURN_TYPES = ("KEYFRAMED_CURVE",) + + @classmethod + def INPUT_TYPES(s): + return { + "required": {"chigozie_string": ("STRING", { + "multiline": True, #True if you want the field to look like the one on the ClipTextEncode node + "default": "0:(1)" + }), + }, + } + + def main(self, chigozie_string): + curve = curve_from_cn_string(chigozie_string) + return (curve,) + + +class KfCurveFromYAML: + CATEGORY=CATEGORY + FUNCTION = 'main' + RETURN_TYPES = ("KEYFRAMED_CURVE",) + + @classmethod + def INPUT_TYPES(s): + return { + "required": {"yaml": ("STRING", { + "multiline": True, #True if you want the field to look like the one on the ClipTextEncode node + # TO DO: replace this with kf.serializaiton.to_dict() # or whatever + "default": """curve: +- - 0 + - 0 + - linear +- - 1 + - 1 +loop: false +bounce: false +duration: 1 +label: foo""" + }), + }, + } + + def main(self, yaml): + curve = kf.serialization.from_yaml(yaml) + return (curve,) + + +class KfEvaluateCurveAtT: + CATEGORY=CATEGORY + FUNCTION = 'main' + RETURN_TYPES = ("FLOAT","INT") + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "curve": ("KEYFRAMED_CURVE",{"forceInput": True,}), + "t": ("INT",{"default":0}) + }, + } + + def main(self, curve, t): + return curve[t], int(curve[t]) + + +# class KfCurveToAcnLatentKeyframe: +# CATEGORY=CATEGORY +# FUNCTION = 'main' +# RETURN_NAMES = ("LATENT_KF", ) +# RETURN_TYPES = ("LATENT_KEYFRAME",) +# """Compatibility with Kosinkadink "Advanced Controlnet" AnimateDiff""" +# @classmethod +# def INPUT_TYPES(s): +# return { +# "required": { +# "curve": ("KEYFRAMED_CURVE",{"forceInput": True,}), +# }, +# } +# def main(self, curve): +# warnings.warn("KfCurveToAcnLatentKeyframe not implemented") +# return (curve,) + + +class KfApplyCurveToCond: + CATEGORY=CATEGORY + FUNCTION = 'main' + #RETURN_TYPES = ("CONDITIONING","LATENT_KEYFRAME",) + RETURN_TYPES = ("CONDITIONING",) + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "curve": ("KEYFRAMED_CURVE", {"forceInput": True,}), + "cond": ("CONDITIONING", {"forceInput": True,}), + }, + "optional":{ + "latents": ("LATENT", {}), + "start_t": ("INT", {"default":0, }), + "n": ("INT", {}), + }, + } + def main(self, curve, cond, latents=None, start_t=0, n=0): + #logger.info(f"latents: {latents}") + logger.info(f"type(latents): {type(latents)}") # Latent is a dict that (presently) has one key, `samples` + #device = 'cpu' # probably should be handling this some other way + #if latents is not None: + if isinstance(latents, dict): + if 'samples' in latents: + n = latents['samples'].shape[0] # batch dimension + #device = latents['samples'].device + #weights = [curve[start_t+i] for i in range(n)] + #weights = torch.tensor(weights, device=device) + cond_out = [] + for c_tensor, c_dict in cond: + #weights.to(c_tensor.device) + m=c_tensor.shape[0] + if c_tensor.shape[0] == 1: + c_tensor = c_tensor.repeat(n, 1, 1) # batch, n_tokens, embeding_dim + m=n + weights = [curve[start_t+i] for i in range(m)] + weights = torch.tensor(weights, device=c_tensor.device) + #logger.info(f"c_tensor.shape:{c_tensor.shape}") + #logger.info(f"weights.shape:{weights.shape}") + #logger.info(f"weights.shape:{weights.view(n,1,1).shape}") + #c_tensor.mul_(weights) + #c_tensor.mul_(weights.view(m,1,1)) # I think these in-place/mutating operations are messing things up + c_tensor = c_tensor * weights.view(m,1,1) + #c_tensor = c_tensor + if "pooled_output" in c_dict: + c_dict = deepcopy(c_dict) # hate this. + pooled = c_dict['pooled_output'] + if pooled.shape[0] == 1: + pooled = pooled.repeat(m, 1) # batch, embeding_dim + #pooled.mul_(weights) + c_dict['pooled_output'] = pooled * weights.view(m,1) + cond_out.append((c_tensor, c_dict)) + return (cond_out,) + + #outv = torch.ones_like(latents) * torch.tensor(weights, device=latents.device) + #return (cond, outv) + + +# TODO: Add Conds +#class ConditioningAverage: +class KfConditioningAdd: + CATEGORY = CATEGORY + FUNCTION = "main" + RETURN_TYPES = ("CONDITIONING",) + + @classmethod + def INPUT_TYPES(s): + return {"required": {"conditioning_1": ("CONDITIONING", ), + "conditioning_2": ("CONDITIONING", ), + }} + + def main(self, conditioning_1, conditioning_2): + assert len(conditioning_1) == len(conditioning_2) + + outv = [] + for i, ((c1_tensor, c1_dict), (c2_tensor, c2_dict) ) in enumerate(zip(conditioning_1, conditioning_2)): + c1_tensor += c2_tensor + if ('pooled_output' in c1_dict) and ('pooled_output' in c2_dict): + c1_dict['pooled_output'] += c2_dict['pooled_output'] + outv.append((c1_tensor, c1_dict)) + return (outv, ) + + +# class KfCurveInverse: +# CATEGORY = CATEGORY +# FUNCTION = "main" +# RETURN_TYPES = ("KEYFRAMED_CURVE",) + +# @classmethod +# def INPUT_TYPES(s): +# return { +# "required": { +# "curve": ("KEYFRAMED_CURVE",{"forceInput": True,}), +# }, +# "hidden": { +# "a": ("FLOAT", {"default": 0.0001}), +# }, +# } + +# def main(self, curve, a=0.0001): +# curve = curve + a +# curve = 1/curve +# return (curve,) + + +class KfCurveDraw: + CATEGORY = f"{CATEGORY}/experimental" + FUNCTION = "main" + RETURN_TYPES = ("IMAGE",) + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "curve": ("KEYFRAMED_CURVE",) + } + } + + def main(self, curve): + """ + + """ + # Create a figure and axes object + fig, ax = plt.subplots() + + # Build the plot using the provided function + #build_plot(ax) + #curve.plot(ax=ax) + curve.plot() + 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') + 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) + + # 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,) + +########################################### + +# curve arithmetic + +# TODO: Add Curves (to compute normalization) +class KfCurvesAdd: + CATEGORY = CATEGORY + FUNCTION = "main" + RETURN_TYPES = ("KEYFRAMED_CURVE",) + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "curve_1": ("KEYFRAMED_CURVE",{"forceInput": True,}), + "curve_2": ("KEYFRAMED_CURVE",{"forceInput": True,}), + }, + } + + def main(self, curve_1, curve_2): + return (curve_1 + curve_2, ) + + +class KfCurvesSubtract: + CATEGORY = CATEGORY + FUNCTION = "main" + RETURN_TYPES = ("KEYFRAMED_CURVE",) + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "curve_1": ("KEYFRAMED_CURVE",{"forceInput": True,}), + "curve_2": ("KEYFRAMED_CURVE",{"forceInput": True,}), + }, + } + + def main(self, curve_1, curve_2): + return (curve_1 - curve_2, ) + + +class KfCurvesMultiply: + CATEGORY = CATEGORY + FUNCTION = "main" + RETURN_TYPES = ("KEYFRAMED_CURVE",) + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "curve_1": ("KEYFRAMED_CURVE",{"forceInput": True,}), + "curve_2": ("KEYFRAMED_CURVE",{"forceInput": True,}), + }, + } + + def main(self, curve_1, curve_2): + return (curve_1 * curve_2, ) + + +class KfCurvesDivide: + CATEGORY = CATEGORY + FUNCTION = "main" + RETURN_TYPES = ("KEYFRAMED_CURVE",) + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "curve_1": ("KEYFRAMED_CURVE",{"forceInput": True,}), + "curve_2": ("KEYFRAMED_CURVE",{"forceInput": True,}), + }, + } + + def main(self, curve_1, curve_2): + return (curve_1 / curve_2, ) + + +################################################################## + +#### 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 + +# get curve from parameter group +## inputs: pgroup, label +## returns curve + +# extract a time slice from the parameter group + + +################################################################## + +NODE_CLASS_MAPPINGS = { + "KfCurveFromString": KfCurveFromString, + "KfCurveFromYAML": KfCurveFromYAML, + "KfEvaluateCurveAtT": KfEvaluateCurveAtT, + "KfApplyCurveToCond": KfApplyCurveToCond, + "KfConditioningAdd": KfConditioningAdd, + #"KfCurveInverse": KfCurveInverse, + "KfCurveDraw": KfCurveDraw, + "KfCurvesAdd": KfCurvesAdd, + "KfCurvesSubtract": KfCurvesSubtract, + "KfCurvesMultiply": KfCurvesMultiply, + "KfCurvesDivide": KfCurvesDivide, + #"KfCurveToAcnLatentKeyframe": KfCurveToAcnLatentKeyframe, +} + +# A dictionary that contains the friendly/humanly readable titles for the nodes +NODE_DISPLAY_NAME_MAPPINGS = { + "KfCurveFromString": "Curve From String", + "KfCurveFromYAML": "Curve From YAML", + "KfEvaluateCurveAtT": "Evaluate Curve At T", + "KfApplyCurveToCond": "Apply Curve to Conditioning", + "KfConditioningAdd": "Add Conditions", + #"KfCurveToAcnLatentKeyframe": "Curve to ACN Latent Keyframe", + "KfCurvesAdd": "Curve_1 + Curve_2", + "KfCurvesSubtract": "Curve_1 - Curve_2", + "KfCurvesMultiply": "Curve_1 * Curve_2", + "KfCurvesDivide": "Curve_1 / Curve_2", +} \ No newline at end of file diff --git a/requirements.txt b/requirements.txt new file mode 100644 index 0000000..b020240 --- /dev/null +++ b/requirements.txt @@ -0,0 +1 @@ +keyframed \ No newline at end of file