File size: 9,854 Bytes
4770b6a | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 | """Flan-T5-Base encoder, used exactly as Supra2-IMG's inference.py uses it -- frozen by
default, optionally with its top blocks trainable.
Supra: tokenizer("google/flan-t5-base"), max_length 128, truncation, fp32 weights run
under bf16 autocast, `last_hidden_state`, and the empty string "" encoded as the
unconditional context for CFG. The only difference is padding="longest" instead of
padding="max_length": T5 masks padded keys, so real-token outputs are unchanged and
typical batches (FLUX-Reason captions: p50 ~48, p90 ~70 tokens) run ~2x cheaper.
train_blocks=N unfreezes the top N of the 12 encoder blocks plus the final layer norm.
The lower blocks (and the relative-position bias, which lives in block 0) stay frozen, so
general language knowledge is kept while the top re-shapes caption embeddings for the
image model. Dropout stays off (eval mode) either way.
"""
from __future__ import annotations
import torch
def unfreeze_top_blocks(t5_encoder_model, n: int) -> list:
"""Freeze everything, then unfreeze the top `n` encoder blocks and the final layer norm;
n < 0 unfreezes the WHOLE encoder (token embeddings and relative-position bias too).
Returns the trainable (name, parameter) pairs."""
t5_encoder_model.requires_grad_(n < 0)
if n > 0:
enc = t5_encoder_model.encoder
if not 0 < n <= len(enc.block):
raise ValueError(f"train_blocks={n}, but the encoder has {len(enc.block)} blocks")
for blk in enc.block[-n:]:
blk.requires_grad_(True)
enc.final_layer_norm.requires_grad_(True)
return [(name, p) for name, p in t5_encoder_model.named_parameters() if p.requires_grad]
class T5TextEncoder:
def __init__(self, path: str, device, max_len: int = 128, train_blocks: int = 0):
from transformers import AutoTokenizer, T5EncoderModel
self.device, self.max_len = torch.device(device), max_len
self.tok = AutoTokenizer.from_pretrained(path)
self.model = T5EncoderModel.from_pretrained(path, torch_dtype=torch.float32).to(self.device)
self.model.eval().requires_grad_(False)
self.dim = self.model.config.d_model
# (name, parameter) of everything that trains; empty when frozen
self.trainable = unfreeze_top_blocks(self.model, train_blocks)
def tokenize(self, captions: list[str], pad_to: int = 1) -> tuple[torch.Tensor, torch.Tensor]:
"""CPU only -- safe to run on the prefetch thread. `pad_to` rounds the length up
(masked padding) so a compiled model sees few distinct shapes."""
t = self.tok(captions, padding="longest", truncation=True, max_length=self.max_len,
return_tensors="pt")
ids, mask = t["input_ids"], t["attention_mask"]
extra = (-ids.shape[1]) % pad_to
if extra:
ids = torch.nn.functional.pad(ids, (0, extra), value=self.tok.pad_token_id)
mask = torch.nn.functional.pad(mask, (0, extra))
return ids, mask
def encode(self, ids: torch.Tensor, mask: torch.Tensor, grad: bool = False
) -> tuple[torch.Tensor, torch.Tensor]:
ids, mask = ids.to(self.device, non_blocking=True), mask.to(self.device, non_blocking=True)
with torch.set_grad_enabled(grad and bool(self.trainable)), \
torch.autocast(self.device.type, dtype=torch.bfloat16, enabled=self.device.type == "cuda"):
h = self.model(input_ids=ids, attention_mask=mask).last_hidden_state
return h.float(), mask
def __call__(self, captions: list[str]) -> tuple[torch.Tensor, torch.Tensor]:
"""-> (ctx (B, L, 768) fp32, mask (B, L) long). Never builds a graph."""
return self.encode(*self.tokenize(captions))
def _load_modernbert(path: str):
"""AutoModel.from_pretrained, or -- where transformers refuses a .bin checkpoint on
torch < 2.6 (Ettin ships only pytorch_model.bin) -- build from the config and load the
state dict with weights_only=True. The MLM head's tensors are dropped."""
from transformers import AutoConfig, AutoModel
try:
return AutoModel.from_pretrained(path, torch_dtype=torch.float32)
except ValueError:
import os
m = AutoModel.from_config(AutoConfig.from_pretrained(path))
f = os.path.join(path, "pytorch_model.bin") if os.path.isdir(path) else None
if f is None:
from huggingface_hub import hf_hub_download
f = hf_hub_download(path, "pytorch_model.bin")
sd = {k.removeprefix("model."): v for k, v in torch.load(f, map_location="cpu", weights_only=True).items()}
missing, _ = m.load_state_dict(sd, strict=False)
if missing:
raise KeyError(f"{path}: encoder tensors missing from the checkpoint: {missing[:5]}")
return m
class EttinTextEncoder:
"""Ettin (jhu-clsp/ettin-encoder-68m, MIT): a ModernBERT encoder, 68M, hidden 512,
byte-level BPE (keeps case, punctuation, accents and non-Latin brand names that
Flan-T5's SentencePiece maps to <unk>). Same interface as T5TextEncoder.
train_blocks != 0 trains the WHOLE encoder (there is no "top blocks" option: it is
small, and the conditioning space has to be learned anew anyway). The empty string
encodes as [CLS][SEP], the unconditional context for CFG."""
def __init__(self, path: str, device, max_len: int = 128, train_blocks: int = 0, grad_ckpt: bool = False):
from transformers import AutoTokenizer
self.device, self.max_len = torch.device(device), max_len
self.tok = AutoTokenizer.from_pretrained(path)
self.model = _load_modernbert(path).to(self.device)
self.model.eval().requires_grad_(bool(train_blocks))
if grad_ckpt and train_blocks:
# HF only checkpoints in train() mode; every Ettin dropout is 0.0 (config), so train()
# changes nothing numerically. Recomputes each layer in the backward: trainable Ettin
# activations were ~22 GB/GPU at 128 x 416 tokens uncompiled (2026-09-25 OOM).
if any(getattr(self.model.config, k, 0) for k in ("attention_dropout", "mlp_dropout", "embedding_dropout")):
raise ValueError("Ettin checkpointing needs train() mode, but this config has dropout")
self.model.gradient_checkpointing_enable(gradient_checkpointing_kwargs={"use_reentrant": False})
self.model.train()
self.dim = self.model.config.hidden_size
self.trainable = [(n, p) for n, p in self.model.named_parameters() if p.requires_grad]
def tokenize(self, captions: list[str], pad_to: int = 1) -> tuple[torch.Tensor, torch.Tensor]:
t = self.tok(captions, padding="longest", truncation=True, max_length=self.max_len,
return_tensors="pt")
ids, mask = t["input_ids"], t["attention_mask"]
extra = (-ids.shape[1]) % pad_to
if extra:
ids = torch.nn.functional.pad(ids, (0, extra), value=self.tok.pad_token_id)
mask = torch.nn.functional.pad(mask, (0, extra))
return ids, mask
def encode(self, ids: torch.Tensor, mask: torch.Tensor, grad: bool = False
) -> tuple[torch.Tensor, torch.Tensor]:
ids, mask = ids.to(self.device, non_blocking=True), mask.to(self.device, non_blocking=True)
with torch.set_grad_enabled(grad and bool(self.trainable)), \
torch.autocast(self.device.type, dtype=torch.bfloat16, enabled=self.device.type == "cuda"):
h = self.model(input_ids=ids, attention_mask=mask).last_hidden_state
return h.float(), mask
def __call__(self, captions: list[str]) -> tuple[torch.Tensor, torch.Tensor]:
return self.encode(*self.tokenize(captions))
def build_text_encoder(kind: str, path: str, device, max_len: int = 128, train_blocks: int = 0,
grad_ckpt: bool = False):
if kind == "t5":
return T5TextEncoder(path, device, max_len=max_len, train_blocks=train_blocks)
if kind == "ettin":
return EttinTextEncoder(path, device, max_len=max_len, train_blocks=train_blocks, grad_ckpt=grad_ckpt)
raise ValueError(f"unknown text encoder {kind!r} (expected t5 or ettin)")
class HashTextEncoder:
"""Stand-in for tests and --debug_tiny: deterministic per-character embeddings, same
interface and masking behaviour, no download. "" gives one token, like T5's </s>.
train_blocks > 0 makes the embedding table trainable (the whole "encoder")."""
def __init__(self, dim: int = 768, device="cpu", max_len: int = 128, seed: int = 0,
train_blocks: int = 0):
g = torch.Generator().manual_seed(seed)
self.table = torch.nn.Parameter(torch.randn(257, dim, generator=g), requires_grad=bool(train_blocks))
self.device, self.max_len, self.dim = torch.device(device), max_len, dim
self.trainable = [("table", self.table)] if train_blocks else []
def tokenize(self, captions: list[str], pad_to: int = 1) -> tuple[torch.Tensor, torch.Tensor]:
ids = [[256] + [b for b in c.encode("utf-8")][: self.max_len - 1] for c in captions]
L = max(len(i) for i in ids)
L += (-L) % pad_to
mask = torch.zeros(len(ids), L, dtype=torch.long)
idx = torch.zeros(len(ids), L, dtype=torch.long)
for r, i in enumerate(ids):
idx[r, : len(i)] = torch.tensor(i)
mask[r, : len(i)] = 1
return idx, mask
def encode(self, ids: torch.Tensor, mask: torch.Tensor, grad: bool = False
) -> tuple[torch.Tensor, torch.Tensor]:
with torch.set_grad_enabled(grad and bool(self.trainable)):
return self.table[ids.cpu()].to(self.device), mask.to(self.device)
def __call__(self, captions: list[str]) -> tuple[torch.Tensor, torch.Tensor]:
return self.encode(*self.tokenize(captions))
|