diff --git a/nodes.py b/nodes.py index bc69772..dc2dfe0 100644 --- a/nodes.py +++ b/nodes.py @@ -7,13 +7,9 @@ import numpy as np import torch.nn.functional as F from comfy.utils import load_torch_file from .utils.convert_unet import convert_iclight_unet -from .utils.patches import calculate_weight_adjust_channel from .utils.image import generate_gradient_image, LightPosition from nodes import MAX_RESOLUTION -from comfy.model_patcher import ModelPatcher -from comfy import lora import model_management -import logging class LoadAndApplyICLightUnet: @classmethod @@ -38,6 +34,8 @@ Used with ICLightConditioning -node def load(self, model, model_path): type_str = str(type(model.model.model_config).__name__) + device = model_management.get_torch_device() + dtype = model_management.unet_dtype() if "SD15" not in type_str: raise Exception(f"Attempted to load {type_str} model, IC-Light is only compatible with SD 1.5 models.") @@ -50,33 +48,30 @@ Used with ICLightConditioning -node model_clone = model.clone() iclight_state_dict = load_torch_file(model_full_path) - + print("LoadAndApplyICLightUnet: Attempting to add patches with IC-Light Unet weights") try: if 'conv_in.weight' in iclight_state_dict: iclight_state_dict = convert_iclight_unet(iclight_state_dict) - in_channels = iclight_state_dict["diffusion_model.input_blocks.0.0.weight"].shape[1] - for key in iclight_state_dict: - model_clone.add_patches({key: (iclight_state_dict[key],)}, 1.0, 1.0) + prefix = "" else: - for key in iclight_state_dict: - model_clone.add_patches({"diffusion_model." + key: (iclight_state_dict[key],)}, 1.0, 1.0) + prefix = "diffusion_model." - in_channels = iclight_state_dict["input_blocks.0.0.weight"].shape[1] + patches={ + (prefix + key): ( + "diff", + [value.to(dtype=dtype, device=device), + {"pad_weight": key == "diffusion_model.input_blocks.0.0.weight" or key == "input_blocks.0.0.weight"},], + ) + for key, value in iclight_state_dict.items() + } + model_clone.add_patches(patches) + except: raise Exception("Could not patch model") print("LoadAndApplyICLightUnet: Added LoadICLightUnet patches") - #Patch ComfyUI's LoRA weight application to accept multi-channel inputs. Thanks @huchenlei - try: - if hasattr(lora, 'calculate_weight'): - lora.calculate_weight = calculate_weight_adjust_channel(lora.calculate_weight) - else: - raise Exception("IC-Light: The 'calculate_weight' function does not exist in 'lora'") - except Exception as e: - raise Exception(f"IC-Light: Could not patch calculate_weight - {str(e)}") - # Mimic the existing IP2P class to enable extra_conds def bound_extra_conds(self, **kwargs): return ICLight.extra_conds(self, **kwargs) @@ -84,7 +79,7 @@ Used with ICLightConditioning -node model_clone.add_object_patch("extra_conds", new_extra_conds) - model_clone.model.model_config.unet_config["in_channels"] = in_channels + #model_clone.model.model_config.unet_config["in_channels"] = in_channels return (model_clone, ) diff --git a/utils/patches.py b/utils/patches.py deleted file mode 100644 index 439d191..0000000 --- a/utils/patches.py +++ /dev/null @@ -1,64 +0,0 @@ - -#credit to huchenlei for this -#from https://github.com/huchenlei/ComfyUI-layerdiffuse/blob/151f7460bbc9d7437d4f0010f21f80178f7a84a6/layered_diffusion.py#L34-L96 - -import torch -import functools -from comfy.model_patcher import ModelPatcher -import comfy.model_management - -def calculate_weight_adjust_channel(func): - """Patches ComfyUI's LoRA weight application to accept multi-channel inputs.""" - - @functools.wraps(func) - def calculate_weight(patches, weight: torch.Tensor, key: str, intermediate_dtype=torch.float32) -> torch.Tensor: - weight = func(patches, weight, key, intermediate_dtype) - - for p in patches: - alpha = p[0] - v = p[1] - - # The recursion call should be handled in the main func call. - if isinstance(v, list): - continue - - if len(v) == 1: - patch_type = "diff" - elif len(v) == 2: - patch_type = v[0] - v = v[1] - - if patch_type == "diff": - w1 = v[0] - if all( - ( - alpha != 0.0, - w1.shape != weight.shape, - w1.ndim == weight.ndim == 4, - ) - ): - new_shape = [max(n, m) for n, m in zip(weight.shape, w1.shape)] - print( - f"IC-Light: Merged with {key} channel changed from {weight.shape} to {new_shape}" - ) - new_diff = alpha * comfy.model_management.cast_to_device( - w1, weight.device, weight.dtype - ) - new_weight = torch.zeros(size=new_shape).to(weight) - new_weight[ - : weight.shape[0], - : weight.shape[1], - : weight.shape[2], - : weight.shape[3], - ] = weight - new_weight[ - : new_diff.shape[0], - : new_diff.shape[1], - : new_diff.shape[2], - : new_diff.shape[3], - ] += new_diff - new_weight = new_weight.contiguous().clone() - weight = new_weight - return weight - - return calculate_weight