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:
Jedrzej Kosinski
2025-08-05 20:12:09 -07:00
committed by GitHub
3 changed files with 61 additions and 3 deletions
+3 -1
View File
@@ -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
+33 -2
View File
@@ -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,)
+25
View File
@@ -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