docuracy/symphonym-v7
057
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 