File size: 2,831 Bytes
d589677
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
#!/usr/bin/env python3
"""Self-contained ONNX inference for the HerBERT Polish legal NER model.
No PyTorch needed.  pip install onnxruntime transformers numpy

Run from the repo root:  python examples/inference_onnx.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  # recall-first: flip a token to PER if summed PER prob >= this

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)
    offsets = 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])      # (seq, num_labels)
    ids = probs.argmax(-1)
    per_b, per_i = label2id["B-PER"], label2id["I-PER"]

    spans, cur = [], None  # cur = [start, end, type]
    for i, (s, e) in enumerate(offsets):
        if s == e:                                    # special token
            if cur:
                spans.append(cur); cur = None
            continue
        label = id2label[int(ids[i])]
        # recall-first override for persons
        if label == "O" and probs[i, per_b] + probs[i, per_i] >= PER_THRESHOLD:
            label = "I-PER" if (cur and cur[2] == "PER") else "B-PER"
        if label == "O":
            if cur:
                spans.append(cur); cur = None
            continue
        tag, etype = label.split("-", 1)
        if tag == "B" or cur is None or cur[2] != etype:
            if cur:
                spans.append(cur)
            cur = [int(s), int(e), etype]
        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]


if __name__ == "__main__":
    samples = [
        "Pozwany Jan Kowalski, zam. ul. Słoneczna 5 w Krakowie, PESEL 02070803628.",
        "Powódka Anna Nowak-Kowalska, e-mail a.nowak@example.pl, tel. 501 234 567.",
    ]
    for t in samples:
        print("\n" + t)
        for ent in predict(t):
            print(f"   {ent['type']:8} [{ent['start']:>3}:{ent['end']:<3}] {ent['text']!r}")