Update to use Comfy API channel padding, better compatibility for current ComfyUI version
This commit is contained in:
@@ -7,13 +7,9 @@ import numpy as np
|
|||||||
import torch.nn.functional as F
|
import torch.nn.functional as F
|
||||||
from comfy.utils import load_torch_file
|
from comfy.utils import load_torch_file
|
||||||
from .utils.convert_unet import convert_iclight_unet
|
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 .utils.image import generate_gradient_image, LightPosition
|
||||||
from nodes import MAX_RESOLUTION
|
from nodes import MAX_RESOLUTION
|
||||||
from comfy.model_patcher import ModelPatcher
|
|
||||||
from comfy import lora
|
|
||||||
import model_management
|
import model_management
|
||||||
import logging
|
|
||||||
|
|
||||||
class LoadAndApplyICLightUnet:
|
class LoadAndApplyICLightUnet:
|
||||||
@classmethod
|
@classmethod
|
||||||
@@ -38,6 +34,8 @@ Used with ICLightConditioning -node
|
|||||||
|
|
||||||
def load(self, model, model_path):
|
def load(self, model, model_path):
|
||||||
type_str = str(type(model.model.model_config).__name__)
|
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:
|
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.")
|
raise Exception(f"Attempted to load {type_str} model, IC-Light is only compatible with SD 1.5 models.")
|
||||||
|
|
||||||
@@ -55,28 +53,25 @@ Used with ICLightConditioning -node
|
|||||||
try:
|
try:
|
||||||
if 'conv_in.weight' in iclight_state_dict:
|
if 'conv_in.weight' in iclight_state_dict:
|
||||||
iclight_state_dict = convert_iclight_unet(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]
|
prefix = ""
|
||||||
for key in iclight_state_dict:
|
|
||||||
model_clone.add_patches({key: (iclight_state_dict[key],)}, 1.0, 1.0)
|
|
||||||
else:
|
else:
|
||||||
for key in iclight_state_dict:
|
prefix = "diffusion_model."
|
||||||
model_clone.add_patches({"diffusion_model." + key: (iclight_state_dict[key],)}, 1.0, 1.0)
|
|
||||||
|
|
||||||
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:
|
except:
|
||||||
raise Exception("Could not patch model")
|
raise Exception("Could not patch model")
|
||||||
print("LoadAndApplyICLightUnet: Added LoadICLightUnet patches")
|
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
|
# Mimic the existing IP2P class to enable extra_conds
|
||||||
def bound_extra_conds(self, **kwargs):
|
def bound_extra_conds(self, **kwargs):
|
||||||
return ICLight.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.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, )
|
return (model_clone, )
|
||||||
|
|
||||||
|
|||||||
@@ -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
|
|
||||||
Reference in New Issue
Block a user