fixed token size issue and WIP Gligen schedule(disabled)

This commit is contained in:
FizzleDorf
2023-10-23 01:44:47 -04:00
parent a8c10cfab8
commit 2b52fdaa7b
5 changed files with 220 additions and 118 deletions
+32 -5
View File
@@ -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 = []
+7 -3
View File
@@ -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
+69 -108
View File
@@ -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
t = batch_get_inbetweens(batch_parse_key_frames(text, max_frames), max_frames)
return (t, list(map(int,t)),)
+108
View File
@@ -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
+4 -2
View File
@@ -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')