microblog-topic-classifier-v5 / topic_classifier.py
haileyok's picture
Model card: charts, MIT licence, a pass over the text; topic_classifier.py: parallel picture resize, one shared GPU worker, optional stage timing
d037570 verified
Raw History Blame Contribute Delete
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
@contextlib.contextmanager
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)
@classmethod
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()
@staticmethod
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"])
@torch.no_grad()
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)
@torch.no_grad()
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
@staticmethod
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)]
@staticmethod
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
@torch.no_grad()
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