Added strength_multival inputs to NaiveReuse (working) and ContextRef (no functionality at all) nodes, some scaffolding code for cleanup
This commit is contained in:
@@ -90,6 +90,7 @@ class ContextOptionsGroup:
|
||||
self._current_index = 0
|
||||
self.step = 0
|
||||
self._set_first_as_current()
|
||||
self.extras.cleanup()
|
||||
|
||||
@property
|
||||
def step(self):
|
||||
|
||||
@@ -2,6 +2,8 @@ from torch import Tensor
|
||||
|
||||
from comfy.model_base import BaseModel
|
||||
|
||||
from .utils_motion import prepare_mask_batch, extend_to_batch_size, get_combined_multival
|
||||
|
||||
|
||||
class ContextExtra:
|
||||
def __init__(self, start_percent: float, end_percent: float):
|
||||
@@ -24,9 +26,12 @@ class ContextExtra:
|
||||
return False
|
||||
return True
|
||||
|
||||
def cleanup(self):
|
||||
pass
|
||||
|
||||
|
||||
################################
|
||||
# Context Ref
|
||||
# ContextRef
|
||||
class ContextRefParams:
|
||||
def __init__(self,
|
||||
attn_style_fidelity=0.0, attn_ref_weight=0.0, attn_atrength=0.0,
|
||||
@@ -54,11 +59,30 @@ class ContextRef(ContextExtra):
|
||||
################################
|
||||
# NaiveReuse
|
||||
class NaiveReuse(ContextExtra):
|
||||
def __init__(self, start_percent: float, end_percent: float, weighted_mean: float, mask_opt: Tensor=None):
|
||||
def __init__(self, start_percent: float, end_percent: float, weighted_mean: float, multival_opt: Tensor=None):
|
||||
super().__init__(start_percent=start_percent, end_percent=end_percent)
|
||||
self.weighted_mean = weighted_mean
|
||||
self.mask_opt = mask_opt
|
||||
self.orig_multival = multival_opt
|
||||
self.mask: Tensor = None
|
||||
|
||||
def cleanup(self):
|
||||
super().cleanup()
|
||||
del self.mask
|
||||
self.mask = None
|
||||
|
||||
def get_effective_weighted_mean(self, x: Tensor, idxs: list[int]):
|
||||
if self.orig_multival is None:
|
||||
return self.weighted_mean
|
||||
# otherwise, is Tensor and should be extended to match dims and size of x;
|
||||
# see if needs to be recalculated
|
||||
if type(self.orig_multival) != Tensor:
|
||||
return self.weighted_mean * self.orig_multival
|
||||
elif self.mask is None or self.mask.shape[0] != x.shape[0] or self.mask.shape[-1] != x.shape[-1] or self.mask.shape[-2] != x.shape[-2]:
|
||||
del self.mask
|
||||
self.mask = prepare_mask_batch(self.orig_multival, x.shape)
|
||||
self.mask = extend_to_batch_size(self.mask, x.shape[0])
|
||||
return self.weighted_mean * self.mask[idxs].to(dtype=x.dtype, device=x.device)
|
||||
|
||||
def should_run(self):
|
||||
to_return = super().should_run()
|
||||
# if weighted_mean is 0.0, then reuse will take no effect anyway
|
||||
@@ -105,6 +129,10 @@ class ContextExtrasGroup:
|
||||
else:
|
||||
raise Exception(f"Unrecognized ContextExtras type: {type(extra)}")
|
||||
|
||||
def cleanup(self):
|
||||
for extra in self.get_extras_list():
|
||||
extra.cleanup()
|
||||
|
||||
def clone(self):
|
||||
cloned = ContextExtrasGroup()
|
||||
cloned.context_ref = self.context_ref
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
import torch
|
||||
from torch import Tensor
|
||||
from typing import Union
|
||||
|
||||
import comfy.samplers
|
||||
from comfy.model_patcher import ModelPatcher
|
||||
@@ -477,11 +478,11 @@ class ContextExtras_NaiveReuse:
|
||||
},
|
||||
"optional": {
|
||||
"prev_extras": ("CONTEXT_EXTRAS",),
|
||||
"mask_opt": ("MASK",),
|
||||
"strength_multival": ("MULTIVAL",),
|
||||
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}),
|
||||
"end_percent": ("FLOAT", {"default": 0.15, "min": 0.0, "max": 1.0, "step": 0.001}),
|
||||
"weighted_mean": ("FLOAT", {"default": 0.95, "min": 0.0, "max": 1.0, "step": 0.001}),
|
||||
"autosize": ("ADEAUTOSIZE", {"padding": 55}),
|
||||
"autosize": ("ADEAUTOSIZE", {"padding": 0}),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -489,12 +490,13 @@ class ContextExtras_NaiveReuse:
|
||||
CATEGORY = "Animate Diff 🎭🅐🅓/context opts/context extras"
|
||||
FUNCTION = "create_context_extra"
|
||||
|
||||
def create_context_extra(self, start_percent=0.0, end_percent=0.1, weighted_mean=0.95, mask_opt: Tensor=None, prev_extras: ContextExtrasGroup=None):
|
||||
def create_context_extra(self, start_percent=0.0, end_percent=0.1, weighted_mean=0.95, strength_multival: Union[float, Tensor]=None,
|
||||
prev_extras: ContextExtrasGroup=None):
|
||||
if prev_extras is None:
|
||||
prev_extras = prev_extras = ContextExtrasGroup()
|
||||
prev_extras = prev_extras.clone()
|
||||
# create extra
|
||||
naive_reuse = NaiveReuse(start_percent=start_percent, end_percent=end_percent, weighted_mean=weighted_mean, mask_opt=mask_opt)
|
||||
naive_reuse = NaiveReuse(start_percent=start_percent, end_percent=end_percent, weighted_mean=weighted_mean, multival_opt=strength_multival)
|
||||
prev_extras.add(naive_reuse)
|
||||
return (prev_extras,)
|
||||
|
||||
@@ -507,10 +509,10 @@ class ContextExtras_ContextRef:
|
||||
},
|
||||
"optional": {
|
||||
"prev_extras": ("CONTEXT_EXTRAS",),
|
||||
"mask_opt": ("MASK",),
|
||||
"strength_multival": ("MULTIVAL",),
|
||||
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}),
|
||||
"end_percent": ("FLOAT", {"default": 0.25, "min": 0.0, "max": 1.0, "step": 0.001}),
|
||||
"autosize": ("ADEAUTOSIZE", {"padding": 55}),
|
||||
"autosize": ("ADEAUTOSIZE", {"padding": 0}),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -518,7 +520,8 @@ class ContextExtras_ContextRef:
|
||||
CATEGORY = "Animate Diff 🎭🅐🅓/context opts/context extras"
|
||||
FUNCTION = "create_context_extra"
|
||||
|
||||
def create_context_extra(self, start_percent=0.0, end_percent=0.1, mask_opt: Tensor=None, prev_extras: ContextExtrasGroup=None):
|
||||
def create_context_extra(self, start_percent=0.0, end_percent=0.1, strength_multival: Union[float, Tensor]=None,
|
||||
prev_extras: ContextExtrasGroup=None):
|
||||
if prev_extras is None:
|
||||
prev_extras = prev_extras = ContextExtrasGroup()
|
||||
prev_extras = prev_extras.clone()
|
||||
|
||||
@@ -4,7 +4,7 @@ from typing import Union
|
||||
import torch
|
||||
from torch import Tensor
|
||||
|
||||
from .utils_motion import linear_conversion, normalize_min_max, extend_to_batch_size, extend_list_to_batch_size
|
||||
from .utils_motion import create_multival_combo, linear_conversion, normalize_min_max, extend_to_batch_size, extend_list_to_batch_size
|
||||
|
||||
|
||||
class ScaleType:
|
||||
@@ -31,39 +31,7 @@ class MultivalDynamicNode:
|
||||
FUNCTION = "create_multival"
|
||||
|
||||
def create_multival(self, float_val: Union[float, list[float]]=1.0, mask_optional: Tensor=None):
|
||||
# first, normalize inputs
|
||||
# if float_val is iterable, treat as a list and assume inputs are floats
|
||||
float_is_iterable = False
|
||||
if isinstance(float_val, Iterable):
|
||||
float_is_iterable = True
|
||||
float_val = list(float_val)
|
||||
# if mask present, make sure float_val list can be applied to list - match lengths
|
||||
if mask_optional is not None:
|
||||
if len(float_val) < mask_optional.shape[0]:
|
||||
# copies last entry enough times to match mask shape
|
||||
float_val = extend_list_to_batch_size(float_val, mask_optional.shape[0])
|
||||
if mask_optional.shape[0] < len(float_val):
|
||||
mask_optional = extend_to_batch_size(mask_optional, len(float_val))
|
||||
float_val = float_val[:mask_optional.shape[0]]
|
||||
float_val: Tensor = torch.tensor(float_val).unsqueeze(-1).unsqueeze(-1)
|
||||
# now that inputs are normalized, figure out what value to actually return
|
||||
if mask_optional is not None:
|
||||
mask_optional = mask_optional.clone()
|
||||
if float_is_iterable:
|
||||
mask_optional = mask_optional[:] * float_val.to(mask_optional.dtype).to(mask_optional.device)
|
||||
else:
|
||||
mask_optional = mask_optional * float_val
|
||||
return (mask_optional,)
|
||||
else:
|
||||
if not float_is_iterable:
|
||||
return (float_val,)
|
||||
# create a dummy mask of b,h,w=float_len,1,1 (sigle pixel)
|
||||
# purpose is for float input to work with mask code, without special cases
|
||||
float_len = float_val.shape[0] if float_is_iterable else 1
|
||||
shape = (float_len,1,1)
|
||||
mask_optional = torch.ones(shape)
|
||||
mask_optional = mask_optional[:] * float_val.to(mask_optional.dtype).to(mask_optional.device)
|
||||
return (mask_optional,)
|
||||
return (create_multival_combo(float_val=float_val, mask_optional=mask_optional),)
|
||||
|
||||
|
||||
class MultivalScaledMaskNode:
|
||||
|
||||
@@ -114,6 +114,7 @@ class AnimateDiffHelper_GlobalState:
|
||||
del self.motion_models
|
||||
self.motion_models = None
|
||||
if self.params is not None:
|
||||
self.params.context_options.reset()
|
||||
del self.params
|
||||
self.params = None
|
||||
if self.sample_settings is not None:
|
||||
@@ -894,7 +895,7 @@ def sliding_calc_conds_batch(model, conds, x_in: Tensor, timestep, model_options
|
||||
new_ctx_idxs = [zz for zz in list(range(z, z+len(cached_naive_ctx_idxs))) if zz < ADGS.params.full_length]
|
||||
# make sure when getting cached_naive idxs, they are adjusted for actual length leftover length
|
||||
adjusted_cnaive_ctx_idxs = cached_naive_ctx_idxs[:len(new_ctx_idxs)]
|
||||
weighted_mean = ADGS.params.context_options.extras.naive_reuse.weighted_mean
|
||||
weighted_mean = ADGS.params.context_options.extras.naive_reuse.get_effective_weighted_mean(x_in, new_ctx_idxs)
|
||||
conds_final[i][new_ctx_idxs] = (weighted_mean * (cached_naive_conds[i][adjusted_cnaive_ctx_idxs]*counts_final[i][new_ctx_idxs])) + ((1.-weighted_mean) * conds_final[i][new_ctx_idxs])
|
||||
#conds_final[i][new_idxs] += (cached_naive_conds[i][cached_naive_full_idxs] / cached_naive_counts[i][cached_naive_full_idxs]) * counts
|
||||
#counts = counts_final[i][new_idxs] * naive_counts_mult# / 2
|
||||
|
||||
@@ -3,6 +3,7 @@ import torch
|
||||
import torch.nn.functional as F
|
||||
from torch import Tensor, nn
|
||||
from abc import ABC, abstractmethod
|
||||
from collections.abc import Iterable
|
||||
|
||||
import comfy.model_management as model_management
|
||||
import comfy.ops
|
||||
@@ -238,6 +239,42 @@ class InputPIA_Multival(InputPIA):
|
||||
return mask * self.multival
|
||||
|
||||
|
||||
def create_multival_combo(float_val: Union[float, list[float]], mask_optional: Tensor=None):
|
||||
# first, normalize inputs
|
||||
# if float_val is iterable, treat as a list and assume inputs are floats
|
||||
float_is_iterable = False
|
||||
if isinstance(float_val, Iterable):
|
||||
float_is_iterable = True
|
||||
float_val = list(float_val)
|
||||
# if mask present, make sure float_val list can be applied to list - match lengths
|
||||
if mask_optional is not None:
|
||||
if len(float_val) < mask_optional.shape[0]:
|
||||
# copies last entry enough times to match mask shape
|
||||
float_val = extend_list_to_batch_size(float_val, mask_optional.shape[0])
|
||||
if mask_optional.shape[0] < len(float_val):
|
||||
mask_optional = extend_to_batch_size(mask_optional, len(float_val))
|
||||
float_val = float_val[:mask_optional.shape[0]]
|
||||
float_val: Tensor = torch.tensor(float_val).unsqueeze(-1).unsqueeze(-1)
|
||||
# now that inputs are normalized, figure out what value to actually return
|
||||
if mask_optional is not None:
|
||||
mask_optional = mask_optional.clone()
|
||||
if float_is_iterable:
|
||||
mask_optional = mask_optional[:] * float_val.to(mask_optional.dtype).to(mask_optional.device)
|
||||
else:
|
||||
mask_optional = mask_optional * float_val
|
||||
return mask_optional
|
||||
else:
|
||||
if not float_is_iterable:
|
||||
return float_val
|
||||
# create a dummy mask of b,h,w=float_len,1,1 (sigle pixel)
|
||||
# purpose is for float input to work with mask code, without special cases
|
||||
float_len = float_val.shape[0] if float_is_iterable else 1
|
||||
shape = (float_len,1,1)
|
||||
mask_optional = torch.ones(shape)
|
||||
mask_optional = mask_optional[:] * float_val.to(mask_optional.dtype).to(mask_optional.device)
|
||||
return mask_optional
|
||||
|
||||
|
||||
def get_combined_multival(multivalA: Union[float, Tensor], multivalB: Union[float, Tensor]) -> Union[float, Tensor]:
|
||||
# if one is None, use the other
|
||||
if multivalA == None:
|
||||
|
||||
Reference in New Issue
Block a user