CoolFace
Modelpublic

docuracy/symphonym-v7

sourceHugging Facecc-by-4.0updated 9d agoView on Hugging Face
0likes57downloads
inference.py436 linesDownload Raw Back to root
1"""2Symphonym v7 — Standalone Inference3====================================4Loads the Student (UniversalEncoder) model and computes phonetic embeddings5for toponyms from any script.  No G2P or IPA transcription required at6inference time.7 8Usage9-----10    from inference import SymphonymModel11 12    model = SymphonymModel()                        # loads from this directory13    emb   = model.embed("London", lang="en")        # (128,) numpy array14    sim   = model.similarity("London", "en",15                             "Лондон", "ru")        # cosine similarity16    pairs = model.batch_embed([17        ("London", "en"),18        ("Лондон", "ru"),19        ("伦敦",   "zh"),20    ])21"""22 23from __future__ import annotations24 25import json26import math27import os28from pathlib import Path29from typing import List, Optional, Tuple, Union30 31import numpy as np32import torch33import torch.nn as nn34import torch.nn.functional as F35 36# ---------------------------------------------------------------------------37# Minimal architecture (copy of UniversalEncoder from models/models.py)38# Keep in sync with the training code if re-training.39# ---------------------------------------------------------------------------40 41class SelfAttention(nn.Module):42    def __init__(self, hidden_dim: int, num_heads: int = 2, dropout: float = 0.1):43        super().__init__()44        assert hidden_dim % num_heads == 045        self.num_heads = num_heads46        self.head_dim  = hidden_dim // num_heads47        self.scale     = math.sqrt(self.head_dim)48        self.q_proj  = nn.Linear(hidden_dim, hidden_dim)49        self.k_proj  = nn.Linear(hidden_dim, hidden_dim)50        self.v_proj  = nn.Linear(hidden_dim, hidden_dim)51        self.out_proj = nn.Linear(hidden_dim, hidden_dim)52        self.dropout  = nn.Dropout(dropout)53 54    def forward(self, x, mask=None):55        B, L, H = x.shape56        def reshape(t):57            return t.view(B, L, self.num_heads, self.head_dim).transpose(1, 2)58        Q, K, V = reshape(self.q_proj(x)), reshape(self.k_proj(x)), reshape(self.v_proj(x))59        scores = torch.matmul(Q, K.transpose(-2, -1)) / self.scale60        if mask is not None:61            scores = scores.masked_fill(~mask[:, None, None, :], float("-inf"))62        w = self.dropout(F.softmax(scores, dim=-1))63        out = torch.matmul(w, V).transpose(1, 2).contiguous().view(B, L, H)64        return self.out_proj(out), w65 66 67class AttentionPooling(nn.Module):68    def __init__(self, hidden_dim: int, dropout: float = 0.2):69        super().__init__()70        self.proj = nn.Sequential(71            nn.Linear(hidden_dim, hidden_dim),72            nn.Tanh(),73            nn.Linear(hidden_dim, 1),74        )75        self.dropout = nn.Dropout(dropout)76 77    def forward(self, x, mask=None):78        scores = self.proj(x).squeeze(-1)79        if mask is not None:80            scores = scores.masked_fill(~mask, float("-inf"))81        w = self.dropout(F.softmax(scores, dim=-1))82        return torch.bmm(w.unsqueeze(1), x).squeeze(1), w83 84 85class UniversalEncoder(nn.Module):86    """Symphonym Student: script-/language-conditioned character encoder."""87 88    def __init__(89        self,90        vocab_size: int        = 113280,91        num_scripts: int       = 25,92        num_langs: int         = 1944,93        char_embed_dim: int    = 64,94        script_embed_dim: int  = 16,95        lang_embed_dim: int    = 16,96        hidden_dim: int        = 128,97        embed_dim: int         = 128,98        num_layers: int        = 2,99        num_attention_heads: int = 2,100        dropout: float         = 0.2,101        lang_dropout: float    = 0.5,102        num_length_buckets: int = 16,103        length_embed_dim: int  = 8,104    ):105        super().__init__()106        self.embed_dim          = embed_dim107        self.lang_dropout_rate  = lang_dropout108        self.num_length_buckets = num_length_buckets109 110        self.char_embed   = nn.Embedding(vocab_size,   char_embed_dim,   padding_idx=0)111        self.script_embed = nn.Embedding(num_scripts,  script_embed_dim)112        self.lang_embed   = nn.Embedding(num_langs,    lang_embed_dim,   padding_idx=0)113        self.length_embed = nn.Embedding(num_length_buckets, length_embed_dim)114 115        input_dim = char_embed_dim + script_embed_dim + lang_embed_dim + length_embed_dim116        self.input_proj = nn.Linear(input_dim, hidden_dim)117        self.input_norm = nn.LayerNorm(hidden_dim)118 119        self.bilstm = nn.LSTM(120            hidden_dim, hidden_dim, num_layers=num_layers,121            batch_first=True, bidirectional=True,122            dropout=dropout if num_layers > 1 else 0,123        )124        self.self_attention = SelfAttention(hidden_dim * 2, num_attention_heads, dropout)125        self.pooling        = AttentionPooling(hidden_dim * 2, dropout)126        self.output_proj    = nn.Sequential(127            nn.Linear(hidden_dim * 2, hidden_dim),128            nn.ReLU(),129            nn.Dropout(dropout),130            nn.Linear(hidden_dim, embed_dim),131            nn.LayerNorm(embed_dim),132        )133 134    def _length_bucket(self, lengths: torch.Tensor) -> torch.Tensor:135        buckets = (lengths.to(torch.long) - 1) // 2136        return buckets.clamp(0, self.num_length_buckets - 1)137 138    def forward(self, char_ids, script_ids, lang_ids, lengths):139        B, L    = char_ids.shape140        device  = char_ids.device141        mask    = torch.arange(L, device=device).unsqueeze(0) < lengths.to(device).unsqueeze(1)142 143        c_emb = self.char_embed(char_ids)144        s_emb = self.script_embed(script_ids).unsqueeze(1).expand(-1, L, -1)145        l_emb = self.lang_embed(lang_ids).unsqueeze(1).expand(-1, L, -1)146        lb    = self._length_bucket(lengths)147        len_emb = self.length_embed(lb.to(device)).unsqueeze(1).expand(-1, L, -1)148 149        x = torch.cat([c_emb, s_emb, l_emb, len_emb], dim=-1)150        x = self.input_norm(self.input_proj(x))151 152        packed   = nn.utils.rnn.pack_padded_sequence(x, lengths.cpu(), batch_first=True, enforce_sorted=False)153        lstm_out, _ = self.bilstm(packed)154        lstm_out, _ = nn.utils.rnn.pad_packed_sequence(lstm_out, batch_first=True, total_length=L)155 156        attended, _ = self.self_attention(lstm_out, mask)157        attended    = attended + lstm_out158        pooled,   _ = self.pooling(attended, mask)159        emb         = self.output_proj(pooled)160        return F.normalize(emb, p=2, dim=-1)161 162 163# ---------------------------------------------------------------------------164# Tokeniser helpers165# ---------------------------------------------------------------------------166 167# Unicode script ranges used during training (deterministic detection)168_SCRIPT_RANGES = [169    ("LATIN",      [(0x0041, 0x007A), (0x00C0, 0x024F), (0x1E00, 0x1EFF)]),170    ("CYRILLIC",   [(0x0400, 0x04FF), (0x0500, 0x052F)]),171    ("ARABIC",     [(0x0600, 0x06FF), (0x0750, 0x077F), (0xFB50, 0xFDFF), (0xFE70, 0xFEFF)]),172    ("CJK",        [(0x4E00, 0x9FFF), (0x3400, 0x4DBF), (0x20000, 0x2A6DF), (0xF900, 0xFAFF)]),173    ("HANGUL",     [(0xAC00, 0xD7AF), (0x1100, 0x11FF), (0x3130, 0x318F)]),174    ("HIRAGANA",   [(0x3041, 0x3096)]),175    ("KATAKANA",   [(0x30A1, 0x30FA), (0x31F0, 0x31FF)]),176    ("DEVANAGARI", [(0x0900, 0x097F)]),177    ("BENGALI",    [(0x0980, 0x09FF)]),178    ("GUJARATI",   [(0x0A80, 0x0AFF)]),179    ("GURMUKHI",   [(0x0A00, 0x0A7F)]),180    ("TAMIL",      [(0x0B80, 0x0BFF)]),181    ("TELUGU",     [(0x0C00, 0x0C7F)]),182    ("KANNADA",    [(0x0C80, 0x0CFF)]),183    ("MALAYALAM",  [(0x0D00, 0x0D7F)]),184    ("THAI",       [(0x0E00, 0x0E7F)]),185    ("GEORGIAN",   [(0x10A0, 0x10FF)]),186    ("ARMENIAN",   [(0x0530, 0x058F)]),187    ("HEBREW",     [(0x0590, 0x05FF), (0xFB1D, 0xFB4F)]),188    ("GREEK",      [(0x0370, 0x03FF), (0x1F00, 0x1FFF)]),189]190 191def _detect_script(text: str) -> str:192    """Return the dominant script name for a text string."""193    counts: dict[str, int] = {}194    for ch in text:195        cp = ord(ch)196        for name, ranges in _SCRIPT_RANGES:197            if any(lo <= cp <= hi for lo, hi in ranges):198                counts[name] = counts.get(name, 0) + 1199                break200        else:201            counts["OTHER"] = counts.get("OTHER", 0) + 1202    if not counts:203        return "OTHER"204    return max(counts, key=counts.__getitem__)205 206 207# ---------------------------------------------------------------------------208# Main model class209# ---------------------------------------------------------------------------210 211class SymphonymModel:212    """213    High-level wrapper for Symphonym v7 inference.214 215    Parameters216    ----------217    model_dir : str or Path, optional218        Directory containing ``model.safetensors`` (or ``final_model.pt``),219        ``vocab/char_vocab.json``, ``vocab/lang_vocab.json``, and220        ``vocab/script_vocab.json``.  Defaults to the directory of this file.221    device : str, optional222        ``"cpu"`` (default) or ``"cuda"``.223 224    Examples225    --------226    >>> model = SymphonymModel()227    >>> model.similarity("London", "en", "Лондон", "ru")228    0.991229    >>> embeddings = model.batch_embed([("London", "en"), ("Лондон", "ru")])230    >>> embeddings.shape231    (2, 128)232    """233 234    def __init__(235        self,236        model_dir: Union[str, Path, None] = None,237        device: str = "cpu",238    ):239        if model_dir is None:240            model_dir = Path(__file__).parent241        model_dir = Path(model_dir)242 243        self.device = torch.device(device)244 245        # --- Load vocabularies ---246        vocab_dir = model_dir / "vocab"247        with open(vocab_dir / "char_vocab.json") as f:248            cv = json.load(f)249        with open(vocab_dir / "lang_vocab.json") as f:250            lv = json.load(f)251        with open(vocab_dir / "script_vocab.json") as f:252            sv = json.load(f)253 254        self._char_to_id:   dict[str, int] = cv.get("char_to_id",   cv)255        self._lang_to_id:   dict[str, int] = lv.get("lang_to_id",   lv)256        self._script_to_id: dict[str, int] = sv.get("script_to_id", sv)257 258        # --- Build model from config ---259        cfg_path = model_dir / "config.json"260        with open(cfg_path) as f:261            cfg = json.load(f)262 263        self._model = UniversalEncoder(264            vocab_size          = cfg.get("vocab_size",          len(self._char_to_id) + 1),265            num_scripts         = cfg.get("num_scripts",         25),266            num_langs           = cfg.get("num_langs",           len(self._lang_to_id) + 1),267            char_embed_dim      = cfg.get("char_embed_dim",      64),268            script_embed_dim    = cfg.get("script_embed_dim",    16),269            lang_embed_dim      = cfg.get("lang_embed_dim",      16),270            hidden_dim          = cfg.get("hidden_dim",          128),271            embed_dim           = cfg.get("embed_dim",           128),272            num_layers          = cfg.get("num_layers",          2),273            num_attention_heads = cfg.get("num_attention_heads", 2),274            dropout             = cfg.get("dropout",             0.2),275            lang_dropout        = cfg.get("lang_dropout",        0.5),276            num_length_buckets  = cfg.get("num_length_buckets",  16),277            length_embed_dim    = cfg.get("length_embed_dim",    8),278        )279 280        # --- Load weights (prefer safetensors, fall back to .pt) ---281        st_path = model_dir / "model.safetensors"282        pt_path = model_dir / "final_model.pt"283        if st_path.exists():284            from safetensors.torch import load_file285            state = load_file(str(st_path), device=str(self.device))286            self._model.load_state_dict(state)287        elif pt_path.exists():288            ckpt = torch.load(str(pt_path), map_location=self.device)289            state = ckpt.get("model_state_dict", ckpt.get("model_state", ckpt))290            self._model.load_state_dict(state)291        else:292            raise FileNotFoundError(293                f"No model weights found in {model_dir}. "294                "Expected model.safetensors or final_model.pt"295            )296 297        self._model.to(self.device).eval()298 299    # ------------------------------------------------------------------300    # Tokenisation301    # ------------------------------------------------------------------302 303    def _tokenise(self, text: str, lang: str) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:304        """Convert a single (text, lang) pair to model inputs."""305        unk_char   = self._char_to_id.get("<UNK>", 1)306        unk_lang   = self._lang_to_id.get("<UNK>", 0)307        script_name = _detect_script(text)308 309        char_ids  = [self._char_to_id.get(ch, unk_char) for ch in text]310        lang_id   = self._lang_to_id.get(lang, unk_lang)311        script_id = self._script_to_id.get(script_name, 0)312        length    = len(char_ids)313 314        return (315            torch.tensor([char_ids],  dtype=torch.long),316            torch.tensor([script_id], dtype=torch.long),317            torch.tensor([lang_id],   dtype=torch.long),318            torch.tensor([length],    dtype=torch.long),319        )320 321    def _pad_batch(322        self,323        items: List[Tuple[str, str]],324    ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:325        """Tokenise and pad a list of (text, lang) pairs."""326        unk_char   = self._char_to_id.get("<UNK>", 1)327        unk_lang   = self._lang_to_id.get("<UNK>", 0)328 329        char_seqs, script_ids, lang_ids, lengths = [], [], [], []330        for text, lang in items:331            script_name = _detect_script(text)332            char_ids    = [self._char_to_id.get(ch, unk_char) for ch in text]333            char_seqs.append(char_ids)334            script_ids.append(self._script_to_id.get(script_name, 0))335            lang_ids.append(self._lang_to_id.get(lang, unk_lang))336            lengths.append(len(char_ids))337 338        max_len = max(lengths)339        padded  = [ids + [0] * (max_len - len(ids)) for ids in char_seqs]340 341        return (342            torch.tensor(padded,      dtype=torch.long),343            torch.tensor(script_ids,  dtype=torch.long),344            torch.tensor(lang_ids,    dtype=torch.long),345            torch.tensor(lengths,     dtype=torch.long),346        )347 348    # ------------------------------------------------------------------349    # Public API350    # ------------------------------------------------------------------351 352    @torch.no_grad()353    def embed(self, text: str, lang: str = "und") -> np.ndarray:354        """355        Compute a 128-dimensional L2-normalised phonetic embedding.356 357        Parameters358        ----------359        text : str360            Toponym in any script.361        lang : str, optional362            ISO 639-1 language code (e.g. ``"en"``, ``"ar"``, ``"zh"``).363            Use ``"und"`` (undetermined) if unknown — the model will fall364            back to script-level generalisation.365 366        Returns367        -------368        numpy.ndarray of shape (128,)369        """370        char_ids, script_ids, lang_ids, lengths = self._tokenise(text, lang)371        char_ids   = char_ids.to(self.device)372        script_ids = script_ids.to(self.device)373        lang_ids   = lang_ids.to(self.device)374        emb = self._model(char_ids, script_ids, lang_ids, lengths)375        return emb.cpu().numpy()[0]376 377    @torch.no_grad()378    def batch_embed(self, items: List[Tuple[str, str]]) -> np.ndarray:379        """380        Compute embeddings for a list of (text, lang) pairs.381 382        Parameters383        ----------384        items : list of (text, lang) tuples385 386        Returns387        -------388        numpy.ndarray of shape (N, 128)389        """390        char_ids, script_ids, lang_ids, lengths = self._pad_batch(items)391        char_ids   = char_ids.to(self.device)392        script_ids = script_ids.to(self.device)393        lang_ids   = lang_ids.to(self.device)394        emb = self._model(char_ids, script_ids, lang_ids, lengths)395        return emb.cpu().numpy()396 397    def similarity(398        self,399        text1: str, lang1: str,400        text2: str, lang2: str,401    ) -> float:402        """403        Cosine similarity between two toponyms.404 405        Returns a float in [-1, 1]; embeddings are L2-normalised so this406        equals the dot product.  Values above 0.75 generally indicate407        phonetically similar names.408        """409        e1 = self.embed(text1, lang1)410        e2 = self.embed(text2, lang2)411        return float(np.dot(e1, e2))412 413 414# ---------------------------------------------------------------------------415# CLI smoke test416# ---------------------------------------------------------------------------417if __name__ == "__main__":418    model = SymphonymModel()419    pairs = [420        ("London",   "en", "Лондон",          "ru"),421        ("London",   "en", "伦敦",              "zh"),422        ("London",   "en", "لندن",              "ar"),423        ("London",   "en", "Londres",           "fr"),424        ("Tokyo",    "en", "東京",              "ja"),425        ("Beijing",  "en", "北京",              "zh"),426        ("Jerusalem","en", "ירושלים",           "he"),427        ("Baghdad",  "en", "بغداد",             "ar"),428        ("Tbilisi",  "en", "თბილისი",           "ka"),429    ]430    print(f"\n{'Name 1':<14} {'Name 2':<16} {'Lang':<6} {'Sim':>6}")431    print("-" * 46)432    for t1, l1, t2, l2 in pairs:433        sim = model.similarity(t1, l1, t2, l2)434        print(f"{t1:<14} {t2:<16} {l1}→{l2:<3}  {sim:>6.3f}")435 436