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)