Skew-Midpoint Neural Dynamics / experiment.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
 1import json, math
 2import torch
 3from skew_midpoint import make_metric, skew, midpoint_step
 4
 5torch.set_default_dtype(torch.float64)
 6torch.manual_seed(7)
 7
 8def midpoint_M(e, dt, M, J):
 9    n=e.numel(); I=torch.eye(n)
10    return torch.linalg.solve(M/dt-J/2, (M/dt+J/2)@e)
11
12def euler(e, dt, M, J):
13    return e + dt*torch.linalg.solve(M, J@e)
14
15def energy(x,M): return 0.5*x@M@x
16
17# Mechanism verification 1: exact quadratic energy conservation for random SPD M, skew J.
18L=torch.tensor([[1.3,0.,0.],[.25,.9,0.],[-.2,.15,1.1]])
19M=make_metric(L, 1e-3)
20A=torch.randn(3,3); J=skew(A)
21x=torch.tensor([.7,-1.1,.4])
22H0=energy(x,M); max_drift=0.;
23for _ in range(1000):
24    x=midpoint_M(x,.17,M,J)
25    max_drift=max(max_drift,abs((energy(x,M)-H0).item()))
26# SPD/skew numerical properties
27spd_min=torch.linalg.eigvalsh(M).min().item(); skew_err=torch.linalg.norm(J+J.T).item()
28
29# Mechanism verification 2: energy growth law for explicit Euler on M=I, J=omega rotation.
30def rot(omega): return torch.tensor([[0.,-omega],[omega,0.]])
31energy_rows=[]
32for omega in [.25, .5, 1., 2.]:
33  for dt in [.05,.1,.2]:
34    N=100
35    Jw=rot(omega); z=torch.tensor([1.,0.]); H=energy(z,torch.eye(2)).item()
36    for _ in range(N): z=euler(z,dt,torch.eye(2),Jw)
37    observed=energy(z,torch.eye(2)).item()/H
38    predicted=(1+(omega*dt)**2)**N
39    energy_rows.append({'omega':omega,'dt':dt,'observed_ratio':observed,'predicted_ratio':predicted,'relative_error':abs(observed-predicted)/predicted})
40
41# Mechanism verification 3: midpoint phase transition/scaling.
42phase_rows=[]
43for omega in [.5,1.,2.,4.]:
44  for dt in [.05,.1,.2]:
45    z=torch.tensor([1.,0.]); Jw=rot(omega)
46    z=midpoint_M(z,dt,torch.eye(2),Jw)
47    # one-step angle, matching positive rotation convention
48    observed=math.atan2(z[1].item(),z[0].item())
49    predicted=2*math.atan(omega*dt/2)
50    exact=omega*dt
51    phase_rows.append({'omega':omega,'dt':dt,'observed_angle':observed,'predicted_angle':predicted,'exact_angle':exact,'midpoint_phase_error':abs(observed-exact)})
52
53# Same-step rollout benchmark against exact oscillator trajectory.
54bench=[]
55for dt in [.05,.1,.2]:
56  omega=2.; N=200; Jw=rot(omega); M2=torch.eye(2)
57  a=torch.tensor([1.,0.]); b=a.clone(); exact=torch.tensor([1.,0.]);
58  for k in range(N):
59    a=midpoint_M(a,dt,M2,Jw); b=euler(b,dt,M2,Jw)
60    t=(k+1)*dt; exact=torch.tensor([math.cos(omega*t),math.sin(omega*t)])
61  bench.append({'dt':dt,'horizon':N*dt,'midpoint_state_error':torch.linalg.norm(a-exact).item(),'euler_state_error':torch.linalg.norm(b-exact).item(),'midpoint_energy':energy(a,M2).item(),'euler_energy':energy(b,M2).item()})
62
63out={'spd_min_eigenvalue':spd_min,'skew_frobenius_error':skew_err,'random_metric_midpoint_max_energy_drift':max_drift,'energy_law_sweep':energy_rows,'phase_sweep':phase_rows,'rollout_benchmark':bench}
64print(json.dumps(out,indent=2))