From 7d36d7176c647b990a540e518c870fe7dab2f96a Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Tue, 9 Jul 2024 12:44:27 -0500 Subject: [PATCH 1/9] start of work on context visualization, custom cfg simple variants --- animatediff/context.py | 31 +++++++++++++++++++++++++ animatediff/motion_module_ad.py | 4 ++-- animatediff/nodes.py | 15 ++++++++---- animatediff/nodes_context.py | 30 ++++++++++++++++++++++++ animatediff/nodes_sample.py | 41 +++++++++++++++++++++++++++++++++ 5 files changed, 115 insertions(+), 6 deletions(-) diff --git a/animatediff/context.py b/animatediff/context.py index ee77cf0..63ddd12 100644 --- a/animatediff/context.py +++ b/animatediff/context.py @@ -1,5 +1,8 @@ from typing import Callable, Optional, Union +import torchvision +import PIL + import numpy as np from torch import Tensor @@ -473,3 +476,31 @@ def shift_window_to_end(window: list[int], num_frames: int): for i in range(len(window)): # 2) add end_delta to each val to slide windows to end window[i] = window[i] + end_delta + + +########################## +# Context Visualization +########################## +class Colors: + BLACK = (0, 0, 0) + WHITE = (255, 255, 255) + RED = (255, 0, 0) + GREEN = (0, 255, 0) + BLUE = (0, 0, 255) + YELLOW = (255, 255, 0) + MAGENTA = (255, 0, 255) + CYAN = (0, 255, 255) + + +class VisualizeSettings: + def __init__(self, img_width, img_height, video_length): + self.img_width = img_width + self.img_height = img_height + self.video_length = video_length + self.grid = img_width // video_length + self.pil_to_tensor = torchvision.transforms.Compose([torchvision.transforms.PILToTensor()]) + + +def generate_context_visualization(context_opts: ContextOptionsGroup, model: BaseModel, width=1440, height=200, video_length=32, start_step=0, end_step=20): + vs = VisualizeSettings(width, height, video_length) + pass diff --git a/animatediff/motion_module_ad.py b/animatediff/motion_module_ad.py index 4ee3b63..5435ad3 100644 --- a/animatediff/motion_module_ad.py +++ b/animatediff/motion_module_ad.py @@ -134,10 +134,10 @@ def has_img_encoder(mm_state_dict: dict[str, Tensor]): def normalize_ad_state_dict(mm_state_dict: dict[str, Tensor], mm_name: str) -> Tuple[dict[str, Tensor], AnimateDiffInfo]: # from pathlib import Path - # with open(Path(__file__).parent.parent.parent / f"keys_{mm_name}.txt", "w") as afile: + # log_name = mm_name.split('\\')[-1] + # with open(Path(__file__).parent.parent.parent / rf"keys_{log_name}.txt", "w") as afile: # for key, value in mm_state_dict.items(): # afile.write(f"{key}:\t{value.shape}\n") - # determine what SD model the motion module is intended for sd_type: str = None down_block_max = get_down_block_max(mm_state_dict) diff --git a/animatediff/nodes.py b/animatediff/nodes.py index 6a8e7cf..ff9c98f 100644 --- a/animatediff/nodes.py +++ b/animatediff/nodes.py @@ -20,10 +20,11 @@ from .nodes_conditioning import (MaskableLoraLoader, MaskableLoraLoaderModelOnly ConditioningTimestepsNode, SetLoraHookKeyframes, CreateLoraHookKeyframe, CreateLoraHookKeyframeInterpolation, CreateLoraHookKeyframeFromStrengthList) from .nodes_sample import (FreeInitOptionsNode, NoiseLayerAddWeightedNode, SampleSettingsNode, NoiseLayerAddNode, NoiseLayerReplaceNode, IterationOptionsNode, - CustomCFGNode, CustomCFGKeyframeNode, NoisedImageInjectionNode, NoisedImageInjectOptionsNode) + CustomCFGNode, CustomCFGSimpleNode, CustomCFGKeyframeNode, CustomCFGKeyframeSimpleNode, + NoisedImageInjectionNode, NoisedImageInjectOptionsNode) from .nodes_sigma_schedule import (SigmaScheduleNode, RawSigmaScheduleNode, WeightedAverageSigmaScheduleNode, InterpolatedWeightedAverageSigmaScheduleNode, SplitAndCombineSigmaScheduleNode) from .nodes_context import (LegacyLoopedUniformContextOptionsNode, LoopedUniformContextOptionsNode, LoopedUniformViewOptionsNode, StandardUniformContextOptionsNode, StandardStaticContextOptionsNode, BatchedContextOptionsNode, - StandardStaticViewOptionsNode, StandardUniformViewOptionsNode, ViewAsContextOptionsNode) + StandardStaticViewOptionsNode, StandardUniformViewOptionsNode, ViewAsContextOptionsNode, VisualizeContextOptionsInt) from .nodes_ad_settings import (AnimateDiffSettingsNode, ManualAdjustPENode, SweetspotStretchPENode, FullStretchPENode, WeightAdjustAllAddNode, WeightAdjustAllMultNode, WeightAdjustIndivAddNode, WeightAdjustIndivMultNode, WeightAdjustIndivAttnAddNode, WeightAdjustIndivAttnMultNode) @@ -56,6 +57,7 @@ NODE_CLASS_MAPPINGS = { "ADE_ViewsOnlyContextOptions": ViewAsContextOptionsNode, "ADE_BatchedContextOptions": BatchedContextOptionsNode, "ADE_AnimateDiffUniformContextOptions": LegacyLoopedUniformContextOptionsNode, # Legacy + "ADE_VisualizeContextOptions": VisualizeContextOptionsInt, # View Opts "ADE_StandardStaticViewOptions": StandardStaticViewOptionsNode, "ADE_StandardUniformViewOptions": StandardUniformViewOptionsNode, @@ -100,7 +102,9 @@ NODE_CLASS_MAPPINGS = { "ADE_AdjustWeightIndivAttnAdd": WeightAdjustIndivAttnAddNode, "ADE_AdjustWeightIndivAttnMult": WeightAdjustIndivAttnMultNode, # Sample Settings + "ADE_CustomCFGSimple": CustomCFGSimpleNode, "ADE_CustomCFG": CustomCFGNode, + "ADE_CustomCFGKeyframeSimple": CustomCFGKeyframeSimpleNode, "ADE_CustomCFGKeyframe": CustomCFGKeyframeNode, "ADE_SigmaSchedule": SigmaScheduleNode, "ADE_RawSigmaSchedule": RawSigmaScheduleNode, @@ -169,6 +173,7 @@ NODE_DISPLAY_NAME_MAPPINGS = { "ADE_ViewsOnlyContextOptions": "Context Optionsβ—†Views Only [VRAMβ‡ˆ] πŸŽ­πŸ…πŸ…“", "ADE_BatchedContextOptions": "Context Optionsβ—†Batched [Non-AD] πŸŽ­πŸ…πŸ…“", "ADE_AnimateDiffUniformContextOptions": "Context Optionsβ—†Looped Uniform πŸŽ­πŸ…πŸ…“", # Legacy + "ADE_VisualizeContextOptions": "Visualize Context Options πŸŽ­πŸ…πŸ…“", # View Opts "ADE_StandardStaticViewOptions": "View Optionsβ—†Standard Static πŸŽ­πŸ…πŸ…“", "ADE_StandardUniformViewOptions": "View Optionsβ—†Standard Uniform πŸŽ­πŸ…πŸ…“", @@ -213,8 +218,10 @@ NODE_DISPLAY_NAME_MAPPINGS = { "ADE_AdjustWeightIndivAttnAdd": "Adjust Weight [Indiv-Attnβ—†Add] πŸŽ­πŸ…πŸ…“", "ADE_AdjustWeightIndivAttnMult": "Adjust Weight [Indiv-Attnβ—†Mult] πŸŽ­πŸ…πŸ…“", # Sample Settings - "ADE_CustomCFG": "Custom CFG πŸŽ­πŸ…πŸ…“", - "ADE_CustomCFGKeyframe": "Custom CFG Keyframe πŸŽ­πŸ…πŸ…“", + "ADE_CustomCFGSimple": "Custom CFG [Simple] πŸŽ­πŸ…πŸ…“", + "ADE_CustomCFG": "Custom CFG [Multival] πŸŽ­πŸ…πŸ…“", + "ADE_CustomCFGKeyframeSimple": "Custom CFG Keyframe [Simple] πŸŽ­πŸ…πŸ…“", + "ADE_CustomCFGKeyframe": "Custom CFG Keyframe [Multival] πŸŽ­πŸ…πŸ…“", "ADE_SigmaSchedule": "Create Sigma Schedule πŸŽ­πŸ…πŸ…“", "ADE_RawSigmaSchedule": "Create Raw Sigma Schedule πŸŽ­πŸ…πŸ…“", "ADE_SigmaScheduleWeightedAverage": "Sigma Schedule Weighted Mean πŸŽ­πŸ…πŸ…“", diff --git a/animatediff/nodes_context.py b/animatediff/nodes_context.py index b4924f0..0940b76 100644 --- a/animatediff/nodes_context.py +++ b/animatediff/nodes_context.py @@ -1,3 +1,8 @@ +import torch +from torch import Tensor + +from comfy.model_patcher import ModelPatcher + from .context import ContextFuseMethod, ContextOptions, ContextOptionsGroup, ContextSchedules from .utils_model import BIGMAX @@ -346,3 +351,28 @@ class LoopedUniformViewOptionsNode: use_on_equal_length=use_on_equal_length, ) return (view_options,) + + +class VisualizeContextOptionsInt: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "model": ("MODEL",), + "context_opts": ("CONTEXT_OPTIONS",), + }, + "optional": { + "latents_length": ("INT", {"min": 1, "max": BIGMAX, "default": 32}), + "start_step": ("INT", {"min": 0, "max": BIGMAX, "default": 0}), + "end_step": ("INT", {"min": 1, "max": BIGMAX, "default": 20}), + } + } + + RETURN_TYPES = ("IMAGE",) + CATEGORY = "Animate Diff πŸŽ­πŸ…πŸ…“/context opts/visualize" + FUNCTION = "visualize" + + def visualize(self, model: ModelPatcher, context_opts: ContextOptionsGroup, + latents_length=32, start_step=0, end_step=20): + images = torch.zeros((latents_length, 256, 256, 3)) + return (images,) diff --git a/animatediff/nodes_sample.py b/animatediff/nodes_sample.py index 0744c66..3fb75b3 100644 --- a/animatediff/nodes_sample.py +++ b/animatediff/nodes_sample.py @@ -231,6 +231,23 @@ class CustomCFGNode: return (cfg_custom,) +class CustomCFGSimpleNode: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "cfg": ("FLOAT", {"default": 8.0, "min": 0.0, "max": 100.0, "step": 0.1}), + }, + } + + RETURN_TYPES = ("CUSTOM_CFG",) + CATEGORY = "Animate Diff πŸŽ­πŸ…πŸ…“/sample settings" + FUNCTION = "create_custom_cfg" + + def create_custom_cfg(self, cfg: float): + return CustomCFGNode.create_custom_cfg(self, cfg_multival=cfg) + + class CustomCFGKeyframeNode: @classmethod def INPUT_TYPES(s): @@ -259,6 +276,30 @@ class CustomCFGKeyframeNode: return (prev_custom_cfg,) +class CustomCFGKeyframeSimpleNode: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "cfg": ("FLOAT", {"default": 8.0, "min": 0.0, "max": 100.0, "step": 0.1}), + "start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}), + "guarantee_steps": ("INT", {"default": 1, "min": 0, "max": BIGMAX}), + }, + "optional": { + "prev_custom_cfg": ("CUSTOM_CFG",), + } + } + + RETURN_TYPES = ("CUSTOM_CFG",) + CATEGORY = "Animate Diff πŸŽ­πŸ…πŸ…“/sample settings" + FUNCTION = "create_custom_cfg" + + def create_custom_cfg(self, cfg: float, start_percent: float=0.0, guarantee_steps: int=1, + prev_custom_cfg: CustomCFGKeyframeGroup=None): + return CustomCFGKeyframeNode.create_custom_cfg(self, cfg_multival=cfg, start_percent=start_percent, + guarantee_steps=guarantee_steps, prev_custom_cfg=prev_custom_cfg) + + class NoisedImageInjectionNode: @classmethod def INPUT_TYPES(s): From e28b67781fbc191671a0e6a1e17d6113ef89eca6 Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Tue, 9 Jul 2024 14:43:50 -0500 Subject: [PATCH 2/9] Added 'comfy [gpu]' and 'auto1111 [gpu]' seed_gen --- animatediff/sample_settings.py | 102 +++++++++++++++++++++++---------- 1 file changed, 72 insertions(+), 30 deletions(-) diff --git a/animatediff/sample_settings.py b/animatediff/sample_settings.py index 818bba7..71086f4 100644 --- a/animatediff/sample_settings.py +++ b/animatediff/sample_settings.py @@ -200,58 +200,100 @@ class NoiseLayerGroup: cloned.add(layer) return cloned + +class RandDevice: + CPU = "cpu" + GPU = "gpu" + NV = "nv" + + +def get_generator(device=RandDevice.CPU, seed: int=None): + generator = None + raw_device = None + if device == RandDevice.CPU: + raw_device = "cpu" + generator = torch.Generator(raw_device) + elif device == RandDevice.GPU: + raw_device = comfy.model_management.get_torch_device() + generator = torch.Generator(raw_device) + # TODO: should I add the NV code from Auto1111? + # It is AGPL licenced, which should be fine since I will not be modifying it. + # elif device == RandDevice.NV: + # pass + else: + raise Exception(f"Unknown noise generator device: '{device}'") + if seed is not None: + generator = generator.manual_seed(seed) + return generator, raw_device + + class SeedNoiseGeneration: COMFY = "comfy" + COMFYGPU = "comfy [gpu]" + #COMFYNV = "comfy [nv]" AUTO1111 = "auto1111" - AUTO1111GPU = "auto1111 [gpu]" # TODO: implement this + AUTO1111GPU = "auto1111 [gpu]" + #AUTO1111NV = "auto1111 [nv]" USE_EXISTING = "use existing" - LIST = [COMFY, AUTO1111] - LIST_WITH_OVERRIDE = [USE_EXISTING, COMFY, AUTO1111] + LIST = [COMFY, COMFYGPU, AUTO1111, AUTO1111GPU] + LIST_WITH_OVERRIDE = [USE_EXISTING, COMFY, COMFYGPU, AUTO1111, AUTO1111GPU] + + _COMFY_GENS = [COMFY, COMFYGPU] + _AUTO1111_GENS = [AUTO1111, AUTO1111GPU] + + _SOURCE_DICT = { + COMFY: RandDevice.CPU, COMFYGPU: RandDevice.GPU, + AUTO1111: RandDevice.CPU, AUTO1111GPU: RandDevice.GPU, + } + + @classmethod + def get_device(cls, seed_gen: str): + return cls._SOURCE_DICT[seed_gen] @classmethod def create_noise(cls, seed: int, latents: Tensor, existing_seed_gen: str=COMFY, seed_gen: str=USE_EXISTING, noise_type: str=NoiseLayerType.DEFAULT, batch_offset: int=0, extra_args: dict={}): # determine if should use existing type if seed_gen == cls.USE_EXISTING: seed_gen = existing_seed_gen - if seed_gen == cls.COMFY: - return cls.create_noise_comfy(seed, latents, noise_type, batch_offset, extra_args) - elif seed_gen in [cls.AUTO1111, cls.AUTO1111GPU]: - return cls.create_noise_auto1111(seed, latents, noise_type, batch_offset, extra_args) + if seed_gen in cls._COMFY_GENS: + return cls.create_noise_comfy(seed, latents, noise_type, batch_offset, extra_args, cls.get_device(seed_gen)) + elif seed_gen in cls._AUTO1111_GENS: + return cls.create_noise_auto1111(seed, latents, noise_type, batch_offset, extra_args, cls.get_device(seed_gen)) raise ValueError(f"Noise seed_gen {seed_gen} is not recognized.") @staticmethod - def create_noise_comfy(seed: int, latents: Tensor, noise_type: str=NoiseLayerType.DEFAULT, batch_offset: int=0, extra_args: dict={}): - common_noise = SeedNoiseGeneration._create_common_noise(seed, latents, noise_type, batch_offset, extra_args) + def create_noise_comfy(seed: int, latents: Tensor, noise_type: str=NoiseLayerType.DEFAULT, batch_offset: int=0, extra_args: dict={}, device=RandDevice.CPU): + common_noise = SeedNoiseGeneration._create_common_noise(seed, latents, noise_type, batch_offset, extra_args, device) if common_noise is not None: return common_noise if noise_type == NoiseLayerType.CONSTANT: - generator = torch.manual_seed(seed) + generator, raw_device = get_generator(device, seed) length = latents.shape[0] single_shape = (1 + batch_offset, latents.shape[1], latents.shape[2], latents.shape[3]) - single_noise = torch.randn(single_shape, dtype=latents.dtype, layout=latents.layout, generator=generator, device="cpu") + single_noise = torch.randn(single_shape, dtype=latents.dtype, layout=latents.layout, generator=generator, device=raw_device).to(device="cpu") return torch.cat([single_noise[batch_offset:]] * length, dim=0) # comfy creates noise with a single seed for the entire shape of the latents batched tensor - generator = torch.manual_seed(seed) + generator, raw_device = get_generator(device, seed) offset_shape = (latents.shape[0] + batch_offset, latents.shape[1], latents.shape[2], latents.shape[3]) - final_noise = torch.randn(offset_shape, dtype=latents.dtype, layout=latents.layout, generator=generator, device="cpu") + final_noise = torch.randn(offset_shape, dtype=latents.dtype, layout=latents.layout, generator=generator, device=raw_device).to(device="cpu") final_noise = final_noise[batch_offset:] # convert to derivative noise type, if needed - derivative_noise = SeedNoiseGeneration._create_derivative_noise(final_noise, noise_type=noise_type, seed=seed, extra_args=extra_args) + derivative_noise = SeedNoiseGeneration._create_derivative_noise(final_noise, noise_type=noise_type, seed=seed, extra_args=extra_args, device=device) if derivative_noise is not None: return derivative_noise return final_noise @staticmethod - def create_noise_auto1111(seed: int, latents: Tensor, noise_type: str=NoiseLayerType.DEFAULT, batch_offset: int=0, extra_args: dict={}): - common_noise = SeedNoiseGeneration._create_common_noise(seed, latents, noise_type, batch_offset, extra_args) + def create_noise_auto1111(seed: int, latents: Tensor, noise_type: str=NoiseLayerType.DEFAULT, batch_offset: int=0, extra_args: dict={}, device=RandDevice.CPU): + common_noise = SeedNoiseGeneration._create_common_noise(seed, latents, noise_type, batch_offset, extra_args, device) if common_noise is not None: return common_noise if noise_type == NoiseLayerType.CONSTANT: - generator = torch.manual_seed(seed+batch_offset) + generator, raw_device = get_generator(device, seed+batch_offset) length = latents.shape[0] single_shape = (1, latents.shape[1], latents.shape[2], latents.shape[3]) - single_noise = torch.randn(single_shape, dtype=latents.dtype, layout=latents.layout, generator=generator, device="cpu") + single_noise = torch.randn(single_shape, dtype=latents.dtype, layout=latents.layout, generator=generator, device=raw_device).to(device="cpu") return torch.cat([single_noise] * length, dim=0) # auto1111 applies growing seeds for a batch length = latents.shape[0] @@ -259,17 +301,17 @@ class SeedNoiseGeneration: all_noises = [] # i starts at 0 for i in range(length): - generator = torch.manual_seed(seed+i+batch_offset) - all_noises.append(torch.randn(single_shape, dtype=latents.dtype, layout=latents.layout, generator=generator, device="cpu")) + generator, raw_device = get_generator(device, seed+i+batch_offset) + all_noises.append(torch.randn(single_shape, dtype=latents.dtype, layout=latents.layout, generator=generator, device=raw_device).to(device="cpu")) final_noise = torch.cat(all_noises, dim=0) # convert to derivative noise type, if needed - derivative_noise = SeedNoiseGeneration._create_derivative_noise(final_noise, noise_type=noise_type, seed=seed, extra_args=extra_args) + derivative_noise = SeedNoiseGeneration._create_derivative_noise(final_noise, noise_type=noise_type, seed=seed, extra_args=extra_args, device=device) if derivative_noise is not None: return derivative_noise return final_noise @staticmethod - def create_noise_individual_seeds(seeds: list[int], latents: Tensor, seed_offset: int=0, extra_args: dict={}): + def create_noise_individual_seeds(seeds: list[int], latents: Tensor, seed_offset: int=0, extra_args: dict={}, device=RandDevice.CPU): length = latents.shape[0] if len(seeds) < length: raise ValueError(f"{len(seeds)} seeds in seed_override were provided, but at least {length} are required to work with the current latents.") @@ -277,25 +319,25 @@ class SeedNoiseGeneration: single_shape = (1, latents.shape[1], latents.shape[2], latents.shape[3]) all_noises = [] for seed in seeds: - generator = torch.manual_seed(seed+seed_offset) - all_noises.append(torch.randn(single_shape, dtype=latents.dtype, layout=latents.layout, generator=generator, device="cpu")) + generator, raw_device = get_generator(device, seed+seed_offset) + all_noises.append(torch.randn(single_shape, dtype=latents.dtype, layout=latents.layout, generator=generator, device=raw_device).to(device="cpu")) return torch.cat(all_noises, dim=0) @staticmethod - def _create_common_noise(seed: int, latents: Tensor, noise_type: str=NoiseLayerType.DEFAULT, batch_offset: int=0, extra_args: dict={}): + def _create_common_noise(seed: int, latents: Tensor, noise_type: str=NoiseLayerType.DEFAULT, batch_offset: int=0, extra_args: dict={}, device=RandDevice.CPU): if noise_type == NoiseLayerType.EMPTY: return torch.zeros_like(latents) return None @staticmethod - def _create_derivative_noise(noise: Tensor, noise_type: str, seed: int, extra_args: dict): + def _create_derivative_noise(noise: Tensor, noise_type: str, seed: int, extra_args: dict, device=RandDevice.CPU): derivative_func = DERIVATIVE_NOISE_FUNC_MAP.get(noise_type, None) if derivative_func is None: return None - return derivative_func(noise=noise, seed=seed, extra_args=extra_args) + return derivative_func(noise=noise, seed=seed, extra_args=extra_args, device=device) @staticmethod - def _convert_to_repeated_context(noise: Tensor, extra_args: dict, **kwargs): + def _convert_to_repeated_context(noise: Tensor, extra_args: dict, device=RandDevice.CPU, **kwargs): # if no context_length, return unmodified noise opts: ContextOptionsGroup = extra_args["context_options"] context_length: int = opts.context_length if not opts.view_options else opts.view_options.context_length @@ -307,7 +349,7 @@ class SeedNoiseGeneration: return torch.cat([noise] * cat_count, dim=0)[:length] @staticmethod - def _convert_to_freenoise(noise: Tensor, seed: int, extra_args: dict, **kwargs): + def _convert_to_freenoise(noise: Tensor, seed: int, extra_args: dict, device=RandDevice.CPU, **kwargs): # if no context_length, return unmodified noise opts: ContextOptionsGroup = extra_args["context_options"] context_length: int = opts.context_length if not opts.view_options else opts.view_options.context_length @@ -316,7 +358,7 @@ class SeedNoiseGeneration: if context_length is None: return noise delta = context_length - context_overlap - generator = torch.manual_seed(seed) + generator, _ = get_generator(RandDevice.CPU, seed) # no point in ever using non-CPU to just shuffle indexes for start_idx in range(0, video_length-context_length, delta): # start_idx corresponds to the beginning of a context window From 109e07a23c7d6dcd7659e48ab0e0d0a8045a68dd Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Tue, 9 Jul 2024 16:26:51 -0500 Subject: [PATCH 3/9] Enabled cfg1 optimization for Custom CFG when current cfg is 1.0 --- animatediff/sampling.py | 10 +++++++--- 1 file changed, 7 insertions(+), 3 deletions(-) diff --git a/animatediff/sampling.py b/animatediff/sampling.py index fb83362..d6be33a 100644 --- a/animatediff/sampling.py +++ b/animatediff/sampling.py @@ -607,11 +607,15 @@ def evolved_sampling_function(model, x: Tensor, timestep: Tensor, uncond, cond, try: cond, uncond = ADGS.perform_special_model_features(model, [cond, uncond], x) - # never use cfg1 optimization if using custom_cfg (since can have timesteps and such) + # only use cfg1_optimization if not using custom_cfg or explicitly set to 1.0 + uncond_ = uncond if ADGS.sample_settings.custom_cfg is None and math.isclose(cond_scale, 1.0) and model_options.get("disable_cfg1_optimization", False) == False: uncond_ = None - else: - uncond_ = uncond + elif ADGS.sample_settings.custom_cfg is not None: + cfg_multival = ADGS.sample_settings.custom_cfg.cfg_multival + if type(cfg_multival) != Tensor and math.isclose(cfg_multival, 1.0) and model_options.get("disable_cfg1_optimization", False) == False: + uncond_ = None + del cfg_multival # add AD/evolved-sampling params to model_options (transformer_options) model_options = model_options.copy() From b7104b0da2aae600182830ae70b216f0ce87fe7f Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Tue, 9 Jul 2024 17:32:44 -0500 Subject: [PATCH 4/9] Refactored custom_cfg code to no longer use sampler_cfg_function patch, allowing that patch to not be overridden when using Custom CFG (some stuff might fail if it expects cond_scale to be a float instead of a tensor) --- animatediff/sample_settings.py | 9 +++++++++ animatediff/sampling.py | 4 ++-- 2 files changed, 11 insertions(+), 2 deletions(-) diff --git a/animatediff/sample_settings.py b/animatediff/sample_settings.py index 71086f4..a2c0764 100644 --- a/animatediff/sample_settings.py +++ b/animatediff/sample_settings.py @@ -583,7 +583,16 @@ class CustomCFGKeyframeGroup: # update steps current context is used self._current_used_steps += 1 + def get_cfg_scale(self, cond: Tensor): + cond_scale = self.cfg_multival + if isinstance(cond_scale, Tensor): + cond_scale = prepare_mask_batch(cond_scale.to(cond.dtype).to(cond.device), cond.shape) + cond_scale = extend_to_batch_size(cond_scale, cond.shape[0]) + return cond_scale + def patch_model(self, model: ModelPatcher) -> ModelPatcher: + # NOTE: no longer used at the moment, as most sampler_cfg_function patches should work with tensor cfg_scales, + # meaning get_cfg_scale is a direct replacement def evolved_custom_cfg(args): cond: Tensor = args["cond"] uncond: Tensor = args["uncond"] diff --git a/animatediff/sampling.py b/animatediff/sampling.py index d6be33a..ee573d4 100644 --- a/animatediff/sampling.py +++ b/animatediff/sampling.py @@ -414,8 +414,6 @@ def motion_sample_factory(orig_comfy_sample: Callable, is_custom: bool=False) -> cached_noise = None function_injections = FunctionInjectionHolder() try: - if model.sample_settings.custom_cfg is not None: - model = model.sample_settings.custom_cfg.patch_model(model) # clone params from model params = model.motion_injection_params.clone() # get amount of latents passed in, and store in params @@ -629,6 +627,8 @@ def evolved_sampling_function(model, x: Tensor, timestep: Tensor, uncond, cond, cond_pred, uncond_pred = sliding_calc_conds_batch(model, [cond, uncond_], x, timestep, model_options) if hasattr(comfy.samplers, "cfg_function"): + if ADGS.sample_settings.custom_cfg is not None: + cond_scale = ADGS.sample_settings.custom_cfg.get_cfg_scale(cond_pred) try: cached_calc_cond_batch = comfy.samplers.calc_cond_batch # support hooks and sliding context for PAG/other sampler_post_cfg_function tech that may use calc_cond_batch From 7ad604a30be6ff634d80af596f1a664981fa84ed Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Tue, 9 Jul 2024 17:56:20 -0500 Subject: [PATCH 5/9] Added PerturbedAttnGuide [Multival] node to allow multival inputs into PerturbedAttentionGuidance --- animatediff/nodes.py | 4 ++- animatediff/nodes_extras.py | 58 +++++++++++++++++++++++++++++++++++++ 2 files changed, 61 insertions(+), 1 deletion(-) diff --git a/animatediff/nodes.py b/animatediff/nodes.py index ff9c98f..7dd70b5 100644 --- a/animatediff/nodes.py +++ b/animatediff/nodes.py @@ -28,7 +28,7 @@ from .nodes_context import (LegacyLoopedUniformContextOptionsNode, LoopedUniform from .nodes_ad_settings import (AnimateDiffSettingsNode, ManualAdjustPENode, SweetspotStretchPENode, FullStretchPENode, WeightAdjustAllAddNode, WeightAdjustAllMultNode, WeightAdjustIndivAddNode, WeightAdjustIndivMultNode, WeightAdjustIndivAttnAddNode, WeightAdjustIndivAttnMultNode) -from .nodes_extras import AnimateDiffUnload, EmptyLatentImageLarge, CheckpointLoaderSimpleWithNoiseSelect +from .nodes_extras import AnimateDiffUnload, EmptyLatentImageLarge, CheckpointLoaderSimpleWithNoiseSelect, PerturbedAttentionGuidanceMultival from .nodes_deprecated import (AnimateDiffLoader_Deprecated, AnimateDiffLoaderAdvanced_Deprecated, AnimateDiffCombine_Deprecated, AnimateDiffModelSettings, AnimateDiffModelSettingsSimple, AnimateDiffModelSettingsAdvanced, AnimateDiffModelSettingsAdvancedAttnStrengths) from .nodes_lora import AnimateDiffLoraLoader @@ -117,6 +117,7 @@ NODE_CLASS_MAPPINGS = { "ADE_AnimateDiffUnload": AnimateDiffUnload, "ADE_EmptyLatentImageLarge": EmptyLatentImageLarge, "CheckpointLoaderSimpleWithNoiseSelect": CheckpointLoaderSimpleWithNoiseSelect, + "ADE_PerturbedAttentionGuidanceMultival": PerturbedAttentionGuidanceMultival, # Gen1 Nodes "ADE_AnimateDiffLoaderGen1": AnimateDiffLoaderGen1, "ADE_AnimateDiffLoaderWithContext": LegacyAnimateDiffLoaderWithContext, @@ -233,6 +234,7 @@ NODE_DISPLAY_NAME_MAPPINGS = { "ADE_AnimateDiffUnload": "AnimateDiff Unload πŸŽ­πŸ…πŸ…“", "ADE_EmptyLatentImageLarge": "Empty Latent Image (Big Batch) πŸŽ­πŸ…πŸ…“", "CheckpointLoaderSimpleWithNoiseSelect": "Load Checkpoint w/ Noise Select πŸŽ­πŸ…πŸ…“", + "ADE_PerturbedAttentionGuidanceMultival": "PerturbedAttnGuide [Multival] πŸŽ­πŸ…πŸ…“", # Gen1 Nodes "ADE_AnimateDiffLoaderGen1": "AnimateDiff Loader πŸŽ­πŸ…πŸ…“β‘ ", "ADE_AnimateDiffLoaderWithContext": "AnimateDiff Loader [Legacy] πŸŽ­πŸ…πŸ…“β‘ ", diff --git a/animatediff/nodes_extras.py b/animatediff/nodes_extras.py index 9b22534..e54dd25 100644 --- a/animatediff/nodes_extras.py +++ b/animatediff/nodes_extras.py @@ -1,12 +1,18 @@ +from typing import Union + import torch +from torch import Tensor import folder_paths import nodes as comfy_nodes from comfy.model_patcher import ModelPatcher +import comfy.model_patcher +import comfy.samplers from comfy.sd import load_checkpoint_guess_config from .logger import logger from .utils_model import BetaSchedules +from .utils_motion import extend_to_batch_size, prepare_mask_batch from .model_injection import get_vanilla_model_patcher @@ -76,3 +82,55 @@ class EmptyLatentImageLarge: def generate(self, width, height, batch_size=1): latent = torch.zeros([batch_size, 4, height // 8, width // 8]) return ({"samples":latent}, ) + + +# this is a modified copy of PerturbedAttentionGuidance node from comfy_extras/nodes_pag.py +class PerturbedAttentionGuidanceMultival: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "model": ("MODEL",), + "scale_multival": ("MULTIVAL",), + } + } + + RETURN_TYPES = ("MODEL",) + FUNCTION = "patch" + + CATEGORY = "Animate Diff πŸŽ­πŸ…πŸ…“/extras" + + def patch(self, model: ModelPatcher, scale_multival: Union[float, Tensor]): + unet_block = "middle" + unet_block_id = 0 + m = model.clone() + + def perturbed_attention(q, k, v, extra_options, mask=None): + return v + + def post_cfg_function(args): + model = args["model"] + cond_pred = args["cond_denoised"] + cond = args["cond"] + cfg_result = args["denoised"] + sigma = args["sigma"] + model_options = args["model_options"].copy() + x = args["input"] + + if type(scale_multival) != Tensor and scale_multival == 0: + return cfg_result + + scale = scale_multival + if isinstance(scale, Tensor): + scale = prepare_mask_batch(scale.to(cond_pred.dtype).to(cond_pred.device), cond_pred.shape) + scale = extend_to_batch_size(scale, cond_pred.shape[0]) + + # Replace Self-attention with PAG + model_options = comfy.model_patcher.set_model_options_patch_replace(model_options, perturbed_attention, "attn1", unet_block, unet_block_id) + (pag,) = comfy.samplers.calc_cond_batch(model, [cond], x, sigma, model_options) + + return cfg_result + (cond_pred - pag) * scale + + m.set_model_sampler_post_cfg_function(post_cfg_function) + + return (m,) From 7e8f521f89d248c795587de58958cbbe4245a12f Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Tue, 9 Jul 2024 20:47:16 -0500 Subject: [PATCH 6/9] Added cfg_extras to Custom CFG nodes to allow cfg patch scheduling (and stacking) --- animatediff/cfg_extras.py | 91 ++++++++++++++++++++++ animatediff/nodes.py | 21 ++++-- animatediff/nodes_extras.py | 54 ++++++------- animatediff/nodes_sample.py | 133 ++++++++++++++++++++++++++++++--- animatediff/sample_settings.py | 39 +++++++++- animatediff/sampling.py | 1 + 6 files changed, 292 insertions(+), 47 deletions(-) create mode 100644 animatediff/cfg_extras.py diff --git a/animatediff/cfg_extras.py b/animatediff/cfg_extras.py new file mode 100644 index 0000000..f2b3d47 --- /dev/null +++ b/animatediff/cfg_extras.py @@ -0,0 +1,91 @@ +from typing import Union + +import inspect +import torch +from torch import Tensor + +import comfy.model_patcher +import comfy.samplers + +from .utils_motion import extend_to_batch_size, prepare_mask_batch + + +################################################################################ +# helpers for modifying model_options to apply cfg function patches; +# taken from comfy/model_patcher.py +def set_model_options_sampler_cfg_function(model_options: dict[str], sampler_cfg_function, disable_cfg1_optimization=False): + if len(inspect.signature(sampler_cfg_function).parameters) == 3: + model_options["sampler_cfg_function"] = lambda args: sampler_cfg_function(args["cond"], args["uncond"], args["cond_scale"]) #Old way + else: + model_options["sampler_cfg_function"] = sampler_cfg_function + if disable_cfg1_optimization: + model_options["disable_cfg1_optimization"] = True + return model_options +#------------------------------------------------------------------------------- + + +# this is a modified version of PerturbedAttentionGuidance from comfy_extras/nodes_pag.py +def perturbed_attention_guidance_patch(scale_multival: Union[float, Tensor]): + unet_block = "middle" + unet_block_id = 0 + + def perturbed_attention(q, k, v, extra_options, mask=None): + return v + + def post_cfg_function(args): + model = args["model"] + cond_pred: Tensor = args["cond_denoised"] + cond = args["cond"] + cfg_result = args["denoised"] + sigma = args["sigma"] + model_options = args["model_options"].copy() + x = args["input"] + + if type(scale_multival) != Tensor and scale_multival == 0: + return cfg_result + + scale = scale_multival + if isinstance(scale, Tensor): + scale = prepare_mask_batch(scale.to(cond_pred.dtype).to(cond_pred.device), cond_pred.shape) + scale = extend_to_batch_size(scale, cond_pred.shape[0]) + + # Replace Self-attention with PAG + model_options = comfy.model_patcher.set_model_options_patch_replace(model_options, perturbed_attention, "attn1", unet_block, unet_block_id) + (pag,) = comfy.samplers.calc_cond_batch(model, [cond], x, sigma, model_options) + + return cfg_result + (cond_pred - pag) * scale + + return post_cfg_function + + +# this is a modified version of RescaleCFG from comfy_extras/nodes_model_advanced.py +def rescale_cfg_patch(multiplier_multival: Union[float, Tensor]): + def cfg_function(args): + cond: Tensor = args["cond"] + uncond = args["uncond"] + cond_scale = args["cond_scale"] + sigma = args["sigma"] + sigma = sigma.view(sigma.shape[:1] + (1,) * (cond.ndim - 1)) + x_orig = args["input"] + + #rescale cfg has to be done on v-pred model output + x = x_orig / (sigma * sigma + 1.0) + cond = ((x - (x_orig - cond)) * (sigma ** 2 + 1.0) ** 0.5) / (sigma) + uncond = ((x - (x_orig - uncond)) * (sigma ** 2 + 1.0) ** 0.5) / (sigma) + + #rescalecfg + x_cfg = uncond + cond_scale * (cond - uncond) + ro_pos = torch.std(cond, dim=(1,2,3), keepdim=True) + ro_cfg = torch.std(x_cfg, dim=(1,2,3), keepdim=True) + + multiplier = multiplier_multival + if isinstance(multiplier, Tensor): + multiplier = prepare_mask_batch(multiplier.to(cond.dtype).to(cond.device), cond.shape) + multiplier = extend_to_batch_size(multiplier, cond.shape[0]) + + x_rescaled = x_cfg * (ro_pos / ro_cfg) + x_final = multiplier * x_rescaled + (1.0 - multiplier) * x_cfg + + return x_orig - (x - x_final * sigma / (sigma * sigma + 1.0) ** 0.5) + + return cfg_function diff --git a/animatediff/nodes.py b/animatediff/nodes.py index 7dd70b5..4481aca 100644 --- a/animatediff/nodes.py +++ b/animatediff/nodes.py @@ -21,6 +21,7 @@ from .nodes_conditioning import (MaskableLoraLoader, MaskableLoraLoaderModelOnly CreateLoraHookKeyframe, CreateLoraHookKeyframeInterpolation, CreateLoraHookKeyframeFromStrengthList) from .nodes_sample import (FreeInitOptionsNode, NoiseLayerAddWeightedNode, SampleSettingsNode, NoiseLayerAddNode, NoiseLayerReplaceNode, IterationOptionsNode, CustomCFGNode, CustomCFGSimpleNode, CustomCFGKeyframeNode, CustomCFGKeyframeSimpleNode, + CFGExtrasPAGNode, CFGExtrasPAGSimpleNode, CFGExtrasRescaleCFGNode, CFGExtrasRescaleCFGSimpleNode, NoisedImageInjectionNode, NoisedImageInjectOptionsNode) from .nodes_sigma_schedule import (SigmaScheduleNode, RawSigmaScheduleNode, WeightedAverageSigmaScheduleNode, InterpolatedWeightedAverageSigmaScheduleNode, SplitAndCombineSigmaScheduleNode) from .nodes_context import (LegacyLoopedUniformContextOptionsNode, LoopedUniformContextOptionsNode, LoopedUniformViewOptionsNode, StandardUniformContextOptionsNode, StandardStaticContextOptionsNode, BatchedContextOptionsNode, @@ -28,7 +29,7 @@ from .nodes_context import (LegacyLoopedUniformContextOptionsNode, LoopedUniform from .nodes_ad_settings import (AnimateDiffSettingsNode, ManualAdjustPENode, SweetspotStretchPENode, FullStretchPENode, WeightAdjustAllAddNode, WeightAdjustAllMultNode, WeightAdjustIndivAddNode, WeightAdjustIndivMultNode, WeightAdjustIndivAttnAddNode, WeightAdjustIndivAttnMultNode) -from .nodes_extras import AnimateDiffUnload, EmptyLatentImageLarge, CheckpointLoaderSimpleWithNoiseSelect, PerturbedAttentionGuidanceMultival +from .nodes_extras import AnimateDiffUnload, EmptyLatentImageLarge, CheckpointLoaderSimpleWithNoiseSelect, PerturbedAttentionGuidanceMultival, RescaleCFGMultival from .nodes_deprecated import (AnimateDiffLoader_Deprecated, AnimateDiffLoaderAdvanced_Deprecated, AnimateDiffCombine_Deprecated, AnimateDiffModelSettings, AnimateDiffModelSettingsSimple, AnimateDiffModelSettingsAdvanced, AnimateDiffModelSettingsAdvancedAttnStrengths) from .nodes_lora import AnimateDiffLoraLoader @@ -113,11 +114,16 @@ NODE_CLASS_MAPPINGS = { "ADE_SigmaScheduleSplitAndCombine": SplitAndCombineSigmaScheduleNode, "ADE_NoisedImageInjection": NoisedImageInjectionNode, "ADE_NoisedImageInjectOptions": NoisedImageInjectOptionsNode, + "ADE_CFGExtrasPAGSimple": CFGExtrasPAGSimpleNode, + "ADE_CFGExtrasPAG": CFGExtrasPAGNode, + "ADE_CFGExtrasRescaleCFGSimple": CFGExtrasRescaleCFGSimpleNode, + "ADE_CFGExtrasRescaleCFG": CFGExtrasRescaleCFGNode, # Extras Nodes "ADE_AnimateDiffUnload": AnimateDiffUnload, "ADE_EmptyLatentImageLarge": EmptyLatentImageLarge, "CheckpointLoaderSimpleWithNoiseSelect": CheckpointLoaderSimpleWithNoiseSelect, "ADE_PerturbedAttentionGuidanceMultival": PerturbedAttentionGuidanceMultival, + "ADE_RescaleCFGMultival": RescaleCFGMultival, # Gen1 Nodes "ADE_AnimateDiffLoaderGen1": AnimateDiffLoaderGen1, "ADE_AnimateDiffLoaderWithContext": LegacyAnimateDiffLoaderWithContext, @@ -163,8 +169,8 @@ NODE_DISPLAY_NAME_MAPPINGS = { "ADE_AnimateDiffSamplingSettings": "Sample Settings πŸŽ­πŸ…πŸ…“", "ADE_AnimateDiffKeyframe": "AnimateDiff Keyframe πŸŽ­πŸ…πŸ…“", # Multival Nodes - "ADE_MultivalDynamic": "Multival Dynamic πŸŽ­πŸ…πŸ…“", - "ADE_MultivalDynamicFloatInput": "Multival Dynamic [Float List] πŸŽ­πŸ…πŸ…“", + "ADE_MultivalDynamic": "Multival πŸŽ­πŸ…πŸ…“", + "ADE_MultivalDynamicFloatInput": "Multival [Float List] πŸŽ­πŸ…πŸ…“", "ADE_MultivalScaledMask": "Multival Scaled Mask πŸŽ­πŸ…πŸ…“", "ADE_MultivalConvertToMask": "Multival to Mask πŸŽ­πŸ…πŸ…“", # Context Opts @@ -219,9 +225,9 @@ NODE_DISPLAY_NAME_MAPPINGS = { "ADE_AdjustWeightIndivAttnAdd": "Adjust Weight [Indiv-Attnβ—†Add] πŸŽ­πŸ…πŸ…“", "ADE_AdjustWeightIndivAttnMult": "Adjust Weight [Indiv-Attnβ—†Mult] πŸŽ­πŸ…πŸ…“", # Sample Settings - "ADE_CustomCFGSimple": "Custom CFG [Simple] πŸŽ­πŸ…πŸ…“", + "ADE_CustomCFGSimple": "Custom CFG πŸŽ­πŸ…πŸ…“", "ADE_CustomCFG": "Custom CFG [Multival] πŸŽ­πŸ…πŸ…“", - "ADE_CustomCFGKeyframeSimple": "Custom CFG Keyframe [Simple] πŸŽ­πŸ…πŸ…“", + "ADE_CustomCFGKeyframeSimple": "Custom CFG Keyframe πŸŽ­πŸ…πŸ…“", "ADE_CustomCFGKeyframe": "Custom CFG Keyframe [Multival] πŸŽ­πŸ…πŸ…“", "ADE_SigmaSchedule": "Create Sigma Schedule πŸŽ­πŸ…πŸ…“", "ADE_RawSigmaSchedule": "Create Raw Sigma Schedule πŸŽ­πŸ…πŸ…“", @@ -230,11 +236,16 @@ NODE_DISPLAY_NAME_MAPPINGS = { "ADE_SigmaScheduleSplitAndCombine": "Sigma Schedule Split Combine πŸŽ­πŸ…πŸ…“", "ADE_NoisedImageInjection": "Image Injection πŸŽ­πŸ…πŸ…“", "ADE_NoisedImageInjectOptions": "Image Injection Options πŸŽ­πŸ…πŸ…“", + "ADE_CFGExtrasPAGSimple": "CFG Extrasβ—†PAG πŸŽ­πŸ…πŸ…“", + "ADE_CFGExtrasPAG": "CFG Extrasβ—†PAG [Multival] πŸŽ­πŸ…πŸ…“", + "ADE_CFGExtrasRescaleCFGSimple": "CFG Extrasβ—†RescaleCFG πŸŽ­πŸ…πŸ…“", + "ADE_CFGExtrasRescaleCFG": "CFG Extrasβ—†RescaleCFG [Multival] πŸŽ­πŸ…πŸ…“", # Extras Nodes "ADE_AnimateDiffUnload": "AnimateDiff Unload πŸŽ­πŸ…πŸ…“", "ADE_EmptyLatentImageLarge": "Empty Latent Image (Big Batch) πŸŽ­πŸ…πŸ…“", "CheckpointLoaderSimpleWithNoiseSelect": "Load Checkpoint w/ Noise Select πŸŽ­πŸ…πŸ…“", "ADE_PerturbedAttentionGuidanceMultival": "PerturbedAttnGuide [Multival] πŸŽ­πŸ…πŸ…“", + "ADE_RescaleCFGMultival": "RescaleCFG [Multival] πŸŽ­πŸ…πŸ…“", # Gen1 Nodes "ADE_AnimateDiffLoaderGen1": "AnimateDiff Loader πŸŽ­πŸ…πŸ…“β‘ ", "ADE_AnimateDiffLoaderWithContext": "AnimateDiff Loader [Legacy] πŸŽ­πŸ…πŸ…“β‘ ", diff --git a/animatediff/nodes_extras.py b/animatediff/nodes_extras.py index e54dd25..9be4b8e 100644 --- a/animatediff/nodes_extras.py +++ b/animatediff/nodes_extras.py @@ -14,6 +14,7 @@ from .logger import logger from .utils_model import BetaSchedules from .utils_motion import extend_to_batch_size, prepare_mask_batch from .model_injection import get_vanilla_model_patcher +from .cfg_extras import perturbed_attention_guidance_patch, rescale_cfg_patch class AnimateDiffUnload: @@ -84,7 +85,6 @@ class EmptyLatentImageLarge: return ({"samples":latent}, ) -# this is a modified copy of PerturbedAttentionGuidance node from comfy_extras/nodes_pag.py class PerturbedAttentionGuidanceMultival: @classmethod def INPUT_TYPES(s): @@ -101,36 +101,28 @@ class PerturbedAttentionGuidanceMultival: CATEGORY = "Animate Diff πŸŽ­πŸ…πŸ…“/extras" def patch(self, model: ModelPatcher, scale_multival: Union[float, Tensor]): - unet_block = "middle" - unet_block_id = 0 m = model.clone() - - def perturbed_attention(q, k, v, extra_options, mask=None): - return v - - def post_cfg_function(args): - model = args["model"] - cond_pred = args["cond_denoised"] - cond = args["cond"] - cfg_result = args["denoised"] - sigma = args["sigma"] - model_options = args["model_options"].copy() - x = args["input"] - - if type(scale_multival) != Tensor and scale_multival == 0: - return cfg_result - - scale = scale_multival - if isinstance(scale, Tensor): - scale = prepare_mask_batch(scale.to(cond_pred.dtype).to(cond_pred.device), cond_pred.shape) - scale = extend_to_batch_size(scale, cond_pred.shape[0]) - - # Replace Self-attention with PAG - model_options = comfy.model_patcher.set_model_options_patch_replace(model_options, perturbed_attention, "attn1", unet_block, unet_block_id) - (pag,) = comfy.samplers.calc_cond_batch(model, [cond], x, sigma, model_options) - - return cfg_result + (cond_pred - pag) * scale - - m.set_model_sampler_post_cfg_function(post_cfg_function) + m.set_model_sampler_post_cfg_function(perturbed_attention_guidance_patch(scale_multival)) return (m,) + + +class RescaleCFGMultival: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "model": ("MODEL",), + "mult_multival": ("MULTIVAL",), + } + } + + RETURN_TYPES = ("MODEL",) + FUNCTION = "patch" + + CATEGORY = "Animate Diff πŸŽ­πŸ…πŸ…“/extras" + + def patch(self, model: ModelPatcher, mult_multival: Union[float, Tensor]): + m = model.clone() + m.set_model_sampler_cfg_function(rescale_cfg_patch(mult_multival)) + return (m, ) diff --git a/animatediff/nodes_sample.py b/animatediff/nodes_sample.py index 3fb75b3..1c5a143 100644 --- a/animatediff/nodes_sample.py +++ b/animatediff/nodes_sample.py @@ -2,13 +2,16 @@ from typing import Union from torch import Tensor from comfy.sd import VAE +from comfy.model_patcher import set_model_options_post_cfg_function from .freeinit import FreeInitFilter from .sample_settings import (FreeInitOptions, IterationOptions, NoiseLayerAdd, NoiseLayerAddWeighted, NoiseLayerGroup, NoiseLayerReplace, NoiseLayerType, - SeedNoiseGeneration, SampleSettings, CustomCFGKeyframeGroup, CustomCFGKeyframe, + SeedNoiseGeneration, SampleSettings, + CustomCFGKeyframeGroup, CustomCFGKeyframe, CFGExtrasGroup, CFGExtras, NoisedImageToInjectGroup, NoisedImageToInject, NoisedImageInjectOptions) from .utils_model import BIGMIN, BIGMAX, MAX_RESOLUTION, SigmaSchedule +from .cfg_extras import perturbed_attention_guidance_patch, rescale_cfg_patch, set_model_options_sampler_cfg_function class SampleSettingsNode: @@ -217,6 +220,9 @@ class CustomCFGNode: return { "required": { "cfg_multival": ("MULTIVAL",), + }, + "optional": { + "cfg_extras": ("CFG_EXTRAS",), } } @@ -224,8 +230,8 @@ class CustomCFGNode: CATEGORY = "Animate Diff πŸŽ­πŸ…πŸ…“/sample settings" FUNCTION = "create_custom_cfg" - def create_custom_cfg(self, cfg_multival: Union[float, Tensor]): - keyframe = CustomCFGKeyframe(cfg_multival=cfg_multival) + def create_custom_cfg(self, cfg_multival: Union[float, Tensor], cfg_extras: CFGExtrasGroup=None): + keyframe = CustomCFGKeyframe(cfg_multival=cfg_multival, cfg_extras=cfg_extras) cfg_custom = CustomCFGKeyframeGroup() cfg_custom.add(keyframe) return (cfg_custom,) @@ -238,14 +244,17 @@ class CustomCFGSimpleNode: "required": { "cfg": ("FLOAT", {"default": 8.0, "min": 0.0, "max": 100.0, "step": 0.1}), }, + "optional": { + "cfg_extras": ("CFG_EXTRAS",), + } } RETURN_TYPES = ("CUSTOM_CFG",) CATEGORY = "Animate Diff πŸŽ­πŸ…πŸ…“/sample settings" FUNCTION = "create_custom_cfg" - def create_custom_cfg(self, cfg: float): - return CustomCFGNode.create_custom_cfg(self, cfg_multival=cfg) + def create_custom_cfg(self, cfg: float, cfg_extras: CFGExtrasGroup=None): + return CustomCFGNode.create_custom_cfg(self, cfg_multival=cfg, cfg_extras=cfg_extras) class CustomCFGKeyframeNode: @@ -259,6 +268,7 @@ class CustomCFGKeyframeNode: }, "optional": { "prev_custom_cfg": ("CUSTOM_CFG",), + "cfg_extras": ("CFG_EXTRAS",), } } @@ -267,11 +277,11 @@ class CustomCFGKeyframeNode: FUNCTION = "create_custom_cfg" def create_custom_cfg(self, cfg_multival: Union[float, Tensor], start_percent: float=0.0, guarantee_steps: int=1, - prev_custom_cfg: CustomCFGKeyframeGroup=None): + prev_custom_cfg: CustomCFGKeyframeGroup=None, cfg_extras: CFGExtrasGroup=None): if not prev_custom_cfg: prev_custom_cfg = CustomCFGKeyframeGroup() prev_custom_cfg = prev_custom_cfg.clone() - keyframe = CustomCFGKeyframe(cfg_multival=cfg_multival, start_percent=start_percent, guarantee_steps=guarantee_steps) + keyframe = CustomCFGKeyframe(cfg_multival=cfg_multival, start_percent=start_percent, guarantee_steps=guarantee_steps, cfg_extras=cfg_extras) prev_custom_cfg.add(keyframe) return (prev_custom_cfg,) @@ -287,6 +297,7 @@ class CustomCFGKeyframeSimpleNode: }, "optional": { "prev_custom_cfg": ("CUSTOM_CFG",), + "cfg_extras": ("CFG_EXTRAS",), } } @@ -295,9 +306,113 @@ class CustomCFGKeyframeSimpleNode: FUNCTION = "create_custom_cfg" def create_custom_cfg(self, cfg: float, start_percent: float=0.0, guarantee_steps: int=1, - prev_custom_cfg: CustomCFGKeyframeGroup=None): + prev_custom_cfg: CustomCFGKeyframeGroup=None, cfg_extras: CFGExtrasGroup=None): return CustomCFGKeyframeNode.create_custom_cfg(self, cfg_multival=cfg, start_percent=start_percent, - guarantee_steps=guarantee_steps, prev_custom_cfg=prev_custom_cfg) + guarantee_steps=guarantee_steps, prev_custom_cfg=prev_custom_cfg, cfg_extras=cfg_extras) + + +class CFGExtrasPAGNode: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "scale_multival": ("MULTIVAL",), + }, + "optional": { + "prev_extras": ("CFG_EXTRAS",), + } + } + + RETURN_TYPES = ("CFG_EXTRAS",) + FUNCTION = "add_cfg_extras" + + CATEGORY = "Animate Diff πŸŽ­πŸ…πŸ…“/sample settings/cfg extras" + + def add_cfg_extras(self, scale_multival: Union[float, Tensor], prev_extras: CFGExtrasGroup=None): + if prev_extras is None: + prev_extras = CFGExtrasGroup() + prev_extras = prev_extras.clone() + + patch = perturbed_attention_guidance_patch(scale_multival) + def call_extras(model_options: dict[str]): + return set_model_options_post_cfg_function(model_options.copy(), patch) + + extra = CFGExtras(call_extras) + prev_extras.add(extra) + return (prev_extras,) + + +class CFGExtrasPAGSimpleNode: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "scale": ("FLOAT", {"default": 3.0, "min": 0.0, "max": 100.0, "step": 0.1, "round": 0.01}), + }, + "optional": { + "prev_extras": ("CFG_EXTRAS",), + } + } + + RETURN_TYPES = ("CFG_EXTRAS",) + FUNCTION = "add_cfg_extras" + + CATEGORY = "Animate Diff πŸŽ­πŸ…πŸ…“/sample settings/cfg extras" + + def add_cfg_extras(self, scale: float, prev_extras: CFGExtrasGroup=None): + return CFGExtrasPAGNode.add_cfg_extras(self, scale_multival=scale, prev_extras=prev_extras) + + +class CFGExtrasRescaleCFGNode: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "mult_multival": ("MULTIVAL",), + }, + "optional": { + "prev_extras": ("CFG_EXTRAS",), + } + } + + RETURN_TYPES = ("CFG_EXTRAS",) + FUNCTION = "add_cfg_extras" + + CATEGORY = "Animate Diff πŸŽ­πŸ…πŸ…“/sample settings/cfg extras" + + def add_cfg_extras(self, mult_multival: Union[float, Tensor], prev_extras: CFGExtrasGroup=None): + if prev_extras is None: + prev_extras = CFGExtrasGroup() + prev_extras = prev_extras.clone() + + patch = rescale_cfg_patch(mult_multival) + def call_extras(model_options: dict[str]): + return set_model_options_sampler_cfg_function(model_options.copy(), patch) + + extra = CFGExtras(call_extras) + prev_extras.add(extra) + return (prev_extras,) + + +class CFGExtrasRescaleCFGSimpleNode: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "multiplier": ("FLOAT", {"default": 0.7, "min": 0.0, "max": 1.0, "step": 0.01}), + }, + "optional": { + "prev_extras": ("CFG_EXTRAS",), + } + } + + RETURN_TYPES = ("CFG_EXTRAS",) + FUNCTION = "add_cfg_extras" + + CATEGORY = "Animate Diff πŸŽ­πŸ…πŸ…“/sample settings/cfg extras" + + def add_cfg_extras(self, multiplier: float, prev_extras: CFGExtrasGroup=None): + return CFGExtrasRescaleCFGNode.add_cfg_extras(self, mult_multival=multiplier, prev_extras=prev_extras) class NoisedImageInjectionNode: diff --git a/animatediff/sample_settings.py b/animatediff/sample_settings.py index a2c0764..e14dc72 100644 --- a/animatediff/sample_settings.py +++ b/animatediff/sample_settings.py @@ -1,5 +1,5 @@ from collections.abc import Iterable -from typing import Union +from typing import Union, Callable import torch from torch import Tensor @@ -503,9 +503,31 @@ class FreeInitOptions(IterationOptions): raise ValueError(f"FreeInit init_type '{self.init_type}' is not recognized.") +class CFGExtras: + def __init__(self, call_fn: Callable): + self.call_fn = call_fn + + +class CFGExtrasGroup: + def __init__(self): + self.extras: list[CFGExtras] = [] + + def add(self, extra: CFGExtras): + self.extras.append(extra) + + def is_empty(self) -> bool: + return len(self.extras) == 0 + + def clone(self): + cloned = CFGExtrasGroup() + cloned.extras = self.extras.copy() + return cloned + + class CustomCFGKeyframe: - def __init__(self, cfg_multival: Union[float, Tensor], start_percent=0.0, guarantee_steps=1): + def __init__(self, cfg_multival: Union[float, Tensor], start_percent=0.0, guarantee_steps=1, cfg_extras: CFGExtrasGroup=None): self.cfg_multival = cfg_multival + self.cfg_extras = cfg_extras # scheduling self.start_percent = float(start_percent) self.start_t = 999999999.9 @@ -589,6 +611,13 @@ class CustomCFGKeyframeGroup: cond_scale = prepare_mask_batch(cond_scale.to(cond.dtype).to(cond.device), cond.shape) cond_scale = extend_to_batch_size(cond_scale, cond.shape[0]) return cond_scale + + def get_model_options(self, model_options: dict[str]): + cfg_extras = self.cfg_extras + if cfg_extras is not None: + for extra in cfg_extras.extras: + model_options = extra.call_fn(model_options) + return model_options def patch_model(self, model: ModelPatcher) -> ModelPatcher: # NOTE: no longer used at the moment, as most sampler_cfg_function patches should work with tensor cfg_scales, @@ -613,6 +642,12 @@ class CustomCFGKeyframeGroup: if self._current_keyframe != None: return self._current_keyframe.cfg_multival return None + + @property + def cfg_extras(self): + if self._current_keyframe != None: + return self._current_keyframe.cfg_extras + return None class NoisedImageInjectOptions: diff --git a/animatediff/sampling.py b/animatediff/sampling.py index ee573d4..02a2267 100644 --- a/animatediff/sampling.py +++ b/animatediff/sampling.py @@ -629,6 +629,7 @@ def evolved_sampling_function(model, x: Tensor, timestep: Tensor, uncond, cond, if hasattr(comfy.samplers, "cfg_function"): if ADGS.sample_settings.custom_cfg is not None: cond_scale = ADGS.sample_settings.custom_cfg.get_cfg_scale(cond_pred) + model_options = ADGS.sample_settings.custom_cfg.get_model_options(model_options) try: cached_calc_cond_batch = comfy.samplers.calc_cond_batch # support hooks and sliding context for PAG/other sampler_post_cfg_function tech that may use calc_cond_batch From 583c61cb662221ca0d057e355330e1b12924fcee Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Tue, 9 Jul 2024 21:02:58 -0500 Subject: [PATCH 7/9] Fixed deprecation prop on ancient animatediff settings nodes --- animatediff/nodes_deprecated.py | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/animatediff/nodes_deprecated.py b/animatediff/nodes_deprecated.py index f96f538..7b6759c 100644 --- a/animatediff/nodes_deprecated.py +++ b/animatediff/nodes_deprecated.py @@ -292,7 +292,7 @@ class AnimateDiffModelSettings: }, "optional": { "mask_motion_scale": ("MASK",), - "optional": {"deprecation_warning": ("ADEWARN", {"text": "Deprecated"})}, + "deprecation_warning": ("ADEWARN", {"text": "Deprecated"}), } } @@ -321,7 +321,7 @@ class AnimateDiffModelSettingsSimple: "mask_motion_scale": ("MASK",), "min_motion_scale": ("FLOAT", {"default": 1.0, "min": 0.0, "step": 0.001}), "max_motion_scale": ("FLOAT", {"default": 1.0, "min": 0.0, "step": 0.001}), - "optional": {"deprecation_warning": ("ADEWARN", {"text": "Deprecated"})}, + "deprecation_warning": ("ADEWARN", {"text": "Deprecated"}), } } @@ -360,7 +360,7 @@ class AnimateDiffModelSettingsAdvanced: "mask_motion_scale": ("MASK",), "min_motion_scale": ("FLOAT", {"default": 1.0, "min": 0.0, "step": 0.001}), "max_motion_scale": ("FLOAT", {"default": 1.0, "min": 0.0, "step": 0.001}), - "optional": {"deprecation_warning": ("ADEWARN", {"text": "Deprecated"})}, + "deprecation_warning": ("ADEWARN", {"text": "Deprecated"}), } } @@ -415,7 +415,7 @@ class AnimateDiffModelSettingsAdvancedAttnStrengths: "mask_motion_scale": ("MASK",), "min_motion_scale": ("FLOAT", {"default": 1.0, "min": 0.0, "step": 0.001}), "max_motion_scale": ("FLOAT", {"default": 1.0, "min": 0.0, "step": 0.001}), - "optional": {"deprecation_warning": ("ADEWARN", {"text": "Deprecated"})}, + "deprecation_warning": ("ADEWARN", {"text": "Deprecated"}), } } From 35c8fdb8073dafd5bdd5ee5c9bd37aafc4c495ff Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Tue, 9 Jul 2024 21:07:31 -0500 Subject: [PATCH 8/9] Hid Visualize Context Options node since I am still working on it --- animatediff/nodes.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/animatediff/nodes.py b/animatediff/nodes.py index 4481aca..363adc5 100644 --- a/animatediff/nodes.py +++ b/animatediff/nodes.py @@ -58,7 +58,7 @@ NODE_CLASS_MAPPINGS = { "ADE_ViewsOnlyContextOptions": ViewAsContextOptionsNode, "ADE_BatchedContextOptions": BatchedContextOptionsNode, "ADE_AnimateDiffUniformContextOptions": LegacyLoopedUniformContextOptionsNode, # Legacy - "ADE_VisualizeContextOptions": VisualizeContextOptionsInt, + #"ADE_VisualizeContextOptions": VisualizeContextOptionsInt, # View Opts "ADE_StandardStaticViewOptions": StandardStaticViewOptionsNode, "ADE_StandardUniformViewOptions": StandardUniformViewOptionsNode, From b74d56c46039ae0ef1eeae0c86360427701a111d Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Tue, 9 Jul 2024 21:07:49 -0500 Subject: [PATCH 9/9] version bump --- pyproject.toml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/pyproject.toml b/pyproject.toml index 7b8d7a4..5bd4395 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,7 +1,7 @@ [project] name = "comfyui-animatediff-evolved" description = "Improved AnimateDiff integration for ComfyUI." -version = "1.0.8" +version = "1.0.9" license = "LICENSE" dependencies = []