Spectrally identifiable phaseless recurrent layer / stage2_bench.py
Mechanism confirmed, baseline not beaten
1import sys,json,random
2import numpy as np
3import torch
4import torch.nn as nn
5sys.path.insert(0,'/home/maxwelhelp/all/math2nn')
6from bench import get_dataset,train_model,evaluate,sweep_baseline,make_report,count_params
7SEEDS=tuple(range(8)); SWEEP_SEEDS=tuple(range(4)); LRS=[1e-3,3e-3,1e-2]; EPOCHS=20
8
9def seed(s):
10 random.seed(s); np.random.seed(s); torch.manual_seed(s)
11 if torch.cuda.is_available(): torch.cuda.manual_seed_all(s)
12
13def op(n=32):
14 L=np.zeros((n,n));
15 for i in range(n-1): L[i,i+1]=L[i+1,i]=1
16 L=np.diag(L.sum(1))-L; rng=np.random.RandomState(1461); best=None
17 for _ in range(30):
18 q=rng.normal(size=n); lam,phi=np.linalg.eigh(L+np.diag(q)); S=phi**2
19 sums=np.array([lam[i]+lam[j] for i in range(n) for j in range(i,n)])
20 gap=np.min(np.abs(sums[:,None]-sums[None,:]+np.eye(len(sums))*1e9)); sm=np.linalg.svd(S,compute_uv=False)[-1]
21 if best is None or gap*sm>best[0]: best=(gap*sm,lam,phi,gap,sm)
22 return best[1].astype('float32'),best[2].astype('float32'),float(best[3]),float(best[4])
23
24class Base(nn.Module):
25 def __init__(self,h=64):
26 super().__init__(); self.r=nn.GRU(3,h,batch_first=True); self.head=nn.Linear(h,1)
27 def forward(self,x): return self.head(self.r(x.view(x.shape[0],-1,3))[1][-1])
28
29class Spectral(nn.Module):
30 def __init__(self,h=64,K=6):
31 super().__init__(); self.K=K; self.inp=nn.Linear(3,h); self.h=h
32 lam,phi,gap,sm=op(h); self.register_buffer('lam',torch.tensor(lam)); self.register_buffer('phi',torch.tensor(phi));
33 self.head=nn.Sequential(nn.Linear(K*h,64),nn.Tanh(),nn.Linear(64,1)); self.gap=gap; self.smin=sm
34 def encode(self,z):
35 c=z if torch.is_complex(z) else torch.complex(z,torch.zeros_like(z)); out=[]
36 for k in range(self.K):
37 out.append(torch.sqrt(c.real*c.real+c.imag*c.imag+1e-8)); c=torch.matmul(torch.exp(-1j*self.lam) * torch.matmul(c, self.phi.to(c.dtype)), self.phi.to(c.dtype).T)
38 return torch.cat(out,1)
39 def forward(self,x):
40 z=self.inp(x.view(x.shape[0],-1,3).mean(1)); c=torch.complex(z,torch.zeros_like(z)); out=[]
41 for k in range(self.K):
42 out.append(torch.sqrt(c.real*c.real+c.imag*c.imag+1e-8)); c=torch.matmul(torch.exp(-1j*self.lam) * torch.matmul(c, self.phi.to(c.dtype)), self.phi.to(c.dtype).T)
43 return self.head(torch.cat(out,1))
44
45def run(model,ds,lr,s):
46 seed(s); model=model(); _,metric,_=train_model(model,ds,epochs=EPOCHS,lr=lr,batch=128,log=lambda *_:None); return metric
47
48def main():
49 def ds(s): return get_dataset('dynamics',s,n_train=400,n_test=200)
50 def mk(c): return lambda s: run(Base,ds(s),c['lr'],s)
51 base=sweep_baseline(mk,[{'lr':x} for x in LRS],seeds=SWEEP_SEEDS)
52 # Full evaluation for every union lr, with the selected best retained as the idea comparison.
53 base_full={}
54 for lr in LRS: base_full[lr]=evaluate(mk({'lr':lr}),SEEDS)
55 best_lr=base['best_cfg']['lr']; idea_cfgs=[best_lr]+[x for x in LRS if x!=best_lr][:2]
56 ir=[]
57 for lr in idea_cfgs: ir.append((lr,evaluate(lambda s:run(lambda:Spectral(K=6),ds(s),lr,s),SEEDS)))
58 idea_lr,idea=min(ir,key=lambda z:z[1]['mean'])
59 d=ds(0); m=Spectral(K=6); seed(0); _,_,_=train_model(m,d,epochs=EPOCHS,lr=idea_lr,batch=128,log=lambda *_:None)
60 m.eval(); dev=next(m.parameters()).device; x=d['xte'][:32].to(dev)
61 with torch.no_grad():
62 z=m.inp(x.view(x.shape[0],-1,3).mean(1)); e1=m.encode(z); e2=m.encode(torch.complex(z,torch.zeros_like(z))*torch.exp(1j*torch.tensor(1.234)))
63 sig={'prediction':'global phase is unobservable by magnitude trajectory','predicted_max_abs_difference':0.0,'observed_mean_abs_difference':float((e1-e2).abs().mean()),'operator_pair_sum_gap':m.gap,'operator_sigma_min_S':m.smin,'confirmed':bool(float((e1-e2).abs().mean())<1e-5)}
64 report=make_report('dynamics','rnn_small',{'best_cfg':base['best_cfg'],'sweep':base['sweep'],'full':base_full[best_lr]},idea,{'mechanism_signature':sig})
65 report['idea_sweep']=[{'cfg':{'lr':lr},'result':r} for lr,r in ir]; report['parameter_counts']={'baseline':count_params(Base()),'idea':count_params(Spectral())}
66 with open('bench_report.json','w') as f: json.dump(report,f,indent=2)
67 print(json.dumps(report,indent=2))
68if __name__=='__main__': main()