i008's picture
HerBERT Polish legal NER (general): weights + ONNX + card + examples
399294e verified
Raw History Blame Contribute Delete
3.42 kB
#!/usr/bin/env python3
"""Hard cases the model handles WELL — the flip side of KNOWN_LIMITATIONS.md.
All inputs are synthetic (fictional names). Outputs are the model's actual
predictions (quantized ONNX, recall-first PER threshold 0.2). Reproduce:
python examples/strong_cases.py
"""
import json
from pathlib import Path
import numpy as np
import onnxruntime as ort
from transformers import AutoTokenizer
ROOT = Path(__file__).resolve().parent.parent
PER_THRESHOLD = 0.2
tok = AutoTokenizer.from_pretrained(str(ROOT))
cfg = json.load(open(ROOT / "config.json", encoding="utf-8"))
id2label = {int(k): v for k, v in cfg["id2label"].items()}
label2id = cfg["label2id"]
sess = ort.InferenceSession(str(ROOT / "onnx" / "model_quantized.onnx"))
in_names = {i.name for i in sess.get_inputs()}
def softmax(x):
e = np.exp(x - x.max(-1, keepdims=True)); return e / e.sum(-1, keepdims=True)
def predict(text):
enc = tok(text, return_offsets_mapping=True, return_tensors="np", truncation=True, max_length=512)
offs = enc["offset_mapping"][0]
feeds = {"input_ids": enc["input_ids"].astype(np.int64),
"attention_mask": enc["attention_mask"].astype(np.int64)}
if "token_type_ids" in in_names:
feeds["token_type_ids"] = np.zeros_like(enc["input_ids"], dtype=np.int64)
probs = softmax(sess.run(None, feeds)[0][0])
ids = probs.argmax(-1)
pb, pi = label2id["B-PER"], label2id["I-PER"]
spans, cur = [], None
for i, (s, e) in enumerate(offs):
if s == e:
if cur: spans.append(cur); cur = None
continue
lab = id2label[int(ids[i])]
if lab == "O" and probs[i, pb] + probs[i, pi] >= PER_THRESHOLD:
lab = "I-PER" if (cur and cur[2] == "PER") else "B-PER"
if lab == "O":
if cur: spans.append(cur); cur = None
continue
tag, et = lab.split("-", 1)
if tag == "B" or cur is None or cur[2] != et:
if cur: spans.append(cur)
cur = [int(s), int(e), et]
else:
cur[1] = int(e)
if cur: spans.append(cur)
return [{"type": t, "start": a, "end": b, "text": text[a:b]} for a, b, t in spans]
CASES = [
("Pozwany Damian Zotadek wniósł odpowiedź na pozew.", "OCR: missing diacritics (Żołądek)"),
("Z powództwa Jiga Bękowiki przeciwko spółce.", "OCR: heavily garbled name (Jaga Bętkowski)"),
("Pełnomocnik: Ptek Gostoiski, adwokat.", "OCR garble in a signature line (Patryk Gostomski)"),
("Pozew skierowano przeciwko Jerzemu Sekule.", "inflected (dative) Polish name"),
("Sprawa dotyczy akt Wojciecha Michalika.", "inflected (genitive) Polish name"),
("Powódka Anna Nowak-Kowalska złożyła wniosek.", "hyphenated double surname (single span)"),
("Stawiła się ALEKSANDRA ŚWIDERSKA, pełnomocnik.", "ALL-CAPS name (header/signature style)"),
("Z poważaniem,\nIga Mełech\nradca prawny", "short female first + rare surname in a signature"),
("Jan Kowalski, zam. ul. Słoneczna 5 w Krakowie, PESEL 02070803628.",
"private address (LOC) vs public city (LOC_PUB) + national ID"),
]
if __name__ == "__main__":
print(f"PER threshold = {PER_THRESHOLD}\n")
for text, note in CASES:
print(f"# {note}")
print(f" input : {text.replace(chr(10), ' / ')}")
for e in predict(text):
print(f" {e['type']:8} {e['text']!r}")
print()