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:
Jedrzej Kosinski
2024-05-29 01:23:21 -05:00
parent 5ef18991ee
commit 949e1acf7e
3 changed files with 24 additions and 15 deletions
+2 -3
View File
@@ -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)
+13 -3
View File
@@ -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)
+9 -9
View File
@@ -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, )