Directed-Path Synchronization Coupling / directed_path_bench.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
 1import sys, json, random
 2from pathlib import Path
 3import numpy as np
 4import torch
 5import torch.nn as nn
 6sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
 7from bench import get_dataset, train_model, sweep_baseline, evaluate, make_report
 8
 9TRACK='dynamics'; MODEL='rnn_small'; N=4; H=16; EPOCHS=12; BATCH=128
10ROOT=Path(__file__).parent
11
12def seed_all(s):
13    random.seed(s); np.random.seed(s); torch.manual_seed(s)
14    if torch.cuda.is_available(): torch.cuda.manual_seed_all(s)
15
16def lap(n=N):
17    a=np.zeros((n,n))
18    for i in range(1,n): a[i,i]=1.; a[i,0]=-1.
19    return a
20
21def math_check():
22    L=lap(); x0=np.linspace(-1,1,N); dt=.001; rows=[]
23    for eta in (.5,1.):
24        x=x0.copy(); v=[]
25        for _ in range(4000):
26            m=x.mean(); v.append(.5*np.sum((x-m)**2)); x += dt*(.8*x-eta*L@x)
27        obs=float(np.polyfit(np.arange(len(v))*dt,np.log(np.maximum(v,1e-300)),1)[0])
28        pred=-2*(eta-.8)
29        rows.append({'eta':eta,'predicted_slope':pred,'observed_slope':obs,'relative_error':abs(obs-pred)/max(abs(pred),1e-9)})
30    return {'gamma_G':1.,'L':.8,'threshold':.8,'rows':rows,'passed':rows[1]['observed_slope']<0 and rows[0]['observed_slope']>0}
31
32class ParallelGRU(nn.Module):
33    def __init__(self, coupled=False, eta=0.):
34        super().__init__(); self.coupled=coupled; self.eta=float(eta)
35        self.grus=nn.ModuleList([nn.GRUCell(3,H) for _ in range(N)])
36        self.head=nn.Linear(H,1)
37    def forward(self,x, return_signature=False):
38        z=x.view(x.shape[0],-1,3); hs=[torch.zeros(x.shape[0],H,device=x.device) for _ in range(N)]
39        for t in range(z.shape[1]):
40            hs=[g(z[:,t],h) for g,h in zip(self.grus,hs)]
41            if self.coupled:
42                root=hs[0]
43                hs=[hs[0]]+[h+self.eta*(root-h) for h in hs[1:]]
44        stack=torch.stack(hs,1); out=self.head(stack.mean(1))
45        if return_signature: return out, float((stack-stack.mean(1,keepdim=True)).pow(2).mean().sqrt().detach().cpu())
46        return out
47
48def run(cfg, seed, coupled):
49    seed_all(seed); ds=get_dataset(TRACK,seed,n_train=400,n_test=400)
50    model=ParallelGRU(coupled=coupled,eta=cfg['eta'])
51    _, metric, _=train_model(model,ds,epochs=EPOCHS,lr=cfg['lr'],batch=BATCH,log=lambda *_:None)
52    return float(metric)
53
54def train_and_measure(cfg,seed,coupled):
55    seed_all(seed); ds=get_dataset(TRACK,seed,n_train=400,n_test=400)
56    model=ParallelGRU(coupled=coupled,eta=cfg['eta'])
57    net,metric,_=train_model(model,ds,epochs=EPOCHS,lr=cfg['lr'],batch=BATCH,log=lambda *_:None)
58    with torch.no_grad():
59        dev=next(net.parameters()).device
60        _, disagreement=net(ds['xte'].to(dev),return_signature=True)
61    return float(metric), disagreement
62
63def main():
64    check=math_check(); lrs=[0.0015,0.003,0.006]; base_grid=[{'lr':lr,'eta':0.0} for lr in lrs]
65    def base_fn(cfg): return lambda s: run(cfg,s,False)
66    baseline=sweep_baseline(lambda c: base_fn(c),base_grid,seeds=(0,1,2,3))
67    best_lr=baseline['best_cfg']['lr']; idea_grid=[{'lr':best_lr,'eta':e} for e in (.02,.08,.20)]
68    # Baseline was evaluated at every union learning rate; idea settings share best lr.
69    idea_records=[]
70    for cfg in idea_grid:
71        r=evaluate(lambda s: run(cfg,s,True),seeds=tuple(range(8)))
72        idea_records.append({'cfg':cfg,'result':r})
73    best=min(idea_records,key=lambda q:q['result']['mean']); idea=best['result']; cfg=best['cfg']
74    base_full=baseline['full']; sig=[]
75    for s in range(8):
76        bm,bd=train_and_measure({'lr':best_lr,'eta':0.},s,False)
77        im,idg=train_and_measure(cfg,s,True); sig.append({'seed':s,'baseline_disagreement':bd,'idea_disagreement':idg})
78    obs=float(np.mean([q['idea_disagreement'] for q in sig])); bobs=float(np.mean([q['baseline_disagreement'] for q in sig]))
79    predicted='negative disagreement trend when eta*gamma_G exceeds L'; confirmed=bool(cfg['eta']>0 and obs < bobs)
80    rep=make_report(TRACK,MODEL,baseline,idea,{'math_check':check,'idea_grid':idea_records,'trained_behavior':sig,'prediction':predicted,'predicted_threshold':.8,'observed_mean_baseline_disagreement':bobs,'observed_mean_idea_disagreement':obs,'confirmed':confirmed})
81    rep['parameter_counts']={'baseline':sum(p.numel() for p in ParallelGRU(False).parameters()),'idea':sum(p.numel() for p in ParallelGRU(True,cfg['eta']).parameters())}
82    rep['protocol_notes']='8 paired seeds; baseline sweep uses 4 seeds and full re-evaluation; n_train=n_test=400; same ParallelGRU base and optimizer.'
83    (ROOT/'bench_report.json').write_text(json.dumps(rep,indent=2))
84    print(json.dumps(rep,indent=2))
85if __name__=='__main__': main()