Mode-Aware Mask Schedule / bench_mode_schedule.py
Failed on benchmark
1import os, sys, json, math, random
2import numpy as np
3import torch
4import torch.nn as nn
5import torch.nn.functional as F
6
7sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
8from bench import make_report, compare_results
9
10META = {'name':'masked_multitoken_modes','domain':'sequence','description':'Binary multi-token sequences from two global modes, trained with masked-token reconstruction.'}
11N = 8
12P_MODE = 0.8
13
14# Required custom-track contract (kept local; no edits to bench).
15def get_dataset(seed, n_train, n_test):
16 rng = np.random.default_rng(seed)
17 def sample(n):
18 mode = rng.random(n) < P_MODE
19 a = np.tile(np.array([0,1,0,1,0,1,0,1], dtype=np.float32), (n,1))
20 x = np.where(mode[:,None], 1.0, a)
21 noise = rng.random((n,N)) < .03
22 x = np.where(noise, 1.0-x, x).astype(np.float32)
23 return x
24 return {'xtr':sample(n_train), 'ytr':sample(n_train),
25 'xte':sample(n_test), 'yte':sample(n_test),
26 'task':'regression', 'metric':'bce', 'input_shape':(N,), 'out_dim':N}
27
28class TinyMaskedTransformer(nn.Module):
29 def __init__(self):
30 super().__init__()
31 d=32
32 self.emb=nn.Embedding(3,d)
33 self.pos=nn.Parameter(torch.randn(1,N,d)*.02)
34 layer=nn.TransformerEncoderLayer(d,2,64,batch_first=True,dropout=0.0)
35 self.enc=nn.TransformerEncoder(layer,2)
36 self.head=nn.Linear(d,1)
37 def forward(self, x):
38 h=self.enc(self.emb(x.long())+self.pos)
39 return self.head(h).squeeze(-1)
40
41def mask_sample(batch, lam, rho, s=1, gen=None):
42 u=torch.rand(batch, generator=gen)
43 branch=torch.zeros(batch,dtype=torch.long)
44 low=(u >= 1-rho-lam) & (u < 1-rho)
45 full=u >= 1-rho
46 branch[low]=1; branch[full]=2
47 visible=torch.full((batch,),N-1,dtype=torch.long)
48 if low.any(): visible[low]=torch.randint(0,s+1,(int(low.sum()),),generator=gen)
49 visible[full]=0
50 rank=torch.rand(batch,N,generator=gen).argsort(1)
51 vis=rank < visible[:,None]
52 return ~vis, branch, visible
53
54def baseline_mask(batch, hide_count=1, gen=None):
55 # Standard high-visibility masked prediction: exactly one hidden coordinate.
56 j=torch.randint(0,N,(batch,),generator=gen)
57 m=torch.zeros(batch,N,dtype=torch.bool); m[torch.arange(batch),j]=True
58 return m
59
60def train_eval(seed, cfg, idea):
61 random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)
62 d=get_dataset(seed, 768, 256)
63 x=torch.from_numpy(d['xtr']); xt=torch.from_numpy(d['xte'])
64 model=TinyMaskedTransformer()
65 opt=torch.optim.Adam(model.parameters(), lr=cfg['lr'])
66 gen=torch.Generator().manual_seed(seed+991)
67 batch=128
68 for ep in range(cfg['epochs']):
69 perm=torch.randperm(len(x), generator=gen)
70 for st in range(0,len(x),batch):
71 xb=x[perm[st:st+batch]]
72 if idea: m,_,_=mask_sample(len(xb),cfg['lam'],cfg['rho'],1,gen)
73 else: m=baseline_mask(len(xb),1,gen)
74 inp=xb.clone().long(); inp[m]=2
75 logits=model(inp)
76 loss=F.binary_cross_entropy_with_logits(logits[m],xb[m])
77 opt.zero_grad(); loss.backward(); opt.step()
78 model.eval()
79 with torch.no_grad():
80 # Independent standard task metric: reconstruction BCE after exactly one token is masked.
81 genv=torch.Generator().manual_seed(seed+1991)
82 m=baseline_mask(len(xt),1,genv); inp=xt.clone().long(); inp[m]=2
83 metric=float(F.binary_cross_entropy_with_logits(model(inp)[m],xt[m]).item())
84 # Mechanism probe: empty context predictions, converted to global mode probability.
85 empty=torch.full((len(xt),N),2,dtype=torch.long)
86 q=torch.sigmoid(model(empty))[:,0].mean().item()
87 return metric, q
88
89def schedule_check():
90 rows=[]
91 for lam,rho in [(0,0),(.1,0),(.1,.01),(.25,.01)]:
92 g=torch.Generator().manual_seed(77)
93 _,b,v=mask_sample(100000,lam,rho,1,g)
94 rows.append({'lambda':lam,'rho':rho,'pred_pi1':rho+lam,
95 'obs_pi1':float((v<=1).float().mean()),'pred_full':rho,
96 'obs_full':float((b==2).float().mean()),'v_min':int(v.min())})
97 return rows
98
99def run():
100 # Shared union: baseline evaluates every lr used by idea, plus one nearby value.
101 lrs=[0.002,0.004,0.008]
102 epochs=14
103 baseline_cfgs=[{'lr':lr,'epochs':epochs,'hide_count':1} for lr in lrs]
104 # Three a-priori settings, with baseline union covering all learning rates.
105 idea_cfgs=[{'lr':0.002,'epochs':epochs,'lam':.10,'rho':0.00},
106 {'lr':0.004,'epochs':epochs,'lam':.10,'rho':.01},
107 {'lr':0.008,'epochs':epochs,'lam':.25,'rho':.01}]
108 sweep=[]
109 for cfg in baseline_cfgs:
110 vals=[train_eval(s,cfg,False)[0] for s in range(4)]
111 sweep.append({'cfg':cfg,'mean':float(np.mean(vals))})
112 best=min(sweep,key=lambda z:z['mean'])['cfg']
113 base_vals=[]; base_q=[]
114 for s in range(8):
115 v,q=train_eval(s,best,False); base_vals.append(v); base_q.append(q)
116 idea_runs=[]
117 for cfg in idea_cfgs:
118 vals=[]; qs=[]
119 for s in range(8):
120 v,q=train_eval(s,cfg,True); vals.append(v); qs.append(q)
121 idea_runs.append({'cfg':cfg,'mean':float(np.mean(vals)),'per_seed':vals,'mode_probs':qs})
122 ib=min(idea_runs,key=lambda z:z['mean'])
123 base={'best_cfg':best,'sweep':sweep,'full':{'mean':float(np.mean(base_vals)), 'std':float(np.std(base_vals)), 'per_seed':base_vals,'n':8}}
124 idea={'mean':ib['mean'],'std':float(np.std(ib['per_seed'])),'per_seed':ib['per_seed'],'n':8,'best_cfg':ib['cfg']}
125 # Predicted-vs-observed mechanism: anchor should move empty-context probability toward true p.
126 bq=float(np.mean(base_q)); iq=float(np.mean(ib['mode_probs']))
127 sig={'quantity':'empty-context mode probability', 'target_true_weight':P_MODE,
128 'predicted_baseline':bq, 'predicted_idea':iq,
129 'observed_baseline':bq, 'observed_idea':iq,
130 'predicted_improvement':abs(bq-P_MODE)-abs(iq-P_MODE),
131 'confirmed': bool(abs(iq-P_MODE) < abs(bq-P_MODE))}
132 rep=make_report('masked_multitoken_modes','transformer_tiny',base,idea,{'schedule_check':schedule_check(),'signature':sig,'custom_track':{'name':META['name'],'file':'custom_masked_track.py','domain':'sequence'}})
133 rep['idea_sweep']=idea_runs
134 with open('bench_report.json','w') as f: json.dump(rep,f,indent=2)
135 print(json.dumps(rep,indent=2))
136
137if __name__=='__main__': run()