diff --git a/control/control.py b/control/control.py index c5812ed..6500af7 100644 --- a/control/control.py +++ b/control/control.py @@ -9,7 +9,7 @@ import comfy.model_detection import comfy.controlnet as comfy_cn from comfy.controlnet import ControlBase, ControlNet, ControlLora, T2IAdapter, broadcast_image_to -from .control_sparsectrl import SparseControlNet, SparseCtrlMotionWrapper, SparseMethod, SparseSettings, SparseSpreadMethod +from .control_sparsectrl import SparseControlNet, SparseCtrlMotionWrapper, SparseMethod, SparseSettings, SparseSpreadMethod, PreprocSparseRGBWrapper from .utils import (AdvancedControlBase, TimestepKeyframeGroup, LatentKeyframeGroup, ControlWeightType, ControlWeights, WeightTypeException, manual_cast_clean_groupnorm, disable_weight_init_clean_groupnorm, prepare_mask_batch, get_properly_arranged_t2i_weights, load_torch_file_with_dict_factory) from .logger import logger @@ -228,6 +228,7 @@ class SparseCtrlAdvanced(ControlNetAdvanced): self.control_model: SparseControlNet = self.control_model # does nothing except help with IDE hints self.sparse_settings = sparse_settings if sparse_settings is not None else SparseSettings.default() self.latent_format = None + self.preprocessed = False def get_control_advanced(self, x_noisy: Tensor, t, cond, batched_number: int): # normal ControlNet stuff @@ -311,6 +312,8 @@ class SparseCtrlAdvanced(ControlNetAdvanced): def pre_run_advanced(self, model, percent_to_timestep_function): super().pre_run_advanced(model, percent_to_timestep_function) + if type(self.cond_hint_original) == PreprocSparseRGBWrapper: + self.cond_hint_original = self.cond_hint_original.condhint self.latent_format = model.latent_format # LatentFormat object, used to process_in latent cond hint if self.control_model.motion_holder is not None: self.control_model.motion_holder.motion_wrapper.reset() diff --git a/control/control_sparsectrl.py b/control/control_sparsectrl.py index c3bb217..41698b4 100644 --- a/control/control_sparsectrl.py +++ b/control/control_sparsectrl.py @@ -79,6 +79,17 @@ class SparseControlNet(ControlNetCLDM): return outs +class PreprocSparseRGBWrapper: + def __init__(self, condhint: Tensor): + self.condhint = condhint + + def movedim(self, *args, **kwargs): + return self + + def __getattr__(self, name): + raise AttributeError("Invalid use of RGB SparseCtrl output. The output of RGB SparseCtrl preprocessor is NOT a usual image, but a latent pretending to be an image - you must connect the output directly to an Apply ControlNet node (advanced or otherwise).") + + class SparseSettings: def __init__(self, sparse_method: 'SparseMethod', use_motion: bool=True, motion_strength=1.0, motion_scale=1.0, merged=False): self.sparse_method = sparse_method diff --git a/control/nodes.py b/control/nodes.py index 5076062..3a071c9 100644 --- a/control/nodes.py +++ b/control/nodes.py @@ -9,7 +9,7 @@ from .utils import StrengthInterpolation as SI from .nodes_weight import (DefaultWeights, ScaledSoftMaskedUniversalWeights, ScaledSoftUniversalWeights, SoftControlNetWeights, CustomControlNetWeights, SoftT2IAdapterWeights, CustomT2IAdapterWeights) from .nodes_latent_keyframe import LatentKeyframeGroupNode, LatentKeyframeInterpolationNode, LatentKeyframeBatchedGroupNode, LatentKeyframeNode -from .nodes_sparsectrl import SparseCtrlMergedLoaderAdvanced, SparseCtrlLoaderAdvanced, SparseIndexMethodNode, SparseSpreadMethodNode, VAEEncodePreprocessor +from .nodes_sparsectrl import SparseCtrlMergedLoaderAdvanced, SparseCtrlLoaderAdvanced, SparseIndexMethodNode, SparseSpreadMethodNode, RgbSparseCtrlPreprocessor from .logger import logger @@ -215,7 +215,7 @@ NODE_CLASS_MAPPINGS = { "CustomT2IAdapterWeights": CustomT2IAdapterWeights, "ACN_DefaultUniversalWeights": DefaultWeights, # SparseCtrl - "ACN_VAEEncodePreprocessor": VAEEncodePreprocessor, + "ACN_SparseCtrlRGBPreprocessor": RgbSparseCtrlPreprocessor, "ACN_SparseCtrlLoaderAdvanced": SparseCtrlLoaderAdvanced, "ACN_SparseCtrlMergedLoaderAdvanced": SparseCtrlMergedLoaderAdvanced, "ACN_SparseCtrlIndexMethodNode": SparseIndexMethodNode, @@ -243,7 +243,7 @@ NODE_DISPLAY_NAME_MAPPINGS = { "CustomT2IAdapterWeights": "T2IAdapter Custom Weights 🛂🅐🅒🅝", "ACN_DefaultUniversalWeights": "Force Default Weights 🛂🅐🅒🅝", # SparseCtrl - "ACN_VAEEncodePreprocessor": "RGB SparseCtrl 🛂🅐🅒🅝", + "ACN_SparseCtrlRGBPreprocessor": "RGB SparseCtrl 🛂🅐🅒🅝", "ACN_SparseCtrlLoaderAdvanced": "Load SparseCtrl Model 🛂🅐🅒🅝", "ACN_SparseCtrlMergedLoaderAdvanced": "Load Merged SparseCtrl Model 🛂🅐🅒🅝", "ACN_SparseCtrlIndexMethodNode": "SparseCtrl Index Method 🛂🅐🅒🅝", diff --git a/control/nodes_sparsectrl.py b/control/nodes_sparsectrl.py index 0e713fe..dca2e3f 100644 --- a/control/nodes_sparsectrl.py +++ b/control/nodes_sparsectrl.py @@ -1,8 +1,11 @@ +from torch import Tensor + import folder_paths from nodes import VAEEncode +import comfy.utils from .utils import TimestepKeyframeGroup -from .control_sparsectrl import SparseMethod, SparseIndexMethod, SparseSettings, SparseSpreadMethod +from .control_sparsectrl import SparseMethod, SparseIndexMethod, SparseSettings, SparseSpreadMethod, PreprocSparseRGBWrapper from .control import load_sparsectrl, load_controlnet, ControlNetAdvanced, SparseCtrlAdvanced @@ -55,7 +58,7 @@ class SparseCtrlMergedLoaderAdvanced: RETURN_TYPES = ("CONTROL_NET", ) FUNCTION = "load_controlnet" - CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/SparseCtrl" + CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/SparseCtrl/experimental" def load_controlnet(self, sparsectrl_name: str, control_net_name: str, use_motion: bool, motion_strength: float, motion_scale: float, sparse_method: SparseMethod=SparseSpreadMethod(), tk_optional: TimestepKeyframeGroup=None): sparsectrl_path = folder_paths.get_full_path("controlnet", sparsectrl_name) @@ -128,24 +131,29 @@ class SparseSpreadMethodNode: return (SparseSpreadMethod(spread=spread),) -class VAEEncodePreprocessor: +class RgbSparseCtrlPreprocessor: @classmethod def INPUT_TYPES(s): return { "required": { "image": ("IMAGE", ), "vae": ("VAE", ), + "latent_size": ("LATENT", ), } } RETURN_TYPES = ("IMAGE",) - RETURN_NAMES = ("latent_IMAGE",) + RETURN_NAMES = ("proc_IMAGE",) FUNCTION = "preprocess_images" CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/SparseCtrl/preprocess" - def preprocess_images(self, vae, image): + def preprocess_images(self, vae, image: Tensor, latent_size: Tensor): + # first, resize image to match latents + image = image.movedim(-1,1) + image = comfy.utils.common_upscale(image, latent_size["samples"].shape[3] * 8, latent_size["samples"].shape[2] * 8, 'nearest-exact', "center") + image = image.movedim(1,-1) + # then, vae encode image = VAEEncode.vae_encode_crop_pixels(image) encoded = vae.encode(image[:,:,:,:3]) - encoded = encoded.movedim(1,-1) - return (encoded,) + return (PreprocSparseRGBWrapper(condhint=encoded),)