hellosindh/indus-script-models
0
1"""2Indus Script — Inference & Generation3======================================4Download models from HuggingFace and run:5 1. Sequence validation — is this inscription valid?6 2. Sign prediction — predict a masked sign7 3. Generate synthetic — generate new Indus sequences8 4. Score any sequence — get ensemble confidence score9 10Install:11 pip install torch transformers huggingface_hub12 13Usage:14 python inference.py --task validate --sequence "T638 T177 T420 T122"15 python inference.py --task predict --sequence "T638 [MASK] T420 T122"16 python inference.py --task generate --count 1017 python inference.py --task score --sequence "T638 T177 T420"18 python inference.py --task demo19"""20 21import argparse22import math23import os24import pickle25import sys26from pathlib import Path27 28import torch29import torch.nn as nn30import torch.nn.functional as F31 32 33# ── Auto-download from HuggingFace ────────────────────────────34HF_REPO = "hellosindh/indus-script-models" # update after upload35 36def download_models(repo_id=HF_REPO, local_dir="indus_models"):37 """Download all model files from HuggingFace."""38 try:39 from huggingface_hub import snapshot_download40 print(f"Downloading models from {repo_id}...")41 path = snapshot_download(repo_id=repo_id, local_dir=local_dir)42 print(f"✓ Downloaded to {path}")43 return path44 except Exception as e:45 print(f"Download failed: {e}")46 print("Manual download: https://huggingface.co/{repo_id}")47 sys.exit(1)48 49 50def get_model_dir():51 """52 Find model directory.53 Priority:54 1. ./models/ (running from cloned HuggingFace repo)55 2. DATA/models/ (running from original indus_script folder)56 3. Auto-download from HuggingFace57 """58 # Running from cloned repo — models/ is right here59 cloned = Path("models")60 if cloned.exists() and (cloned / "nanogpt_indus.pt").exists():61 data = Path("data") if Path("data").exists() else Path(".")62 return cloned, data63 # Running from original indus_script folder64 local = Path("DATA/models")65 if local.exists():66 return local, Path("DATA")67 # Auto-download from HuggingFace68 path = download_models()69 return Path(path) / "models", Path(path) / "data"70 71 72# ── Device ─────────────────────────────────────────────────────73device = torch.device("cuda" if torch.cuda.is_available() else "cpu")74 75BOS_ID = 81476EOS_ID = 81577PAD_ID = 81678 79 80# ── Load helpers ───────────────────────────────────────────────81def load_tokenizer(data_dir):82 from transformers import PreTrainedTokenizerFast83 # Try data/indus_tokenizer first, then just data_dir itself84 tok_path = data_dir / "indus_tokenizer"85 if not tok_path.exists():86 tok_path = data_dir87 return PreTrainedTokenizerFast.from_pretrained(str(tok_path))88 89 90def load_bert_mlm(model_dir):91 from transformers import BertForMaskedLM92 return BertForMaskedLM.from_pretrained(93 str(model_dir / "mlm")).to(device).eval()94 95 96def load_bert_cls(model_dir):97 from transformers import BertForSequenceClassification98 return BertForSequenceClassification.from_pretrained(99 str(model_dir / "cls")).to(device).eval()100 101 102def load_ngram(model_dir):103 # indus_ngram.py must be importable104 sys.path.insert(0, str(Path(__file__).parent))105 with open(model_dir / "ngram_model.pkl", "rb") as f:106 return pickle.load(f)107 108 109def load_electra(model_dir):110 from transformers import BertModel, BertConfig, PreTrainedTokenizerFast111 import json112 113 class ElectraDisc(nn.Module):114 def __init__(self, cfg):115 super().__init__()116 self.bert = BertModel(cfg)117 self.classifier = nn.Linear(cfg.hidden_size, 2)118 self.dropout = nn.Dropout(0.1)119 120 def forward(self, input_ids, attention_mask):121 out = self.bert(input_ids=input_ids,122 attention_mask=attention_mask)123 return self.classifier(self.dropout(out.last_hidden_state))124 125 p = model_dir / "electra"126 with open(p / "discriminator_config.json") as f:127 cfg = json.load(f)128 m = ElectraDisc(BertConfig(**cfg))129 m.load_state_dict(torch.load(p / "discriminator.pt",130 map_location=device, weights_only=True))131 tok = PreTrainedTokenizerFast.from_pretrained(str(p))132 return tok, m.to(device).eval()133 134 135def load_nanogpt(model_dir):136 ckpt = torch.load(model_dir / "nanogpt_indus.pt",137 map_location=device, weights_only=False)138 cfg = ckpt["cfg"]139 140 class CSA(nn.Module):141 def __init__(self, c):142 super().__init__()143 self.nh = c["n_head"]; self.ne = c["n_embd"]144 self.hd = c["n_embd"] // c["n_head"]145 self.qkv = nn.Linear(c["n_embd"], 3*c["n_embd"], bias=False)146 self.proj = nn.Linear(c["n_embd"], c["n_embd"], bias=False)147 self.drop = nn.Dropout(c["dropout"])148 ml = c["block_size"]149 self.register_buffer("mask",150 torch.tril(torch.ones(ml, ml)).view(1, 1, ml, ml))151 152 def forward(self, x):153 B, T, C = x.shape154 q, k, v = self.qkv(x).split(self.ne, dim=2)155 sh = lambda t: t.view(B, T, self.nh, self.hd).transpose(1, 2)156 q, k, v = sh(q), sh(k), sh(v)157 a = (q @ k.transpose(-2, -1)) / math.sqrt(self.hd)158 a = a.masked_fill(self.mask[:,:,:T,:T] == 0, float("-inf"))159 return self.proj(160 (self.drop(F.softmax(a, dim=-1)) @ v)161 .transpose(1, 2).contiguous().view(B, T, C))162 163 class TB(nn.Module):164 def __init__(self, c):165 super().__init__()166 self.ln1 = nn.LayerNorm(c["n_embd"]); self.attn = CSA(c)167 self.ln2 = nn.LayerNorm(c["n_embd"])168 self.ffn = nn.Sequential(169 nn.Linear(c["n_embd"], 4*c["n_embd"]), nn.GELU(),170 nn.Linear(4*c["n_embd"], c["n_embd"]), nn.Dropout(c["dropout"]))171 def forward(self, x):172 return x + self.ffn(self.ln2(x + self.attn(self.ln1(x))))173 174 class GPT(nn.Module):175 def __init__(self, c):176 super().__init__()177 self.cfg = c178 self.tok_emb = nn.Embedding(c["vocab_size"], c["n_embd"])179 self.pos_emb = nn.Embedding(c["block_size"], c["n_embd"])180 self.drop = nn.Dropout(c["dropout"])181 self.blocks = nn.ModuleList([TB(c) for _ in range(c["n_layer"])])182 self.ln_f = nn.LayerNorm(c["n_embd"])183 self.head = nn.Linear(c["n_embd"], c["vocab_size"], bias=False)184 self.tok_emb.weight = self.head.weight185 186 def forward(self, idx):187 B, T = idx.shape188 x = self.drop(self.tok_emb(idx) + self.pos_emb(189 torch.arange(T, device=idx.device).unsqueeze(0)))190 for b in self.blocks: x = b(x)191 return self.head(self.ln_f(x))192 193 @torch.no_grad()194 def generate(self, temperature=0.85, top_k=40, max_len=15):195 self.eval()196 idx = torch.tensor([[BOS_ID]], device=device)197 gen = []198 for _ in range(max_len):199 logits = self(idx[:, -self.cfg["block_size"]:])[: ,-1, :] / temperature200 logits[:, PAD_ID] = logits[:, BOS_ID] = logits[:, EOS_ID] = float("-inf")201 if top_k > 0:202 v, _ = torch.topk(logits, min(top_k, logits.size(-1)))203 logits[logits < v[:, [-1]]] = float("-inf")204 nxt = torch.multinomial(F.softmax(logits, dim=-1), 1)205 if nxt.item() == EOS_ID: break206 gen.append(nxt.item())207 idx = torch.cat([idx, nxt], dim=1)208 return list(reversed(gen)) # RTL→LTR209 210 m = GPT(cfg)211 m.load_state_dict(ckpt["model_state"])212 return m.to(device).eval()213 214 215# ── Scoring functions ──────────────────────────────────────────216def parse_sequence(seq_str):217 """Parse 'T638 T177 T420' or '638 177 420' into list of ints."""218 tokens = seq_str.strip().split()219 ids = []220 for t in tokens:221 if t.upper() == "[MASK]":222 ids.append(None)223 else:224 t = t.upper().lstrip("T")225 ids.append(int(t))226 return ids227 228 229def bert_validity_score(seq, tok, cls_model):230 text = " ".join(f"T{t}" for t in seq)231 enc = tok(text, return_tensors="pt", truncation=True,232 max_length=32).to(device)233 with torch.no_grad():234 return float(torch.softmax(cls_model(**enc).logits, dim=-1)[0][1])235 236 237def bert_predict_mask(seq_with_none, tok, mlm_model, top_k=5):238 parts = ["[MASK]" if t is None else f"T{t}" for t in seq_with_none]239 enc = tok(" ".join(parts), return_tensors="pt",240 truncation=True, max_length=32).to(device)241 with torch.no_grad():242 logits = mlm_model(**enc).logits243 results = {}244 for pos, val in enumerate(seq_with_none):245 if val is not None: continue246 tp, ti = torch.softmax(logits[0, pos+1], dim=-1).topk(top_k)247 preds = []248 for p, tid in zip(tp.tolist(), ti.tolist()):249 ts = tok.convert_ids_to_tokens([tid])[0]250 if ts.startswith("T") and ts[1:].isdigit():251 preds.append((int(ts[1:]), round(p, 4)))252 results[pos] = preds253 return results254 255 256def electra_score(seq, tok, disc):257 enc = tok(" ".join(f"T{t}" for t in seq), return_tensors="pt",258 truncation=True, max_length=32).to(device)259 with torch.no_grad():260 logits = disc(enc["input_ids"], enc["attention_mask"])261 probs = torch.softmax(logits[0], dim=-1)262 n = min(len(seq), probs.shape[0]-1)263 return float(probs[1:n+1, 0].mean())264 265 266def ensemble_score(seq, tok, cls, ngram, elec_tok, elec_disc):267 b = bert_validity_score(seq, tok, cls)268 n = ngram.validity_score(seq)269 e = electra_score(seq, elec_tok, elec_disc)270 return 0.50*b + 0.25*n + 0.25*e, b, n, e271 272 273def load_glyph_map(data_dir):274 import json275 p = data_dir / "id_to_glyph.json"276 if p.exists():277 with open(p, encoding="utf-8") as f:278 return json.load(f)279 return {}280 281 282# ── Tasks ──────────────────────────────────────────────────────283def task_validate(seq_str, models):284 tok, cls, ngram, elec_tok, elec_disc, glyph_map = models285 seq = parse_sequence(seq_str)286 if any(t is None for t in seq):287 print("Use --task predict for sequences with [MASK]")288 return289 ens, b, n, e = ensemble_score(seq, tok, cls, ngram, elec_tok, elec_disc)290 glyphs = "".join(glyph_map.get(str(t), f"[{t}]") for t in seq)291 print(f"\n Sequence : {' '.join(f'T{t}' for t in seq)}")292 print(f" Glyphs : {glyphs}")293 print(f" BERT : {b:.4f}")294 print(f" N-gram : {n:.4f}")295 print(f" ELECTRA : {e:.4f}")296 print(f" Ensemble : {ens:.4f}")297 print(f" Verdict : {'✅ VALID (≥85%)' if ens >= 0.85 else '⚠ UNCERTAIN (≥70%)' if ens >= 0.70 else '❌ INVALID (<70%)'}")298 299 300def task_predict(seq_str, models):301 tok, cls, ngram, elec_tok, elec_disc, glyph_map = models302 model_dir, data_dir = get_model_dir()303 mlm = load_bert_mlm(model_dir)304 seq = parse_sequence(seq_str)305 preds = bert_predict_mask(seq, tok, mlm, top_k=5)306 print(f"\n Input: {seq_str}")307 for pos, candidates in preds.items():308 print(f"\n Position {pos} predictions:")309 for sign_id, prob in candidates:310 g = glyph_map.get(str(sign_id), "?")311 print(f" T{sign_id:<5} {g} {prob*100:>6.2f}%")312 313 314def task_generate(count, models, threshold=0.85):315 tok, cls, ngram, elec_tok, elec_disc, glyph_map = models316 model_dir, data_dir = get_model_dir()317 gpt = load_nanogpt(model_dir)318 kept = []319 seen = set()320 attempts = 0321 322 print(f"\n Generating (threshold={threshold:.0%})...\n")323 temps = [0.85, 0.90, 1.00, 1.10]324 topks = [40, 50, 60, 80 ]325 326 while len(kept) < count and attempts < count * 100:327 i = attempts % len(temps)328 seq = gpt.generate(temperature=temps[i], top_k=topks[i])329 attempts += 1330 if len(seq) < 2 or tuple(seq) in seen: continue331 seen.add(tuple(seq))332 ens, b, n, e = ensemble_score(seq, tok, cls, ngram, elec_tok, elec_disc)333 if ens >= threshold:334 glyphs = "".join(glyph_map.get(str(t), "?") for t in seq)335 kept.append((seq, ens, glyphs))336 seq_str = " ".join(f"T{t}" for t in seq)337 print(f" {len(kept):>3}. {glyphs} | {seq_str} | score={ens:.3f}")338 339 print(f"\n Generated {len(kept)} sequences in {attempts} attempts")340 return kept341 342 343def task_score(seq_str, models):344 task_validate(seq_str, models)345 346 347def task_demo(models, glyph_map):348 print("\n" + "="*60)349 print(" INDUS SCRIPT — INFERENCE DEMO")350 print("="*60)351 352 examples = [353 ("T638 T177 T420 T122", "Known valid sequence"),354 ("T604 T123 T609", "Known formula (appears on 80+ seals)"),355 ("T406 T638 T243", "Known formula (appears on 37 seals)"),356 ("T122 T638 T177", "Reversed — should score lower"),357 ("T999 T888 T777", "Invalid token IDs"),358 ]359 360 tok, cls, ngram, elec_tok, elec_disc, glyph_map = models361 print(f"\n {'Sequence':<35} {'Ensemble':>9} Verdict")362 print(" " + "─"*58)363 for seq_str, label in examples:364 try:365 seq = [int(t.lstrip("T")) for t in seq_str.split()]366 ens, b, n, e = ensemble_score(seq, tok, cls, ngram, elec_tok, elec_disc)367 g = "".join(glyph_map.get(str(t),"?") for t in seq)368 verdict = "✅" if ens>=0.85 else "⚠" if ens>=0.70 else "❌"369 print(f" {seq_str:<35} {ens:>8.3f} {verdict} {label}")370 except Exception:371 print(f" {seq_str:<35} {'—':>9} ❌ {label}")372 373 374# ── Main ───────────────────────────────────────────────────────375def main():376 parser = argparse.ArgumentParser(description="Indus Script Inference")377 parser.add_argument("--task", choices=["validate","predict","generate","score","demo"],378 default="demo")379 parser.add_argument("--sequence", type=str, default="T638 T177 T420 T122",380 help="Sequence like 'T638 T177 T420' or 'T638 [MASK] T420'")381 parser.add_argument("--count", type=int, default=10,382 help="Number of sequences to generate")383 parser.add_argument("--threshold",type=float, default=0.85)384 parser.add_argument("--download", action="store_true",385 help="Force re-download from HuggingFace")386 args = parser.parse_args()387 388 if args.download:389 download_models()390 391 print("Loading models...")392 model_dir, data_dir = get_model_dir()393 394 tok = load_tokenizer(data_dir)395 cls = load_bert_cls(model_dir); print(" ✓ TinyBERT")396 ngram = load_ngram(model_dir); print(" ✓ N-gram")397 elec_tok, elec_disc = load_electra(model_dir); print(" ✓ ELECTRA")398 glyph_map = load_glyph_map(data_dir)399 400 models = (tok, cls, ngram, elec_tok, elec_disc, glyph_map)401 402 if args.task == "validate": task_validate(args.sequence, models)403 elif args.task == "predict": task_predict(args.sequence, models)404 elif args.task == "generate": task_generate(args.count, models, args.threshold)405 elif args.task == "score": task_score(args.sequence, models)406 elif args.task == "demo": task_demo(models, glyph_map)407 408 409if __name__ == "__main__":410 main()