Download code/baseline_clean_blocks.py from sayed125/lhc-0-brain: direct link, hf CLI and curl.
- Browser
- Download file 5.83 kB
-
https://huggingface.co/sayed125/lhc-0-brain/resolve/main/code/baseline_clean_blocks.py
- Command line
-
hf download hf://sayed125/lhc-0-brain/code/baseline_clean_blocks.py
-
curl -L -o baseline_clean_blocks.py https://huggingface.co/sayed125/lhc-0-brain/resolve/main/code/baseline_clean_blocks.py
5.83 kB
| #!/usr/bin/env python3 | |
| # -*- coding: utf-8 -*- | |
| """baseline_clean_blocks.py — e5-small على نفس كتل التقييم (بذرة 41) لنفس العناصر: | |
| - مقارنة مقترنة LHC ↔ e5 على العناصر النظيفة (McNemar) | |
| - اختبار تحكم: فجوة «التسرب» لـe5 يجب أن تكون ≈0 (لم يتدرب على هذه المجموعات) | |
| - ملف مضغوط: إصابات e5 لكل عنصر (لإعادة الاستخدام) | |
| """ | |
| import os, sys, json, time | |
| import numpy as np | |
| sys.path.insert(0,"/home/user/lhc/code"); sys.path.insert(0,"/home/user/lhc/recon") | |
| from phase_w1_data import download_corpus | |
| from scipy.stats import binomtest | |
| import torch | |
| from transformers import AutoTokenizer, AutoModel | |
| torch.set_num_threads(2) | |
| T0=time.time() | |
| def log(*a): print(f"[{time.time()-T0:6.1f}s]", *a, flush=True) | |
| R="/home/user/lhc/results"; BS=300; KBLOCKS=4 | |
| M=json.load(open(f"{R}/masks_v2.json")); masks={cn:set(M["masks_strong"][cn]) for cn in M["masks_strong"]} | |
| DD="/home/user/lhc/data/opus" | |
| OPUS=[("Tanzil","https://object.pouta.csc.fi/OPUS-Tanzil/v1/moses/ar-en.txt.zip",20000), | |
| ("Bible","https://object.pouta.csc.fi/OPUS-bible-uedin/v1/moses/ar-en.txt.zip",20000), | |
| ("NeuLab-TED","https://object.pouta.csc.fi/OPUS-NeuLab-TedTalks/v1/moses/ar-en.txt.zip",20000)] | |
| pairs={} | |
| for name,url,lim in OPUS: pairs[name]=download_corpus(name,url,lim,DD) | |
| DEV="/home/user/lhc/data/flores101_dataset/devtest" | |
| ld=lambda p:[l.strip() for l in open(p,encoding="utf-8") if l.strip()] | |
| pairs["FLORES"]=[{"ar":a,"en":e} for a,e in zip(ld(f"{DEV}/ara.devtest"),ld(f"{DEV}/eng.devtest"))] | |
| def blocks(n, seed=41): | |
| rng=np.random.RandomState(seed); idx=rng.permutation(n) | |
| return [idx[i:i+BS] for i in range(0,n-BS+1,BS)] | |
| tok=AutoTokenizer.from_pretrained("intfloat/multilingual-e5-small") | |
| model=AutoModel.from_pretrained("intfloat/multilingual-e5-small").eval() | |
| def e5(texts, prefix, bs=16): | |
| embs=[] | |
| with torch.no_grad(): | |
| for i in range(0,len(texts),bs): | |
| t=tok([prefix+x for x in texts[i:i+bs]],padding=True,truncation=True,max_length=512,return_tensors="pt") | |
| h=model(**t).last_hidden_state | |
| m=t["attention_mask"].unsqueeze(-1).float() | |
| embs.append(((h*m).sum(1)/m.sum(1)).numpy()) | |
| return np.vstack(embs) | |
| zh=np.load(f"{R}/block_hits.npz") | |
| out={"blocks_per_set":KBLOCKS, "block_size":BS, "subsets":{}} | |
| store={} | |
| for cn in ("Tanzil","Bible","NeuLab-TED","FLORES"): | |
| P=pairs[cn]; bl=blocks(len(P))[:KBLOCKS] | |
| items=[int(i) for b in bl for i in b] | |
| ar=[P[i]["ar"] for i in items]; en=[P[i]["en"] for i in items] | |
| inmask=np.array([i in masks[cn] for i in items], bool) | |
| log(f"{cn}: {len(items)} عنصرًا ({KBLOCKS} كتل) — ترميز e5 ...") | |
| qa=e5(ar,"query: "); pe=e5(en,"passage: "); qe=e5(en,"query: "); pa=e5(ar,"passage: ") | |
| def norm(x): return x/np.maximum(np.linalg.norm(x,axis=1,keepdims=True),1e-9) | |
| sim=norm(qa)@norm(pe).T; pred=sim.argmax(axis=1) | |
| hit_ae=np.array([P[items[i]]["en"]==en[pred[i]] for i in range(len(items))],bool) | |
| sim2=norm(qe)@norm(pa).T; pred2=sim2.argmax(axis=0) | |
| hit_ea=np.array([P[items[j]]["ar"]==ar[pred2[j]] for j in range(len(items))],bool) | |
| # إصابات LHC لنفس المواضع (من block_hits: أول KBLOCKS*300 من كل مجموعة) | |
| n_used=KBLOCKS*BS | |
| lhc_ae=zh[f"{cn}|W4|ar_en"][:n_used]; lhc_ea=zh[f"{cn}|W4|en_ar"][:n_used] | |
| lhc_m =zh[f"{cn}|W4|mask"][:n_used] | |
| assert np.array_equal(lhc_m, inmask), "عدم تطابق قناع!" | |
| c=~inmask; n_clean=int(c.sum()) | |
| res={ | |
| "n":len(items),"n_clean":n_clean,"n_masked":int(inmask.sum()), | |
| "e5": {"r1_ar_en":round(float(hit_ae.mean()),4),"r1_en_ar":round(float(hit_ea.mean()),4), | |
| "clean_ar_en":round(float(hit_ae[c].mean()),4) if c.any() else None, | |
| "clean_en_ar":round(float(hit_ea[c].mean()),4) if c.any() else None, | |
| "masked_ar_en":round(float(hit_ae[inmask].mean()),4) if inmask.any() else None, | |
| "masked_en_ar":round(float(hit_ea[inmask].mean()),4) if inmask.any() else None}, | |
| "lhc_w4": {"r1_ar_en":round(float(lhc_ae.mean()),4),"r1_en_ar":round(float(lhc_ea.mean()),4), | |
| "clean_ar_en":round(float(lhc_ae[c].mean()),4) if c.any() else None, | |
| "clean_en_ar":round(float(lhc_ea[c].mean()),4) if c.any() else None, | |
| "masked_ar_en":round(float(lhc_ae[inmask].mean()),4) if inmask.any() else None, | |
| "masked_en_ar":round(float(lhc_ea[inmask].mean()),4) if inmask.any() else None}, | |
| } | |
| # مقارنة مقترنة على العناصر النظيفة فقط | |
| for dname, a, b in (("ar_en", lhc_ae[c], hit_ae[c]), ("en_ar", lhc_ea[c], hit_ea[c])): | |
| b01=int(np.sum(~a&b)); b10=int(np.sum(a&~b)); n=b01+b10 | |
| p=float(binomtest(b01,n,0.5).pvalue) if n>0 else 1.0 | |
| res[f"mcnemar_clean_{dname}"]={"delta_lhc_minus_e5":round(float(a.mean()-b.mean()),4), | |
| "b01_e5_only":b01,"b10_lhc_only":b10,"p":round(p,6) if n>0 else None,"n_discordant":n} | |
| out["subsets"][cn]=res | |
| store[f"{cn}_e5_ar_en"]=hit_ae; store[f"{cn}_e5_en_ar"]=hit_ea; store[f"{cn}_items"]=np.array(items) | |
| log(f" e5 كامل: {res['e5']['r1_ar_en']}/{res['e5']['r1_en_ar']} | نظيف: {res['e5']['clean_ar_en']}/{res['e5']['clean_en_ar']} | LHC نظيف: {res['lhc_w4']['clean_ar_en']}/{res['lhc_w4']['clean_en_ar']}") | |
| for dname in ("ar_en","en_ar"): | |
| m=res[f"mcnemar_clean_{dname}"] | |
| log(f" McNemar-نظيف [{dname}]: Δ={m['delta_lhc_minus_e5']:+.4f} p={m['p']} (b01={m['b01_e5_only']},b10={m['b10_lhc_only']})") | |
| np.savez_compressed(f"{R}/baseline_blocks_e5.npz", **store) | |
| json.dump(out, open(f"{R}/baseline_clean_blocks.json","w"), ensure_ascii=False, indent=1) | |
| log("SAVED baseline_clean_blocks.json + baseline_blocks_e5.npz") | |