Carleman-Lifted Polynomial State Space / carleman_experiment.py
Mechanism failed
1import json, math, random
2from pathlib import Path
3import numpy as np
4from scipy.integrate import solve_ivp
5
6SEED = 186
7np.random.seed(SEED); random.seed(SEED)
8
9
10def apply_level(z, j, d, F0, F1, F2, include_quad=True):
11 """Derivative of tensor level j, using z[j-1],z[j],z[j+1]."""
12 out = np.zeros((d,) * j)
13 # z is a list indexed by tensor level; z[0] is scalar
14 for oi in np.ndindex(out.shape):
15 val = 0.0
16 for r in range(j):
17 # linear term
18 for a in range(d):
19 ii = list(oi); ii[r] = a
20 val += F1[oi[r], a] * z[j][tuple(ii)]
21 # forcing term (for j=1, z[0]=1)
22 if j == 1:
23 val += F0[oi[r]]
24 else:
25 ii = tuple(oi[k] for k in range(j) if k != r)
26 val += F0[oi[r]] * z[j-1][ii]
27 # quadratic term: split the source factor into two indices
28 if include_quad:
29 for a in range(d):
30 for b in range(d):
31 ii = list(oi); ii[r:r+1] = [a, b]
32 val += F2[oi[r], a*d+b] * z[j+1][tuple(ii)]
33 out[oi] = val
34 return out
35
36
37def unpack(y, d, N):
38 levels=[np.array(1.0)]
39 p=0
40 for j in range(1,N+1):
41 n=d**j; levels.append(y[p:p+n].reshape((d,)*j)); p += n
42 return levels
43
44
45def pack(levels): return np.concatenate([np.asarray(x).ravel() for x in levels[1:]])
46def pack_derivs(levels): return np.concatenate([np.asarray(x).ravel() for x in levels])
47
48
49def rhs_truncated(t, y, d, N, F0, F1, F2):
50 z=unpack(y,d,N); dz=[]
51 for j in range(1,N+1):
52 # append a harmless zero high level: omitted interaction is precisely excluded
53 zz=z + [np.zeros((d,)*(N+1))]
54 dz.append(apply_level(zz,j,d,F0,F1,F2,include_quad=(j < N)))
55 return pack_derivs(dz)
56
57
58def rhs_full(t, y, d, Nfull, F0, F1, F2):
59 z=unpack(y,d,Nfull); dz=[]
60 for j in range(1,Nfull+1):
61 zz=z + [np.zeros((d,)*(Nfull+1))]
62 dz.append(apply_level(zz,j,d,F0,F1,F2,include_quad=(j < Nfull)))
63 return pack_derivs(dz)
64
65
66def verify_math():
67 d=2; N=2; full=3
68 F1=np.array([[-0.45,0.12],[-0.08,-0.30]])
69 F2=np.array([[0.10,-0.035,0.02,0.015],[-0.04,0.025,-0.02,0.03]])
70 F0=np.array([0.08,-0.04]); u0=np.array([0.35,-0.25])
71 init=[np.array(1.),u0,np.einsum('i,j->ij',u0,u0),np.einsum('i,j,k->ijk',u0,u0,u0)]
72 ts=np.linspace(0,4,401)
73 sol=solve_ivp(lambda t,y: rhs_full(t,y,d,full,F0,F1,F2),[0,4],pack(init),t_eval=ts,rtol=1e-10,atol=1e-12)
74 tr=solve_ivp(lambda t,y: rhs_truncated(t,y,d,N,F0,F1,F2),[0,4],pack(init[:N+1]),t_eval=ts,rtol=1e-10,atol=1e-12)
75 exact=sol.y.T; approx=tr.y.T
76 errors=[]; defects=[]; residuals=[]
77 dt=ts[1]-ts[0]
78 for k in range(len(ts)):
79 zfull=unpack(exact[k],d,full); zn=unpack(approx[k],d,N)
80 errors.append(np.linalg.norm(exact[k][:d+d*d]-approx[k]))
81 defects.append(np.linalg.norm(apply_level(zfull,N,d,F0,F1,F2,include_quad=True)-apply_level(zfull,N,d,F0,F1,F2,include_quad=False)))
82 if 1 <= k < len(ts)-1:
83 eta=(exact[k+1][:d+d*d]-exact[k-1][:d+d*d])/(2*dt)
84 Aeta=rhs_truncated(ts[k], exact[k][:d+d*d],d,N,F0,F1,F2)-rhs_truncated(ts[k],approx[k],d,N,F0,F1,F2)
85 residuals.append(np.linalg.norm(eta-Aeta-defects[k]*np.r_[np.zeros(d), np.ones(d*d)]*0))
86 # residual above cannot embed the defect into all coordinates; compute exact vector defect
87 residuals=[]
88 for k in range(len(ts)):
89 zfull=unpack(exact[k],d,full); zn=unpack(approx[k],d,N)
90 full_proj=rhs_full(ts[k],exact[k],d,full,F0,F1,F2)[:d+d*d]
91 trunc_on_exact=rhs_truncated(ts[k],exact[k][:d+d*d],d,N,F0,F1,F2)
92 trunc_on_approx=rhs_truncated(ts[k],approx[k],d,N,F0,F1,F2)
93 high=apply_level(zfull,N,d,F0,F1,F2,include_quad=True)-apply_level(zfull,N,d,F0,F1,F2,include_quad=False)
94 defect=np.r_[np.zeros(d),high.ravel()]
95 residuals.append(np.linalg.norm((full_proj-trunc_on_approx)-((trunc_on_exact-trunc_on_approx)+defect)))
96 return dict(max_lift_error=float(max(errors)), max_defect=float(max(defects)),
97 mean_defect=float(np.mean(defects)), residual_rms=float(np.sqrt(np.mean(np.square(residuals)))),
98 final_lift_error=float(errors[-1]))
99
100# Small practical check: learn a forced quadratic sequence, comparing GRU and a lifted Euler cell.
101import torch
102import torch.nn as nn
103
104def make_data(n, T, d):
105 rng=np.random.RandomState(SEED+3)
106 xs=rng.randn(n,T,1).astype('float32')
107 ys=np.zeros((n,T,d),dtype='float32'); ys[:,0]=rng.randn(n,d)*.2
108 F1=np.array([[.82,.10],[-.08,.76]],dtype='float32')[:d,:d]
109 F2=np.array([[.16,-.05,.04,.02],[-.07,.05,-.03,.04]],dtype='float32')[:d,:d*d]
110 for t in range(T-1):
111 u=ys[:,t]; q=np.einsum('bi,bj->bij',u,u).reshape(n,d*d)
112 ys[:,t+1]=u@F1.T + q@F2.T + xs[:,t,:,None].squeeze(1)*np.array([.12,-.09],dtype='float32')[:d]
113 return torch.tensor(xs),torch.tensor(ys)
114
115class GRUModel(nn.Module):
116 def __init__(self,d=2,h=8):
117 super().__init__(); self.r=nn.GRU(1,h,batch_first=True); self.o=nn.Linear(h,d)
118 def forward(self,x): return self.o(self.r(x)[0])
119
120class LiftModel(nn.Module):
121 """N=2 Carleman cell; z2 is maintained explicitly and contractions are factor-wise."""
122 def __init__(self,d=2,dt=0.02):
123 super().__init__(); self.d=d; self.dt=dt
124 self.f0=nn.Linear(1,d); self.f1=nn.Linear(1,d*d); self.f2=nn.Linear(1,d*d*d)
125 for layer in (self.f0,self.f1,self.f2):
126 nn.init.normal_(layer.weight, 0.0, 0.08); nn.init.zeros_(layer.bias)
127 def forward(self,x):
128 B,T,_=x.shape; d=self.d; h=torch.zeros(B,d,device=x.device)
129 z2=h[:,:,None]*h[:,None,:]; outs=[]; defects=[]
130 for t in range(T):
131 f0=self.f0(x[:,t]); f1=self.f1(x[:,t]).view(B,d,d); f2=self.f2(x[:,t]).view(B,d,d,d)
132 z3=h[:,:,None,None]*z2[:,None,:,:]
133 # omitted A_3^2 z3, summed over the two tensor-factor positions
134 quad2=torch.einsum('b i a c,b a c j->b i j',f2,z3)
135 quad2=quad2+torch.einsum('b j a c,b i a c->b i j',f2,z3)
136 # equivalent first-level quadratic contribution F2 z2
137 dh=torch.einsum('b i a,b a->b i',f1,h)+f0+torch.einsum('b i a c,b a c->b i',f2,z2)
138 dz2=quad2+torch.einsum('b i a,b a j->b i j',f1,z2)+torch.einsum('b j a,b i a->b i j',f1,z2)
139 dz2=dz2+f0[:,:,None]*h[:,None,:]+h[:,:,None]*f0[:,None,:]
140 defects.append(torch.linalg.vector_norm(quad2.reshape(B,-1),dim=1))
141 h=h+self.dt*dh; z2=z2+self.dt*dz2; outs.append(h)
142 return torch.stack(outs,1),torch.stack(defects,1)
143
144def train(model,x,y,steps=350):
145 requested = "cuda" if torch.cuda.is_available() else "cpu"
146 def run(device):
147 model.to(device); xx=x.to(device); yy=y.to(device)
148 opt=torch.optim.Adam(model.parameters(),lr=1e-3)
149 model.train()
150 for _ in range(steps):
151 opt.zero_grad()
152 out=model(xx)[0] if isinstance(model,LiftModel) else model(xx)
153 loss=((out-yy)**2).mean()
154 if not torch.isfinite(loss): raise FloatingPointError('non-finite training loss')
155 loss.backward()
156 torch.nn.utils.clip_grad_norm_(model.parameters(),5); opt.step()
157 model.eval()
158 with torch.no_grad():
159 out=model(xx)[0] if isinstance(model,LiftModel) else model(xx)
160 return float(((out-yy)**2).mean().cpu())
161 try:
162 return run(requested)
163 except Exception as exc:
164 if requested != "cuda": raise
165 print("CUDA failed; retrying on CPU:", repr(exc))
166 return run("cpu")
167
168def main():
169 check=verify_math(); torch.manual_seed(SEED)
170 x,y=make_data(96,32,2)
171 b=train(GRUModel(),x,y); torch.manual_seed(SEED); i=train(LiftModel(),x,y)
172 result={'math_check':check,'sequence_train_mse':{'baseline_gru':b,'carleman_lift':i}}
173 Path('results.json').write_text(json.dumps(result,indent=2)); print(json.dumps(result,indent=2))
174if __name__=='__main__': main()