failed09/bashkir-lid
055
1"""Portable Unicode char_wb preprocessing and sparse ONNX inference."""2import hashlib3import json4import re5from collections import Counter6from pathlib import Path7 8import numpy as np9 10ROOT = Path(__file__).resolve().parent11 12 13def sha256(path):14 with open(path, "rb") as stream:15 return hashlib.file_digest(stream, "sha256").hexdigest()16 17 18def char_wb(text, ngram_range):19 # Match sklearn's _char_wb_ngrams, including short-word handling.20 text = re.sub(r"\s\s+", " ", text.lower())21 for word in text.split():22 word = " " + word + " "23 for n in range(ngram_range[0], ngram_range[1] + 1):24 offset = 025 yield word[offset:offset + n]26 while offset + n < len(word):27 offset += 128 yield word[offset:offset + n]29 if offset == 0:30 break31 32 33WORD_RE = re.compile(r"(?u)\b\w\w+\b")34 35 36def word_ngrams(text, ngram_range):37 tokens = WORD_RE.findall(text.lower())38 for n in range(ngram_range[0], ngram_range[1] + 1):39 for i in range(len(tokens) - n + 1):40 yield " ".join(tokens[i:i + n])41 42 43class LanguageIdentifier:44 def __init__(self, model_dir=None, batch_size=256):45 import onnxruntime as ort46 47 if model_dir is not None:48 self.model_dir = Path(model_dir)49 else:50 self.model_dir = ROOT / "model" if (ROOT / "model" / "META.json").is_file() else ROOT51 config_path = self.model_dir / "config.json"52 self.config = json.loads(config_path.read_text(encoding="utf-8")) if config_path.is_file() else None53 self.meta = json.loads((self.model_dir / "META.json").read_text(encoding="utf-8"))54 for name in ("model.onnx", "vectorizer.json"):55 if name in self.meta.get("sha256", {}):56 digest = self.meta["sha256"][name]57 if sha256(self.model_dir / name) != digest:58 raise ValueError(f"LID artifact checksum mismatch: {name}")59 cfg = json.loads((self.model_dir / "vectorizer.json").read_text(encoding="utf-8"))60 self.char_vocab = cfg.get("char_vocabulary", cfg.get("vocabulary", {}))61 self.char_ngram_range = tuple(cfg.get("char_ngram_range", cfg.get("ngram_range", (2, 5))))62 self.word_vocab = cfg.get("word_vocabulary", {})63 self.word_ngram_range = tuple(cfg.get("word_ngram_range", (1, 2)))64 self.classes = np.asarray(self.meta["classes"])65 if self.config is not None:66 expected = {"classes": self.classes.tolist(), "model_file": "model.onnx",67 "vectorizer_file": "vectorizer.json", "metadata_file": "META.json"}68 if any(self.config.get(k) != v for k, v in expected.items()):69 raise ValueError("LID config does not match model artifacts")70 if batch_size < 1:71 raise ValueError("batch_size must be positive")72 self.batch_size = batch_size73 options = ort.SessionOptions()74 options.intra_op_num_threads = 175 self.session = ort.InferenceSession(str(self.model_dir / "model.onnx"), options,76 providers=["CPUExecutionProvider"])77 78 def features(self, texts):79 char_rows = []80 word_rows = []81 for text in texts:82 if not isinstance(text, str):83 raise TypeError("LID input must contain strings")84 c_counts = Counter(self.char_vocab[g] for g in char_wb(text, self.char_ngram_range)85 if g in self.char_vocab)86 char_rows.append(c_counts)87 if self.word_vocab:88 w_counts = Counter(self.word_vocab[w] for w in word_ngrams(text, self.word_ngram_range)89 if w in self.word_vocab)90 word_rows.append(w_counts)91 92 w_c = max(1, max(map(len, char_rows), default=0))93 c_ids = np.zeros((len(char_rows), w_c), dtype=np.int64)94 c_counts = np.zeros((len(char_rows), w_c), dtype=np.float32)95 for i, row in enumerate(char_rows):96 if row:97 c_ids[i, :len(row)] = list(row.keys())98 c_counts[i, :len(row)] = list(row.values())99 100 if self.word_vocab:101 w_w = max(1, max(map(len, word_rows), default=0))102 w_ids = np.zeros((len(word_rows), w_w), dtype=np.int64)103 w_counts = np.zeros((len(word_rows), w_w), dtype=np.float32)104 for i, row in enumerate(word_rows):105 if row:106 w_ids[i, :len(row)] = list(row.keys())107 w_counts[i, :len(row)] = list(row.values())108 return {"char_ids": c_ids, "char_counts": c_counts, "word_ids": w_ids, "word_counts": w_counts}109 return {"char_ids": c_ids, "char_counts": c_counts}110 111 def predict_proba(self, texts):112 texts = list(texts)113 if not texts:114 return np.empty((0, len(self.classes)), dtype=np.float32)115 results = []116 for start in range(0, len(texts), self.batch_size):117 batch_texts = texts[start:start + self.batch_size]118 results.append(self.session.run(["probabilities"], self.features(batch_texts))[0])119 return np.concatenate(results) if results else np.empty((0, len(self.classes)), dtype=np.float32)120 121 def predict(self, texts):122 texts = list(texts)123 if not texts:124 return np.empty(0, dtype=str)125 return self.classes[self.predict_proba(texts).argmax(axis=1)]126 127 def describe(self):128 return {"name": self.meta["id"], "version": self.meta["version"],129 "sha256": self.meta["sha256"], "backend": "onnxruntime",130 "decision_rule": "argmax probabilities; classes ordered as in model passport"}131 132 