Gramian-Regularized Latent State Models / gramian_experiment.py

Failed on benchmark

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