| 211 | |
| 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) |
| 231 | ttoken = ttoken + self.cross_attn(query=self.normq(ttoken), key=self.normk(text.permute(0,2,1)), mask=pad_mask[...,0]) |
| 232 | tokens = torch.cat((ltoken, ttoken), dim=1) |
| 233 | return tokens |
| 234 | |
| 235 | class LoadLayer(nn.Module): |
| 236 | def __init__(self, token_dim, drop, bias=False, pe_shape=None) -> None: |