import json, math import torch from skew_midpoint import make_metric, skew, midpoint_step torch.set_default_dtype(torch.float64) torch.manual_seed(7) def midpoint_M(e, dt, M, J): n=e.numel(); I=torch.eye(n) return torch.linalg.solve(M/dt-J/2, (M/dt+J/2)@e) def euler(e, dt, M, J): return e + dt*torch.linalg.solve(M, J@e) def energy(x,M): return 0.5*x@M@x # Mechanism verification 1: exact quadratic energy conservation for random SPD M, skew J. L=torch.tensor([[1.3,0.,0.],[.25,.9,0.],[-.2,.15,1.1]]) M=make_metric(L, 1e-3) A=torch.randn(3,3); J=skew(A) x=torch.tensor([.7,-1.1,.4]) H0=energy(x,M); max_drift=0.; for _ in range(1000): x=midpoint_M(x,.17,M,J) max_drift=max(max_drift,abs((energy(x,M)-H0).item())) # SPD/skew numerical properties spd_min=torch.linalg.eigvalsh(M).min().item(); skew_err=torch.linalg.norm(J+J.T).item() # Mechanism verification 2: energy growth law for explicit Euler on M=I, J=omega rotation. def rot(omega): return torch.tensor([[0.,-omega],[omega,0.]]) energy_rows=[] for omega in [.25, .5, 1., 2.]: for dt in [.05,.1,.2]: N=100 Jw=rot(omega); z=torch.tensor([1.,0.]); H=energy(z,torch.eye(2)).item() for _ in range(N): z=euler(z,dt,torch.eye(2),Jw) observed=energy(z,torch.eye(2)).item()/H predicted=(1+(omega*dt)**2)**N energy_rows.append({'omega':omega,'dt':dt,'observed_ratio':observed,'predicted_ratio':predicted,'relative_error':abs(observed-predicted)/predicted}) # Mechanism verification 3: midpoint phase transition/scaling. phase_rows=[] for omega in [.5,1.,2.,4.]: for dt in [.05,.1,.2]: z=torch.tensor([1.,0.]); Jw=rot(omega) z=midpoint_M(z,dt,torch.eye(2),Jw) # one-step angle, matching positive rotation convention observed=math.atan2(z[1].item(),z[0].item()) predicted=2*math.atan(omega*dt/2) exact=omega*dt phase_rows.append({'omega':omega,'dt':dt,'observed_angle':observed,'predicted_angle':predicted,'exact_angle':exact,'midpoint_phase_error':abs(observed-exact)}) # Same-step rollout benchmark against exact oscillator trajectory. bench=[] for dt in [.05,.1,.2]: omega=2.; N=200; Jw=rot(omega); M2=torch.eye(2) a=torch.tensor([1.,0.]); b=a.clone(); exact=torch.tensor([1.,0.]); for k in range(N): a=midpoint_M(a,dt,M2,Jw); b=euler(b,dt,M2,Jw) t=(k+1)*dt; exact=torch.tensor([math.cos(omega*t),math.sin(omega*t)]) 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()}) out={'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} print(json.dumps(out,indent=2))