File size: 2,516 Bytes
2495418
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
import os, sys, 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()
EOS=tok.eos_token_id
@torch.no_grad()
def gen(prompt,k=50,rep=1.8,top_k=40,top_p=0.95,temp=0.8,nrn=3):
    ids=[151830]+tok(prompt,add_special_tokens=False)["input_ids"]  # prepend [Y2000]=151830 (v11 训练格式)
    start=len(ids)
    for _ in range(k):
        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)
        lg=(o[0,-1] if o.shape[0]==1 else o[-1,0]).float()
        for t in set(ids): lg[t]=lg[t]/rep if lg[t]>0 else lg[t]*rep   # rep penalty
        if len(ids)>=nrn-1:                                            # no_repeat_ngram=3
            pref=tuple(ids[-(nrn-1):])
            for i in range(len(ids)-nrn+1):
                if tuple(ids[i:i+nrn-1])==pref: lg[ids[i+nrn-1]]=-1e9
        lg=lg/temp
        v,i=lg.topk(min(top_k,lg.numel()))                            # top-k
        sp=torch.softmax(v,-1); cum=sp.cumsum(-1); keep=cum<=top_p; keep[0]=True  # top-p
        v=v[keep]; i=i[keep]; pr=torch.softmax(v,-1); nxt=int(i[torch.multinomial(pr,1)])
        ids.append(nxt)
        if EOS is not None and nxt==EOS: break
    return tok.decode(ids[start:])
FIN=["The Company's total revenue for the fiscal year","Net income increased primarily due to",
     "The consolidated financial statements","Our principal risk factors include"]
GEN=["The capital of France is","Water is made of hydrogen and"]
print("== FINANCE (rep1.8/topk40/topp0.95/temp0.8/no_repeat_ngram3) ==",flush=True)
for p in FIN: print(f"[fin] {p!r} -> {gen(p)!r}",flush=True)
print("== GENERAL ==",flush=True)
for p in GEN: print(f"[gen] {p!r} -> {gen(p)!r}",flush=True)
print("TALK2_DONE",flush=True)