File size: 2,531 Bytes
59aae3e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
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)