| 1569 | |
| 1570 | class ImageProjection(nn.Module): |
| 1571 | def __init__( |
| 1572 | self, |
| 1573 | image_embed_dim: int = 768, |
| 1574 | cross_attention_dim: int = 768, |
| 1575 | num_image_text_embeds: int = 32, |
| 1576 | ): |
| 1577 | super().__init__() |
| 1578 | |
| 1579 | self.num_image_text_embeds = num_image_text_embeds |
| 1580 | self.image_embeds = nn.Linear(image_embed_dim, self.num_image_text_embeds * cross_attention_dim) |
| 1581 | self.norm = nn.LayerNorm(cross_attention_dim) |
| 1582 | |
| 1583 | def forward(self, image_embeds: torch.Tensor): |
| 1584 | batch_size = image_embeds.shape[0] |