Dissipative Softmax Latent Layer / official_bench.py

Failed on benchmark

Raw ⬇ ZIP
  1import json, random
  2import numpy as np
  3import torch
  4from torch import nn
  5import sys
  6sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
  7from bench import get_dataset, make_model, train_model, evaluate, sweep_baseline, make_report
  8
  9TRACK='token_expert_sequence'
 10MODEL='mlp_tiny'
 11SEEDS=tuple(range(8))
 12EPOCHS=25
 13
 14class DissipativeNet(nn.Module):
 15    """Same mlp_tiny backbone, with a differentiable reset/cycle latent readout."""
 16    def __init__(self, input_shape, out_dim=2, reset_rate=2.0, drive=.12, steps=16):
 17        super().__init__()
 18        # Exact mlp_tiny: 24 -> 64 -> 64 -> 2.
 19        dim=int(np.prod(input_shape))
 20        self.body=nn.Sequential(nn.Linear(dim,64),nn.ReLU(),nn.Linear(64,64),nn.ReLU(),nn.Linear(64,out_dim))
 21        self.r=float(reset_rate); self.drive=float(drive); self.steps=int(steps)
 22    def forward(self,x):
 23        X=self.body(x.reshape(x.shape[0],-1)); K=X.shape[1]
 24        p=torch.softmax(X,dim=1)
 25        B=x.shape[0]; rates=torch.zeros(B,K+1,K+1,device=x.device,dtype=x.dtype)
 26        rates[:,0,1:]=self.r*p
 27        rates[:,1:,0]=self.r
 28        d=(X[:,None,:]-X[:,:,None])/2
 29        W=torch.exp(d)*(1-torch.eye(K,device=x.device,dtype=x.dtype)[None])
 30        rates[:,1:,1:]=W
 31        # Directed reset-sector cycle 0 -> 1 -> ... -> K-1 -> 0.
 32        for i in range(K-1): rates[:,i,i+1]+=self.drive
 33        rates[:,K-1,0]+=self.drive
 34        total=rates.sum(2)
 35        dt=.45/(total.max().detach()+1e-6)
 36        P=torch.eye(K+1,device=x.device,dtype=x.dtype)[None]-dt*torch.diag_embed(total)+dt*rates
 37        q=torch.zeros(B,K+1,device=x.device,dtype=x.dtype); q[:,0]=1.
 38        for _ in range(self.steps): q=torch.bmm(q[:,None,:],P).squeeze(1).clamp_min(1e-8); q=q/q.sum(1,keepdim=True)
 39        return torch.log(q[:,1:].clamp_min(1e-8))
 40
 41def set_seed(s):
 42    random.seed(s); np.random.seed(s); torch.manual_seed(s)
 43
 44def baseline_fn(cfg):
 45    def run(seed):
 46        set_seed(seed); d=get_dataset(TRACK,seed,n_train=400,n_test=400)
 47        net=nn.Sequential(nn.Flatten(), make_model(MODEL,d['input_shape'],d['out_dim']))
 48        _,metric,_=train_model(net,d,epochs=cfg['epochs'],lr=cfg['lr'],batch=128,log=lambda *_:None)
 49        return metric
 50    return run
 51
 52def idea_fn(cfg):
 53    def run(seed):
 54        set_seed(seed); d=get_dataset(TRACK,seed,n_train=400,n_test=400)
 55        net=DissipativeNet(d['input_shape'],d['out_dim'],cfg['r'])
 56        _,metric,_=train_model(net,d,epochs=cfg['epochs'],lr=cfg['lr'],batch=128,log=lambda *_:None)
 57        return metric
 58    return run
 59
 60def signature(cfg):
 61    vals=[]; currents=[]
 62    for seed in SEEDS:
 63        set_seed(seed); d=get_dataset(TRACK,seed,n_train=400,n_test=400)
 64        net=DissipativeNet(d['input_shape'],d['out_dim'],cfg['r'])
 65        net,_,_=train_model(net,d,epochs=cfg['epochs'],lr=cfg['lr'],batch=128,log=lambda *_:None)
 66        with torch.no_grad():
 67            dev=next(net.parameters()).device; xt=torch.as_tensor(d['xte'],dtype=torch.float32,device=dev)
 68            X=net.body(xt.reshape(len(d['xte']),-1)); p=torch.softmax(X,1).cpu().numpy()
 69            q=torch.exp(net(xt)).cpu().numpy()
 70        vals.append(float(np.abs(q-p).mean()))
 71        currents.append(float(np.abs(q[:,0]-q[:,1]).mean()))
 72    # A model-behavior check, not an analytical identity: rapid reset should
 73    # reduce occupation mismatch relative to the nearby low-reset setting.
 74    low=[]
 75    for seed in SEEDS:
 76        set_seed(seed); d=get_dataset(TRACK,seed,n_train=400,n_test=400)
 77        net=DissipativeNet(d['input_shape'],d['out_dim'],.5)
 78        net,_,_=train_model(net,d,epochs=EPOCHS,lr=cfg['lr'],batch=128,log=lambda *_:None)
 79        with torch.no_grad():
 80            dev=next(net.parameters()).device; xt=torch.as_tensor(d['xte'],dtype=torch.float32,device=dev)
 81            X=net.body(xt.reshape(len(d['xte']),-1)); p=torch.softmax(X,1); q=torch.exp(net(xt))
 82        low.append(float(torch.abs(q-p).mean()))
 83    return {'predicted_effect':'occupation mismatch decreases as reset_rate increases', 'mismatch_at_r':float(np.mean(vals)), 'mismatch_at_r_0.5':float(np.mean(low)), 'observed_cycle_proxy':float(np.mean(currents)), 'confirmed':bool(np.mean(vals)<np.mean(low))}
 84
 85def main():
 86    # Union parity: all idea learning rates are also baseline candidates.
 87    grid=[{'lr':lr,'epochs':EPOCHS} for lr in (1e-3,3e-3,1e-2)]
 88    base=sweep_baseline(baseline_fn,grid,seeds=(0,1,2,3))
 89    # Three idea settings at the selected baseline lr; reset is the method knob.
 90    lr=base['best_cfg']['lr']
 91    idea_cfgs=[{'lr':lr,'epochs':EPOCHS,'r':r} for r in (.5,2.,8.)]
 92    idea_runs=[(cfg,evaluate(idea_fn(cfg),seeds=SEEDS)) for cfg in idea_cfgs]
 93    best_cfg,idea_res=min(idea_runs,key=lambda z:z[1]['mean'])
 94    extra=signature(best_cfg)
 95    rep=make_report(TRACK,MODEL,base,idea_res,extra=extra)
 96    rep['idea_sweep']=[{'cfg':cfg,'result':res} for cfg,res in idea_runs]
 97    rep['custom_track']={'name':TRACK,'file':'/home/maxwelhelp/all/math2nn/bench/custom_tracks/token_expert_sequence.py','domain':'moe-routing'}
 98    with open('bench_report.json','w') as f: json.dump(rep,f,indent=2)
 99    print(json.dumps(rep,indent=2))
100if __name__=='__main__': main()