Gramian-Regularized Latent State Models / gramian_experiment.py
Failed on benchmark
1import json, random
2from pathlib import Path
3import numpy as np
4import torch
5
6SEED = 137
7np.random.seed(SEED); random.seed(SEED); torch.manual_seed(SEED)
8
9def discrete_gramians(A, B, C, T):
10 n = A.shape[0]
11 Wo = np.zeros((n, n)); Phi = np.eye(n)
12 for _ in range(T):
13 Wo += Phi.T @ C.T @ C @ Phi
14 Phi = A @ Phi
15 Wr = np.zeros((n, n))
16 for k in range(T):
17 P = np.linalg.matrix_power(A, T - 1 - k)
18 Wr += P @ B @ B.T @ P.T
19 return (Wo + Wo.T) / 2, (Wr + Wr.T) / 2
20
21def norm_min(W):
22 tr = np.trace(W)
23 return float(np.min(np.linalg.eigvalsh(W)) / (tr / W.shape[0])) if tr > 1e-12 else 0.0
24
25def mechanism_checks():
26 # With A=aI and C=diag(1,r), Wo is proportional to diag(1,r^2).
27 # Therefore normalized lambda_min is exactly 2r^2/(1+r^2), a quantitative
28 # parameter-scaling prediction. The same construction applies to Wr.
29 A = np.diag([.82, .82]); T = 30
30 ratios = [0.0, .05, .1, .2, .4, .7, 1.0]
31 obs = []
32 reach = []
33 for r in ratios:
34 C = np.diag([1., r]); Wo, _ = discrete_gramians(A, np.eye(2), C, T)
35 _, Wr = discrete_gramians(A, np.diag([1., r]), C, T)
36 pred = 2*r*r/(1+r*r) if r else 0.0
37 obs.append({'coupling':r, 'normalized_lambda_min':norm_min(Wo), 'predicted':pred})
38 reach.append({'coupling':r, 'normalized_lambda_min':norm_min(Wr), 'predicted':pred})
39 # Direct quadrature verifies the continuous-time formula for a scalar mode.
40 a=.18; c=.7; b=.55; dt=.002; horizon=2.
41 times=np.arange(0, horizon, dt)
42 wo_num=float(np.sum(dt*(c*np.exp(-a*times))**2))
43 wr_num=float(np.sum(dt*(b*np.exp(-a*(horizon-times)))**2))
44 wo_exact=c*c*(1-np.exp(-2*a*horizon))/(2*a)
45 wr_exact=b*b*(1-np.exp(-2*a*horizon))/(2*a)
46 # Measurement noise prediction: least-squares covariance is sigma^2 Wo^-1.
47 noise=.03; samples=12000; est=[]; rng=np.random.default_rng(SEED)
48 for r in [.08,.12,.2,.3,.5]:
49 C=np.diag([1., r]); Wo,_=discrete_gramians(A,np.zeros((2,2)),C,T)
50 H=np.vstack([C @ np.linalg.matrix_power(A,k) for k in range(T)])
51 z=rng.standard_normal((2,samples)); ys=H@z+noise*rng.standard_normal((2*T,samples))
52 zhat=np.linalg.solve(H.T@H, H.T@ys).T
53 mse=float(np.mean((zhat-z.T)**2)); lam=float(np.min(np.linalg.eigvalsh(Wo)))
54 est.append({'coupling':r, 'lambda_min':lam, 'mse':mse, 'predicted_sigma2_over_lambda':noise*noise/lam})
55 slope=float(np.polyfit(np.log([x['lambda_min'] for x in est]), np.log([x['mse'] for x in est]), 1)[0])
56 return {'observation_scaling':obs, 'reachability_scaling':reach,
57 'scaling_max_abs_error_observation':max(abs(x['normalized_lambda_min']-x['predicted']) for x in obs),
58 'scaling_max_abs_error_reachability':max(abs(x['normalized_lambda_min']-x['predicted']) for x in reach),
59 'integral_check':{'wo_numeric':wo_num,'wo_exact':wo_exact,'wr_numeric':wr_num,'wr_exact':wr_exact},
60 'noise_inverse_scaling':est, 'log_mse_vs_log_lambda_slope':slope}
61
62class LinearSSM(torch.nn.Module):
63 def __init__(self):
64 super().__init__()
65 self.A=torch.nn.Parameter(torch.tensor([[.75,.02],[.01,.65]]))
66 self.B=torch.nn.Parameter(torch.randn(2,1)*.1)
67 self.C=torch.nn.Parameter(torch.randn(1,2)*.1)
68 def forward(self,u):
69 z=torch.zeros(u.shape[0],2,device=u.device); ys=[]
70 for k in range(u.shape[1]):
71 z=z@self.A.T+u[:,k,0:1]*self.B.T; ys.append(z@self.C.T)
72 return torch.stack(ys,1)
73 def grams(self,T):
74 phi=torch.eye(2,device=self.A.device); wo=torch.zeros((2,2),device=self.A.device)
75 for _ in range(T): wo=wo+phi.T@[email protected]@phi; phi=self.A@phi
76 wr=torch.zeros((2,2),device=self.A.device)
77 for k in range(T):
78 p=torch.linalg.matrix_power(self.A,T-1-k); wr=wr+p@[email protected]@p.T
79 return (wo+wo.T)/2,(wr+wr.T)/2
80
81def train(reg, device):
82 torch.manual_seed(SEED + int(reg*1000)); n=256; T=20
83 A0=torch.tensor([[.90,0.],[0.,.72]]); B0=torch.tensor([[.8],[.35]]); C0=torch.tensor([[1.,.12]])
84 u=torch.randn(n,T,1); z=torch.zeros(n,2); ys=[]
85 for k in range(T): z=z@A0.T+u[:,k,0:1]*B0.T; ys.append(z@C0.T)
86 y=torch.stack(ys,1).to(device); u=u.to(device); model=LinearSSM().to(device)
87 opt=torch.optim.Adam(model.parameters(),lr=.025)
88 for _ in range(500):
89 pred=model(u); task=((pred-y)**2).mean(); wo,wr=model.grams(T)
90 def penalty(w):
91 ev=torch.linalg.eigvalsh(w); ratio=ev[0]/(torch.trace(w)/2+1e-8)
92 return torch.relu(torch.tensor(.08,device=device)-ratio), ratio
93 po,ro=penalty(wo); pr,rr=penalty(wr); loss=task+reg*(po+pr)
94 opt.zero_grad(); loss.backward(); torch.nn.utils.clip_grad_norm_(model.parameters(),2.); opt.step()
95 return {'task_mse':float(task.detach().cpu()), 'normalized_observability':float(ro.detach().cpu()), 'normalized_reachability':float(rr.detach().cpu())}
96
97def main():
98 checks=mechanism_checks(); device='cuda' if torch.cuda.is_available() else 'cpu'
99 try:
100 baseline=train(0.,device); idea=train(.08,device)
101 except Exception:
102 device='cpu'; baseline=train(0.,device); idea=train(.08,device)
103 out={'seed':SEED,'device':device,'checks':checks,'training':{'baseline':baseline,'gramian_regularized':idea}}
104 Path('results.json').write_text(json.dumps(out,indent=2)); print(json.dumps(out,indent=2))
105if __name__=='__main__': main()