MiniAI-TITS2 / text_encoder2.py
Codeminute's picture
T.I.T.S.2 — 93M DiT, rectified flow, 256px, 24 epochs on cleaned CC12M
2ca14fe verified
Raw History Blame Contribute Delete
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
@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