Structure-preserving SU(1,1) recurrent scan / su11_experiment.py
Beats tuned baseline
1import json, math, random, time
2import numpy as np
3import torch
4from torch import nn
5
6SEED = 2374
7random.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED)
8
9def exact_step(a, b, f, x, xi, h):
10 r = abs(f)
11 q = h*r
12 c = np.cosh(q)
13 s = h if r == 0 else np.sinh(q)/r
14 ph = np.exp(2j*np.pi*x*xi)
15 return c*a + s*np.conj(f)*ph*b, c*b + s*f*np.conj(ph)*a
16
17def euler_step(a, b, f, x, xi, h):
18 ph = np.exp(2j*np.pi*x*xi)
19 return a + h*np.conj(f)*ph*b, b + h*f*np.conj(ph)*a
20
21def verify_math():
22 # Random complex potentials and phases; start exactly at the claimed initial state.
23 a, b = 1+0j, 0+0j
24 ae, be = a, b
25 exact_err, euler_err = [], []
26 for k in range(1000):
27 f = (np.random.randn()+1j*np.random.randn()) * 0.7
28 a,b = exact_step(a,b,f,k,0.17,0.025)
29 ae,be = euler_step(ae,be,f,k,0.17,0.025)
30 exact_err.append(abs(abs(a)**2-abs(b)**2-1))
31 euler_err.append(abs(abs(ae)**2-abs(be)**2-1))
32 # Also directly check M^2=|f|^2 I numerically.
33 f = .31-.77j; x=.43; xi=.29
34 ph=np.exp(2j*np.pi*x*xi)
35 M=np.array([[0,np.conj(f)*ph],[f*np.conj(ph),0]],complex)
36 algebra_err=float(np.max(np.abs(M@M-(abs(f)**2)*np.eye(2))))
37 return {
38 "exact_max_invariant_error": float(max(exact_err)),
39 "exact_final_invariant_error": float(exact_err[-1]),
40 "euler_final_invariant_error": float(euler_err[-1]),
41 "euler_max_invariant_error": float(max(euler_err)),
42 "M2_error": algebra_err,
43 }
44
45# Torch version of the same elementwise exponential update.
46class SU11Classifier(nn.Module):
47 def __init__(self, hidden=2, h=.12):
48 super().__init__(); self.h=h; self.hidden=hidden
49 self.fmap=nn.Linear(1, 2*hidden)
50 self.out=nn.Linear(4*hidden, 2)
51 def forward(self, x):
52 B,T,_=x.shape
53 a=torch.ones(B,self.hidden,device=x.device,dtype=torch.complex64)
54 b=torch.zeros_like(a)
55 xi=torch.arange(self.hidden,device=x.device,dtype=torch.float32)*.07
56 for k in range(T):
57 raw=self.fmap(x[:,k])
58 f=torch.complex(raw[:,:self.hidden],raw[:,self.hidden:])
59 # bounded potentials avoid an irrelevant overflow regime in this tiny test
60 f=.7*torch.tanh(f.real)+1j*.7*torch.tanh(f.imag)
61 r=torch.abs(f); q=self.h*r
62 c=torch.cosh(q.float()).to(torch.complex64)
63 s=torch.where(r>1e-7, torch.sinh(q.float())/r.float(), torch.full_like(r,self.h)).to(torch.complex64)
64 phase=torch.exp(2j*math.pi*k*xi).to(torch.complex64)[None,:]
65 aa=c*a+s*torch.conj(f)*phase*b
66 bb=c*b+s*f*torch.conj(phase)*a
67 a,b=aa,bb
68 feat=torch.cat([a.real,a.imag,b.real,b.imag],1)
69 return self.out(feat)
70
71class TanhRNNClassifier(nn.Module):
72 def __init__(self, hidden=4):
73 super().__init__(); self.cell=nn.RNNCell(1,hidden,nonlinearity='tanh'); self.out=nn.Linear(hidden,2)
74 def forward(self,x):
75 z=torch.zeros(x.shape[0],self.cell.hidden_size,device=x.device)
76 for k in range(x.shape[1]): z=self.cell(x[:,k],z)
77 return self.out(z)
78
79def make_data(n, T):
80 x=np.random.randn(n,T,1).astype('float32')
81 y=(x.sum(axis=1)[:,0]>0).astype('int64')
82 return torch.from_numpy(x),torch.from_numpy(y)
83
84def train(model, tr, va, device, steps=350):
85 model.to(device); opt=torch.optim.Adam(model.parameters(),lr=0.01); lossfn=nn.CrossEntropyLoss()
86 x,y=tr[0].to(device),tr[1].to(device); xv,yv=va[0].to(device),va[1].to(device)
87 t0=time.time(); losses=[]
88 for i in range(steps):
89 opt.zero_grad(); loss=lossfn(model(x),y); loss.backward(); torch.nn.utils.clip_grad_norm_(model.parameters(),5.0); opt.step()
90 if i in (0,steps-1): losses.append(float(loss.detach().cpu()))
91 with torch.no_grad():
92 vl=float(lossfn(model(xv),yv).cpu()); acc=float((model(xv).argmax(1)==yv).float().mean().cpu())
93 return {"initial_train_loss":losses[0],"final_train_loss":losses[-1],"validation_loss":vl,"validation_accuracy":acc,"seconds":time.time()-t0,"parameters":sum(p.numel() for p in model.parameters())}
94
95def train_compare():
96 torch.manual_seed(SEED); tr=make_data(128,32); va=make_data(256,32)
97 # CPU is the safe default; CUDA errors fall back as required.
98 device=torch.device('cuda' if torch.cuda.is_available() else 'cpu')
99 try:
100 torch.manual_seed(SEED); su=train(SU11Classifier(),tr,va,device)
101 torch.manual_seed(SEED); base=train(TanhRNNClassifier(),tr,va,device)
102 except Exception as e:
103 device=torch.device('cpu'); torch.manual_seed(SEED); su=train(SU11Classifier(),tr,va,device)
104 torch.manual_seed(SEED); base=train(TanhRNNClassifier(),tr,va,device)
105 su['fallback_error']=str(e)
106 return {"device":str(device),"su11":su,"tanh_rnn":base}
107
108def main():
109 out={"math":verify_math(),"training":train_compare()}
110 with open('results.json','w') as f: json.dump(out,f,indent=2)
111 print(json.dumps(out,indent=2))
112if __name__=='__main__': main()