Team Ai
Modelpublic

hellosindh/indus-script-models

sourceHugging Facecc-by-4.0updated 6mo agoView on Hugging Face
0likes
inference.py410 linesDownload Raw Back to root
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()