Structure-preserving SU(1,1) recurrent scan / su11_experiment.py

✓✓ Beats tuned baseline

Raw ⬇ ZIP
  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()