Download text_encoder2.py from M1n1A1/MiniAI-TITS2: direct link, hf CLI and curl.
- Browser
- Download file 1.35 kB
-
https://huggingface.co/M1n1A1/MiniAI-TITS2/resolve/main/text_encoder2.py
- Command line
-
hf download hf://M1n1A1/MiniAI-TITS2/text_encoder2.py
-
curl -L -o text_encoder2.py https://huggingface.co/M1n1A1/MiniAI-TITS2/resolve/main/text_encoder2.py
1.35 kB
| """ | |
| Frozen CLIP ViT-L/14 text encoder for T.I.T.S.2. | |
| Unlike T.I.T.S.1, which squashed the whole prompt into one 512-dim vector, this returns | |
| the per-token hidden states (77 x 768) so the model can cross-attend word by word, plus | |
| a pooled vector for the global conditioning path, plus the attention mask so padding | |
| tokens are ignored. | |
| """ | |
| from typing import List, Tuple | |
| import torch | |
| import torch.nn as nn | |
| from transformers import CLIPTextModel, CLIPTokenizer | |
| class FrozenCLIPTextEncoder(nn.Module): | |
| def __init__(self, model_name: str = "openai/clip-vit-large-patch14", device: str = "cuda"): | |
| super().__init__() | |
| self.tokenizer = CLIPTokenizer.from_pretrained(model_name) | |
| self.model = CLIPTextModel.from_pretrained(model_name).to(device).eval() | |
| self.model.requires_grad_(False) | |
| self.device = device | |
| self.embed_dim = self.model.config.hidden_size # 768 | |
| def encode(self, prompts: List[str]) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: | |
| """-> (tokens (b,77,768), pooled (b,768), mask (b,77))""" | |
| tok = self.tokenizer( | |
| prompts, padding="max_length", truncation=True, max_length=77, return_tensors="pt" | |
| ).to(self.device) | |
| out = self.model(**tok) | |
| return out.last_hidden_state, out.pooler_output, tok.attention_mask | |