Turned context_aware into a dropdown in case I add other context_aware types in future, renamed sparse_mask_strength to sparse_mask_mult
This commit is contained in:
@@ -303,7 +303,7 @@ class SparseCtrlAdvanced(ControlNetAdvanced):
|
||||
self.cond_hint = None
|
||||
# first, figure out which cond idxs are relevant, and where they fit in
|
||||
cond_idxs, hint_order = self.sparse_settings.sparse_method.get_indexes(hint_length=self.cond_hint_original.size(0), full_length=full_length,
|
||||
sub_idxs=self.sub_idxs if self.sparse_settings.context_aware else None)
|
||||
sub_idxs=self.sub_idxs if self.sparse_settings.is_context_aware() else None)
|
||||
range_idxs = list(range(full_length)) if self.sub_idxs is None else self.sub_idxs
|
||||
hint_idxs = [] # idxs in cond_idxs
|
||||
local_idxs = [] # idx to put in final cond_hint
|
||||
@@ -345,8 +345,7 @@ class SparseCtrlAdvanced(ControlNetAdvanced):
|
||||
# prepare cond_mask (b, 1, h, w)
|
||||
cond_shape[1] = 1
|
||||
cond_mask = torch.zeros(cond_shape).to(dtype).to(self.device)
|
||||
cond_mask[local_idxs] = self.sparse_settings.sparse_mask_strength * self.weights.extras.get(SparseConst.MASK_STRENGTH, 1.0)
|
||||
#cond_mask[local_idxs] = 2.5
|
||||
cond_mask[local_idxs] = self.sparse_settings.sparse_mask_mult * self.weights.extras.get(SparseConst.MASK_MULT, 1.0)
|
||||
# combine cond_hint and cond_mask into (b, c+1, h, w)
|
||||
if not self.sparse_settings.merged:
|
||||
self.cond_hint = torch.cat([self.cond_hint, cond_mask], dim=1)
|
||||
|
||||
@@ -57,7 +57,7 @@ else:
|
||||
class SparseConst:
|
||||
HINT_MULT = "sparse_hint_mult"
|
||||
NONHINT_MULT = "sparse_nonhint_mult"
|
||||
MASK_STRENGTH = "sparse_mask_strength"
|
||||
MASK_MULT = "sparse_mask_mult"
|
||||
|
||||
|
||||
class SparseControlNet(ControlNetCLDM):
|
||||
@@ -185,19 +185,29 @@ class PreprocSparseRGBWrapper:
|
||||
raise AttributeError(self.error_msg)
|
||||
|
||||
|
||||
class SparseContextAware:
|
||||
NEAREST_HINT = "nearest_hint"
|
||||
OFF = "off"
|
||||
|
||||
LIST = [NEAREST_HINT, OFF]
|
||||
|
||||
|
||||
class SparseSettings:
|
||||
def __init__(self, sparse_method: 'SparseMethod', use_motion: bool=True, motion_strength=1.0, motion_scale=1.0, merged=False,
|
||||
sparse_mask_strength=1.0, sparse_hint_mult=1.0, sparse_nonhint_mult=1.0, context_aware=True):
|
||||
sparse_mask_mult=1.0, sparse_hint_mult=1.0, sparse_nonhint_mult=1.0, context_aware=SparseContextAware.NEAREST_HINT):
|
||||
self.sparse_method = sparse_method
|
||||
self.use_motion = use_motion
|
||||
self.motion_strength = motion_strength
|
||||
self.motion_scale = motion_scale
|
||||
self.merged = merged
|
||||
self.sparse_mask_strength = float(sparse_mask_strength)
|
||||
self.sparse_mask_mult = float(sparse_mask_mult)
|
||||
self.sparse_hint_mult = float(sparse_hint_mult)
|
||||
self.sparse_nonhint_mult = float(sparse_nonhint_mult)
|
||||
self.context_aware = context_aware
|
||||
|
||||
def is_context_aware(self):
|
||||
return self.context_aware != SparseContextAware.OFF
|
||||
|
||||
@classmethod
|
||||
def default(cls):
|
||||
return SparseSettings(sparse_method=SparseSpreadMethod(), use_motion=True)
|
||||
|
||||
@@ -6,7 +6,7 @@ import comfy.utils
|
||||
from comfy.sd import VAE
|
||||
|
||||
from .utils import TimestepKeyframeGroup
|
||||
from .control_sparsectrl import SparseMethod, SparseIndexMethod, SparseSettings, SparseSpreadMethod, PreprocSparseRGBWrapper, SparseConst
|
||||
from .control_sparsectrl import SparseMethod, SparseIndexMethod, SparseSettings, SparseSpreadMethod, PreprocSparseRGBWrapper, SparseConst, SparseContextAware
|
||||
from .control import load_sparsectrl, load_controlnet, ControlNetAdvanced, SparseCtrlAdvanced
|
||||
|
||||
|
||||
@@ -24,10 +24,10 @@ class SparseCtrlLoaderAdvanced:
|
||||
"optional": {
|
||||
"sparse_method": ("SPARSE_METHOD", ),
|
||||
"tk_optional": ("TIMESTEP_KEYFRAME", ),
|
||||
"context_aware": ("BOOLEAN", {"default": True}, ),
|
||||
"context_aware": (SparseContextAware.LIST, ),
|
||||
"sparse_hint_mult": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"sparse_nonhint_mult": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"sparse_mask_strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"sparse_mask_mult": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -37,11 +37,11 @@ class SparseCtrlLoaderAdvanced:
|
||||
CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/SparseCtrl"
|
||||
|
||||
def load_controlnet(self, sparsectrl_name: str, use_motion: bool, motion_strength: float, motion_scale: float, sparse_method: SparseMethod=SparseSpreadMethod(), tk_optional: TimestepKeyframeGroup=None,
|
||||
context_aware=True, sparse_hint_mult=1.0, sparse_nonhint_mult=1.0, sparse_mask_strength=1.0):
|
||||
context_aware=SparseContextAware.NEAREST_HINT, sparse_hint_mult=1.0, sparse_nonhint_mult=1.0, sparse_mask_mult=1.0):
|
||||
sparsectrl_path = folder_paths.get_full_path("controlnet", sparsectrl_name)
|
||||
sparse_settings = SparseSettings(sparse_method=sparse_method, use_motion=use_motion, motion_strength=motion_strength, motion_scale=motion_scale,
|
||||
context_aware=context_aware,
|
||||
sparse_mask_strength=sparse_mask_strength, sparse_hint_mult=sparse_hint_mult, sparse_nonhint_mult=sparse_nonhint_mult)
|
||||
sparse_mask_mult=sparse_mask_mult, sparse_hint_mult=sparse_hint_mult, sparse_nonhint_mult=sparse_nonhint_mult)
|
||||
sparsectrl = load_sparsectrl(sparsectrl_path, timestep_keyframe=tk_optional, sparse_settings=sparse_settings)
|
||||
return (sparsectrl,)
|
||||
|
||||
@@ -178,7 +178,7 @@ class SparseWeightExtras:
|
||||
"extras": ("CN_WEIGHTS_EXTRAS",),
|
||||
"sparse_hint_mult": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"sparse_nonhint_mult": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"sparse_mask_strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"sparse_mask_mult": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -186,11 +186,11 @@ class SparseWeightExtras:
|
||||
RETURN_NAMES = ("extras", )
|
||||
FUNCTION = "create_weight_extras"
|
||||
|
||||
CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/SparseCtrl"
|
||||
CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/SparseCtrl/extras"
|
||||
|
||||
def create_weight_extras(self, extras: dict[str]={}, sparse_hint_mult=1.0, sparse_nonhint_mult=1.0, sparse_mask_strength=1.0):
|
||||
def create_weight_extras(self, extras: dict[str]={}, sparse_hint_mult=1.0, sparse_nonhint_mult=1.0, sparse_mask_mult=1.0):
|
||||
extras = extras.copy()
|
||||
extras[SparseConst.HINT_MULT] = sparse_hint_mult
|
||||
extras[SparseConst.NONHINT_MULT] = sparse_nonhint_mult
|
||||
extras[SparseConst.MASK_STRENGTH] = sparse_mask_strength
|
||||
extras[SparseConst.MASK_MULT] = sparse_mask_mult
|
||||
return (extras, )
|
||||
|
||||
Reference in New Issue
Block a user