TimmaJ's picture
v2: four-class race head, per-domain calibration, ad-domain benchmark
324be87 verified
Raw History Blame Contribute Delete
11.3 kB
"""
Perceived-demographic prediction (age, gender, race) reproducing the exact
preprocessing and decision rule the models were trained and validated with.
from predict import DemographicPredictor
p = DemographicPredictor()
p.predict("advert.jpg", domain="real") # photographs
p.predict("generated.png", domain="ai") # anything a generator made
Pipeline, in order:
1. InsightFace buffalo_l detection; keep the largest face, det_score >= 0.10
2. YOLOv11n-face fallback (imgsz 1280, conf >= 0.10) when InsightFace finds nothing
3. Expand the box by FACE_MARGIN (0.35) of the box size on every side, clamped
4. If neither detector fires, use the whole image
5. Square resize to the checkpoint's own size, ImageNet normalisation
6. Race: temperature-scale the logits, add the per-domain offsets, argmax
Step 3 is not cosmetic. Tightening the margin from 0.35 to 0.25 costs ~2.8 points
of race accuracy and ~5.9 points of macro-F1 on our validation set.
Step 6 is not optional either, and `domain` is not a convenience flag: the
classifier behaves differently on photographs and on generated images, so the
offsets were fitted separately on each. Passing the wrong one shifts the race
distribution by several points. See calibration.py.
Licence: MIT (code). Weights are non-commercial research use only.
"""
from __future__ import annotations
import os
from dataclasses import dataclass
from typing import Iterable, Optional
import numpy as np
import torch
from PIL import Image, ImageFile
from calibration import RACE4, calibrated_probs
ImageFile.LOAD_TRUNCATED_IMAGES = True
# --------------------------------------------------------------------------
# Configuration
# --------------------------------------------------------------------------
REPO = os.environ.get("DEMOG_REPO", "TimmaJ/age-gender-race-prediction")
SUBFOLDERS = {"age": "age", "gender": "gender", "race": "race"}
GENDER_LABELS = ["male", "female"]
RACE_LABELS = RACE4 # White, Black, Asian, Other
# Fitted on the 1,000 human-labelled real ads and applied before the offsets.
# Monotonic on its own, but it scales the logits against the fixed offsets, so
# it is part of the decision rule rather than a cosmetic rescaling.
RACE_TEMPERATURE = 0.909039
# Half-open bands: lo <= age < hi + 1.
AGE_BANDS = [("0-17", 0, 17), ("18-24", 18, 24), ("25-34", 25, 34),
("35-44", 35, 44), ("45-54", 45, 54), ("55-64", 55, 64),
("65+", 65, 200)]
FACE_MARGIN = 0.35
MIN_INSIGHTFACE_SCORE = 0.10
MIN_YOLO_CONF = 0.10
YOLO_IMGSZ = 1280
YOLO_REPO = "AdamCodd/YOLOv11n-face-detection"
def age_to_band(age: Optional[float]) -> Optional[str]:
if age is None or (isinstance(age, float) and np.isnan(age)):
return None
a = float(age)
for name, lo, hi in AGE_BANDS:
if lo <= a < hi + 1:
return name
return AGE_BANDS[-1][0] if a >= AGE_BANDS[-1][1] else AGE_BANDS[0][0]
@dataclass
class FaceBox:
box: Optional[tuple] # (x1, y1, x2, y2) or None
detector: str # "insightface" | "yolo" | "none"
score: Optional[float]
# --------------------------------------------------------------------------
# Detection
# --------------------------------------------------------------------------
class FaceCropper:
"""InsightFace primary, YOLOv11n-face fallback, whole image as last resort."""
def __init__(self, device: str = "cuda", use_yolo_fallback: bool = True):
self.device = device
self.use_yolo_fallback = use_yolo_fallback
self._insight = None
self._yolo = None
@property
def insight(self):
if self._insight is None:
from insightface.app import FaceAnalysis
providers = (["CUDAExecutionProvider", "CPUExecutionProvider"]
if self.device.startswith("cuda") else ["CPUExecutionProvider"])
app = FaceAnalysis(name="buffalo_l", allowed_modules=["detection"],
providers=providers)
app.prepare(ctx_id=0 if self.device.startswith("cuda") else -1,
det_size=(640, 640))
self._insight = app
return self._insight
@property
def yolo(self):
if self._yolo is None and self.use_yolo_fallback:
from huggingface_hub import hf_hub_download
from ultralytics import YOLO
path = hf_hub_download(repo_id=YOLO_REPO, filename="model.pt")
# torch >= 2.6 defaults weights_only=True; ultralytics does not pass it.
orig = torch.load
def _compat(*a, **k):
k.setdefault("weights_only", False)
return orig(*a, **k)
torch.load = _compat
try:
self._yolo = YOLO(path)
finally:
torch.load = orig
return self._yolo
def detect(self, img_rgb: np.ndarray) -> FaceBox:
faces = self.insight.get(img_rgb)
if faces:
best = max(faces, key=lambda f: float(getattr(f, "det_score", 0.0)))
score = float(getattr(best, "det_score", 0.0))
if score >= MIN_INSIGHTFACE_SCORE:
x1, y1, x2, y2 = [float(v) for v in best.bbox]
return FaceBox((x1, y1, x2, y2), "insightface", score)
if self.use_yolo_fallback and self.yolo is not None:
res = self.yolo.predict(img_rgb, verbose=False, conf=0.01,
imgsz=YOLO_IMGSZ, max_det=50)[0]
if res.boxes is not None and len(res.boxes):
xyxy = res.boxes.xyxy.cpu().numpy()
conf = res.boxes.conf.cpu().numpy()
i = int(np.argmax(conf))
if float(conf[i]) >= MIN_YOLO_CONF:
return FaceBox(tuple(float(v) for v in xyxy[i]), "yolo", float(conf[i]))
return FaceBox(None, "none", None)
@staticmethod
def crop(img: Image.Image, fb: FaceBox, margin: float = FACE_MARGIN) -> Image.Image:
if fb.box is None:
return img
w, h = img.size
x1, y1, x2, y2 = fb.box
mx, my = (x2 - x1) * margin, (y2 - y1) * margin
nx1 = max(0, int(round(x1 - mx))); ny1 = max(0, int(round(y1 - my)))
nx2 = min(w, int(round(x2 + mx))); ny2 = min(h, int(round(y2 + my)))
if nx2 <= nx1 + 1 or ny2 <= ny1 + 1:
return img
return img.crop((nx1, ny1, nx2, ny2))
# --------------------------------------------------------------------------
# Prediction
# --------------------------------------------------------------------------
class DemographicPredictor:
"""
Parameters
----------
repo : Hugging Face repo id holding the three checkpoints, or a local path
device : "cuda" | "cpu"
detect : run face detection. Set False only if inputs are already
padded face crops in the training convention.
domain : "real" for photographs, "ai" for generated imagery. Selects the
race offsets; can also be overridden per call.
"""
def __init__(self, repo: str = REPO, device: Optional[str] = None,
detect: bool = True, use_yolo_fallback: bool = True,
domain: str = "real"):
from transformers import (AutoImageProcessor as _Proc,
AutoModelForImageClassification as _Model)
self.device = device or ("cuda" if torch.cuda.is_available() else "cpu")
self.detect = detect
self.domain = domain
self.cropper = FaceCropper(self.device, use_yolo_fallback) if detect else None
local = os.path.isdir(repo)
self.proc, self.models = {}, {}
for task, sub in SUBFOLDERS.items():
src = os.path.join(repo, sub) if local else repo
kw = {} if local else {"subfolder": sub}
self.proc[task] = _Proc.from_pretrained(src, **kw)
self.models[task] = _Model.from_pretrained(src, **kw).to(self.device).eval()
# -- internals ---------------------------------------------------------
@torch.no_grad()
def _forward(self, pils: list[Image.Image], domain: str) -> list[dict]:
out = [{} for _ in pils]
inp = self.proc["age"](pils, return_tensors="pt").to(self.device)
ages = self.models["age"](**inp).logits.squeeze(-1).cpu().numpy().reshape(-1)
for i, a in enumerate(ages):
a = float(np.clip(a, 0, 120))
out[i]["age"] = round(a, 1)
out[i]["age_group"] = age_to_band(a)
inp = self.proc["gender"](pils, return_tensors="pt").to(self.device)
gp = torch.softmax(self.models["gender"](**inp).logits, -1).cpu().numpy()
for i, p in enumerate(gp):
out[i]["gender"] = GENDER_LABELS[int(p.argmax())]
out[i]["p_female"] = round(float(p[1]), 5)
inp = self.proc["race"](pils, return_tensors="pt").to(self.device)
logits = self.models["race"](**inp).logits.double().cpu().numpy()
probs = torch.softmax(torch.tensor(logits) / RACE_TEMPERATURE, dim=1).numpy()
cal = calibrated_probs(probs, domain=domain)
for i in range(len(pils)):
out[i]["race"] = RACE_LABELS[int(cal[i].argmax())]
out[i]["race_probs"] = {c: round(float(probs[i, j]), 5)
for j, c in enumerate(RACE_LABELS)}
out[i]["domain"] = domain
return out
def _load(self, image) -> Image.Image:
if isinstance(image, Image.Image):
return image.convert("RGB")
return Image.open(image).convert("RGB")
def _prepare(self, image):
img = self._load(image)
if not self.detect:
return img, FaceBox(None, "disabled", None)
fb = self.cropper.detect(np.asarray(img))
return self.cropper.crop(img, fb), fb
# -- public ------------------------------------------------------------
def predict(self, image, domain: Optional[str] = None) -> dict:
"""image: path, file object, or PIL.Image."""
crop, fb = self._prepare(image)
res = self._forward([crop], domain or self.domain)[0]
res["face_found"] = fb.box is not None
res["detector"] = fb.detector
return res
def predict_batch(self, images: Iterable, batch_size: int = 32,
domain: Optional[str] = None) -> list[dict]:
images = list(images)
results: list[dict] = []
for start in range(0, len(images), batch_size):
chunk = images[start:start + batch_size]
prepared = [self._prepare(im) for im in chunk]
preds = self._forward([c for c, _ in prepared], domain or self.domain)
for pred, (_, fb) in zip(preds, prepared):
pred["face_found"] = fb.box is not None
pred["detector"] = fb.detector
results.extend(preds)
return results
if __name__ == "__main__":
import sys, json
if len(sys.argv) < 2:
print("usage: python predict.py IMAGE [IMAGE ...]")
raise SystemExit(1)
p = DemographicPredictor()
for path, res in zip(sys.argv[1:], p.predict_batch(sys.argv[1:])):
print(f"{path}: {json.dumps(res)}")