| 107 | |
| 108 | |
| 109 | class FusionFaceId(ModelMixin): |
| 110 | def __init__(self, cross_attention_dim=768, id_embeddings_dim=512, clip_embeddings_dim=1024, num_tokens=4): |
| 111 | super().__init__() |
| 112 | self.cross_attention_dim = cross_attention_dim |
| 113 | self.num_tokens = num_tokens |
| 114 | |
| 115 | self.proj = torch.nn.Sequential( |
| 116 | torch.nn.Linear(id_embeddings_dim, id_embeddings_dim*2), |
| 117 | torch.nn.GELU(), |
| 118 | torch.nn.Linear(id_embeddings_dim*2, cross_attention_dim*num_tokens), |
| 119 | ) |
| 120 | |
| 121 | self.norm = torch.nn.LayerNorm(cross_attention_dim) |
| 122 | |
| 123 | self.fusion_model = FacePerceiver( |
| 124 | dim=cross_attention_dim, |
| 125 | depth=4, |
| 126 | dim_head=64, |
| 127 | heads=cross_attention_dim // 64, |
| 128 | embedding_dim=clip_embeddings_dim, |
| 129 | output_dim=cross_attention_dim, |
| 130 | ff_mult=4, |
| 131 | ) |
| 132 | |
| 133 | |
| 134 | def forward(self, id_embeds, clip_embeds, shortcut=False, scale=1.0): |
| 135 | x = self.proj(id_embeds) |
| 136 | x = x.reshape(-1, self.num_tokens, self.cross_attention_dim) |
| 137 | x = self.norm(x) |
| 138 | out = self.fusion_model(x, clip_embeds) |
| 139 | if shortcut: |
| 140 | out = x + scale * out |
| 141 | return out |
no outgoing calls
no test coverage detected