CoolFace
Modelpublic

failed09/bashkir-lid

sourceHugging Faceapache-2.0updated 8d agoView on Hugging Face
0likes55downloads
lid.py132 linesDownload Raw Back to root
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