failed09's picture
Release update
1cad660 verified
Raw History Blame Contribute Delete
5.71 kB
"""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]