fluxdev/stable-diffusion-webui-forge
1
1import os2import torch3import ldm_patched.modules.utils4import ldm_patched.controlnet5 6from ldm_patched.modules.controlnet import ControlLora, ControlNet, load_t2i_adapter7from modules_forge.controlnet import apply_controlnet_advanced8from modules_forge.shared import add_supported_control_model9 10 11class ControlModelPatcher:12 @staticmethod13 def try_build_from_state_dict(state_dict, ckpt_path):14 return None15 16 def __init__(self, model_patcher=None):17 self.model_patcher = model_patcher18 self.strength = 1.019 self.start_percent = 0.020 self.end_percent = 1.021 self.positive_advanced_weighting = None22 self.negative_advanced_weighting = None23 self.advanced_frame_weighting = None24 self.advanced_sigma_weighting = None25 self.advanced_mask_weighting = None26 27 def process_after_running_preprocessors(self, process, params, *args, **kwargs):28 return29 30 def process_before_every_sampling(self, process, cond, mask, *args, **kwargs):31 return32 33 def process_after_every_sampling(self, process, params, *args, **kwargs):34 return35 36 37class ControlNetPatcher(ControlModelPatcher):38 @staticmethod39 def try_build_from_state_dict(controlnet_data, ckpt_path):40 if "lora_controlnet" in controlnet_data:41 return ControlNetPatcher(ControlLora(controlnet_data))42 43 controlnet_config = None44 if "controlnet_cond_embedding.conv_in.weight" in controlnet_data: # diffusers format45 unet_dtype = ldm_patched.modules.model_management.unet_dtype()46 controlnet_config = ldm_patched.modules.model_detection.unet_config_from_diffusers_unet(controlnet_data,47 unet_dtype)48 diffusers_keys = ldm_patched.modules.utils.unet_to_diffusers(controlnet_config)49 diffusers_keys["controlnet_mid_block.weight"] = "middle_block_out.0.weight"50 diffusers_keys["controlnet_mid_block.bias"] = "middle_block_out.0.bias"51 52 count = 053 loop = True54 while loop:55 suffix = [".weight", ".bias"]56 for s in suffix:57 k_in = "controlnet_down_blocks.{}{}".format(count, s)58 k_out = "zero_convs.{}.0{}".format(count, s)59 if k_in not in controlnet_data:60 loop = False61 break62 diffusers_keys[k_in] = k_out63 count += 164 65 count = 066 loop = True67 while loop:68 suffix = [".weight", ".bias"]69 for s in suffix:70 if count == 0:71 k_in = "controlnet_cond_embedding.conv_in{}".format(s)72 else:73 k_in = "controlnet_cond_embedding.blocks.{}{}".format(count - 1, s)74 k_out = "input_hint_block.{}{}".format(count * 2, s)75 if k_in not in controlnet_data:76 k_in = "controlnet_cond_embedding.conv_out{}".format(s)77 loop = False78 diffusers_keys[k_in] = k_out79 count += 180 81 new_sd = {}82 for k in diffusers_keys:83 if k in controlnet_data:84 new_sd[diffusers_keys[k]] = controlnet_data.pop(k)85 86 leftover_keys = controlnet_data.keys()87 if len(leftover_keys) > 0:88 print("leftover keys:", leftover_keys)89 controlnet_data = new_sd90 91 pth_key = 'control_model.zero_convs.0.0.weight'92 pth = False93 key = 'zero_convs.0.0.weight'94 if pth_key in controlnet_data:95 pth = True96 key = pth_key97 prefix = "control_model."98 elif key in controlnet_data:99 prefix = ""100 else:101 net = load_t2i_adapter(controlnet_data)102 if net is None:103 return None104 return ControlNetPatcher(net)105 106 if controlnet_config is None:107 unet_dtype = ldm_patched.modules.model_management.unet_dtype()108 controlnet_config = ldm_patched.modules.model_detection.model_config_from_unet(controlnet_data, prefix,109 unet_dtype, True).unet_config110 load_device = ldm_patched.modules.model_management.get_torch_device()111 manual_cast_dtype = ldm_patched.modules.model_management.unet_manual_cast(unet_dtype, load_device)112 if manual_cast_dtype is not None:113 controlnet_config["operations"] = ldm_patched.modules.ops.manual_cast114 controlnet_config.pop("out_channels")115 controlnet_config["hint_channels"] = controlnet_data["{}input_hint_block.0.weight".format(prefix)].shape[1]116 control_model = ldm_patched.controlnet.cldm.ControlNet(**controlnet_config)117 118 if pth:119 if 'difference' in controlnet_data:120 print("WARNING: Your controlnet model is diff version rather than official float16 model. "121 "Please use an official float16/float32 model for robust performance.")122 123 class WeightsLoader(torch.nn.Module):124 pass125 126 w = WeightsLoader()127 w.control_model = control_model128 missing, unexpected = w.load_state_dict(controlnet_data, strict=False)129 else:130 missing, unexpected = control_model.load_state_dict(controlnet_data, strict=False)131 print(missing, unexpected)132 133 global_average_pooling = False134 filename = os.path.splitext(ckpt_path)[0]135 if filename.endswith("_shuffle") or filename.endswith("_shuffle_fp16"):136 # TODO: smarter way of enabling global_average_pooling137 global_average_pooling = True138 139 control = ControlNet(control_model, global_average_pooling=global_average_pooling, load_device=load_device,140 manual_cast_dtype=manual_cast_dtype)141 return ControlNetPatcher(control)142 143 def __init__(self, model_patcher):144 super().__init__(model_patcher)145 146 def process_before_every_sampling(self, process, cond, mask, *args, **kwargs):147 unet = process.sd_model.forge_objects.unet148 149 unet = apply_controlnet_advanced(150 unet=unet,151 controlnet=self.model_patcher,152 image_bchw=cond,153 strength=self.strength,154 start_percent=self.start_percent,155 end_percent=self.end_percent,156 positive_advanced_weighting=self.positive_advanced_weighting,157 negative_advanced_weighting=self.negative_advanced_weighting,158 advanced_frame_weighting=self.advanced_frame_weighting,159 advanced_sigma_weighting=self.advanced_sigma_weighting,160 advanced_mask_weighting=self.advanced_mask_weighting161 )162 163 process.sd_model.forge_objects.unet = unet164 return165 166 167add_supported_control_model(ControlNetPatcher)168 