Skip to main content

ferromotion_learn/
hnn.rs

1//! **Hamiltonian Neural Network (HNN)** — physics injected into the **architecture**. Instead of predicting
2//! the dynamics directly, the network learns a single scalar function, the Hamiltonian `H_θ(q, p)` (the total
3//! energy), and the dynamics are *derived* from it by Hamilton's equations: `q̇ = ∂H/∂p`, `ṗ = −∂H/∂q`. A
4//! vector field obtained this way is symplectic by construction, so trajectories conserve the learned `H` —
5//! the network cannot leak or inject energy no matter how long you roll it out. This is the structural answer
6//! to the integrator-drift problem from the energy lesson: conservation is baked in, not hoped for.
7//!
8//! Training uses the network's derivatives w.r.t. its **inputs** `(q, p)` as the predicted velocities, matched
9//! to observed `(q̇, ṗ)`. Those input-derivatives are computed exactly by propagating a first-order two-input
10//! jet `(value, ∂/∂q, ∂/∂p)` through the network where each slot is a reverse-mode [`Var`], so one backward
11//! pass gives the exact parameter gradient of the matching loss. Verified on the ideal mass–spring
12//! (`H = ½p² + ½q²`): the learned vector field matches the truth, and a long rollout conserves energy where a
13//! direct black-box predictor of the same data drifts.
14
15use crate::autodiff::{Tape, Var};
16
17/// A Hamiltonian network `H_θ(q, p)` (two inputs, one scalar output).
18pub struct Hnn {
19    sizes: Vec<usize>,
20    params: Vec<f64>,
21    m: Vec<f64>,
22    v: Vec<f64>,
23    t: u64,
24}
25
26fn splitmix64(state: &mut u64) -> u64 {
27    *state = state.wrapping_add(0x9E3779B97F4A7C15);
28    let mut z = *state;
29    z = (z ^ (z >> 30)).wrapping_mul(0xBF58476D1CE4E5B9);
30    z = (z ^ (z >> 27)).wrapping_mul(0x94D049BB133111EB);
31    z ^ (z >> 31)
32}
33
34// A first-order jet in two inputs (q, p): (value, ∂/∂q, ∂/∂p), each a tape Var.
35type Jet2<'t> = (Var<'t>, Var<'t>, Var<'t>);
36
37fn j_add<'t>(a: Jet2<'t>, b: Jet2<'t>) -> Jet2<'t> {
38    (a.0 + b.0, a.1 + b.1, a.2 + b.2)
39}
40fn j_scale<'t>(w: Var<'t>, a: Jet2<'t>) -> Jet2<'t> {
41    (w * a.0, w * a.1, w * a.2)
42}
43fn j_tanh<'t>(g: Jet2<'t>) -> Jet2<'t> {
44    let f0 = g.0.tanh();
45    let fp = f0 * f0 * (-1.0) + 1.0; // 1 − tanh²
46    (f0, fp * g.1, fp * g.2)
47}
48
49impl Hnn {
50    /// A Hamiltonian network with the given hidden width.
51    pub fn new(hidden: usize, seed: u64) -> Self {
52        let sizes = vec![2, hidden, hidden, 1];
53        let mut state = seed ^ 0x0F1E_2D3C_4B5A_6978;
54        let mut params = Vec::new();
55        for l in 0..sizes.len() - 1 {
56            let (ind, outd) = (sizes[l], sizes[l + 1]);
57            let r = (6.0 / (ind + outd) as f64).sqrt();
58            for _ in 0..ind * outd {
59                let u = (splitmix64(&mut state) as f64 / u64::MAX as f64) * 2.0 - 1.0;
60                params.push(u * r);
61            }
62            params.extend(std::iter::repeat_n(0.0, outd));
63        }
64        let n = params.len();
65        Hnn { sizes, params, m: vec![0.0; n], v: vec![0.0; n], t: 0 }
66    }
67
68    // Forward as a two-input jet; returns (H, ∂H/∂q, ∂H/∂p) as tape Vars.
69    fn forward_jet<'t>(&self, tape: &'t Tape, pv: &[Var<'t>], q: f64, p: f64, zero: Var<'t>, one: Var<'t>) -> Jet2<'t> {
70        let mut a: Vec<Jet2<'t>> = vec![(tape.constant(q), one, zero), (tape.constant(p), zero, one)];
71        let mut off = 0;
72        let layers = self.sizes.len() - 1;
73        for l in 0..layers {
74            let (ind, outd) = (self.sizes[l], self.sizes[l + 1]);
75            let mut z: Vec<Jet2<'t>> = Vec::with_capacity(outd);
76            for o in 0..outd {
77                let mut s: Jet2<'t> = (pv[off + ind * outd + o], zero, zero);
78                for (i, &ai) in a.iter().enumerate() {
79                    s = j_add(s, j_scale(pv[off + o * ind + i], ai));
80                }
81                z.push(if l + 1 < layers { j_tanh(s) } else { s });
82            }
83            off += ind * outd + outd;
84            a = z;
85        }
86        a[0]
87    }
88
89    // Plain-f64 two-input jet: (H, ∂H/∂q, ∂H/∂p) with no tape (for inference / rollout).
90    fn forward_deriv(&self, q: f64, p: f64) -> (f64, f64, f64) {
91        let mut a: Vec<(f64, f64, f64)> = vec![(q, 1.0, 0.0), (p, 0.0, 1.0)];
92        let mut off = 0;
93        let layers = self.sizes.len() - 1;
94        for l in 0..layers {
95            let (ind, outd) = (self.sizes[l], self.sizes[l + 1]);
96            let mut z = Vec::with_capacity(outd);
97            for o in 0..outd {
98                let (mut v0, mut v1, mut v2) = (self.params[off + ind * outd + o], 0.0, 0.0);
99                for (i, ai) in a.iter().enumerate() {
100                    let w = self.params[off + o * ind + i];
101                    v0 += w * ai.0;
102                    v1 += w * ai.1;
103                    v2 += w * ai.2;
104                }
105                if l + 1 < layers {
106                    let t = v0.tanh();
107                    let fp = 1.0 - t * t;
108                    z.push((t, fp * v1, fp * v2));
109                } else {
110                    z.push((v0, v1, v2));
111                }
112            }
113            off += ind * outd + outd;
114            a = z;
115        }
116        a[0]
117    }
118
119    /// The learned Hamiltonian (energy) at a phase point.
120    pub fn hamiltonian(&self, q: f64, p: f64) -> f64 {
121        self.forward_deriv(q, p).0
122    }
123
124    /// The learned dynamics `(q̇, ṗ) = (∂H/∂p, −∂H/∂q)` at a phase point.
125    pub fn field(&self, q: f64, p: f64) -> (f64, f64) {
126        let (_, dq, dp) = self.forward_deriv(q, p);
127        (dp, -dq)
128    }
129
130    /// One RK4 rollout step of the learned dynamics.
131    pub fn step_rk4(&self, q: f64, p: f64, dt: f64) -> (f64, f64) {
132        let f = |q: f64, p: f64| self.field(q, p);
133        let (k1q, k1p) = f(q, p);
134        let (k2q, k2p) = f(q + 0.5 * dt * k1q, p + 0.5 * dt * k1p);
135        let (k3q, k3p) = f(q + 0.5 * dt * k2q, p + 0.5 * dt * k2p);
136        let (k4q, k4p) = f(q + dt * k3q, p + dt * k3p);
137        (q + dt / 6.0 * (k1q + 2.0 * k2q + 2.0 * k3q + k4q), p + dt / 6.0 * (k1p + 2.0 * k2p + 2.0 * k3p + k4p))
138    }
139
140    /// One Adam step matching the learned field `(∂H/∂p, −∂H/∂q)` to the observed `(q̇, ṗ)`. Returns the loss.
141    pub fn train_step(&mut self, q: &[f64], p: &[f64], qdot: &[f64], pdot: &[f64], lr: f64) -> f64 {
142        let tape = Tape::new();
143        let pv: Vec<Var> = self.params.iter().map(|&x| tape.var(x)).collect();
144        let zero = tape.constant(0.0);
145        let one = tape.constant(1.0);
146        let mut loss = tape.constant(0.0);
147        for i in 0..q.len() {
148            let (_h, hq, hp) = self.forward_jet(&tape, &pv, q[i], p[i], zero, one);
149            let r_q = hp - qdot[i]; //  q̇_pred − q̇ = ∂H/∂p − q̇
150            let r_p = (-hq) - pdot[i]; //  ṗ_pred − ṗ = −∂H/∂q − ṗ
151            loss = loss + r_q * r_q + r_p * r_p;
152        }
153        loss = loss * (1.0 / q.len() as f64);
154        let g = loss.backward();
155        let grad: Vec<f64> = pv.iter().map(|&x| g.wrt(x)).collect();
156
157        self.t += 1;
158        let (b1, b2, eps) = (0.9_f64, 0.999_f64, 1e-8);
159        let bc1 = 1.0 - b1.powi(self.t as i32);
160        let bc2 = 1.0 - b2.powi(self.t as i32);
161        for (i, &gi) in grad.iter().enumerate() {
162            self.m[i] = b1 * self.m[i] + (1.0 - b1) * gi;
163            self.v[i] = b2 * self.v[i] + (1.0 - b2) * gi * gi;
164            self.params[i] -= lr * (self.m[i] / bc1) / ((self.v[i] / bc2).sqrt() + eps);
165        }
166        loss.value()
167    }
168
169    /// Train for `epochs` steps; returns the final loss.
170    pub fn train(&mut self, q: &[f64], p: &[f64], qdot: &[f64], pdot: &[f64], epochs: usize, lr: f64) -> f64 {
171        let mut l = f64::INFINITY;
172        for _ in 0..epochs {
173            l = self.train_step(q, p, qdot, pdot, lr);
174        }
175        l
176    }
177}
178
179#[cfg(test)]
180mod tests {
181    use super::*;
182
183    // Ideal mass-spring H = ½p² + ½q²: true field q̇ = p, ṗ = −q; energy ½(q²+p²) is conserved.
184    fn mass_spring_data() -> (Vec<f64>, Vec<f64>, Vec<f64>, Vec<f64>) {
185        let (mut q, mut p, mut qd, mut pd) = (vec![], vec![], vec![], vec![]);
186        let mut s = 12345u64;
187        for _ in 0..80 {
188            let qi = (splitmix64(&mut s) as f64 / u64::MAX as f64) * 4.0 - 2.0;
189            let pi = (splitmix64(&mut s) as f64 / u64::MAX as f64) * 4.0 - 2.0;
190            q.push(qi);
191            p.push(pi);
192            qd.push(pi); // q̇ = p
193            pd.push(-qi); // ṗ = −q
194        }
195        (q, p, qd, pd)
196    }
197
198    #[test]
199    fn hnn_learns_the_field_and_conserves_energy() {
200        // Train once, then check both properties. THE HEADLINE: from (q,p)→(q̇,ṗ) samples the HNN recovers the
201        // true vector field via ∂H/∂p, −∂H/∂q. THE STRUCTURAL PAYOFF: because the field derives from a scalar
202        // H, rolling it out conserves the true energy over a long horizon — no drift.
203        let (q, p, qd, pd) = mass_spring_data();
204        let mut hnn = Hnn::new(16, 3);
205        hnn.train(&q, &p, &qd, &pd, 2500, 5e-3);
206
207        let mut err = 0.0;
208        let mut n = 0;
209        for &(tq, tp) in &[(1.0, 0.0), (0.0, 1.0), (1.0, 1.0), (-1.5, 0.5)] {
210            let (fq, fp) = hnn.field(tq, tp);
211            err += (fq - tp).abs() + (fp - (-tq)).abs(); // truth (p, −q)
212            n += 2;
213        }
214        assert!(err / (n as f64) < 0.15, "learned field should match (p,−q): mean abs err {}", err / (n as f64));
215
216        let (mut qq, mut pp) = (1.5, 0.0);
217        let e0 = 0.5 * (qq * qq + pp * pp);
218        let mut max_dev: f64 = 0.0;
219        for _ in 0..2000 {
220            let (nq, np) = hnn.step_rk4(qq, pp, 0.02);
221            qq = nq;
222            pp = np;
223            max_dev = max_dev.max((0.5 * (qq * qq + pp * pp) - e0).abs());
224        }
225        assert!(max_dev < 0.1, "HNN rollout should conserve energy: max deviation {max_dev}");
226    }
227}