Download topic_classifier.py from haileyok/microblog-topic-classifier-v5: direct link, hf CLI and curl.
- Browser
- Download file 20.4 kB
-
https://huggingface.co/haileyok/microblog-topic-classifier-v5/resolve/main/topic_classifier.py
- Command line
-
hf download hf://haileyok/microblog-topic-classifier-v5/topic_classifier.py
-
curl -L -o topic_classifier.py https://huggingface.co/haileyok/microblog-topic-classifier-v5/resolve/main/topic_classifier.py
20.4 kB
| """Standalone inference for microblog-topic-classifier-v5. | |
| The model reads one post (its text, plus up to two pictures) and returns probabilities for | |
| broad 25 broad topics ("sports", "us_politics", ... "unclear") | |
| paths 118 subtopics ("sports/baseball", ...); a broad topic with no subtopics has one entry named like itself | |
| signals 11 scores between 0 and 1 (substance, news, promo, general_interest, sentiment, critical, ad, | |
| engagement_bait, spam, self_promo, meme) | |
| tone 6 tones (informative, humorous, personal, outraged, supportive, other) | |
| Architecture: the post text goes through a fine-tuned Ettin-150M text encoder (mean pooled). Each picture goes through a | |
| frozen SigLIP 2 so400m (512 px) image encoder; the pictures' embeddings are projected and averaged (one learned vector | |
| stands in for "no picture"). Text vector, picture vector and their product go through a small network to the heads. | |
| Only the Ettin weights, the projection and the heads are in this repository (model.safetensors); the SigLIP 2 image | |
| encoder is downloaded from google/siglip2-so400m-patch16-512 the first time it is needed. | |
| from topic_classifier import TopicClassifier, render_post | |
| clf = TopicClassifier.from_pretrained("haileyok/microblog-topic-classifier-v5") # or a local folder | |
| post = {"text": "Great win for the Braves tonight", "tags": [], "media_kinds": ["image"]} | |
| out = clf.predict([{"text": render_post(post, with_pictures=True), "pictures": ["braves.jpg"]}]) | |
| print(out[0]["top_broad"], out[0]["broad"]["sports"]) | |
| Tested with torch 2.14.0 and transformers 5.17.0 (an NVIDIA GPU is strongly recommended for posts with pictures). | |
| Probabilities for the broad topic and subtopic heads are temperature scaled (config.json: temperature_broad, | |
| temperature_path), fitted on held-out validation posts so that the confidence is not over- or under-stated. | |
| """ | |
| import concurrent.futures as cf | |
| import contextlib | |
| import json | |
| import os | |
| import threading | |
| import time | |
| from pathlib import Path | |
| import numpy as np | |
| import torch | |
| import torch.nn as nn | |
| MAX_CHARS = 1800 # the post text is cut here, before tokenizing (as in training) | |
| MAX_LINK_DESCRIPTION, MAX_QUOTE, MAX_IMAGE_TEXT = 300, 500, 300 | |
| VIDEO_KIND = {"image": "image", "video_thumb": "video"} | |
| # --- turning a post into the text the model reads ---------------------------------------------------------------- | |
| def _one_line(s): | |
| return " ".join(s.split()) | |
| def _truncate(s, n): | |
| return s if len(s) <= n else s[: n - 1].rstrip() + "…" | |
| def _lines(xs, n=0): | |
| out = [] | |
| for s in xs or []: | |
| s = _one_line(s) | |
| if s: | |
| out.append(_truncate(s, n) if n else s) | |
| return out | |
| def _media_summary(kinds): | |
| order, count = [], {} | |
| for k in kinds or []: | |
| k = _one_line(k) | |
| if not k: | |
| continue | |
| if k not in count: | |
| order.append(k) | |
| count[k] = count.get(k, 0) + 1 | |
| return ", ".join(f"1 {k}" if count[k] == 1 else f"{count[k]} {k}s" for k in order) | |
| def render_post(post, with_pictures=False): | |
| """The text the model reads for a post. | |
| `post` is a dict with any of these keys (all optional): | |
| text the post text | |
| tags list of hashtags | |
| media_kinds list of "image" / "video", one per attachment | |
| labels list of moderation labels on the post | |
| media_alts list of alt texts | |
| link_domain, link_title, link_description the link card, if any | |
| quote_text the quoted post's text, if any | |
| image_texts, image_text_sources text found in the pictures by OCR ("ocr") or a picture-description service | |
| ("luna"), one source per text | |
| image_kinds "image" / "video_thumb" per picture (only used to fill in media_kinds, see below) | |
| With `with_pictures=True` (you will pass the pictures themselves) any text found in the pictures is dropped, as | |
| it was in training for picture posts, and an empty media_kinds is filled in from image_kinds. | |
| """ | |
| row = dict(post) | |
| if with_pictures: | |
| row["image_texts"], row["image_text_sources"] = [], [] | |
| if not row.get("media_kinds"): | |
| row["media_kinds"] = [VIDEO_KIND[k] for k in row.get("image_kinds") or [] if k in VIDEO_KIND] | |
| ocr, desc = [], [] | |
| for t, src in zip(row.get("image_texts") or [], row.get("image_text_sources") or []): | |
| if src == "ocr": | |
| ocr.append(t) | |
| elif src == "luna": | |
| desc.append(t) | |
| labels = [] | |
| for l in row.get("labels") or []: | |
| l = _one_line(l) | |
| if l and not l.startswith("!") and l != "needs-review" and l not in labels: | |
| labels.append(l) | |
| link = { | |
| "domain": _one_line(row.get("link_domain") or ""), | |
| "title": _one_line(row.get("link_title") or ""), | |
| "description": _truncate(_one_line(row.get("link_description") or ""), MAX_LINK_DESCRIPTION), | |
| } | |
| tags, seen = [], set() | |
| for t in row.get("tags") or []: | |
| t = _one_line(t).lstrip("#") | |
| if t and t.lower() not in seen: | |
| seen.add(t.lower()) | |
| tags.append(t) | |
| text = (row.get("text") or "").strip() | |
| alt, ocr, desc = _lines(row.get("media_alts")), _lines(ocr, MAX_IMAGE_TEXT), _lines(desc, MAX_IMAGE_TEXT) | |
| media = _media_summary(row.get("media_kinds")) | |
| quote = _truncate(_one_line(row.get("quote_text") or ""), MAX_QUOTE) | |
| out = [] | |
| if text: | |
| out.append(text) | |
| if tags: | |
| out.append("[tags] " + " ".join("#" + t for t in tags)) | |
| if media: | |
| out.append("[media] " + media) | |
| if labels: | |
| out.append("[labels] " + ", ".join(labels)) | |
| if alt: | |
| out.append("[alt] " + " | ".join(alt)) | |
| if ocr: | |
| out.append("[image text] " + " | ".join(ocr)) | |
| if desc: | |
| out.append("[image description] " + " | ".join(desc)) | |
| if any(link.values()): | |
| out.append("[link] " + " | ".join(p for p in (link["domain"], link["title"], link["description"]) if p)) | |
| if quote: | |
| out.append("[quote] " + quote) | |
| return "\n".join(out) | |
| # --- the model --------------------------------------------------------------------------------------------------- | |
| class FusionModel(nn.Module): | |
| def __init__(self, text_config, img_dim, n_broad, n_paths, n_signals, n_tones): | |
| super().__init__() | |
| from transformers import AutoModel | |
| self.text = AutoModel.from_config(text_config) | |
| h = self.text.config.hidden_size | |
| self.img = nn.Sequential(nn.LayerNorm(img_dim), nn.Linear(img_dim, h), nn.GELU()) | |
| self.no_img = nn.Parameter(torch.zeros(h)) | |
| self.fuse = nn.Sequential(nn.LayerNorm(3 * h), nn.Linear(3 * h, h), nn.GELU(), nn.Dropout(0.1)) | |
| self.broad, self.path = nn.Linear(h, n_broad), nn.Linear(h, n_paths) | |
| self.signals, self.tone = nn.Linear(h, n_signals), nn.Linear(h, n_tones) | |
| def forward(self, input_ids, attention_mask, img, img_mask): | |
| hs = self.text(input_ids=input_ids, attention_mask=attention_mask).last_hidden_state | |
| m = attention_mask.unsqueeze(-1).to(hs.dtype) | |
| t = (hs * m).sum(1) / m.sum(1).clamp(min=1) | |
| v = self.img(img.to(t.dtype)) # [batch, pictures, hidden] | |
| w = img_mask.unsqueeze(-1).to(v.dtype) | |
| have = w.sum(1) | |
| pooled = (v * w).sum(1) / have.clamp(min=1) | |
| v = torch.where(have > 0, pooled, self.no_img.to(pooled.dtype).expand_as(pooled)) | |
| z = self.fuse(torch.cat([t, v, t * v], -1)) | |
| return self.broad(z), self.path(z), self.signals(z), self.tone(z) | |
| class TopicClassifier: | |
| def __init__(self, folder, device=None, image_encoder=None, decode_threads=None, prefetch=True): | |
| """decode_threads: threads that decode pictures (default the TOPIC_DECODE_THREADS environment variable, else 8; | |
| 1 decodes on the calling thread). prefetch: decode the next batch while the current one runs on the GPU.""" | |
| from safetensors.torch import load_file | |
| from transformers import AutoConfig, AutoModel, AutoProcessor, AutoTokenizer | |
| folder = Path(folder) | |
| self.cfg = json.load(open(folder / "config.json")) | |
| self.device = torch.device(device or ("cuda" if torch.cuda.is_available() else "cpu")) | |
| self.tok = AutoTokenizer.from_pretrained(folder / "tokenizer") | |
| self.model = FusionModel(AutoConfig.from_pretrained(folder / "text_encoder"), self.cfg["image_features"]["dim"], | |
| len(self.cfg["broad"]), len(self.cfg["paths"]), len(self.cfg["signals"]), len(self.cfg["tones"])) | |
| self.model.load_state_dict(load_file(folder / "model.safetensors"), strict=True) | |
| self.model.to(self.device).eval() | |
| self._image_encoder_id = image_encoder or self.cfg["image_features"]["model"] | |
| self._image_model = self._image_processor = None | |
| self._image_lock = threading.Lock() | |
| self._AutoModel, self._AutoProcessor = AutoModel, AutoProcessor | |
| self.t_broad, self.t_path = self.cfg["temperature_broad"], self.cfg["temperature_path"] | |
| n = decode_threads if decode_threads is not None else int(os.environ.get("TOPIC_DECODE_THREADS", "8")) | |
| self._decode_pool = cf.ThreadPoolExecutor(n, thread_name_prefix="decode") if n > 1 else None | |
| self._prefetch = prefetch | |
| # Optional timing hook, called as stage_timer(name, seconds, count=1) from whichever thread ran the stage | |
| # (the service sets it to log where the time goes). None costs nothing. | |
| self.stage_timer = None | |
| def _stage(self, name): | |
| if self.stage_timer is None: | |
| yield | |
| return | |
| t0 = time.perf_counter() | |
| try: | |
| yield | |
| finally: | |
| self.stage_timer(name, time.perf_counter() - t0) | |
| def _count(self, name, n): | |
| if self.stage_timer is not None: | |
| self.stage_timer(name, 0.0, n) | |
| def from_pretrained(cls, path_or_repo, device=None, image_encoder=None, **kwargs): | |
| """A local folder with the repository's files, or a Hugging Face repository id.""" | |
| folder = path_or_repo | |
| if not os.path.isdir(folder): | |
| from huggingface_hub import snapshot_download | |
| folder = snapshot_download(path_or_repo) | |
| return cls(folder, device, image_encoder, **kwargs) | |
| # -- pictures | |
| def _load_image_encoder(self): | |
| if self._image_model is not None: | |
| return | |
| with self._image_lock: # several threads may prepare batches at once | |
| if self._image_model is None: | |
| dtype = torch.bfloat16 if self.device.type == "cuda" else torch.float32 | |
| m = self._AutoModel.from_pretrained(self._image_encoder_id, dtype=dtype) | |
| if hasattr(m, "text_model"): | |
| del m.text_model # only the picture tower is used | |
| processor = self._AutoProcessor.from_pretrained(self._image_encoder_id) | |
| self._image_processor = processor | |
| self._image_model = m.to(self.device).eval() # set last: other threads treat it as "loaded" | |
| def load_image_encoder(self): | |
| """Load the picture encoder now (it is otherwise loaded when the first picture arrives).""" | |
| self._load_image_encoder() | |
| def _open(p): | |
| """A PIL RGB image from encoded bytes, a file path or a PIL image; None when it cannot be read.""" | |
| import io | |
| from PIL import Image | |
| try: | |
| if isinstance(p, (bytes, bytearray, memoryview)): | |
| with Image.open(io.BytesIO(p)) as im: | |
| return im.convert("RGB") | |
| if isinstance(p, (str, os.PathLike)): | |
| with Image.open(p) as im: | |
| return im.convert("RGB") | |
| return p.convert("RGB") | |
| except Exception: # noqa: BLE001 broken, truncated or hostile files must not fail a whole batch | |
| return None | |
| def embed_pictures(self, ims): | |
| """One embedding per PIL RGB image, as float16 (the encoder's pooled output, as stored in training).""" | |
| self._load_image_encoder() | |
| return self._embed_pixels(self._image_processor(images=ims, return_tensors="pt")["pixel_values"]) | |
| def _embed_pixels(self, pixel_values): | |
| """The picture encoder on the GPU, for pictures the image processor has already resized and normalised.""" | |
| px = pixel_values.to(self.device, self._image_model.dtype) | |
| out = self._image_model.get_image_features(pixel_values=px) | |
| emb = out if torch.is_tensor(out) else out.pooler_output | |
| return emb.float().cpu().numpy().astype(np.float16) | |
| # -- predictions | |
| def _pixels(self, p): | |
| """One picture decoded, resized and normalised by the image processor: a float32 tensor [3, H, W], or None | |
| when the file cannot be read. Run on the decode threads; the result is exactly what the processor gives for | |
| the same picture inside a batch.""" | |
| im = self._open(p) | |
| if im is None: | |
| return None | |
| with self._stage("prep.resize"): # summed over threads | |
| return self._image_processor(images=[im], return_tensors="pt")["pixel_values"][0] | |
| def prepare_batch(self, items): | |
| """The CPU part of a batch: tokenize the texts and decode, resize and normalise the pictures (one picture per | |
| task on the decode threads). Thread safe; it can run for the next batches while the current one is on the GPU.""" | |
| max_images = self.cfg["max_images"] | |
| t_prep = time.perf_counter() | |
| with self._stage("prep.tokenize"): | |
| enc = self.tok([(it["text"] or "")[:MAX_CHARS] for it in items], truncation=True, max_length=self.cfg["max_len"])["input_ids"] | |
| refs, sources = [], [] | |
| for b, it in enumerate(items): | |
| for p in list(it.get("pictures") or [])[:max_images]: | |
| refs.append(b) | |
| sources.append(p) | |
| px, ok = None, [] | |
| if sources: | |
| self._load_image_encoder() | |
| with self._stage("prep.pictures"): # wall time of decoding and resizing the batch's pictures | |
| if self._decode_pool is not None and len(sources) > 1: | |
| tensors = list(self._decode_pool.map(self._pixels, sources)) | |
| else: | |
| tensors = [self._pixels(p) for p in sources] | |
| ok = [t is not None for t in tensors] | |
| readable = [t for t in tensors if t is not None] | |
| if readable: | |
| px = torch.stack(readable) | |
| if self.stage_timer is not None: | |
| self.stage_timer("prep", time.perf_counter() - t_prep) | |
| self._count("pictures", sum(ok)) | |
| return enc, refs, ok, px | |
| def _finish(self, n, prepared): | |
| """The GPU part: embed the decoded pictures and build the model's input tensors.""" | |
| max_images, dim = self.cfg["max_images"], self.cfg["image_features"]["dim"] | |
| enc, refs, ok, px = prepared | |
| L = max(len(x) for x in enc) | |
| ids = torch.full((n, L), self.tok.pad_token_id, dtype=torch.long) | |
| att = torch.zeros((n, L), dtype=torch.long) | |
| for b, x in enumerate(enc): | |
| ids[b, : len(x)] = torch.tensor(x) | |
| att[b, : len(x)] = 1 | |
| img = np.zeros((n, max_images, dim), np.float16) | |
| have = np.zeros((n, max_images), bool) | |
| where, slot = [], [0] * n | |
| for b, readable in zip(refs, ok): | |
| if not readable: # a picture that cannot be read is skipped, as in training | |
| continue | |
| where.append((b, slot[b])) | |
| slot[b] += 1 | |
| if px is not None: | |
| with self._stage("gpu.embed"): # copy to the GPU, the picture encoder, and the copy back (waits for the GPU) | |
| emb = self._embed_pixels(px) | |
| for (b, k), e in zip(where, emb): | |
| img[b, k] = e | |
| have[b, k] = True | |
| d = self.device | |
| with self._stage("gpu.upload"): | |
| out = (ids.to(d), att.to(d), torch.from_numpy(img).to(d), torch.from_numpy(have).to(d)) | |
| return out, have.sum(1) | |
| def run_batch(self, n, prepared): | |
| """The GPU part for one batch of `n` posts that prepare_batch has prepared: ([broad, paths, signals, tone] | |
| probabilities as numpy arrays, pictures_used [n]). Call it from one thread at a time (the GPU).""" | |
| with self._stage("finish"): | |
| inputs, n_used = self._finish(n, prepared) | |
| with self._stage("gpu.forward"): # text model and heads, up to the copy of the probabilities back | |
| with torch.autocast(self.device.type, dtype=torch.bfloat16, enabled=self.device.type == "cuda"): | |
| lb, lp, ls, lt = self.model(*inputs) | |
| probs = (torch.softmax(lb.float() / self.t_broad, -1), torch.softmax(lp.float() / self.t_path, -1), | |
| torch.sigmoid(ls.float()), torch.softmax(lt.float(), -1)) | |
| arrays = [p.cpu().numpy() for p in probs] | |
| self._count("batches", 1) | |
| self._count("batch.posts", n) | |
| return arrays, n_used | |
| def plan_batches(items, batch_size=32): | |
| """Splits the items into batches of similar text length: (order, batches), where batches holds the items in | |
| that order, `batch_size` at a time.""" | |
| order = np.argsort([len(it["text"] or "") for it in items], kind="stable") # similar lengths together | |
| return order, [[items[i] for i in order[s:s + batch_size]] for s in range(0, len(items), batch_size)] | |
| def assemble(order, parts, used): | |
| """The predict_arrays result from per-batch results (parts[k] = run_batch's arrays, used[k] = pictures_used), | |
| in batch order, put back in the input order.""" | |
| inv = np.empty(len(order), int) | |
| inv[order] = np.arange(len(order)) | |
| out = {name: np.concatenate([p[j] for p in parts])[inv] for j, name in enumerate(("broad", "paths", "signals", "tone"))} | |
| out["pictures_used"] = np.concatenate(used)[inv] | |
| return out | |
| def predict_arrays(self, items, batch_size=32): | |
| """Arrays in the input order: broad [n, 25], paths [n, 118], signals [n, 11], tone [n, 6] (all | |
| probabilities) and pictures_used [n] (how many of each post's pictures could be read and were used).""" | |
| order, batches = self.plan_batches(items, batch_size) | |
| parts, used = [], [] | |
| ahead = cf.ThreadPoolExecutor(1, thread_name_prefix="prefetch") if self._prefetch and len(batches) > 1 else None | |
| try: | |
| nxt = ahead.submit(self.prepare_batch, batches[0]) if ahead else None | |
| for k, batch in enumerate(batches): | |
| with self._stage("wait_prep.first" if k == 0 else "wait_prep"): # the GPU is idle while this waits | |
| prepared = nxt.result() if ahead else self.prepare_batch(batch) | |
| nxt = ahead.submit(self.prepare_batch, batches[k + 1]) if ahead and k + 1 < len(batches) else None | |
| arrays, n_used = self.run_batch(len(batch), prepared) | |
| parts.append(arrays) | |
| used.append(n_used) | |
| finally: | |
| if ahead: | |
| ahead.shutdown(wait=True) | |
| return self.assemble(order, parts, used) | |
| def predict(self, items, batch_size=32, top_k=None): | |
| """items: list of {"text": str, "pictures": [PIL image or file path, ...]} ("pictures" optional). | |
| Returns one dict per item: broad, paths, signals, tone (name -> probability, broad/paths sorted high to low, | |
| cut to top_k when given) and top_broad, top_path. | |
| """ | |
| a = self.predict_arrays(items, batch_size) | |
| out = [] | |
| for i in range(len(items)): | |
| r = {} | |
| for key, names in (("broad", self.cfg["broad"]), ("paths", self.cfg["paths"])): | |
| pairs = sorted(zip(names, a[key][i].tolist()), key=lambda kv: -kv[1]) | |
| r[key] = dict(pairs[:top_k] if top_k else pairs) | |
| r["signals"] = dict(zip(self.cfg["signals"], a["signals"][i].tolist())) | |
| r["tone"] = dict(zip(self.cfg["tones"], a["tone"][i].tolist())) | |
| r["top_broad"], r["top_path"] = next(iter(r["broad"])), next(iter(r["paths"])) | |
| out.append(r) | |
| return out | |