Added basic NaiveReuse Keyframe nodes
This commit is contained in:
+128
-11
@@ -1,8 +1,12 @@
|
||||
from typing import Union
|
||||
import math
|
||||
import torch
|
||||
from torch import Tensor
|
||||
|
||||
from comfy.model_base import BaseModel
|
||||
|
||||
from .utils_motion import prepare_mask_batch, extend_to_batch_size, get_combined_multival
|
||||
from .utils_motion import (prepare_mask_batch, extend_to_batch_size, get_combined_multival, resize_multival,
|
||||
get_sorted_list_via_attr)
|
||||
|
||||
|
||||
class ContextExtra:
|
||||
@@ -93,34 +97,147 @@ class ContextRef(ContextExtra):
|
||||
|
||||
|
||||
################################
|
||||
# NaiveReuse
|
||||
# NaiveReuse
|
||||
class NaiveReuseKeyframe:
|
||||
def __init__(self, mult_multival: Union[float, Tensor], start_percent=0.0, guarantee_steps=1):
|
||||
self.mult_multival = mult_multival
|
||||
# scheduling
|
||||
self.start_percent = float(start_percent)
|
||||
self.start_t = 999999999.9
|
||||
self.guarantee_steps = guarantee_steps
|
||||
|
||||
def clone(self):
|
||||
c = NaiveReuseKeyframe(mult_multival=self.mult_multival,
|
||||
start_percent=self.start_percent, guarantee_steps=self.guarantee_steps)
|
||||
c.start_t = self.start_t
|
||||
return c
|
||||
|
||||
class NaiveReuseKeyframeGroup:
|
||||
def __init__(self):
|
||||
self.keyframes: list[NaiveReuseKeyframe] = []
|
||||
self._current_keyframe: NaiveReuseKeyframe = None
|
||||
self._current_used_steps: int = 0
|
||||
self._current_index: int = 0
|
||||
self._previous_t = -1
|
||||
|
||||
def reset(self):
|
||||
self._current_keyframe = None
|
||||
self._current_used_steps = 0
|
||||
self._current_index = 0
|
||||
self._set_first_as_current()
|
||||
|
||||
def add(self, keyframe: NaiveReuseKeyframe):
|
||||
# add to end of list, then sort
|
||||
self.keyframes.append(keyframe)
|
||||
self.keyframes = get_sorted_list_via_attr(self.keyframes, "start_percent")
|
||||
self._set_first_as_current()
|
||||
|
||||
def _set_first_as_current(self):
|
||||
if len(self.keyframes) > 0:
|
||||
self._current_keyframe = self.keyframes[0]
|
||||
else:
|
||||
self._current_keyframe = None
|
||||
|
||||
def has_index(self, index: int) -> int:
|
||||
return index >=0 and index < len(self.keyframes)
|
||||
|
||||
def is_empty(self) -> bool:
|
||||
return len(self.keyframes) == 0
|
||||
|
||||
def clone(self):
|
||||
cloned = NaiveReuseKeyframeGroup()
|
||||
for keyframe in self.keyframes:
|
||||
cloned.keyframes.append(keyframe)
|
||||
cloned._set_first_as_current()
|
||||
return cloned
|
||||
|
||||
def initialize_timesteps(self, model: BaseModel):
|
||||
for keyframe in self.keyframes:
|
||||
keyframe.start_t = model.model_sampling.percent_to_sigma(keyframe.start_percent)
|
||||
|
||||
def prepare_current_keyframe(self, t: Tensor):
|
||||
curr_t: float = t[0]
|
||||
# if curr_t same as before, do nothing as step already accounted for
|
||||
if curr_t == self._previous_t:
|
||||
return
|
||||
prev_index = self._current_index
|
||||
# if met guaranteed steps, look for next keyframe in case need to switch
|
||||
if self._current_used_steps >= self._current_keyframe.guarantee_steps:
|
||||
# if has next index, loop through and see if need t oswitch
|
||||
if self.has_index(self._current_index+1):
|
||||
for i in range(self._current_index+1, len(self.keyframes)):
|
||||
eval_c = self.keyframes[i]
|
||||
# check if start_t is greater or equal to curr_t
|
||||
# NOTE: t is in terms of sigmas, not percent, so bigger number = earlier step in sampling
|
||||
if eval_c.start_t >= curr_t:
|
||||
self._current_index = i
|
||||
self._current_keyframe = eval_c
|
||||
self._current_used_steps = 0
|
||||
# if guarantee_steps greater than zero, stop searching for other keyframes
|
||||
if self._current_keyframe.guarantee_steps > 0:
|
||||
break
|
||||
# if eval_c is outside the percent range, stop looking further
|
||||
else: break
|
||||
# update steps current context is used
|
||||
self._current_used_steps += 1
|
||||
# update previous_t
|
||||
self._previous_t = curr_t
|
||||
|
||||
# properties shadow those of NaiveReuseKeyframe
|
||||
@property
|
||||
def mult_multival(self):
|
||||
if self._current_keyframe != None:
|
||||
return self._current_keyframe.mult_multival
|
||||
return None
|
||||
|
||||
class NaiveReuse(ContextExtra):
|
||||
def __init__(self, start_percent: float, end_percent: float, weighted_mean: float, multival_opt: Tensor=None):
|
||||
def __init__(self, start_percent: float, end_percent: float, weighted_mean: float, multival_opt: Union[float, Tensor]=None, naivereuse_kf: NaiveReuseKeyframeGroup=None):
|
||||
super().__init__(start_percent=start_percent, end_percent=end_percent)
|
||||
self.weighted_mean = weighted_mean
|
||||
self.orig_multival = multival_opt
|
||||
self.mask: Tensor = None
|
||||
self.keyframe = naivereuse_kf if naivereuse_kf else NaiveReuseKeyframeGroup()
|
||||
self._prev_keyframe = None
|
||||
|
||||
def cleanup(self):
|
||||
super().cleanup()
|
||||
del self.mask
|
||||
self.mask = None
|
||||
self._prev_keyframe = None
|
||||
self.keyframe.reset()
|
||||
|
||||
def initialize_timesteps(self, model: BaseModel):
|
||||
super().initialize_timesteps(model)
|
||||
self.keyframe.initialize_timesteps(model)
|
||||
|
||||
def prepare_current(self, t: Tensor):
|
||||
super().prepare_current(t)
|
||||
self.keyframe.prepare_current_keyframe(t)
|
||||
|
||||
def get_effective_weighted_mean(self, x: Tensor, idxs: list[int]):
|
||||
if self.orig_multival is None:
|
||||
if self.orig_multival is None and self.keyframe.mult_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]:
|
||||
# check if keyframe changed
|
||||
keyframe_changed = False
|
||||
if self.keyframe._current_keyframe != self._prev_keyframe:
|
||||
keyframe_changed = True
|
||||
self._prev_keyframe = self.keyframe._current_keyframe
|
||||
|
||||
if type(self.orig_multival) != Tensor and type(self.keyframe.mult_multival) != Tensor:
|
||||
return self.weighted_mean * get_combined_multival(self.orig_multival, self.keyframe.mult_multival)
|
||||
|
||||
if self.mask is None or keyframe_changed 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])
|
||||
real_mult_multival = resize_multival(self.keyframe.mult_multival, batch_size=x.shape[0], height=x.shape[-1], width=x.shape[-2])
|
||||
self.mask = resize_multival(self.orig_multival, batch_size=x.shape[0], height=x.shape[-1], width=x.shape[-2])
|
||||
self.mask = get_combined_multival(self.mask, real_mult_multival)
|
||||
return self.weighted_mean * self.mask[idxs].to(dtype=x.dtype, device=x.device)
|
||||
|
||||
def should_run(self):
|
||||
to_return = super().should_run()
|
||||
# if keyframe has 0.0 val, should not run
|
||||
if self.keyframe.mult_multival is not None and type(self.keyframe.mult_multival) != Tensor and math.isclose(self.keyframe.mult_multival, 0.0):
|
||||
return False
|
||||
# if weighted_mean is 0.0, then reuse will take no effect anyway
|
||||
return to_return and self.weighted_mean > 0.0
|
||||
#--------------------------------
|
||||
|
||||
@@ -31,7 +31,8 @@ from .nodes_context import (LegacyLoopedUniformContextOptionsNode, LoopedUniform
|
||||
VisualizeContextOptionsK, VisualizeContextOptionsKAdv, VisualizeContextOptionsSCustom,
|
||||
SetContextExtrasOnContextOptions, ContextExtras_NaiveReuse, ContextExtras_ContextRef,
|
||||
ContextRef_ModeFirst, ContextRef_ModeSliding, ContextRef_ModeIndexes,
|
||||
ContextRef_TuneAttn, ContextRef_TuneAttnAdain)
|
||||
ContextRef_TuneAttn, ContextRef_TuneAttnAdain,
|
||||
NaiveReuse_KeyframeNode, NaiveReuse_KeyframeMultivalNode)
|
||||
from .nodes_ad_settings import (AnimateDiffSettingsNode, ManualAdjustPENode, SweetspotStretchPENode, FullStretchPENode,
|
||||
WeightAdjustAllAddNode, WeightAdjustAllMultNode, WeightAdjustIndivAddNode, WeightAdjustIndivMultNode,
|
||||
WeightAdjustIndivAttnAddNode, WeightAdjustIndivAttnMultNode)
|
||||
@@ -76,12 +77,14 @@ NODE_CLASS_MAPPINGS = {
|
||||
# Context Extras
|
||||
"ADE_ContextExtras_Set": SetContextExtrasOnContextOptions,
|
||||
"ADE_ContextExtras_ContextRef": ContextExtras_ContextRef,
|
||||
"ADE_ContextExtras_NaiveReuse": ContextExtras_NaiveReuse,
|
||||
"ADE_ContextExtras_ContextRef_ModeFirst": ContextRef_ModeFirst,
|
||||
"ADE_ContextExtras_ContextRef_ModeSliding": ContextRef_ModeSliding,
|
||||
"ADE_ContextExtras_ContextRef_ModeIndexes": ContextRef_ModeIndexes,
|
||||
"ADE_ContextExtras_ContextRef_TuneAttn": ContextRef_TuneAttn,
|
||||
"ADE_ContextExtras_ContextRef_TuneAttnAdain": ContextRef_TuneAttnAdain,
|
||||
"ADE_ContextExtras_NaiveReuse": ContextExtras_NaiveReuse,
|
||||
"ADE_ContextExtras_NaiveReuse_Keyframe": NaiveReuse_KeyframeNode,
|
||||
"ADE_ContextExtras_NaiveReuse_KeyframeMultival": NaiveReuse_KeyframeMultivalNode,
|
||||
#------------------------------------------------------------------------------
|
||||
###############################################################################
|
||||
# Iteration Opts
|
||||
@@ -219,12 +222,14 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
# Context Extras
|
||||
"ADE_ContextExtras_Set": "Set Context Extras 🎭🅐🅓",
|
||||
"ADE_ContextExtras_ContextRef": "Context Extras◆ContextRef 🎭🅐🅓",
|
||||
"ADE_ContextExtras_NaiveReuse": "Context Extras◆NaiveReuse 🎭🅐🅓",
|
||||
"ADE_ContextExtras_ContextRef_ModeFirst": "ContextRef Mode◆First 🎭🅐🅓",
|
||||
"ADE_ContextExtras_ContextRef_ModeSliding": "ContextRef Mode◆Sliding 🎭🅐🅓",
|
||||
"ADE_ContextExtras_ContextRef_ModeIndexes": "ContextRef Mode◆Indexes 🎭🅐🅓",
|
||||
"ADE_ContextExtras_ContextRef_TuneAttn": "ContextRef Tune◆Attn 🎭🅐🅓",
|
||||
"ADE_ContextExtras_ContextRef_TuneAttnAdain": "ContextRef Tune◆Attn+Adain 🎭🅐🅓",
|
||||
"ADE_ContextExtras_NaiveReuse": "Context Extras◆NaiveReuse 🎭🅐🅓",
|
||||
"ADE_ContextExtras_NaiveReuse_Keyframe": "NaiveReuse Keyframe 🎭🅐🅓",
|
||||
"ADE_ContextExtras_NaiveReuse_KeyframeMultival": "NaiveReuse Keyframe [Multival] 🎭🅐🅓",
|
||||
#------------------------------------------------------------------------------
|
||||
###############################################################################
|
||||
# Iteration Opts
|
||||
|
||||
@@ -7,7 +7,8 @@ from comfy.model_patcher import ModelPatcher
|
||||
|
||||
from .context import (ContextFuseMethod, ContextOptions, ContextOptionsGroup, ContextSchedules,
|
||||
generate_context_visualization)
|
||||
from .context_extras import ContextExtrasGroup, ContextRef, ContextRefParams, ContextRefMode, NaiveReuse
|
||||
from .context_extras import (ContextExtrasGroup, ContextRef, ContextRefParams, ContextRefMode, NaiveReuse,
|
||||
NaiveReuseKeyframe, NaiveReuseKeyframeGroup)
|
||||
from .utils_model import BIGMAX, MAX_RESOLUTION
|
||||
from .utils_scheduling import convert_str_to_indexes
|
||||
|
||||
@@ -480,6 +481,7 @@ class ContextExtras_NaiveReuse:
|
||||
"optional": {
|
||||
"prev_extras": ("CONTEXT_EXTRAS",),
|
||||
"strength_multival": ("MULTIVAL",),
|
||||
"naivereuse_kf": ("NAIVEREUSE_KEYFRAME",),
|
||||
"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}),
|
||||
@@ -492,16 +494,73 @@ class ContextExtras_NaiveReuse:
|
||||
FUNCTION = "create_context_extra"
|
||||
|
||||
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):
|
||||
naivereuse_kf: NaiveReuseKeyframeGroup=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, multival_opt=strength_multival)
|
||||
naive_reuse = NaiveReuse(start_percent=start_percent, end_percent=end_percent, weighted_mean=weighted_mean, multival_opt=strength_multival,
|
||||
naivereuse_kf=naivereuse_kf)
|
||||
prev_extras.add(naive_reuse)
|
||||
return (prev_extras,)
|
||||
|
||||
|
||||
class NaiveReuse_KeyframeMultivalNode:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"mult_multival": ("MULTIVAL",),
|
||||
},
|
||||
"optional": {
|
||||
"prev_kf": ("NAIVEREUSE_KEYFRAME",),
|
||||
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}),
|
||||
"guarantee_steps": ("INT", {"default": 1, "min": 0, "max": BIGMAX}),
|
||||
"autosize": ("ADEAUTOSIZE", {"padding": 80}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("NAIVEREUSE_KEYFRAME",)
|
||||
RETURN_NAMES = ("NAIVEREUSE_KF",)
|
||||
CATEGORY = "Animate Diff 🎭🅐🅓/context opts/context extras/naivereuse"
|
||||
FUNCTION = "create_keyframe"
|
||||
|
||||
def create_keyframe(self, prev_kf=None, mult_multival=1.0,
|
||||
start_percent=0.0, guarantee_steps=1):
|
||||
if prev_kf is None:
|
||||
prev_kf = NaiveReuseKeyframeGroup()
|
||||
prev_kf = prev_kf.clone()
|
||||
kf = NaiveReuseKeyframe(mult_multival=mult_multival, start_percent=start_percent, guarantee_steps=guarantee_steps)
|
||||
prev_kf.add(kf)
|
||||
return (prev_kf,)
|
||||
|
||||
|
||||
class NaiveReuse_KeyframeNode:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
},
|
||||
"optional": {
|
||||
"prev_kf": ("NAIVEREUSE_KEYFRAME",),
|
||||
"mult": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.001}),
|
||||
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}),
|
||||
"guarantee_steps": ("INT", {"default": 1, "min": 0, "max": BIGMAX}),
|
||||
"autosize": ("ADEAUTOSIZE", {"padding": 10}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("NAIVEREUSE_KEYFRAME",)
|
||||
RETURN_NAMES = ("NAIVEREUSE_KF",)
|
||||
CATEGORY = "Animate Diff 🎭🅐🅓/context opts/context extras/naivereuse"
|
||||
FUNCTION = "create_keyframe"
|
||||
|
||||
def create_keyframe(self, prev_kf=None, mult=1.0,
|
||||
start_percent=0.0, guarantee_steps=1):
|
||||
return NaiveReuse_KeyframeMultivalNode.create_keyframe(self, prev_kf=prev_kf, mult_multival=float(mult),
|
||||
start_percent=start_percent, guarantee_steps=guarantee_steps)
|
||||
|
||||
|
||||
class ContextExtras_ContextRef:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
@@ -541,6 +600,32 @@ class ContextExtras_ContextRef:
|
||||
return (prev_extras,)
|
||||
|
||||
|
||||
class ContextRef_KeyframeNode:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
},
|
||||
"optional": {
|
||||
"prev_keyframe": ("CONTEXTREF_KEYFRAME",),
|
||||
"mult_multival": ("MULTIVAL",),
|
||||
"mode_replace": ("CONTEXTREF_MODE",),
|
||||
"tune_replace": ("CONTEXTREF_TUNE",),
|
||||
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}),
|
||||
"guarantee_steps": ("INT", {"default": 1, "min": 0, "max": BIGMAX}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("CONTEXTREF_KEYFRAME",)
|
||||
RETURN_NAMES = ("CONTEXTREF_KF",)
|
||||
CATEGORY = "Animate Diff 🎭🅐🅓/context opts/context extras/contextref"
|
||||
FUNCTION = "create_keyframe"
|
||||
|
||||
def create_keyframe(self, prev_keyframe=None, mult_multival=1.0, mode_replace=None, tune_replace=None,
|
||||
start_percent=1.0, guarantee_steps=1):
|
||||
pass
|
||||
|
||||
|
||||
class ContextRef_ModeFirst:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
@@ -553,7 +638,7 @@ class ContextRef_ModeFirst:
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("CONTEXTREF_MODE",)
|
||||
CATEGORY = "Animate Diff 🎭🅐🅓/context opts/context extras/ContextRef"
|
||||
CATEGORY = "Animate Diff 🎭🅐🅓/context opts/context extras/contextref"
|
||||
FUNCTION = "create_contextref_mode"
|
||||
|
||||
def create_contextref_mode(self):
|
||||
@@ -574,7 +659,7 @@ class ContextRef_ModeSliding:
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("CONTEXTREF_MODE",)
|
||||
CATEGORY = "Animate Diff 🎭🅐🅓/context opts/context extras/ContextRef"
|
||||
CATEGORY = "Animate Diff 🎭🅐🅓/context opts/context extras/contextref"
|
||||
FUNCTION = "create_contextref_mode"
|
||||
|
||||
def create_contextref_mode(self, sliding_width):
|
||||
@@ -596,7 +681,7 @@ class ContextRef_ModeIndexes:
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("CONTEXTREF_MODE",)
|
||||
CATEGORY = "Animate Diff 🎭🅐🅓/context opts/context extras/ContextRef"
|
||||
CATEGORY = "Animate Diff 🎭🅐🅓/context opts/context extras/contextref"
|
||||
FUNCTION = "create_contextref_mode"
|
||||
|
||||
def create_contextref_mode(self, switch_on_idxs: str, always_include_0: bool):
|
||||
@@ -625,7 +710,7 @@ class ContextRef_TuneAttnAdain:
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("CONTEXTREF_TUNE",)
|
||||
CATEGORY = "Animate Diff 🎭🅐🅓/context opts/context extras/ContextRef"
|
||||
CATEGORY = "Animate Diff 🎭🅐🅓/context opts/context extras/contextref"
|
||||
FUNCTION = "create_contextref_tune"
|
||||
|
||||
def create_contextref_tune(self, attn_style_fidelity=1.0, attn_ref_weight=1.0, attn_strength=1.0,
|
||||
@@ -651,7 +736,7 @@ class ContextRef_TuneAttn:
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("CONTEXTREF_TUNE",)
|
||||
CATEGORY = "Animate Diff 🎭🅐🅓/context opts/context extras/ContextRef"
|
||||
CATEGORY = "Animate Diff 🎭🅐🅓/context opts/context extras/contextref"
|
||||
FUNCTION = "create_contextref_tune"
|
||||
|
||||
def create_contextref_tune(self, attn_style_fidelity=1.0, attn_ref_weight=1.0, attn_strength=1.0):
|
||||
|
||||
@@ -275,7 +275,7 @@ def create_multival_combo(float_val: Union[float, list[float]], mask_optional: T
|
||||
return mask_optional
|
||||
|
||||
|
||||
def get_combined_multival(multivalA: Union[float, Tensor], multivalB: Union[float, Tensor]) -> Union[float, Tensor]:
|
||||
def get_combined_multival(multivalA: Union[float, Tensor], multivalB: Union[float, Tensor], force_leader_A=False) -> Union[float, Tensor]:
|
||||
# if one is None, use the other
|
||||
if multivalA == None:
|
||||
return multivalB
|
||||
@@ -284,14 +284,18 @@ def get_combined_multival(multivalA: Union[float, Tensor], multivalB: Union[floa
|
||||
# both have a value - combine them based on type
|
||||
# if both are Tensors, make dims match before multiplying
|
||||
if type(multivalA) == Tensor and type(multivalB) == Tensor:
|
||||
areaA = multivalA.shape[1]*multivalA.shape[2]
|
||||
areaB = multivalB.shape[1]*multivalB.shape[2]
|
||||
# match height/width to mask with larger area
|
||||
leader,follower = (multivalA,multivalB) if areaA >= areaB else (multivalB,multivalA)
|
||||
batch_size = multivalA.shape[0] if multivalA.shape[0] >= multivalB.shape[0] else multivalB.shape[0]
|
||||
if force_leader_A:
|
||||
leader,follower = (multivalA,multivalB)
|
||||
batch_size = multivalA.shape[0]
|
||||
else:
|
||||
areaA = multivalA.shape[1]*multivalA.shape[2]
|
||||
areaB = multivalB.shape[1]*multivalB.shape[2]
|
||||
# match height/width to mask with larger area
|
||||
leader,follower = (multivalA,multivalB) if areaA >= areaB else (multivalB,multivalA)
|
||||
batch_size = multivalA.shape[0] if multivalA.shape[0] >= multivalB.shape[0] else multivalB.shape[0]
|
||||
# make follower same dimensions as leader
|
||||
follower = torch.unsqueeze(follower, 1)
|
||||
follower = comfy.utils.common_upscale(follower, leader.shape[2], leader.shape[1], "bilinear", "center")
|
||||
follower = comfy.utils.common_upscale(follower, leader.shape[-1], leader.shape[-2], "bilinear", "center")
|
||||
follower = torch.squeeze(follower, 1)
|
||||
# make sure batch size will match
|
||||
leader = extend_to_batch_size(leader, batch_size)
|
||||
@@ -301,6 +305,18 @@ def get_combined_multival(multivalA: Union[float, Tensor], multivalB: Union[floa
|
||||
return multivalA * multivalB
|
||||
|
||||
|
||||
def resize_multival(multival: Union[float, Tensor], batch_size: int, height: int, width: int):
|
||||
if multival == None:
|
||||
return 1.0
|
||||
if type(multival) != Tensor:
|
||||
return multival
|
||||
multival = torch.unsqueeze(multival, 1)
|
||||
multival = comfy.utils.common_upscale(multival, height, width, "bilinear", "center")
|
||||
multival = torch.squeeze(multival, 1)
|
||||
multival = extend_to_batch_size(multival, batch_size)
|
||||
return multival
|
||||
|
||||
|
||||
def get_combined_input(inputA: Union[InputPIA, None], inputB: Union[InputPIA, None], x: Tensor):
|
||||
if inputA is None:
|
||||
inputA = InputPIA_Multival(1.0)
|
||||
|
||||
Reference in New Issue
Block a user