(
self,
image_embed_dim: int = 768,
cross_attention_dim: int = 768,
num_image_text_embeds: int = 32,
)
| 969 | |
| 970 | class ImageProjection(nn.Module): |
| 971 | def __init__( |
| 972 | self, |
| 973 | image_embed_dim: int = 768, |
| 974 | cross_attention_dim: int = 768, |
| 975 | num_image_text_embeds: int = 32, |
| 976 | ): |
| 977 | super().__init__() |
| 978 | |
| 979 | self.num_image_text_embeds = num_image_text_embeds |
| 980 | self.image_embeds = nn.Linear(image_embed_dim, self.num_image_text_embeds * cross_attention_dim) |
| 981 | self.norm = nn.LayerNorm(cross_attention_dim) |
| 982 | |
| 983 | def forward(self, image_embeds: torch.Tensor): |
| 984 | batch_size = image_embeds.shape[0] |