| 388 | return self.visual(image.type(self.dtype)) |
| 389 | |
| 390 | def encode_text(self, text): |
| 391 | x = self.token_embedding(text).type(self.dtype) # [batch_size, n_ctx, d_model] |
| 392 | |
| 393 | x = x + self.positional_embedding.type(self.dtype) |
| 394 | x = x.permute(1, 0, 2) # NLD -> LND |
| 395 | x = self.transformer(x) |
| 396 | x = x.permute(1, 0, 2) # LND -> NLD |
| 397 | x = self.ln_final(x).type(self.dtype) |
| 398 | |
| 399 | # x.shape = [batch_size, n_ctx, transformer.width] |
| 400 | # take features from the eot embedding (eot_token is the highest number in each sequence) |
| 401 | x = x[torch.arange(x.shape[0]), text.argmax(dim=-1)] @ self.text_projection |
| 402 | |
| 403 | return x |
| 404 | |
| 405 | def forward(self, image, text): |
| 406 | image_features = self.encode_image(image) |