File size: 5,706 Bytes
1cad660
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
"""Minimal safe ONNX runtime for Bashkir diacritics restoration.

Use ``BashkirDiacriticsRestorer(Path("."), use_lexicon=True).restore(text)``.
The release intentionally has no KenLM dependency: ambiguous words are left to
the neural prediction rather than guessed from a separate language model.
"""
import json
import re
from pathlib import Path

import numpy as np
import onnxruntime as ort

BA_SPEC = set("ғҙҡңөҫүһәҒҘҠҢӨҪҮҺӘ")
WORD_RE = re.compile(r"[A-Za-zА-Яа-яЁёӘәҒғҘҙҠҡҢңӨөҪҫҮүҺһ-]+")
LATIN_RE = re.compile(r"[A-Za-z]")
CYRILLIC_RE = re.compile(r"[А-Яа-яЁёӘәҒғҘҙҠҡҢңӨөҪҫҮүҺһ]")
ROMAN_RE = re.compile(r"^[IVXLCDM]+$", re.I)
QUOTES_RE = re.compile(r"«[^»]*»|“[^”]*”|\"[^\"]*\"")
RUSSIAN_INERT = frozenset("а без бы был была были в во вот вы да для до же и из или их к как когда ли мне мы на над не него нет но о об он она они от по под при про с со так то ты у уже чем что чтобы это я".split())
BASE = str.maketrans({"ғ":"г", "ҙ":"з", "ҡ":"к", "ң":"н", "ө":"о", "ҫ":"с", "ү":"у", "һ":"х", "ә":"э", "Ғ":"г", "Ҙ":"з", "Ҡ":"к", "Ң":"н", "Ө":"о", "Ҫ":"с", "Ү":"у", "Һ":"х", "Ә":"э", "h":"х", "H":"х"})
LATIN = str.maketrans({"a":"а", "b":"б", "c":"с", "d":"д", "e":"е", "f":"ф", "g":"г", "h":"х", "i":"и", "j":"й", "k":"к", "l":"л", "m":"м", "n":"н", "o":"о", "p":"п", "q":"к", "r":"р", "s":"с", "t":"т", "u":"у", "v":"в", "w":"в", "x":"х", "y":"ы", "z":"з", "A":"а", "B":"б", "C":"с", "D":"д", "E":"е", "F":"ф", "G":"г", "H":"х", "I":"и", "J":"й", "K":"к", "L":"л", "M":"м", "N":"н", "O":"о", "P":"п", "Q":"к", "R":"р", "S":"с", "T":"т", "U":"у", "V":"в", "W":"в", "X":"х", "Y":"ы", "Z":"з"})


def _preserve_case(source, target):
    if source.isupper(): return target.upper()
    if source.istitle(): return target[:1].upper() + target[1:].lower()
    return target


class BashkirDiacriticsRestorer:
    def __init__(self, model_dir=".", num_threads=2, use_lexicon=True):
        self.dir = Path(model_dir)
        vocab = json.loads((self.dir / "vocab.json").read_text(encoding="utf-8"))
        self.char2id = vocab["char2id"]
        self.id2char = {int(k): v for k, v in vocab["id2char"].items()}
        self.unk_id, self.pad_id = self.char2id.get("<UNK>", 1), self.char2id.get("<PAD>", 0)
        rules = json.loads((self.dir / "substitution_map.json").read_text(encoding="utf-8"))
        self.allowed = {key: set(value) for key, value in rules["allowed"].items()}
        lex_path = self.dir / "lexicon.json"
        # A separately supplied lexicon may improve quality, but the public
        # package is deliberately self-contained and must start without one.
        lex = json.loads(lex_path.read_text(encoding="utf-8")) if use_lexicon and lex_path.is_file() else {}
        self.lexicon, self.ambiguous = lex.get("unambiguous", {}), lex.get("ambiguous", {})
        opts = ort.SessionOptions(); opts.intra_op_num_threads = num_threads
        self.session = ort.InferenceSession(str(self.dir / "model.onnx"), opts, providers=["CPUExecutionProvider"])

    def _keys(self, word):
        plain = word.lower()
        cyr = plain.translate(LATIN)
        return (plain.translate(BASE), cyr.translate(BASE), cyr, plain)

    def _lookup(self, word):
        for key in self._keys(word):
            if key in self.lexicon: return self.lexicon[key]
        return None

    def _safe_word(self, word, quoted=False):
        if not word or ROMAN_RE.fullmatch(word) or any(c in BA_SPEC for c in word): return False
        known = self._lookup(word) is not None or any(key in self.ambiguous for key in self._keys(word))
        latin, cyrillic = bool(LATIN_RE.search(word)), bool(CYRILLIC_RE.search(word))
        if latin:
            target = self._lookup(word)
            return known and (cyrillic or (len(word) >= 4 and target and any(c in BA_SPEC for c in target)))
        return cyrillic and word.lower() not in RUSSIAN_INERT and (known or not quoted)

    def restore(self, text, min_confidence=0.40):
        if not isinstance(text, str) or not text: return text
        quoted = {i for m in QUOTES_RE.finditer(text) for i in range(*m.span())}
        ids = np.full((1, len(text)), self.pad_id, dtype=np.int64)
        for i, char in enumerate(text): ids[0, i] = self.char2id.get(char, self.unk_id)
        probs = self.session.run(["probabilities"], {"input_ids": ids})[0][0]
        permitted = set()
        for match in WORD_RE.finditer(text):
            if self._safe_word(match.group(), match.start() in quoted): permitted.update(range(*match.span()))
        neural = list(text)
        for i, char in enumerate(text):
            allowed = self.allowed.get(char)
            if i not in permitted or not allowed: continue
            pred = int(probs[i].argmax()); candidate = self.id2char.get(pred, char)
            if candidate in allowed and probs[i, pred] >= min_confidence: neural[i] = candidate
        neural = "".join(neural)
        chunks, last = [], 0
        for match in WORD_RE.finditer(text):
            chunks.append(text[last:match.start()]); word = match.group()
            target = self._lookup(word) if self._safe_word(word, match.start() in quoted) else None
            chunks.append(_preserve_case(word, target) if target else neural[match.start():match.end()]); last = match.end()
        return "".join(chunks) + text[last:]

    def restore_batch(self, texts, min_confidence=0.40):
        return [self.restore(text, min_confidence) for text in texts]