""" 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 @torch.no_grad() 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