Mode-Aware Mask Schedule / bench_mode_schedule.py

Failed on benchmark

Raw ⬇ ZIP
  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()