Team Ai
Apppublic

hugging-apps/echo-memory

sourceHugging Faceupdated 3mo agoView on Hugging Face
0likes
flux_controlnet.py332 linesDownload Raw Back to models
1import torch2from einops import rearrange, repeat3from .flux_dit import RoPEEmbedding, TimestepEmbeddings, FluxJointTransformerBlock, FluxSingleTransformerBlock, RMSNorm4from .utils import hash_state_dict_keys, init_weights_on_device5 6 7 8class FluxControlNet(torch.nn.Module):9    def __init__(self, disable_guidance_embedder=False, num_joint_blocks=5, num_single_blocks=10, num_mode=0, mode_dict={}, additional_input_dim=0):10        super().__init__()11        self.pos_embedder = RoPEEmbedding(3072, 10000, [16, 56, 56])12        self.time_embedder = TimestepEmbeddings(256, 3072)13        self.guidance_embedder = None if disable_guidance_embedder else TimestepEmbeddings(256, 3072)14        self.pooled_text_embedder = torch.nn.Sequential(torch.nn.Linear(768, 3072), torch.nn.SiLU(), torch.nn.Linear(3072, 3072))15        self.context_embedder = torch.nn.Linear(4096, 3072)16        self.x_embedder = torch.nn.Linear(64, 3072)17 18        self.blocks = torch.nn.ModuleList([FluxJointTransformerBlock(3072, 24) for _ in range(num_joint_blocks)])19        self.single_blocks = torch.nn.ModuleList([FluxSingleTransformerBlock(3072, 24) for _ in range(num_single_blocks)])20 21        self.controlnet_blocks = torch.nn.ModuleList([torch.nn.Linear(3072, 3072) for _ in range(num_joint_blocks)])22        self.controlnet_single_blocks = torch.nn.ModuleList([torch.nn.Linear(3072, 3072) for _ in range(num_single_blocks)])23        24        self.mode_dict = mode_dict25        self.controlnet_mode_embedder = torch.nn.Embedding(num_mode, 3072) if len(mode_dict) > 0 else None26        self.controlnet_x_embedder = torch.nn.Linear(64 + additional_input_dim, 3072)27 28 29    def prepare_image_ids(self, latents):30        batch_size, _, height, width = latents.shape31        latent_image_ids = torch.zeros(height // 2, width // 2, 3)32        latent_image_ids[..., 1] = latent_image_ids[..., 1] + torch.arange(height // 2)[:, None]33        latent_image_ids[..., 2] = latent_image_ids[..., 2] + torch.arange(width // 2)[None, :]34 35        latent_image_id_height, latent_image_id_width, latent_image_id_channels = latent_image_ids.shape36 37        latent_image_ids = latent_image_ids[None, :].repeat(batch_size, 1, 1, 1)38        latent_image_ids = latent_image_ids.reshape(39            batch_size, latent_image_id_height * latent_image_id_width, latent_image_id_channels40        )41        latent_image_ids = latent_image_ids.to(device=latents.device, dtype=latents.dtype)42 43        return latent_image_ids44    45 46    def patchify(self, hidden_states):47        hidden_states = rearrange(hidden_states, "B C (H P) (W Q) -> B (H W) (C P Q)", P=2, Q=2)48        return hidden_states49    50 51    def align_res_stack_to_original_blocks(self, res_stack, num_blocks, hidden_states):52        if len(res_stack) == 0:53            return [torch.zeros_like(hidden_states)] * num_blocks54        interval = (num_blocks + len(res_stack) - 1) // len(res_stack)55        aligned_res_stack = [res_stack[block_id // interval] for block_id in range(num_blocks)]56        return aligned_res_stack57 58 59    def forward(60        self,61        hidden_states,62        controlnet_conditioning,63        timestep, prompt_emb, pooled_prompt_emb, guidance, text_ids, image_ids=None,64        processor_id=None,65        tiled=False, tile_size=128, tile_stride=64,66        **kwargs67    ):68        if image_ids is None:69            image_ids = self.prepare_image_ids(hidden_states)70 71        conditioning = self.time_embedder(timestep, hidden_states.dtype) + self.pooled_text_embedder(pooled_prompt_emb)72        if self.guidance_embedder is not None:73            guidance = guidance * 100074            conditioning = conditioning + self.guidance_embedder(guidance, hidden_states.dtype)75        prompt_emb = self.context_embedder(prompt_emb)76        if self.controlnet_mode_embedder is not None: # Different from FluxDiT77            processor_id = torch.tensor([self.mode_dict[processor_id]], dtype=torch.int)78            processor_id = repeat(processor_id, "D -> B D", B=1).to(text_ids.device)79            prompt_emb = torch.concat([self.controlnet_mode_embedder(processor_id), prompt_emb], dim=1)80            text_ids = torch.cat([text_ids[:, :1], text_ids], dim=1)81        image_rotary_emb = self.pos_embedder(torch.cat((text_ids, image_ids), dim=1))82 83        hidden_states = self.patchify(hidden_states)84        hidden_states = self.x_embedder(hidden_states)85        controlnet_conditioning = self.patchify(controlnet_conditioning) # Different from FluxDiT86        hidden_states = hidden_states + self.controlnet_x_embedder(controlnet_conditioning) # Different from FluxDiT87 88        controlnet_res_stack = []89        for block, controlnet_block in zip(self.blocks, self.controlnet_blocks):90            hidden_states, prompt_emb = block(hidden_states, prompt_emb, conditioning, image_rotary_emb)91            controlnet_res_stack.append(controlnet_block(hidden_states))92 93        controlnet_single_res_stack = []94        hidden_states = torch.cat([prompt_emb, hidden_states], dim=1)95        for block, controlnet_block in zip(self.single_blocks, self.controlnet_single_blocks):96            hidden_states, prompt_emb = block(hidden_states, prompt_emb, conditioning, image_rotary_emb)97            controlnet_single_res_stack.append(controlnet_block(hidden_states[:, prompt_emb.shape[1]:]))98 99        controlnet_res_stack = self.align_res_stack_to_original_blocks(controlnet_res_stack, 19, hidden_states[:, prompt_emb.shape[1]:])100        controlnet_single_res_stack = self.align_res_stack_to_original_blocks(controlnet_single_res_stack, 38, hidden_states[:, prompt_emb.shape[1]:])101 102        return controlnet_res_stack, controlnet_single_res_stack103 104 105    @staticmethod106    def state_dict_converter():107        return FluxControlNetStateDictConverter()108    109    def quantize(self):110        def cast_to(weight, dtype=None, device=None, copy=False):111            if device is None or weight.device == device:112                if not copy:113                    if dtype is None or weight.dtype == dtype:114                        return weight115                return weight.to(dtype=dtype, copy=copy)116 117            r = torch.empty_like(weight, dtype=dtype, device=device)118            r.copy_(weight)119            return r120 121        def cast_weight(s, input=None, dtype=None, device=None):122            if input is not None:123                if dtype is None:124                    dtype = input.dtype125                if device is None:126                    device = input.device127            weight = cast_to(s.weight, dtype, device)128            return weight129 130        def cast_bias_weight(s, input=None, dtype=None, device=None, bias_dtype=None):131            if input is not None:132                if dtype is None:133                    dtype = input.dtype134                if bias_dtype is None:135                    bias_dtype = dtype136                if device is None:137                    device = input.device138            bias = None139            weight = cast_to(s.weight, dtype, device)140            bias = cast_to(s.bias, bias_dtype, device)141            return weight, bias142 143        class quantized_layer:144            class QLinear(torch.nn.Linear):145                def __init__(self, *args, **kwargs):146                    super().__init__(*args, **kwargs)147                    148                def forward(self,input,**kwargs):149                    weight,bias= cast_bias_weight(self,input)150                    return torch.nn.functional.linear(input,weight,bias)151            152            class QRMSNorm(torch.nn.Module):153                def __init__(self, module):154                    super().__init__()155                    self.module = module156                    157                def forward(self,hidden_states,**kwargs):158                    weight= cast_weight(self.module,hidden_states)159                    input_dtype = hidden_states.dtype160                    variance = hidden_states.to(torch.float32).square().mean(-1, keepdim=True)161                    hidden_states = hidden_states * torch.rsqrt(variance + self.module.eps)162                    hidden_states = hidden_states.to(input_dtype) * weight163                    return hidden_states164            165            class QEmbedding(torch.nn.Embedding):166                def __init__(self, *args, **kwargs):167                    super().__init__(*args, **kwargs)168                    169                def forward(self,input,**kwargs):170                    weight= cast_weight(self,input)171                    return torch.nn.functional.embedding(172                        input, weight, self.padding_idx, self.max_norm,173                        self.norm_type, self.scale_grad_by_freq, self.sparse)174            175        def replace_layer(model):176            for name, module in model.named_children():177                if isinstance(module,quantized_layer.QRMSNorm):178                    continue179                if isinstance(module, torch.nn.Linear):180                    with init_weights_on_device():181                        new_layer = quantized_layer.QLinear(module.in_features,module.out_features)182                    new_layer.weight = module.weight183                    if module.bias is not None:184                        new_layer.bias = module.bias185                    setattr(model, name, new_layer)186                elif isinstance(module, RMSNorm):187                    if hasattr(module,"quantized"):188                        continue189                    module.quantized= True190                    new_layer = quantized_layer.QRMSNorm(module)191                    setattr(model, name, new_layer)192                elif isinstance(module,torch.nn.Embedding):193                    rows, cols = module.weight.shape194                    new_layer = quantized_layer.QEmbedding(195                        num_embeddings=rows,196                        embedding_dim=cols,197                        _weight=module.weight,198                        # _freeze=module.freeze,199                        padding_idx=module.padding_idx,200                        max_norm=module.max_norm,201                        norm_type=module.norm_type,202                        scale_grad_by_freq=module.scale_grad_by_freq,203                        sparse=module.sparse)204                    setattr(model, name, new_layer)205                else:206                    replace_layer(module)207 208        replace_layer(self)209    210 211 212class FluxControlNetStateDictConverter:213    def __init__(self):214        pass215 216    def from_diffusers(self, state_dict):217        hash_value = hash_state_dict_keys(state_dict)218        global_rename_dict = {219            "context_embedder": "context_embedder",220            "x_embedder": "x_embedder",221            "time_text_embed.timestep_embedder.linear_1": "time_embedder.timestep_embedder.0",222            "time_text_embed.timestep_embedder.linear_2": "time_embedder.timestep_embedder.2",223            "time_text_embed.guidance_embedder.linear_1": "guidance_embedder.timestep_embedder.0",224            "time_text_embed.guidance_embedder.linear_2": "guidance_embedder.timestep_embedder.2",225            "time_text_embed.text_embedder.linear_1": "pooled_text_embedder.0",226            "time_text_embed.text_embedder.linear_2": "pooled_text_embedder.2",227            "norm_out.linear": "final_norm_out.linear",228            "proj_out": "final_proj_out",229        }230        rename_dict = {231            "proj_out": "proj_out",232            "norm1.linear": "norm1_a.linear",233            "norm1_context.linear": "norm1_b.linear",234            "attn.to_q": "attn.a_to_q",235            "attn.to_k": "attn.a_to_k",236            "attn.to_v": "attn.a_to_v",237            "attn.to_out.0": "attn.a_to_out",238            "attn.add_q_proj": "attn.b_to_q",239            "attn.add_k_proj": "attn.b_to_k",240            "attn.add_v_proj": "attn.b_to_v",241            "attn.to_add_out": "attn.b_to_out",242            "ff.net.0.proj": "ff_a.0",243            "ff.net.2": "ff_a.2",244            "ff_context.net.0.proj": "ff_b.0",245            "ff_context.net.2": "ff_b.2",246            "attn.norm_q": "attn.norm_q_a",247            "attn.norm_k": "attn.norm_k_a",248            "attn.norm_added_q": "attn.norm_q_b",249            "attn.norm_added_k": "attn.norm_k_b",250        }251        rename_dict_single = {252            "attn.to_q": "a_to_q",253            "attn.to_k": "a_to_k",254            "attn.to_v": "a_to_v",255            "attn.norm_q": "norm_q_a",256            "attn.norm_k": "norm_k_a",257            "norm.linear": "norm.linear",258            "proj_mlp": "proj_in_besides_attn",259            "proj_out": "proj_out",260        }261        state_dict_ = {}262        for name, param in state_dict.items():263            if name.endswith(".weight") or name.endswith(".bias"):264                suffix = ".weight" if name.endswith(".weight") else ".bias"265                prefix = name[:-len(suffix)]266                if prefix in global_rename_dict:267                    state_dict_[global_rename_dict[prefix] + suffix] = param268                elif prefix.startswith("transformer_blocks."):269                    names = prefix.split(".")270                    names[0] = "blocks"271                    middle = ".".join(names[2:])272                    if middle in rename_dict:273                        name_ = ".".join(names[:2] + [rename_dict[middle]] + [suffix[1:]])274                        state_dict_[name_] = param275                elif prefix.startswith("single_transformer_blocks."):276                    names = prefix.split(".")277                    names[0] = "single_blocks"278                    middle = ".".join(names[2:])279                    if middle in rename_dict_single:280                        name_ = ".".join(names[:2] + [rename_dict_single[middle]] + [suffix[1:]])281                        state_dict_[name_] = param282                    else:283                        state_dict_[name] = param284                else:285                    state_dict_[name] = param286        for name in list(state_dict_.keys()):287            if ".proj_in_besides_attn." in name:288                name_ = name.replace(".proj_in_besides_attn.", ".to_qkv_mlp.")289                param = torch.concat([290                    state_dict_[name.replace(".proj_in_besides_attn.", f".a_to_q.")],291                    state_dict_[name.replace(".proj_in_besides_attn.", f".a_to_k.")],292                    state_dict_[name.replace(".proj_in_besides_attn.", f".a_to_v.")],293                    state_dict_[name],294                ], dim=0)295                state_dict_[name_] = param296                state_dict_.pop(name.replace(".proj_in_besides_attn.", f".a_to_q."))297                state_dict_.pop(name.replace(".proj_in_besides_attn.", f".a_to_k."))298                state_dict_.pop(name.replace(".proj_in_besides_attn.", f".a_to_v."))299                state_dict_.pop(name)300        for name in list(state_dict_.keys()):301            for component in ["a", "b"]:302                if f".{component}_to_q." in name:303                    name_ = name.replace(f".{component}_to_q.", f".{component}_to_qkv.")304                    param = torch.concat([305                        state_dict_[name.replace(f".{component}_to_q.", f".{component}_to_q.")],306                        state_dict_[name.replace(f".{component}_to_q.", f".{component}_to_k.")],307                        state_dict_[name.replace(f".{component}_to_q.", f".{component}_to_v.")],308                    ], dim=0)309                    state_dict_[name_] = param310                    state_dict_.pop(name.replace(f".{component}_to_q.", f".{component}_to_q."))311                    state_dict_.pop(name.replace(f".{component}_to_q.", f".{component}_to_k."))312                    state_dict_.pop(name.replace(f".{component}_to_q.", f".{component}_to_v."))313        if hash_value == "78d18b9101345ff695f312e7e62538c0":314            extra_kwargs = {"num_mode": 10, "mode_dict": {"canny": 0, "tile": 1, "depth": 2, "blur": 3, "pose": 4, "gray": 5, "lq": 6}}315        elif hash_value == "b001c89139b5f053c715fe772362dd2a":316            extra_kwargs = {"num_single_blocks": 0}317        elif hash_value == "52357cb26250681367488a8954c271e8":318            extra_kwargs = {"num_joint_blocks": 6, "num_single_blocks": 0, "additional_input_dim": 4}319        elif hash_value == "0cfd1740758423a2a854d67c136d1e8c":320            extra_kwargs = {"num_joint_blocks": 4, "num_single_blocks": 1}321        elif hash_value == "7f9583eb8ba86642abb9a21a4b2c9e16":322            extra_kwargs = {"num_joint_blocks": 4, "num_single_blocks": 10}323        elif hash_value == "43ad5aaa27dd4ee01b832ed16773fa52":324            extra_kwargs = {"num_joint_blocks": 6, "num_single_blocks": 0}325        else:326            extra_kwargs = {}327        return state_dict_, extra_kwargs328    329 330    def from_civitai(self, state_dict):331        return self.from_diffusers(state_dict)332