Merge pull request #537 from DrJKL/drjkl/feat/extract_condition_step
[Feat] Add node to pull a single condition out of a scheduled prompt condition
This commit is contained in:
@@ -39,7 +39,7 @@ from .nodes_context_extras import (SetContextExtrasOnContextOptions, ContextExtr
|
||||
from .nodes_ad_settings import (AnimateDiffSettingsNode, ManualAdjustPENode, SweetspotStretchPENode, FullStretchPENode,
|
||||
WeightAdjustAllAddNode, WeightAdjustAllMultNode, WeightAdjustIndivAddNode, WeightAdjustIndivMultNode,
|
||||
WeightAdjustIndivAttnAddNode, WeightAdjustIndivAttnMultNode)
|
||||
from .nodes_scheduling import (PromptSchedulingNode, PromptSchedulingLatentsNode, ValueSchedulingNode, ValueSchedulingLatentsNode,
|
||||
from .nodes_scheduling import (ConditionExtractionNode, PromptSchedulingNode, PromptSchedulingLatentsNode, ValueSchedulingNode, ValueSchedulingLatentsNode,
|
||||
AddValuesReplaceNode, FloatToFloatsNode)
|
||||
from .nodes_per_block import (ADBlockComboNode, ADBlockIndivNode, PerBlockHighLevelNode,
|
||||
PerBlock_SD15_LowLevelNode, PerBlock_SD15_MidLevelNode, PerBlock_SD15_FromFloatsNode,
|
||||
@@ -165,6 +165,7 @@ NODE_CLASS_MAPPINGS = {
|
||||
PromptSchedulingLatentsNode.NodeID: PromptSchedulingLatentsNode,
|
||||
ValueSchedulingNode.NodeID: ValueSchedulingNode,
|
||||
ValueSchedulingLatentsNode.NodeID: ValueSchedulingLatentsNode,
|
||||
ConditionExtractionNode.NodeID: ConditionExtractionNode,
|
||||
AddValuesReplaceNode.NodeID: AddValuesReplaceNode,
|
||||
FloatToFloatsNode.NodeID: FloatToFloatsNode,
|
||||
# Per-Block
|
||||
@@ -345,6 +346,7 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
PromptSchedulingLatentsNode.NodeID: PromptSchedulingLatentsNode.NodeName,
|
||||
ValueSchedulingNode.NodeID: ValueSchedulingNode.NodeName,
|
||||
ValueSchedulingLatentsNode.NodeID: ValueSchedulingLatentsNode.NodeName,
|
||||
ConditionExtractionNode.NodeID: ConditionExtractionNode.NodeName,
|
||||
AddValuesReplaceNode.NodeID: AddValuesReplaceNode.NodeName,
|
||||
FloatToFloatsNode.NodeID:FloatToFloatsNode.NodeName,
|
||||
# Per-Block
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
from typing import Union
|
||||
|
||||
from .documentation import register_description, short_desc, coll, DocHelper
|
||||
from .scheduling import (evaluate_prompt_schedule, evaluate_value_schedule, TensorInterp, PromptOptions,
|
||||
from .scheduling import (evaluate_prompt_schedule, evaluate_value_schedule, extract_cond_from_schedule, TensorInterp, PromptOptions,
|
||||
verify_key_value)
|
||||
from .utils_model import BIGMAX
|
||||
from .logger import logger
|
||||
@@ -24,6 +24,10 @@ desc_FLOAT = {'FLOAT': 'Float (or list of floats) to convert to FLOATS type.'}
|
||||
desc_value_key = {'value_key': 'Key to use for value schedule in Prompt Scheduling node. Can only contain a-z, A-Z, 0-9, and _ characters. In Prompt Scheduling, keys can be referred to as `some_key`, where the key is surrounded by ` characters.'}
|
||||
desc_prev_replace = {'prev_replace': 'OPTIONAL, other values_replace can be chained.'}
|
||||
|
||||
desc_input_conditioning = {'conditioning': 'Encoded prompts. The output of a Prompt Scheduling node.'}
|
||||
desc_index = {'index': 'The index to extract. Must be within the range [0,N] where N is the length of scheduled prompts.'}
|
||||
desc_output_conditioning_single = {'CONDITIONING': 'The single step conditioning from the schedule.'}
|
||||
|
||||
desc_output_conditioning = {'CONDITIONING': 'Encoded prompts.'}
|
||||
desc_output_latent = {'LATENT': 'Unmodified input latents; can be used as pipe, or can be ignored.'}
|
||||
|
||||
@@ -280,4 +284,31 @@ class FloatToFloatsNode:
|
||||
floats = [float(FLOAT)]
|
||||
else:
|
||||
floats = list(FLOAT)
|
||||
return (floats,)
|
||||
return (floats,)
|
||||
|
||||
class ConditionExtractionNode:
|
||||
NodeID = 'ADE_ConditionExtraction'
|
||||
NodeName = 'Condition Step Extraction 🎭🅐🅓'
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"conditioning": ("CONDITIONING",),
|
||||
"index": ("INT", {"default": 0, "min": 0, "step": 1})
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("CONDITIONING",)
|
||||
CATEGORY = "Animate Diff 🎭🅐🅓/scheduling"
|
||||
FUNCTION = "extract_conditioning"
|
||||
|
||||
Desc = [
|
||||
short_desc('Extract a single conditioning step from a schedule of prompts.'),
|
||||
{coll('Inputs'): DocHelper.combine(desc_input_conditioning, desc_index)},
|
||||
{coll('Outputs'): DocHelper.combine(desc_output_conditioning)}
|
||||
]
|
||||
register_description(NodeID, Desc)
|
||||
|
||||
def extract_conditioning(self, conditioning, index: int=0):
|
||||
conditioning_step = extract_cond_from_schedule(conditioning, index)
|
||||
return (conditioning_step,)
|
||||
|
||||
@@ -112,6 +112,8 @@ class PromptOptions:
|
||||
print_schedule: bool = False
|
||||
add_dict: dict[str] = None
|
||||
|
||||
IndividualConditioning = tuple[torch.Tensor, dict[str, torch.Tensor]]
|
||||
Conditioning = list[IndividualConditioning]
|
||||
|
||||
def evaluate_prompt_schedule(text: str, length: int, clip: CLIP, options: PromptOptions):
|
||||
text = strip_input(text)
|
||||
@@ -445,6 +447,29 @@ def _handle_prompt_interpolation(pairs: list[InputPair], length: int, clip: CLIP
|
||||
clip.add_hooks_to_dict(final_pooled_dict)
|
||||
return [[final_cond, final_pooled_dict]]
|
||||
|
||||
def extract_cond_from_schedule(conditioning: Conditioning, index: int) -> Conditioning:
|
||||
return [_extract_single_cond(t, index) for t in conditioning]
|
||||
|
||||
def _extract_single_cond(single_cond: IndividualConditioning, index:int) -> IndividualConditioning:
|
||||
if index < 0:
|
||||
return single_cond
|
||||
|
||||
cond, kwargs = single_cond[0], single_cond[1].copy()
|
||||
original_pooled = kwargs["pooled_output"]
|
||||
|
||||
cond_schedules = cond.shape[0]
|
||||
pooled_schedules = original_pooled.shape[0]
|
||||
|
||||
if cond_schedules <= index or pooled_schedules <= index:
|
||||
logger.warning(f"Trying to get index {index}, only have {cond_schedules} items")
|
||||
return single_cond
|
||||
|
||||
cond_chunks = cond.chunk(cond_schedules)
|
||||
chosen_cond = cond_chunks[index]
|
||||
|
||||
pool_chunks = original_pooled.chunk(pooled_schedules)
|
||||
kwargs["pooled_output"] = pool_chunks[index]
|
||||
return [chosen_cond, kwargs]
|
||||
|
||||
def pad_cond(cond: Tensor, target_length: int):
|
||||
# FizzNodes-style cond padding
|
||||
|
||||
Reference in New Issue
Block a user