"""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("", 1), self.char2id.get("", 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]