From 2b52fdaa7b8754859c27cf7b6e7c46a3f912e4ef Mon Sep 17 00:00:00 2001 From: FizzleDorf <1fizzledorf@gmail.com> Date: Mon, 23 Oct 2023 01:44:47 -0400 Subject: [PATCH] fixed token size issue and WIP Gligen schedule(disabled) --- BatchFuncs.py | 37 ++++++++-- ScheduleFuncs.py | 10 ++- ScheduledNodes.py | 177 ++++++++++++++++++---------------------------- ValueFuncs.py | 108 ++++++++++++++++++++++++++++ __init__.py | 6 +- 5 files changed, 220 insertions(+), 118 deletions(-) create mode 100644 ValueFuncs.py diff --git a/BatchFuncs.py b/BatchFuncs.py index 1f0c2ae..6649f16 100644 --- a/BatchFuncs.py +++ b/BatchFuncs.py @@ -7,7 +7,7 @@ import numpy as np import pandas as pd import re -from .ScheduleFuncs import addWeighted, check_is_number, parse_weight, prepare_prompt, SDXLencode +from .ScheduleFuncs import addWeighted, check_is_number, parse_weight, prepare_prompt, SDXLencode, reverseConcatenation def prepare_batch_prompt(prompt_series, max_frames, frame_idx, prompt_weight_1=0, prompt_weight_2=0, prompt_weight_3=0, prompt_weight_4=0): # calculate expressions from the text input and return a string @@ -123,15 +123,16 @@ def interpolate_prompt_series(animation_prompts, max_frames, pre_text, app_text, # Evaluate the current and next prompt's expressions for i in range(len(cur_prompt_series)): + print(len(cur_prompt_series)) cur_prompt_series[i] = prepare_batch_prompt(cur_prompt_series[i], max_frames, i, prompt_weight_1[i], prompt_weight_2[i], prompt_weight_3[i], prompt_weight_4[i]) nxt_prompt_series[i] = prepare_batch_prompt(nxt_prompt_series[i], max_frames, i, prompt_weight_1[i], prompt_weight_2[i], prompt_weight_3[i], prompt_weight_4[i]) # Show the to/from prompts with evaluated expressions for transparency. - for i in range(len(cur_prompt_series)): - print("\n", "Max Frames: ", max_frames, "\n", "Current Prompt: ", cur_prompt_series[i], "\n", - "Next Prompt: ", nxt_prompt_series[i], "\n", "Strength : ", weight_series[i], "\n") + #for i in range(len(cur_prompt_series)): + # print("\n", "Max Frames: ", max_frames, "\n", "Current Prompt: ", cur_prompt_series[i], "\n", + # "Next Prompt: ", nxt_prompt_series[i], "\n", "Strength : ", weight_series[i], "\n") # Output methods depending if the prompts are the same or if the current frame is a keyframe. # if it is an in-between frame and the prompts differ, composable diffusion will be performed. @@ -160,10 +161,36 @@ def BatchPoolAnimConditioning(cur_prompt_series, nxt_prompt_series, weight_serie cond_out.append(interpolated_cond) final_pooled_output = torch.cat(pooled_out, dim=0) - final_conditioning = torch.cat(cond_out, dim=0) + final_conditioning = torch.cat(cond_out, dim=1) return [[final_conditioning, {"pooled_output": final_pooled_output}]] +def BatchGLIGENConditioning(cur_prompt_series, nxt_prompt_series, weight_series, clip): + pooled_out = [] + cond_out = [] + + for i in range(len(cur_prompt_series)): + tokens = clip.tokenize(str(cur_prompt_series[i])) + cond_to, pooled_to = clip.encode_from_tokens(tokens, return_pooled=True) + + tokens = clip.tokenize(str(nxt_prompt_series[i])) + cond_from, pooled_from = clip.encode_from_tokens(tokens, return_pooled=True) + + interpolated_conditioning = addWeighted([[cond_to, {"pooled_output": pooled_to}]], + [[cond_from, {"pooled_output": pooled_from}]], + weight_series[i]) + + interpolated_cond = interpolated_conditioning[0][0] + interpolated_pooled = interpolated_conditioning[0][1].get("pooled_output", pooled_from) + + pooled_out.append(interpolated_pooled) + cond_out.append(interpolated_cond) + + final_pooled_output = torch.cat(pooled_out, dim=0) + final_conditioning = torch.cat(cond_out, dim=0) + + return cond_out, pooled_out + def BatchPoolAnimConditioningSDXL(cur_prompt_series, nxt_prompt_series, weight_series, clip): pooled_out = [] cond_out = [] diff --git a/ScheduleFuncs.py b/ScheduleFuncs.py index 1501fc6..15a6dce 100644 --- a/ScheduleFuncs.py +++ b/ScheduleFuncs.py @@ -38,13 +38,17 @@ def addWeighted(conditioning_to, conditioning_from, conditioning_to_strength): out.append(n) return out - # used by both nodes +def reverseConcatenation(final_conditioning, final_pooled_output, max_frames): + # Split the final_conditioning and final_pooled_output tensors into their original components + cond_out = torch.split(final_conditioning, max_frames) + pooled_out = torch.split(final_pooled_output, max_frames) + + return cond_out, pooled_out + def check_is_number(value): float_pattern = r'^(?=.)([+-]?([0-9]*)(\.([0-9]+))?)$' return re.match(float_pattern, value) - - def parse_weight(match, frame=0, max_frames=0) -> float: #calculate weight steps for in-betweens w_raw = match.group("weight") max_f = max_frames # this line has to be left intact as it's in use by numexpr even though it looks like it doesn't diff --git a/ScheduledNodes.py b/ScheduledNodes.py index 1e8a164..b663a28 100644 --- a/ScheduledNodes.py +++ b/ScheduledNodes.py @@ -9,8 +9,9 @@ import re import json -from .ScheduleFuncs import check_is_number, interpolate_prompts, interpolate_prompts_SDXL, PoolAnimConditioning, interpolate_string -from .BatchFuncs import interpolate_prompt_series, BatchPoolAnimConditioning, BatchInterpolatePromptsSDXL +from .ScheduleFuncs import check_is_number, interpolate_prompts, interpolate_prompts_SDXL, PoolAnimConditioning, interpolate_string, addWeighted, reverseConcatenation +from .BatchFuncs import interpolate_prompt_series, BatchPoolAnimConditioning, BatchInterpolatePromptsSDXL #, BatchGLIGENConditioning +from .ValueFuncs import batch_get_inbetweens, batch_parse_key_frames, parse_key_frames, get_inbetweens, sanitize_value #Max resolution value for Gligen area calculation. MAX_RESOLUTION=8192 @@ -52,7 +53,7 @@ class PromptSchedule: "clip": ("CLIP", ), "max_frames": ("INT", {"default": 120.0, "min": 1.0, "max": 9999.0, "step": 1.0}), "current_frame": ("INT", {"default": 0.0, "min": 0.0, "max": 9999.0, "step": 1.0,})},# "forceInput": True}),}, - "optional": {"pre_text": ("STRING", {"multiline": False,}),# "forceInput": True}), + "optional": {"pre_text": ("STRING", {"multiline": False,}),# "forceInput": True}), "app_text": ("STRING", {"multiline": False,}),# "forceInput": True}), "pw_a": ("FLOAT", {"default": 0.0, "min": -9999.0, "max": 9999.0, "step": 0.1,}), #"forceInput": True }), "pw_b": ("FLOAT", {"default": 0.0, "min": -9999.0, "max": 9999.0, "step": 0.1,}), #"forceInput": True }), @@ -311,6 +312,68 @@ class PromptScheduleNodeFlowEnd: animation_prompts = json.loads(inputText.strip()) return (interpolate_prompts(animation_prompts, max_frames, current_frame, clip, pre_text, app_text, pw_a, pw_b, pw_c, pw_d, ),) #return a conditioning value +class BatchGLIGENSchedule: + @classmethod + def INPUT_TYPES(s): + return {"required": {"conditioning_to": ("CONDITIONING",), + "clip": ("CLIP",), + "gligen_textbox_model": ("GLIGEN",), + "text": ("STRING", {"multiline": True, "default":defaultPrompt}), + "width": ("INT", {"default": 64, "min": 8, "max": MAX_RESOLUTION, "step": 8}), + "height": ("INT", {"default": 64, "min": 8, "max": MAX_RESOLUTION, "step": 8}), + "x": ("INT", {"default": 0, "min": 0, "max": MAX_RESOLUTION, "step": 8}), + "y": ("INT", {"default": 0, "min": 0, "max": MAX_RESOLUTION, "step": 8}), + "max_frames": ("INT", {"default": 120.0, "min": 1.0, "max": 9999.0, "step": 1.0}),}, + # "forceInput": True}),}, + "optional": {"pre_text": ("STRING", {"multiline": False, }), # "forceInput": True}), + "app_text": ("STRING", {"multiline": False, }), # "forceInput": True}), + "pw_a": ("FLOAT", {"default": 0.0, "min": -9999.0, "max": 9999.0, "step": 0.1, }), + # "forceInput": True }), + "pw_b": ("FLOAT", {"default": 0.0, "min": -9999.0, "max": 9999.0, "step": 0.1, }), + # "forceInput": True }), + "pw_c": ("FLOAT", {"default": 0.0, "min": -9999.0, "max": 9999.0, "step": 0.1, }), + # "forceInput": True }), + "pw_d": ("FLOAT", {"default": 0.0, "min": -9999.0, "max": 9999.0, "step": 0.1, }), + # "forceInput": True }), + }} + + RETURN_TYPES = ("CONDITIONING",) + FUNCTION = "animate" + + CATEGORY = "FizzNodes/BatchScheduleNodes" + + def animate(self, conditioning_to, clip, gligen_textbox_model, text, width, height, x, y, max_frames, pw_a, pw_b, pw_c, pw_d, pre_text='', app_text=''): + inputText = str("{" + text + "}") + animation_prompts = json.loads(inputText.strip()) + cur_series, nxt_series, weight_series = interpolate_prompt_series(animation_prompts, max_frames, pre_text, app_text, pw_a, pw_b, pw_c, pw_d) + out = [] + for i in range(0, max_frames - 1): + # Calculate changes in x and y here, based on your logic + x_change = 8 + y_change = 0 + + # Update x and y values + x += x_change + y += y_change + print(x) + print(y) + out.append(self.append(conditioning_to, clip, gligen_textbox_model, pre_text, width, height, x, y)) + + return (out,) + + def append(self, conditioning_to, clip, gligen_textbox_model, text, width, height, x, y): + c = [] + cond, cond_pooled = clip.encode_from_tokens(clip.tokenize(text), return_pooled=True) + for t in range(0, len(conditioning_to)): + n = [conditioning_to[t][0], conditioning_to[t][1].copy()] + position_params = [(cond_pooled, height // 8, width // 8, y // 8, x // 8)] + prev = [] + if "gligen" in n[1]: + prev = n[1]['gligen'][2] + + n[1]['gligen'] = ("position", gligen_textbox_model, prev + position_params) + c.append(n) + return c #This node parses the user's test input into #interpolated floats. Expressions can be input @@ -328,60 +391,10 @@ class ValueSchedule: CATEGORY = "FizzNodes/ScheduleNodes" def animate(self, text, max_frames, current_frame,): - t = self.get_inbetweens(self.parse_key_frames(text, max_frames), max_frames) + t = get_inbetweens(parse_key_frames(text, max_frames), max_frames) cFrame = current_frame return (t[cFrame],int(t[cFrame]),) - def sanitize_value(self, value): - return value.replace("'","").replace('"',"").replace('(',"").replace(')',"") - - def get_inbetweens(self, key_frames, max_frames, integer=False, interp_method='Linear', is_single_string = False): - key_frame_series = pd.Series([np.nan for a in range(max_frames)]) - max_f = max_frames -1 #needed for numexpr even though it doesn't look like it's in use. - value_is_number = False - for i in range(0, max_frames): - if i in key_frames: - value = key_frames[i] - value_is_number = check_is_number(self.sanitize_value(value)) - if value_is_number: # if it's only a number, leave the rest for the default interpolation - key_frame_series[i] = self.sanitize_value(value) - if not value_is_number: - t = i - # workaround for values formatted like 0:("I am test") //used for sampler schedules - key_frame_series[i] = numexpr.evaluate(value) if not is_single_string else self.sanitize_value(value) - elif is_single_string:# take previous string value and replicate it - key_frame_series[i] = key_frame_series[i-1] - key_frame_series = key_frame_series.astype(float) if not is_single_string else key_frame_series # as string - - if interp_method == 'Cubic' and len(key_frames.items()) <= 3: - interp_method = 'Quadratic' - if interp_method == 'Quadratic' and len(key_frames.items()) <= 2: - interp_method = 'Linear' - - key_frame_series[0] = key_frame_series[key_frame_series.first_valid_index()] - key_frame_series[max_frames-1] = key_frame_series[key_frame_series.last_valid_index()] - key_frame_series = key_frame_series.interpolate(method=interp_method.lower(), limit_direction='both') - - if integer: - return key_frame_series.astype(int) - return key_frame_series - - def parse_key_frames(self, string, max_frames): - # because math functions (i.e. sin(t)) can utilize brackets - # it extracts the value in form of some stuff - # which has previously been enclosed with brackets and - # with a comma or end of line existing after the closing one - frames = dict() - for match_object in string.split(","): - frameParam = match_object.split(":") - max_f = max_frames -1 #needed for numexpr even though it doesn't look like it's in use. - frame = int(self.sanitize_value(frameParam[0])) if check_is_number(self.sanitize_value(frameParam[0].strip())) else int(numexpr.evaluate(frameParam[0].strip().replace("'","",1).replace('"',"",1)[::-1].replace("'","",1).replace('"',"",1)[::-1])) - frames[frame] = frameParam[1].strip() - if frames == {} and len(string) != 0: - raise RuntimeError('Key Frame string not correctly formatted') - return frames - - class BatchValueSchedule: @classmethod def INPUT_TYPES(s): @@ -395,57 +408,5 @@ class BatchValueSchedule: CATEGORY = "FizzNodes/BatchScheduleNodes" def animate(self, text, max_frames, ): - t = self.get_inbetweens(self.parse_key_frames(text, max_frames), max_frames) - return (t, list(map(int,t)),) - - def sanitize_value(self, value): - return value.replace("'","").replace('"',"").replace('(',"").replace(')',"") - def get_inbetweens(self, key_frames, max_frames, integer=False, interp_method='Linear', is_single_string=False): - key_frame_series = pd.Series([np.nan for a in range(max_frames)]) - max_f = max_frames - 1 # needed for numexpr even though it doesn't look like it's in use. - value_is_number = False - for i in range(0, max_frames): - if i in key_frames: - value = key_frames[i] - value_is_number = check_is_number(self.sanitize_value(value)) - if value_is_number: # if it's only a number, leave the rest for the default interpolation - key_frame_series[i] = self.sanitize_value(value) - if not value_is_number: - t = i - # workaround for values formatted like 0:("I am test") //used for sampler schedules - key_frame_series[i] = numexpr.evaluate(value) if not is_single_string else self.sanitize_value(value) - elif is_single_string: # take previous string value and replicate it - key_frame_series[i] = key_frame_series[i - 1] - key_frame_series = key_frame_series.astype(float) if not is_single_string else key_frame_series # as string - - if interp_method == 'Cubic' and len(key_frames.items()) <= 3: - interp_method = 'Quadratic' - if interp_method == 'Quadratic' and len(key_frames.items()) <= 2: - interp_method = 'Linear' - - key_frame_series[0] = key_frame_series[key_frame_series.first_valid_index()] - key_frame_series[max_frames - 1] = key_frame_series[key_frame_series.last_valid_index()] - key_frame_series = key_frame_series.interpolate(method=interp_method.lower(), limit_direction='both') - - if integer: - return key_frame_series.astype(int) - return key_frame_series - - def parse_key_frames(self, string, max_frames): - # because math functions (i.e. sin(t)) can utilize brackets - # it extracts the value in form of some stuff - # which has previously been enclosed with brackets and - # with a comma or end of line existing after the closing one - frames = dict() - for match_object in string.split(","): - frameParam = match_object.split(":") - max_f = max_frames - 1 # needed for numexpr even though it doesn't look like it's in use. - frame = int(self.sanitize_value(frameParam[0])) if check_is_number( - self.sanitize_value(frameParam[0].strip())) else int(numexpr.evaluate( - frameParam[0].strip().replace("'", "", 1).replace('"', "", 1)[::-1].replace("'", "", 1).replace('"', "", - 1)[ - ::-1])) - frames[frame] = frameParam[1].strip() - if frames == {} and len(string) != 0: - raise RuntimeError('Key Frame string not correctly formatted') - return frames \ No newline at end of file + t = batch_get_inbetweens(batch_parse_key_frames(text, max_frames), max_frames) + return (t, list(map(int,t)),) \ No newline at end of file diff --git a/ValueFuncs.py b/ValueFuncs.py new file mode 100644 index 0000000..c9eda24 --- /dev/null +++ b/ValueFuncs.py @@ -0,0 +1,108 @@ +import numexpr +import torch +import numpy as np +import pandas as pd +import re +import json + +from .ScheduleFuncs import check_is_number +def sanitize_value(value): + return value.replace("'", "").replace('"', "").replace('(', "").replace(')', "") + + +def get_inbetweens(key_frames, max_frames, integer=False, interp_method='Linear', is_single_string=False): + key_frame_series = pd.Series([np.nan for a in range(max_frames)]) + max_f = max_frames - 1 # needed for numexpr even though it doesn't look like it's in use. + value_is_number = False + for i in range(0, max_frames): + if i in key_frames: + value = key_frames[i] + value_is_number = check_is_number(sanitize_value(value)) + if value_is_number: # if it's only a number, leave the rest for the default interpolation + key_frame_series[i] = sanitize_value(value) + if not value_is_number: + t = i + # workaround for values formatted like 0:("I am test") //used for sampler schedules + key_frame_series[i] = numexpr.evaluate(value) if not is_single_string else sanitize_value(value) + elif is_single_string: # take previous string value and replicate it + key_frame_series[i] = key_frame_series[i - 1] + key_frame_series = key_frame_series.astype(float) if not is_single_string else key_frame_series # as string + + if interp_method == 'Cubic' and len(key_frames.items()) <= 3: + interp_method = 'Quadratic' + if interp_method == 'Quadratic' and len(key_frames.items()) <= 2: + interp_method = 'Linear' + + key_frame_series[0] = key_frame_series[key_frame_series.first_valid_index()] + key_frame_series[max_frames - 1] = key_frame_series[key_frame_series.last_valid_index()] + key_frame_series = key_frame_series.interpolate(method=interp_method.lower(), limit_direction='both') + + if integer: + return key_frame_series.astype(int) + return key_frame_series + + +def parse_key_frames(string, max_frames): + # because math functions (i.e. sin(t)) can utilize brackets + # it extracts the value in form of some stuff + # which has previously been enclosed with brackets and + # with a comma or end of line existing after the closing one + frames = dict() + for match_object in string.split(","): + frameParam = match_object.split(":") + max_f = max_frames - 1 # needed for numexpr even though it doesn't look like it's in use. + frame = int(sanitize_value(frameParam[0])) if check_is_number( + sanitize_value(frameParam[0].strip())) else int(numexpr.evaluate( + frameParam[0].strip().replace("'", "", 1).replace('"', "", 1)[::-1].replace("'", "", 1).replace('"', "", 1)[::-1])) + frames[frame] = frameParam[1].strip() + if frames == {} and len(string) != 0: + raise RuntimeError('Key Frame string not correctly formatted') + return frames + +def batch_get_inbetweens(key_frames, max_frames, integer=False, interp_method='Linear', is_single_string=False): + key_frame_series = pd.Series([np.nan for a in range(max_frames)]) + max_f = max_frames - 1 # needed for numexpr even though it doesn't look like it's in use. + value_is_number = False + for i in range(0, max_frames): + if i in key_frames: + value = key_frames[i] + value_is_number = check_is_number(sanitize_value(value)) + if value_is_number: # if it's only a number, leave the rest for the default interpolation + key_frame_series[i] = sanitize_value(value) + if not value_is_number: + t = i + # workaround for values formatted like 0:("I am test") //used for sampler schedules + key_frame_series[i] = numexpr.evaluate(value) if not is_single_string else sanitize_value(value) + elif is_single_string: # take previous string value and replicate it + key_frame_series[i] = key_frame_series[i - 1] + key_frame_series = key_frame_series.astype(float) if not is_single_string else key_frame_series # as string + + if interp_method == 'Cubic' and len(key_frames.items()) <= 3: + interp_method = 'Quadratic' + if interp_method == 'Quadratic' and len(key_frames.items()) <= 2: + interp_method = 'Linear' + + key_frame_series[0] = key_frame_series[key_frame_series.first_valid_index()] + key_frame_series[max_frames - 1] = key_frame_series[key_frame_series.last_valid_index()] + key_frame_series = key_frame_series.interpolate(method=interp_method.lower(), limit_direction='both') + + if integer: + return key_frame_series.astype(int) + return key_frame_series + +def batch_parse_key_frames(string, max_frames): + # because math functions (i.e. sin(t)) can utilize brackets + # it extracts the value in form of some stuff + # which has previously been enclosed with brackets and + # with a comma or end of line existing after the closing one + frames = dict() + for match_object in string.split(","): + frameParam = match_object.split(":") + max_f = max_frames - 1 # needed for numexpr even though it doesn't look like it's in use. + frame = int(sanitize_value(frameParam[0])) if check_is_number( + sanitize_value(frameParam[0].strip())) else int(numexpr.evaluate( + frameParam[0].strip().replace("'", "", 1).replace('"', "", 1)[::-1].replace("'", "", 1).replace('"', "",1)[::-1])) + frames[frame] = frameParam[1].strip() + if frames == {} and len(string) != 0: + raise RuntimeError('Key Frame string not correctly formatted') + return frames \ No newline at end of file diff --git a/__init__.py b/__init__.py index b39c56c..e68ec3e 100644 --- a/__init__.py +++ b/__init__.py @@ -54,7 +54,7 @@ def is_installed(package, package_overwrite=None): print(f"Couldn't install\nCommand: {command}\nError code: {result.returncode}") from .WaveNodes import Lerp, SinWave, InvSinWave, CosWave, InvCosWave, SquareWave, SawtoothWave, TriangleWave, AbsCosWave, AbsSinWave -from .ScheduledNodes import ValueSchedule, PromptSchedule, PromptScheduleNodeFlow, PromptScheduleNodeFlowEnd, PromptScheduleEncodeSDXL, StringSchedule, BatchPromptSchedule, BatchValueSchedule, BatchPromptScheduleEncodeSDXL +from .ScheduledNodes import ValueSchedule, PromptSchedule, PromptScheduleNodeFlow, PromptScheduleNodeFlowEnd, PromptScheduleEncodeSDXL, StringSchedule, BatchPromptSchedule, BatchValueSchedule, BatchPromptScheduleEncodeSDXL #, BatchGLIGENSchedule NODE_CLASS_MAPPINGS = { "Lerp": Lerp, @@ -75,7 +75,9 @@ NODE_CLASS_MAPPINGS = { "StringSchedule":StringSchedule, "BatchPromptSchedule": BatchPromptSchedule, "BatchValueSchedule": BatchValueSchedule, - "BatchPromptScheduleEncodeSDXL": BatchPromptScheduleEncodeSDXL + "BatchPromptScheduleEncodeSDXL": BatchPromptScheduleEncodeSDXL, + #"BatchGLIGENSchedule": BatchGLIGENSchedule, + } print('\033[34mFizzleDorf Custom Nodes: \033[92mLoaded\033[0m') \ No newline at end of file