Lie-Poisson Hamiltonian latent block / experiment.py

Failed on benchmark

Raw ⬇ ZIP
  1import json
  2from pathlib import Path
  3import numpy as np
  4from scipy.optimize import root
  5
  6SEED = 1255
  7np.random.seed(SEED)
  8I = np.array([1.0, 2.0, 4.0])
  9Iinv = 1.0 / I
 10# Levi-Civita symbols: eps[0,1,2] = +1.
 11E = np.zeros((3, 3, 3))
 12E[0, 1, 2] = E[1, 2, 0] = E[2, 0, 1] = 1.0
 13E[0, 2, 1] = E[2, 1, 0] = E[1, 0, 2] = -1.0
 14
 15
 16def H(z):
 17    return 0.5 * np.sum(Iinv * z[3:] ** 2)
 18
 19
 20def grad_H(z):
 21    g = np.zeros(6)
 22    g[3:] = Iinv * z[3:]
 23    return g
 24
 25
 26def poisson_J(z):
 27    p = z[3:]
 28    J = np.zeros((6, 6))
 29    J[:3, 3:] = np.eye(3)
 30    J[3:, :3] = -np.eye(3)
 31    J[3:, 3:] = np.einsum('a,aij->ij', p, E)
 32    return J
 33
 34
 35def vector_field(z, use_coadjoint=True):
 36    p = z[3:]
 37    v = Iinv * p
 38    # p_a c^a_ij dH/dp_j; this is the explicit Lie-Poisson term.
 39    coad = np.einsum('a,aij,j->i', p, E, v) if use_coadjoint else np.zeros(3)
 40    return np.r_[v, coad]
 41
 42
 43def euler_step(z, dt, use_coadjoint=True):
 44    return z + dt * vector_field(z, use_coadjoint)
 45
 46
 47def midpoint_step(z, dt, use_coadjoint=True):
 48    def residual(zn):
 49        zm = 0.5 * (z + zn)
 50        return zn - z - dt * vector_field(zm, use_coadjoint)
 51    sol = root(residual, z + dt * vector_field(z, use_coadjoint), method='hybr')
 52    if not sol.success or np.linalg.norm(residual(sol.x), ord=np.inf) > 1e-9:
 53        raise RuntimeError('implicit midpoint failed: ' + sol.message)
 54    return sol.x
 55
 56
 57def rollout(z0, dt, steps, method='midpoint', use_coadjoint=True):
 58    z = z0.copy()
 59    out = [z.copy()]
 60    for _ in range(steps):
 61        if method == 'euler':
 62            z = euler_step(z, dt, use_coadjoint)
 63        else:
 64            z = midpoint_step(z, dt, use_coadjoint)
 65        out.append(z.copy())
 66    return np.asarray(out)
 67
 68
 69def loglog_slope(xs, ys):
 70    return float(np.polyfit(np.log(xs), np.log(np.maximum(ys, 1e-30)), 1)[0])
 71
 72
 73def main():
 74    z0 = np.r_[np.array([0.2, -0.1, 0.3]), np.array([0.7, 1.1, 1.5])]
 75    # Core algebraic check: J is antisymmetric and hence grad H^T J grad H = 0.
 76    zcheck = np.r_[np.array([0.4, -0.2, 0.1]), np.array([0.8, -1.3, 1.7])]
 77    J = poisson_J(zcheck)
 78    g = grad_H(zcheck)
 79    antisym = float(np.max(np.abs(J + J.T)))
 80    energy_derivative = float(abs(g @ J @ g))
 81
 82    # Prediction 1: coadjoint effect vanishes exactly at p=0.
 83    zzero = np.r_[z0[:3], np.zeros(3)]
 84    a0 = rollout(zzero, 0.05, 100, 'midpoint', True)
 85    b0 = rollout(zzero, 0.05, 100, 'midpoint', False)
 86    zero_momentum_difference = float(np.max(np.abs(a0 - b0)))
 87
 88    # Prediction 2: at fixed time, the coadjoint-vs-ablated difference is O(alpha^2)
 89    # for small momentum scale alpha, because the extra vector field is quadratic in p.
 90    alphas = np.array([0.125, 0.25, 0.5, 1.0])
 91    alpha_errors = []
 92    for alpha in alphas:
 93        za = np.r_[z0[:3], alpha * z0[3:]]
 94        full = rollout(za, 0.01, 100, 'midpoint', True)
 95        ablated = rollout(za, 0.01, 100, 'midpoint', False)
 96        alpha_errors.append(np.linalg.norm(full[-1, 3:] - ablated[-1, 3:]))
 97    alpha_errors = np.asarray(alpha_errors)
 98    alpha_slope = loglog_slope(alphas, alpha_errors)
 99
100    # Prediction 3: midpoint energy error scales quadratically with dt over fixed T,
101    # while Euler has first-order drift.
102    T = 10.0
103    dts = np.array([0.1, 0.05, 0.025, 0.0125])
104    midpoint_energy_errors, euler_energy_errors = [], []
105    for dt in dts:
106        n = int(round(T / dt))
107        zm = rollout(z0, dt, n, 'midpoint', True)[-1]
108        ze = rollout(z0, dt, n, 'euler', True)[-1]
109        midpoint_energy_errors.append(abs(H(zm) - H(z0)))
110        euler_energy_errors.append(abs(H(ze) - H(z0)))
111    midpoint_energy_errors = np.asarray(midpoint_energy_errors)
112    euler_energy_errors = np.asarray(euler_energy_errors)
113    midpoint_slope = loglog_slope(dts, midpoint_energy_errors)
114    euler_slope = loglog_slope(dts, euler_energy_errors)
115
116    # Secondary comparison: long-rollout energy stability and ablation phase error.
117    long_steps, long_dt = 1000, 0.02
118    mid = rollout(z0, long_dt, long_steps, 'midpoint', True)
119    eu = rollout(z0, long_dt, long_steps, 'euler', True)
120    abl = rollout(z0, long_dt, long_steps, 'midpoint', False)
121    long_mid_energy = float(np.max(np.abs(np.array([H(x) for x in mid]) - H(z0))))
122    long_euler_energy = float(np.max(np.abs(np.array([H(x) for x in eu]) - H(z0))))
123    ablation_terminal_error = float(np.linalg.norm(mid[-1, 3:] - abl[-1, 3:]))
124
125    result = {
126        'seed': SEED,
127        'inertia': I.tolist(),
128        'math_check': {
129            'max_J_plus_JT': antisym,
130            'abs_gradH_J_gradH': energy_derivative,
131            'tolerance': 1e-12,
132        },
133        'predictions': {
134            'zero_momentum_effect_predicted': 0.0,
135            'zero_momentum_observed_max_difference': zero_momentum_difference,
136            'momentum_scaling_predicted_exponent': 2.0,
137            'momentum_scaling_observed_exponent': alpha_slope,
138            'momentum_scales': alphas.tolist(),
139            'momentum_effects': alpha_errors.tolist(),
140            'midpoint_energy_dt_predicted_exponent': 2.0,
141            'midpoint_energy_dt_observed_exponent': midpoint_slope,
142            'euler_energy_dt_predicted_exponent': 1.0,
143            'euler_energy_dt_observed_exponent': euler_slope,
144            'dts': dts.tolist(),
145            'midpoint_energy_errors': midpoint_energy_errors.tolist(),
146            'euler_energy_errors': euler_energy_errors.tolist(),
147        },
148        'secondary_comparison': {
149            'midpoint_max_energy_error_T10': long_mid_energy,
150            'euler_max_energy_error_T10': long_euler_energy,
151            'midpoint_full_vs_no_coadjoint_terminal_momentum_error': ablation_terminal_error,
152        },
153    }
154    Path('results.json').write_text(json.dumps(result, indent=2))
155    print(json.dumps(result, indent=2))
156
157
158if __name__ == '__main__':
159    main()