CoolFace
Apppublic

fluxdev/stable-diffusion-webui-forge

sourceHugging Faceupdated 2y agoView on Hugging Face
1likes
supported_controlnet.py168 linesDownload Raw Back to modules_forge
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