Mode-Aware Mask Schedule / experiment.py

Failed on benchmark

Raw ⬇ ZIP
 1import json, math, random
 2import numpy as np
 3import torch
 4import torch.nn as nn
 5import torch.nn.functional as F
 6SEED=1669
 7random.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED)
 8torch.set_num_threads(4)
 9try:
10    device="cuda" if torch.cuda.is_available() else "cpu"
11    if device=="cuda": torch.zeros(1,device="cuda")
12except Exception:
13    device="cpu"
14
15class Schedule:
16    def __init__(self,n,lam=0.,rho=0.,s=1):
17        self.n,self.lam,self.rho,self.s=n,lam,rho,s
18        assert lam+rho<=1
19    def sample(self,batch,generator=None):
20        u=torch.rand(batch,generator=generator); branch=torch.zeros(batch,dtype=torch.long)
21        low=(u>=1-self.rho-self.lam)&(u<1-self.rho); full=u>=1-self.rho
22        branch[low]=1; branch[full]=2
23        vis=torch.full((batch,),self.n-1,dtype=torch.long)
24        if low.any(): vis[low]=torch.randint(0,self.s+1,(int(low.sum()),),generator=generator)
25        vis[full]=0
26        rank=torch.rand(batch,self.n,generator=generator).argsort(1)
27        visible=rank < vis[:,None]
28        return ~visible,branch,vis
29    def analytic(self):
30        return {"v_min":0 if self.lam+self.rho>0 else self.n-1,
31                "pi_s":self.rho+self.lam if self.s>=1 else self.rho+self.lam/(self.s+1),
32                "full_atom":self.rho}
33
34def schedule_check():
35    rows=[]
36    for lam,rho in [(0,0),(.1,0),(.1,.01),(.25,.01)]:
37        sc=Schedule(8,lam,rho,1); g=torch.Generator().manual_seed(SEED+7)
38        _,b,v=sc.sample(100000,g)
39        rows.append({"lambda":lam,"rho":rho,"pred_pi1":sc.analytic()["pi_s"],"obs_pi1":float((v<=1).float().mean()),"pred_full":rho,"obs_full":float((b==2).float().mean()),"v_min":int(v.min())})
40    return rows
41
42def anchor_check():
43    p=.8; eta=1.; steps=100; rows=[]
44    for rho in [.01,.03,.1,.2]:
45        b=0.
46        for _ in range(steps): b-=eta*rho*(1/(1+math.exp(-b))-p)
47        q=1/(1+math.exp(-b)); obs=abs(q-p); pred=.3*(1-eta*rho*p*(1-p))**steps
48        rows.append({"rho":rho,"observed_error":obs,"linear_prediction":pred,"ratio":obs/pred})
49    return rows
50
51class TinyMaskedPredictor(nn.Module):
52    def __init__(self,n=8):
53        super().__init__(); self.n=n; self.emb=nn.Embedding(3,12); self.pos=nn.Embedding(n,12)
54        self.net=nn.Sequential(nn.Linear(n*12,40),nn.Tanh(),nn.Linear(40,n*2))
55    def forward(self,x):
56        z=self.emb(x.long())+self.pos(torch.arange(self.n,device=x.device))[None]
57        return self.net(z.reshape(x.shape[0],-1)).reshape(x.shape[0],self.n,2)
58
59def make_data(num,n=8,p=.8):
60    mode=(torch.rand(num)<p).long(); alt=torch.tensor([j%2 for j in range(n)]).long()
61    x=torch.where(mode[:,None].bool(),torch.ones(num,n,dtype=torch.long),alt[None].expand(num,n)).clone()
62    noise=torch.rand(num,n)<.04; return torch.where(noise,1-x,x),mode
63
64def train_eval(lam,rho,steps=250):
65    torch.manual_seed(SEED); n=8; x,_=make_data(1200,n); xv,_=make_data(500,n)
66    model=TinyMaskedPredictor(n).to(device); opt=torch.optim.Adam(model.parameters(),lr=.005); g=torch.Generator().manual_seed(SEED+99); sc=Schedule(n,lam,rho,1)
67    for _ in range(steps):
68        ix=torch.randint(0,len(x),(64,),generator=g); xb=x[ix].to(device); masks,_,_=sc.sample(len(ix),g); masks=masks.to(device)
69        inp=xb.clone(); inp[masks]=2; loss=F.cross_entropy(model(inp)[masks],xb[masks])
70        opt.zero_grad(); loss.backward(); opt.step()
71    model.eval()
72    with torch.no_grad():
73        xvd=xv.to(device); inp=xvd.clone(); mask=torch.zeros_like(inp,dtype=torch.bool); mask[:,0]=True; inp[mask]=2
74        cond=float((model(inp).argmax(-1)[mask]==xvd[mask]).float().mean().cpu())
75        empty=torch.full((len(xv),n),2,dtype=torch.long,device=device); q=float(model(empty).softmax(-1)[:,0,1].mean().cpu())
76    return {"lambda":lam,"rho":rho,"conditional_accuracy":cond,"empty_mode_prob":q,"mode_abs_error":abs(q-.8)}
77
78def main():
79    out={"device":device,"schedule_checks":schedule_check(),"anchor_checks":anchor_check(),"runs":[train_eval(*z) for z in [(0,0),(.1,0),(.1,.01),(.25,.01)]]}
80    open("results.json","w").write(json.dumps(out,indent=2)); print(json.dumps(out,indent=2))
81if __name__=="__main__": main()