lfqian's picture
Upload eval_scripts/align_check.py with huggingface_hub
59aae3e verified
Raw
History Blame Contribute Delete
2.53 kB
import os, sys, json, glob, torch
os.environ.setdefault("ANNULUS_YEAR_OUTPUT","0")
os.environ["ANNULUS_ROUTED_EXPERTS"]="32"; os.environ["ANNULUS_SHARED_EXPERTS"]="1"
os.environ["ANNULUS_SHARED_FFN"]="2048"; os.environ.setdefault("ANNULUS_LAYERS","24")
os.environ.setdefault("ANNULUS_TOPK","8"); os.environ["ANNULUS_GROUPED_GEMM"]="0"; os.environ["ANNULUS_GROUP_AUX"]="1"
S="/gpfs/radev/scratch/xu_hua/lq62/annulus_v4"; _CV7=S+"/code_v7"; _REPO=os.path.expanduser("~/Annulus")
for p in [_REPO+"/eval",_REPO+"/nemo/src",_CV7]:
if os.path.isdir(p):
if p in sys.path: sys.path.remove(p)
sys.path.insert(0,p)
import icl_eval_v5 as V
core,tok=V.build_v5_model_and_tokenizer(os.environ["CKPT"],os.environ["TOK"]); core.eval()
# find real text
cands=glob.glob(os.path.expanduser("~/Annulus")+"/**/heldout_2001_samesource.jsonl",recursive=True)+\
glob.glob("/gpfs/radev/scratch/xu_hua/lq62/annulus_v5/delta_test/ppl_to2005.jsonl")+\
glob.glob(S+"/**/ppl_to2005.jsonl",recursive=True)
texts=[]
for f in cands:
for ln in open(f):
try: texts.append(json.loads(ln)["text"])
except: pass
if texts: print("data:",f,"docs:",len(texts),flush=True); break
if not texts: texts=["The Company reported total revenue of 4.2 billion dollars for the fiscal year ended December, an increase driven by higher product sales and improved margins across all operating segments."]*5
@torch.no_grad()
def acc(doc):
ids=tok(doc,add_special_tokens=False)["input_ids"][:512]
if len(ids)<8: return None
s=len(ids); inp=torch.tensor([ids],device="cuda"); pos=torch.arange(s,device="cuda")[None]
m=torch.triu(torch.ones(s,s,dtype=torch.bool,device="cuda"),1)[None,None]
o=core(input_ids=inp,position_ids=pos,attention_mask=m)
L=(o[0] if o.shape[0]==1 else o[:,0]).float() # [seq,V]
am=L.argmax(-1) # top1 per pos
idt=torch.tensor(ids,device="cuda")
std=(am[:-1]==idt[1:]).float().mean().item() # logits[i]->ids[i+1]
off=(am==idt).float().mean().item() # logits[i]->ids[i]
lp=torch.log_softmax(L[:-1],-1); nll=(-lp.gather(1,idt[1:,None]).squeeze(1)).mean().item()
return std,off,nll
import statistics as st
S_,O_,N_=[],[],[]
for d in texts[:30]:
r=acc(d)
if r: S_.append(r[0]);O_.append(r[1]);N_.append(r[2])
print(f"[ALIGN] n={len(S_)} std_acc(logits[i]->ids[i+1])={st.mean(S_):.3f} offby1_acc(logits[i]->ids[i])={st.mean(O_):.3f} NLL={st.mean(N_):.3f}",flush=True)
print("ALIGN_DONE",flush=True)