Automatically scale image before RGB Preproc does vae encoding, make sure RGB Preproc image isn't attempted to be used outside of Apply ControlNet nodes

This commit is contained in:
Jedrzej Kosinski
2023-12-22 11:12:23 -06:00
parent e674b995d0
commit fa0a3dba4c
4 changed files with 33 additions and 11 deletions
+4 -1
View File
@@ -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()
+11
View File
@@ -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
+3 -3
View File
@@ -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 🛂🅐🅒🅝",
+15 -7
View File
@@ -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),)