(self, tokens, text, pad_mask)
| 206 | self.normk = nn.LayerNorm(768) |
| 207 | |
| 208 | def forward(self, tokens, text, pad_mask): |
| 209 | tokens = tokens + self.cross_attn(query=self.normq(tokens), key=self.normk(text.permute(0,2,1)), mask=pad_mask[...,0]) |
| 210 | return tokens |
| 211 | |
| 212 | # updated version |
| 213 | class LoadToken(nn.Module): |
nothing calls this directly
no outgoing calls
no test coverage detected