| 212 | # updated version |
| 213 | class LoadToken(nn.Module): |
| 214 | def __init__(self, token_dim, bias, drop) -> None: |
| 215 | super().__init__() |
| 216 | self.cross_attn = CrossAttn( |
| 217 | q_dim=token_dim, |
| 218 | kv_dim=768, |
| 219 | hidden_dim=token_dim, |
| 220 | num_heads=1, |
| 221 | out_dim=token_dim, |
| 222 | qkv_bias=bias, |
| 223 | attn_drop=drop, |
| 224 | proj_drop=drop, |
| 225 | ) |
| 226 | self.normq = nn.LayerNorm(token_dim) |
| 227 | self.normk = nn.LayerNorm(768) |
| 228 | |
| 229 | def forward(self, tokens, text, pad_mask): |
| 230 | ltoken, ttoken = torch.split(tokens, [tokens.shape[1]-1,1], dim=1) |