Team Ai
Apppublic

hugging-apps/echo-memory

sourceHugging Faceupdated 3mo agoView on Hugging Face
0likes
flux_ipadapter.py95 linesDownload Raw Back to models
1from .svd_image_encoder import SVDImageEncoder2from .sd3_dit import RMSNorm3from transformers import CLIPImageProcessor4import torch5 6 7class MLPProjModel(torch.nn.Module):8    def __init__(self, cross_attention_dim=768, id_embeddings_dim=512, num_tokens=4):9        super().__init__()10        11        self.cross_attention_dim = cross_attention_dim12        self.num_tokens = num_tokens13        14        self.proj = torch.nn.Sequential(15            torch.nn.Linear(id_embeddings_dim, id_embeddings_dim*2),16            torch.nn.GELU(),17            torch.nn.Linear(id_embeddings_dim*2, cross_attention_dim*num_tokens),18        )19        self.norm = torch.nn.LayerNorm(cross_attention_dim)20        21    def forward(self, id_embeds):22        x = self.proj(id_embeds)23        x = x.reshape(-1, self.num_tokens, self.cross_attention_dim)24        x = self.norm(x)25        return x26 27class IpAdapterModule(torch.nn.Module):28    def __init__(self, num_attention_heads, attention_head_dim, input_dim):29        super().__init__()30        self.num_heads = num_attention_heads31        self.head_dim = attention_head_dim32        output_dim = num_attention_heads * attention_head_dim33        self.to_k_ip = torch.nn.Linear(input_dim, output_dim, bias=False)34        self.to_v_ip = torch.nn.Linear(input_dim, output_dim, bias=False)35        self.norm_added_k = RMSNorm(attention_head_dim, eps=1e-5, elementwise_affine=False)36        37 38    def forward(self, hidden_states):39        batch_size = hidden_states.shape[0]40        # ip_k41        ip_k = self.to_k_ip(hidden_states)42        ip_k = ip_k.view(batch_size, -1, self.num_heads, self.head_dim).transpose(1, 2)43        ip_k = self.norm_added_k(ip_k)44        # ip_v45        ip_v = self.to_v_ip(hidden_states)46        ip_v = ip_v.view(batch_size, -1, self.num_heads, self.head_dim).transpose(1, 2)47        return ip_k, ip_v48 49 50class FluxIpAdapter(torch.nn.Module):51    def __init__(self, num_attention_heads=24, attention_head_dim=128, cross_attention_dim=4096, num_tokens=128, num_blocks=57):52        super().__init__()53        self.ipadapter_modules = torch.nn.ModuleList([IpAdapterModule(num_attention_heads, attention_head_dim, cross_attention_dim) for _ in range(num_blocks)])54        self.image_proj = MLPProjModel(cross_attention_dim=cross_attention_dim, id_embeddings_dim=1152, num_tokens=num_tokens)55        self.set_adapter()56 57    def set_adapter(self):58        self.call_block_id = {i:i for i in range(len(self.ipadapter_modules))}59 60    def forward(self, hidden_states, scale=1.0):61        hidden_states = self.image_proj(hidden_states)62        hidden_states = hidden_states.view(1, -1, hidden_states.shape[-1])63        ip_kv_dict = {}64        for block_id in self.call_block_id:65            ipadapter_id = self.call_block_id[block_id]66            ip_k, ip_v = self.ipadapter_modules[ipadapter_id](hidden_states)67            ip_kv_dict[block_id] = {68                "ip_k": ip_k,69                "ip_v": ip_v,70                "scale": scale71            }72        return ip_kv_dict73 74    @staticmethod75    def state_dict_converter():76        return FluxIpAdapterStateDictConverter()77 78 79class FluxIpAdapterStateDictConverter:80    def __init__(self):81        pass82 83    def from_diffusers(self, state_dict):84        state_dict_ = {}85        for name in state_dict["ip_adapter"]:86            name_ = 'ipadapter_modules.' + name87            state_dict_[name_] = state_dict["ip_adapter"][name]88        for name in state_dict["image_proj"]:89            name_ = "image_proj." + name90            state_dict_[name_] = state_dict["image_proj"][name]91        return state_dict_92    93    def from_civitai(self, state_dict):94        return self.from_diffusers(state_dict)95