diff --git a/animatediff/nodes.py b/animatediff/nodes.py index ffca1fc..4a8e5d4 100644 --- a/animatediff/nodes.py +++ b/animatediff/nodes.py @@ -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 diff --git a/animatediff/nodes_scheduling.py b/animatediff/nodes_scheduling.py index 96cd96c..d423a6a 100644 --- a/animatediff/nodes_scheduling.py +++ b/animatediff/nodes_scheduling.py @@ -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,) \ No newline at end of file + 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,) diff --git a/animatediff/scheduling.py b/animatediff/scheduling.py index 4b56119..647ba09 100644 --- a/animatediff/scheduling.py +++ b/animatediff/scheduling.py @@ -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