diff --git a/adv_control/control.py b/adv_control/control.py index 787277f..ce00c74 100644 --- a/adv_control/control.py +++ b/adv_control/control.py @@ -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) diff --git a/adv_control/control_sparsectrl.py b/adv_control/control_sparsectrl.py index c70e3e2..e10e08c 100644 --- a/adv_control/control_sparsectrl.py +++ b/adv_control/control_sparsectrl.py @@ -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) diff --git a/adv_control/nodes_sparsectrl.py b/adv_control/nodes_sparsectrl.py index cfbfd2a..70eadc5 100644 --- a/adv_control/nodes_sparsectrl.py +++ b/adv_control/nodes_sparsectrl.py @@ -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, )