Added strength_multival inputs to NaiveReuse (working) and ContextRef (no functionality at all) nodes, some scaffolding code for cleanup

This commit is contained in:
Jedrzej Kosinski
2024-07-25 00:52:15 -05:00
parent 3a8744860d
commit bd514925de
6 changed files with 83 additions and 45 deletions
+1
View File
@@ -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):
+31 -3
View File
@@ -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
+10 -7
View File
@@ -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()
+2 -34
View File
@@ -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:
+2 -1
View File
@@ -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
+37
View File
@@ -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: