| 2013 | |
| 2014 | class MochiAttentionPool(nn.Module): |
| 2015 | def __init__( |
| 2016 | self, |
| 2017 | num_attention_heads: int, |
| 2018 | embed_dim: int, |
| 2019 | output_dim: Optional[int] = None, |
| 2020 | ) -> None: |
| 2021 | super().__init__() |
| 2022 | |
| 2023 | self.output_dim = output_dim or embed_dim |
| 2024 | self.num_attention_heads = num_attention_heads |
| 2025 | |
| 2026 | self.to_kv = nn.Linear(embed_dim, 2 * embed_dim) |
| 2027 | self.to_q = nn.Linear(embed_dim, embed_dim) |
| 2028 | self.to_out = nn.Linear(embed_dim, self.output_dim) |
| 2029 | |
| 2030 | @staticmethod |
| 2031 | def pool_tokens(x: torch.Tensor, mask: torch.Tensor, *, keepdim=False) -> torch.Tensor: |