Merge PR #422 from Kosinkadink/develop - Custom CFG Improvements + GPU Noise
Custom CFG Improvements + GPU Noise
This commit is contained in:
@@ -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
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
+27
-7
@@ -20,14 +20,16 @@ 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,
|
||||
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,
|
||||
StandardStaticViewOptionsNode, StandardUniformViewOptionsNode, ViewAsContextOptionsNode)
|
||||
StandardStaticViewOptionsNode, StandardUniformViewOptionsNode, ViewAsContextOptionsNode, VisualizeContextOptionsInt)
|
||||
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, RescaleCFGMultival
|
||||
from .nodes_deprecated import (AnimateDiffLoader_Deprecated, AnimateDiffLoaderAdvanced_Deprecated, AnimateDiffCombine_Deprecated,
|
||||
AnimateDiffModelSettings, AnimateDiffModelSettingsSimple, AnimateDiffModelSettingsAdvanced, AnimateDiffModelSettingsAdvancedAttnStrengths)
|
||||
from .nodes_lora import AnimateDiffLoraLoader
|
||||
@@ -56,6 +58,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 +103,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,
|
||||
@@ -109,10 +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,
|
||||
@@ -158,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
|
||||
@@ -169,6 +180,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 +225,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 🎭🅐🅓",
|
||||
"ADE_CustomCFG": "Custom CFG [Multival] 🎭🅐🅓",
|
||||
"ADE_CustomCFGKeyframeSimple": "Custom CFG Keyframe 🎭🅐🅓",
|
||||
"ADE_CustomCFGKeyframe": "Custom CFG Keyframe [Multival] 🎭🅐🅓",
|
||||
"ADE_SigmaSchedule": "Create Sigma Schedule 🎭🅐🅓",
|
||||
"ADE_RawSigmaSchedule": "Create Raw Sigma Schedule 🎭🅐🅓",
|
||||
"ADE_SigmaScheduleWeightedAverage": "Sigma Schedule Weighted Mean 🎭🅐🅓",
|
||||
@@ -222,10 +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] 🎭🅐🅓①",
|
||||
|
||||
@@ -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,)
|
||||
|
||||
@@ -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"}),
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -1,13 +1,20 @@
|
||||
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
|
||||
from .cfg_extras import perturbed_attention_guidance_patch, rescale_cfg_patch
|
||||
|
||||
|
||||
class AnimateDiffUnload:
|
||||
@@ -76,3 +83,46 @@ class EmptyLatentImageLarge:
|
||||
def generate(self, width, height, batch_size=1):
|
||||
latent = torch.zeros([batch_size, 4, height // 8, width // 8])
|
||||
return ({"samples":latent}, )
|
||||
|
||||
|
||||
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]):
|
||||
m = model.clone()
|
||||
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, )
|
||||
|
||||
+161
-5
@@ -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,13 +230,33 @@ 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,)
|
||||
|
||||
|
||||
class CustomCFGSimpleNode:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"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, cfg_extras: CFGExtrasGroup=None):
|
||||
return CustomCFGNode.create_custom_cfg(self, cfg_multival=cfg, cfg_extras=cfg_extras)
|
||||
|
||||
|
||||
class CustomCFGKeyframeNode:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
@@ -242,6 +268,7 @@ class CustomCFGKeyframeNode:
|
||||
},
|
||||
"optional": {
|
||||
"prev_custom_cfg": ("CUSTOM_CFG",),
|
||||
"cfg_extras": ("CFG_EXTRAS",),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -250,15 +277,144 @@ 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,)
|
||||
|
||||
|
||||
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",),
|
||||
"cfg_extras": ("CFG_EXTRAS",),
|
||||
}
|
||||
}
|
||||
|
||||
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, 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, 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:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
|
||||
+118
-32
@@ -1,5 +1,5 @@
|
||||
from collections.abc import Iterable
|
||||
from typing import Union
|
||||
from typing import Union, Callable
|
||||
import torch
|
||||
from torch import Tensor
|
||||
|
||||
@@ -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
|
||||
@@ -461,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
|
||||
@@ -541,7 +605,23 @@ 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 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,
|
||||
# meaning get_cfg_scale is a direct replacement
|
||||
def evolved_custom_cfg(args):
|
||||
cond: Tensor = args["cond"]
|
||||
uncond: Tensor = args["uncond"]
|
||||
@@ -562,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:
|
||||
|
||||
+10
-5
@@ -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
|
||||
@@ -607,11 +605,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()
|
||||
@@ -625,6 +627,9 @@ 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)
|
||||
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
|
||||
|
||||
+1
-1
@@ -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 = []
|
||||
|
||||
|
||||
Reference in New Issue
Block a user