Self-Supervised Amortized Mean-Field Controller / bench_stage2.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
  1import sys, json, random, time
  2import numpy as np
  3import torch
  4from torch import nn
  5sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
  6from bench import get_dataset, make_model, train_model, evaluate, sweep_baseline, make_report
  7
  8# Matched dynamics/control track. The intervention is a training loss, so a local loop is required.
  9SEEDS = tuple(range(8))
 10SWEEP_SEEDS = (0,1,2,3)
 11EPOCHS = 12
 12BATCH = 128
 13LRS = [1e-3, 3e-3, 6e-3]
 14WEIGHTS = [0.0, 0.03, 0.1]  # idea sweep; zero is included only as an honest nearby setting
 15
 16
 17def seed_all(seed):
 18    random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)
 19    if torch.cuda.is_available(): torch.cuda.manual_seed_all(seed)
 20
 21
 22def device_safe():
 23    if torch.cuda.is_available():
 24        try:
 25            torch.zeros(1, device='cuda')
 26            return 'cuda'
 27        except Exception:
 28            pass
 29    return 'cpu'
 30
 31
 32def tensors(ds, dev):
 33    return (torch.as_tensor(ds['xtr'], dtype=torch.float32, device=dev),
 34            torch.as_tensor(ds['ytr'], dtype=torch.float32, device=dev).reshape(-1),
 35            torch.as_tensor(ds['xte'], dtype=torch.float32, device=dev),
 36            torch.as_tensor(ds['yte'], dtype=torch.float32, device=dev).reshape(-1))
 37
 38
 39def train_idea(net, ds, lr, score_weight, seed):
 40    """Direct supervised benchmark target plus a probability-flow-inspired consistency term.
 41    For each history, estimate local score s=-z/Var(z) over its 8 temporal states and
 42    require the predicted next angle to be stable under the deterministic flow z+dt*(-gamma*s).
 43    This is a self-supervised structural regularizer; no oracle trajectory is introduced."""
 44    dev = device_safe(); net = net.to(dev)
 45    x,y,xe,ye = tensors(ds, dev)
 46    opt = torch.optim.Adam(net.parameters(), lr=lr)
 47    loss_fn = nn.MSELoss()
 48    n=x.shape[0]
 49    for ep in range(EPOCHS):
 50        g=torch.Generator(device=dev); g.manual_seed(seed+1000+ep)
 51        perm=torch.randperm(n, generator=g, device=dev)
 52        net.train()
 53        for st in range(0,n,BATCH):
 54            ix=perm[st:st+BATCH]; xb=x[ix]; yb=y[ix]
 55            pred=net(xb).reshape(-1)
 56            task=loss_fn(pred,yb)
 57            if score_weight:
 58                z=xb.view(-1,8,3)
 59                # prompt-free empirical score of the particle/time cloud; centered to avoid
 60                # changing the mean, as in probability-flow u=v-gamma grad log p.
 61                theta=z[:,:,0]; var=theta.var(1,keepdim=True,unbiased=False).clamp_min(1e-3)
 62                score=-(theta-theta.mean(1,keepdim=True))/var
 63                dt=0.05; gamma=0.08
 64                flow_theta=theta + dt*(-gamma*score)
 65                flow_x=xb.clone(); flow_x.view(-1,8,3)[:,:,0]=flow_theta
 66                # same controller should be insensitive to the infinitesimal deterministic
 67                # probability-flow transport, a finite NN-scale testable prediction.
 68                pred_flow=net(flow_x).reshape(-1)
 69                consistency=((pred_flow-pred).square()).mean()
 70                loss=task+score_weight*consistency
 71            else: loss=task
 72            opt.zero_grad(set_to_none=True); loss.backward(); opt.step()
 73    net.eval()
 74    with torch.no_grad(): metric=float(((net(xe).reshape(-1)-ye)**2).mean().cpu())
 75    return metric, net
 76
 77
 78def run(cfg, seed, return_model=False):
 79    seed_all(seed); ds=get_dataset('dynamics',seed,n_train=400,n_test=400)
 80    # Use bench constructor; identical model and input/data for both methods.
 81    net=make_model('rnn_small', ds['input_shape'], ds['out_dim'])
 82    metric, net=train_idea(net,ds,cfg['lr'],cfg.get('score_weight',0.0),seed)
 83    if return_model: return metric,net,ds
 84    return metric
 85
 86
 87def baseline_factory(cfg):
 88    return lambda seed: run({'lr':cfg['lr'],'score_weight':0.0},seed)
 89
 90def idea_factory(cfg):
 91    return lambda seed: run(cfg,seed)
 92
 93
 94def signature(cfg, base_cfg):
 95    # Behavioural NN-scale signature: measured change in predictions under the
 96    # deterministic probability-flow perturbation, on trained models.
 97    vals=[]
 98    for seed in (0,1,2,3):
 99        bm, bnet, ds=run(base_cfg,seed,True)
100        im, inet, _=run(cfg,seed,True)
101        dev=next(inet.parameters()).device
102        x=torch.as_tensor(ds['xte'][:128],dtype=torch.float32,device=dev)
103        z=x.view(-1,8,3); theta=z[:,:,0]; var=theta.var(1,keepdim=True,unbiased=False).clamp_min(1e-3)
104        xf=x.clone(); xf.view(-1,8,3)[:,:,0]=theta+0.05*(-0.08)*(-(theta-theta.mean(1,keepdim=True))/var)
105        with torch.no_grad():
106            db=(bnet(xf).reshape(-1)-bnet(x).reshape(-1)).abs().mean().item()
107            di=(inet(xf).reshape(-1)-inet(x).reshape(-1)).abs().mean().item()
108        vals.append((db,di))
109    b=np.array([v[0] for v in vals]); i=np.array([v[1] for v in vals])
110    ratio=float(i.mean()/(b.mean()+1e-12))
111    return {'prediction':'probability-flow consistency reduces prediction sensitivity to score transport',
112            'baseline_abs_sensitivity_mean':float(b.mean()),'idea_abs_sensitivity_mean':float(i.mean()),
113            'ratio_idea_over_baseline':ratio,'n_models':8,
114            'confirmed':bool(i.mean() < b.mean())}
115
116
117def main():
118    # Baseline grid includes union of every idea lr; central baseline knob is lr.
119    grid=[{'lr':v,'score_weight':0.0} for v in LRS]
120    t=time.time(); base=sweep_baseline(baseline_factory,grid,seeds=SWEEP_SEEDS)
121    # Idea has same lr union and two nearby regularizer strengths.
122    idea_cfgs=[{'lr':base['best_cfg']['lr'],'score_weight':w} for w in [0.03,0.1]]
123    idea_cfgs += [{'lr':v,'score_weight':0.03} for v in LRS if v!=base['best_cfg']['lr']]
124    ir=[]
125    for cfg in idea_cfgs:
126        r=evaluate(idea_factory(cfg),SEEDS); ir.append({'cfg':cfg,'result':r})
127    best=min(ir,key=lambda q:q['result']['mean']); idea=best['result']; cfg=best['cfg']
128    rep=make_report('dynamics','rnn_small',base,idea,{'idea_cfg':cfg,'signature':signature(cfg,base['best_cfg'])})
129    rep['runtime_seconds']=time.time()-t; rep['idea_sweep']=ir
130    rep['protocol_notes']='Dynamics chosen because controlled pendulum rollout is explicitly a stability/control task. Both systems use rnn_small, same datasets, epochs and Adam; only the self-supervised score-transport consistency loss differs.'
131    with open('bench_report.json','w') as f: json.dump(rep,f,indent=2)
132    print(json.dumps(rep,indent=2))
133
134if __name__=='__main__': main()