| 233 | return tokens |
| 234 | |
| 235 | class LoadLayer(nn.Module): |
| 236 | def __init__(self, token_dim, drop, bias=False, pe_shape=None) -> None: |
| 237 | super().__init__() |
| 238 | if pe_shape >30: |
| 239 | self.loadtoken = LoadToken( |
| 240 | token_dim=token_dim, |
| 241 | bias=bias, |
| 242 | drop=drop |
| 243 | ) |
| 244 | self.norm = nn.LayerNorm(token_dim) |
| 245 | self.mlp = Mlp(token_dim, token_dim*2, token_dim) |
| 246 | self.positional_embedding = nn.Parameter(torch.randn(pe_shape**2, token_dim) / token_dim ** 0.5) |
| 247 | self.pe_shape = pe_shape |
| 248 | |
| 249 | def forward(self, tokens, text, pad_mask): |
| 250 | if self.pe_shape > 30: |
| 251 | tokens = self.loadtoken(tokens, text, pad_mask) |
| 252 | tokens = self.mlp(self.norm(tokens)) |
| 253 | return tokens, self.positional_embedding |
| 254 | |
| 255 | |
| 256 | class CGAttention(nn.Module): |