diff --git a/nodes.py b/nodes.py index a01b181..b58a791 100644 --- a/nodes.py +++ b/nodes.py @@ -4,6 +4,8 @@ import os import types from comfy.utils import load_torch_file from .utils.convert_unet import convert_iclight_unet +from .utils.patches import calculate_weight_adjust_channel +from comfy.model_patcher import ModelPatcher class LoadAndApplyICLightUnet: @classmethod @@ -58,26 +60,28 @@ Used with ICLightConditioning -node raise Exception("Could not patch model") print("LoadAndApplyICLightUnet: Added LoadICLightUnet patches") - # Create a new Conv2d layer with 8 or 12 input channels - original_conv_layer = model_clone.model.diffusion_model.input_blocks[0][0] + # # Create a new Conv2d layer with 8 or 12 input channels + # original_conv_layer = model_clone.model.diffusion_model.input_blocks[0][0] - print(f"LoadAndApplyICLightUnet: Input channels in currently loaded model: {original_conv_layer.in_channels}") + # print(f"LoadAndApplyICLightUnet: Input channels in currently loaded model: {original_conv_layer.in_channels}") - print("LoadAndApplyICLightUnet: Settings in_channels to: ", in_channels) + # print("LoadAndApplyICLightUnet: Settings in_channels to: ", in_channels) - if model_clone.model.diffusion_model.input_blocks[0][0].in_channels != in_channels: - num_channels_to_copy = min(in_channels, original_conv_layer.in_channels) - new_conv_layer = torch.nn.Conv2d(in_channels, original_conv_layer.out_channels, kernel_size=original_conv_layer.kernel_size, stride=original_conv_layer.stride, padding=original_conv_layer.padding) - new_conv_layer.weight.zero_() - new_conv_layer.weight[:, :num_channels_to_copy, :, :].copy_(original_conv_layer.weight[:, :num_channels_to_copy, :, :]) - new_conv_layer.bias = original_conv_layer.bias - new_conv_layer = new_conv_layer.to(model_clone.model.diffusion_model.dtype) - original_conv_layer.conv_in = new_conv_layer - # Replace the old layer with the new one - model_clone.model.diffusion_model.input_blocks[0][0] = new_conv_layer - # Verify the change - print(f"LoadAndApplyICLightUnet: New number of input channels: {model_clone.model.diffusion_model.input_blocks[0][0].in_channels}") - + # if model_clone.model.diffusion_model.input_blocks[0][0].in_channels != in_channels: + # num_channels_to_copy = min(in_channels, original_conv_layer.in_channels) + # new_conv_layer = torch.nn.Conv2d(in_channels, original_conv_layer.out_channels, kernel_size=original_conv_layer.kernel_size, stride=original_conv_layer.stride, padding=original_conv_layer.padding) + # new_conv_layer.weight.zero_() + # new_conv_layer.weight[:, :num_channels_to_copy, :, :].copy_(original_conv_layer.weight[:, :num_channels_to_copy, :, :]) + # new_conv_layer.bias = original_conv_layer.bias + # new_conv_layer = new_conv_layer.to(model_clone.model.diffusion_model.dtype) + # original_conv_layer.conv_in = new_conv_layer + # # Replace the old layer with the new one + # model_clone.model.diffusion_model.input_blocks[0][0] = new_conv_layer + # # Verify the change + # print(f"LoadAndApplyICLightUnet: New number of input channels: {model_clone.model.diffusion_model.input_blocks[0][0].in_channels}") + + #Patch ComfyUI's LoRA weight application to accept multi-channel inputs. Thanks @huchenlei + ModelPatcher.calculate_weight = calculate_weight_adjust_channel(ModelPatcher.calculate_weight) # Mimic the existing IP2P class to enable extra_conds def bound_extra_conds(self, **kwargs): return ICLight.extra_conds(self, **kwargs) diff --git a/utils/patches.py b/utils/patches.py new file mode 100644 index 0000000..10429b6 --- /dev/null +++ b/utils/patches.py @@ -0,0 +1,66 @@ + +#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( + self: ModelPatcher, patches, weight: torch.Tensor, key: str + ) -> torch.Tensor: + weight = func(self, patches, weight, key) + + 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"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 \ No newline at end of file