| 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() |
| |
| 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() |
| am=L.argmax(-1) |
| idt=torch.tensor(ids,device="cuda") |
| std=(am[:-1]==idt[1:]).float().mean().item() |
| off=(am==idt).float().mean().item() |
| 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) |
|
|