hugging-apps/echo-memory
0
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 