genderize / inference.py
textpie's picture
Genderize v2: fused two-network model (v2.2), script-aware routing, country prior — honest model card; v1 preserved in tag v1
5e6f6db verified
Raw History Blame Contribute Delete
6.83 kB
#!/usr/bin/env python3
"""genderize (cora / ultra) — standalone inference for the released weights.
What the model is
-----------------
A byte-level dual-head convolutional classifier. A personal name is normalised
(Unicode NFC, lowercase, whitespace collapsed), UTF-8 encoded and truncated to 48 bytes for legacy weights, or to the configured
``maxlen``. Each byte is embedded (256 x 64), projected to `ch` channels and passed
through four residual 1-D convolution blocks (kernels 3, 5, 7, 3; BatchNorm +
GELU). Masked mean-pooling and max-pooling over the byte positions are
concatenated and fed to a shared 512-unit layer, from which two linear heads
read out: gender (2 classes, M/F) and country (226 ISO-3166 alpha-2 codes).
`cora` uses ch=160 (~0.77M parameters), `ultra` uses ch=384 (~3.2M).
Logits are divided by a per-head temperature (calibration.json) before softmax.
Files expected next to this script (or pass --models-dir):
genderize_<variant>.pt state_dict (PyTorch)
genderize_<variant>.config.json {"ch": 160 | 384}
genderize_<variant>.calibration.json {"temp_gender": t, "temp_country": t}
maps.json {"gender": {"M":0,"F":1}, "country": {"AD":0, ...}}
Usage
-----
python genderize_infer.py --variant ultra "Andrea Rossi" "Yuki Tanaka" "María García"
python genderize_infer.py --variant cora --top 3 --json "Chen Wei"
Requires only torch and numpy. CPU is enough.
"""
from __future__ import annotations
import argparse
import json
import sys
import unicodedata
from pathlib import Path
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
MAXLEN = 48
class _GBlock(nn.Module):
def __init__(self, ch: int, k: int):
super().__init__()
self.conv = nn.Conv1d(ch, ch, k, padding=k // 2)
self.norm = nn.BatchNorm1d(ch)
def forward(self, x):
return x + F.gelu(self.norm(self.conv(x)))
class NameModel(nn.Module):
"""Byte/codepoint classifier: gender (2) + country (n_countries)."""
def __init__(self, n_countries: int, emb: int = 64, ch: int = 160,
vocab_size: int = 256):
super().__init__()
self.emb = nn.Embedding(vocab_size, emb, padding_idx=0)
self.proj = nn.Conv1d(emb, ch, 1)
self.blocks = nn.Sequential(_GBlock(ch, 3), _GBlock(ch, 5), _GBlock(ch, 7), _GBlock(ch, 3))
self.shared = nn.Sequential(nn.Linear(ch * 2, 512), nn.GELU(), nn.Dropout(0.15))
self.head_g = nn.Linear(512, 2)
self.head_c = nn.Linear(512, n_countries)
def forward(self, x):
mask = (x != 0).float().unsqueeze(1)
h = self.proj(self.emb(x.long()).transpose(1, 2))
h = self.blocks(h) * mask
mean = h.sum(-1) / mask.sum(-1).clamp(min=1)
mx = h.masked_fill(mask == 0, -1e9).max(-1).values
z = self.shared(torch.cat([mean, mx], -1))
return self.head_g(z), self.head_c(z)
def normalise(name: str) -> str:
"""NFC + lowercase + collapse whitespace. Must match what the weights saw."""
return " ".join(unicodedata.normalize("NFC", (name or "").lower()).split())
def encode(texts: list[str], maxlen: int = MAXLEN,
vocab_codepoint: dict[str, int] | None = None) -> torch.Tensor:
"""Codifica byte legacy oppure codepoint (0=pad, 1=unk)."""
dtype = np.uint8 if vocab_codepoint is None else np.int64
X = np.zeros((len(texts), maxlen), dtype=dtype)
for i, t in enumerate(texts):
if vocab_codepoint is None:
b = t.encode("utf-8")[:maxlen]
X[i, : len(b)] = np.frombuffer(b, dtype=np.uint8)
else:
ids = [vocab_codepoint.get(c, 1) for c in t[:maxlen]]
X[i, : len(ids)] = ids
return torch.from_numpy(X)
class Genderize:
def __init__(self, variant: str = "ultra", models_dir: str | Path | None = None):
md = Path(models_dir) if models_dir else Path(__file__).resolve().parent
maps = json.loads((md / "maps.json").read_text())
self.vocab_codepoint = maps.get("vocab_codepoint")
self.inv_country = {v: k for k, v in maps["country"].items()}
self.idx_m = maps["gender"].get("M", 0)
cfg = json.loads((md / f"genderize_{variant}.config.json").read_text())
self.maxlen = int(cfg.get("maxlen", maps.get("maxlen", MAXLEN)))
cal_path = md / f"genderize_{variant}.calibration.json"
self.cal = json.loads(cal_path.read_text()) if cal_path.exists() else {}
vocab_size = len(self.vocab_codepoint) + 2 if self.vocab_codepoint is not None else 256
self.model = NameModel(len(maps["country"]), ch=int(cfg.get("ch", 160)),
vocab_size=vocab_size)
self.model.load_state_dict(torch.load(md / f"genderize_{variant}.pt", map_location="cpu"))
self.model.eval()
self.variant = variant
@torch.no_grad()
def predict(self, names: list[str], top: int = 5) -> list[dict]:
texts = [normalise(n) for n in names]
lg, lc = self.model(encode(texts, maxlen=self.maxlen,
vocab_codepoint=self.vocab_codepoint))
pg = F.softmax(lg / self.cal.get("temp_gender", 1.0), dim=-1)
pc = F.softmax(lc / self.cal.get("temp_country", 1.0), dim=-1)
vals, idx = pc.topk(min(top, pc.shape[-1]), dim=-1)
out = []
for i, name in enumerate(names):
pm = float(pg[i, self.idx_m])
gender = "male" if pm >= 0.5 else "female"
out.append({
"name": name,
"gender": gender,
"probability": round(pm if gender == "male" else 1.0 - pm, 4),
"countries": [{"code": self.inv_country.get(int(j), "??"), "probability": round(float(v), 4)}
for v, j in zip(vals[i].tolist(), idx[i].tolist())],
"model": self.variant,
})
return out
def main() -> int:
ap = argparse.ArgumentParser(description="genderize cora/ultra — gender + nationality from a name")
ap.add_argument("names", nargs="+")
ap.add_argument("--variant", choices=["cora", "ultra"], default="ultra")
ap.add_argument("--models-dir", default=None)
ap.add_argument("--top", type=int, default=5)
ap.add_argument("--json", action="store_true", help="print JSON instead of a table")
a = ap.parse_args()
g = Genderize(a.variant, a.models_dir)
rows = g.predict(a.names, top=a.top)
if a.json:
print(json.dumps(rows, ensure_ascii=False, indent=1))
else:
for r in rows:
cs = ", ".join(f"{c['code']} {c['probability']:.2f}" for c in r["countries"])
print(f"{r['name']:<28} {r['gender']:<7} {r['probability']:.3f} {cs}")
return 0
if __name__ == "__main__":
sys.exit(main())