use crate::autodiff::{Tape, Var};
pub struct Pinn {
sizes: Vec<usize>,
params: Vec<f64>,
m: Vec<f64>,
v: Vec<f64>,
t: u64,
omega: f64,
colloc: Vec<f64>,
ic_weight: f64,
}
fn splitmix64(state: &mut u64) -> u64 {
*state = state.wrapping_add(0x9E3779B97F4A7C15);
let mut z = *state;
z = (z ^ (z >> 30)).wrapping_mul(0xBF58476D1CE4E5B9);
z = (z ^ (z >> 27)).wrapping_mul(0x94D049BB133111EB);
z ^ (z >> 31)
}
type Jet<'t> = (Var<'t>, Var<'t>, Var<'t>);
fn jet_add<'t>(a: Jet<'t>, b: Jet<'t>) -> Jet<'t> {
(a.0 + b.0, a.1 + b.1, a.2 + b.2)
}
fn jet_scale<'t>(w: Var<'t>, a: Jet<'t>) -> Jet<'t> {
(w * a.0, w * a.1, w * a.2)
}
fn jet_tanh<'t>(g: Jet<'t>) -> Jet<'t> {
let f0 = g.0.tanh();
let fp = f0 * f0 * (-1.0) + 1.0; let fpp = (f0 * fp) * (-2.0); (f0, fp * g.1, fpp * (g.1 * g.1) + fp * g.2)
}
impl Pinn {
pub fn new(hidden: usize, omega: f64, t_max: f64, n_colloc: usize, seed: u64) -> Self {
let sizes = vec![1, hidden, hidden, 1];
let mut state = seed ^ 0x1234_5678_9ABC_DEF0;
let mut params = Vec::new();
for l in 0..sizes.len() - 1 {
let (ind, outd) = (sizes[l], sizes[l + 1]);
let r = (6.0 / (ind + outd) as f64).sqrt();
for _ in 0..ind * outd {
let u = (splitmix64(&mut state) as f64 / u64::MAX as f64) * 2.0 - 1.0;
params.push(u * r);
}
params.extend(std::iter::repeat_n(0.0, outd));
}
let n = params.len();
let colloc: Vec<f64> = (0..n_colloc).map(|i| t_max * (i as f64 + 0.5) / n_colloc as f64).collect();
Pinn { sizes, params, m: vec![0.0; n], v: vec![0.0; n], t: 0, omega, colloc, ic_weight: 12.0 }
}
pub fn exact(&self, t: f64) -> f64 {
(self.omega * t).cos()
}
pub fn forward(&self, t: f64) -> f64 {
let mut a = vec![t];
let mut off = 0;
let layers = self.sizes.len() - 1;
for l in 0..layers {
let (ind, outd) = (self.sizes[l], self.sizes[l + 1]);
let mut z = vec![0.0; outd];
for (o, zo) in z.iter_mut().enumerate() {
let mut s = self.params[off + ind * outd + o];
for (i, &ai) in a.iter().enumerate() {
s += self.params[off + o * ind + i] * ai;
}
*zo = if l + 1 < layers { s.tanh() } else { s };
}
off += ind * outd + outd;
a = z;
}
a[0]
}
fn forward_jet<'t>(&self, tape: &'t Tape, pv: &[Var<'t>], t: f64, zero: Var<'t>, one: Var<'t>) -> Jet<'t> {
let mut a: Vec<Jet<'t>> = vec![(tape.constant(t), one, zero)];
let mut off = 0;
let layers = self.sizes.len() - 1;
for l in 0..layers {
let (ind, outd) = (self.sizes[l], self.sizes[l + 1]);
let mut z: Vec<Jet<'t>> = Vec::with_capacity(outd);
for o in 0..outd {
let mut s: Jet<'t> = (pv[off + ind * outd + o], zero, zero); for (i, &ai) in a.iter().enumerate() {
s = jet_add(s, jet_scale(pv[off + o * ind + i], ai));
}
z.push(if l + 1 < layers { jet_tanh(s) } else { s });
}
off += ind * outd + outd;
a = z;
}
a[0]
}
fn loss_and_grad(&self) -> (f64, Vec<f64>) {
let tape = Tape::new();
let pv: Vec<Var> = self.params.iter().map(|&p| tape.var(p)).collect();
let zero = tape.constant(0.0);
let one = tape.constant(1.0);
let w2 = self.omega * self.omega;
let mut loss = tape.constant(0.0);
for &t in &self.colloc {
let (u, _ut, utt) = self.forward_jet(&tape, &pv, t, zero, one);
let r = utt + u * w2;
loss = loss + r * r;
}
loss = loss * (1.0 / self.colloc.len() as f64);
let (u0, ut0, _) = self.forward_jet(&tape, &pv, 0.0, zero, one);
let e1 = u0 - 1.0;
loss = loss + (e1 * e1 + ut0 * ut0) * self.ic_weight;
let g = loss.backward();
(loss.value(), pv.iter().map(|&p| g.wrt(p)).collect())
}
pub fn train_step(&mut self, lr: f64) -> f64 {
let (loss, grad) = self.loss_and_grad();
self.t += 1;
let (b1, b2, eps) = (0.9_f64, 0.999_f64, 1e-8);
let bc1 = 1.0 - b1.powi(self.t as i32);
let bc2 = 1.0 - b2.powi(self.t as i32);
for (i, &gi) in grad.iter().enumerate() {
self.m[i] = b1 * self.m[i] + (1.0 - b1) * gi;
self.v[i] = b2 * self.v[i] + (1.0 - b2) * gi * gi;
self.params[i] -= lr * (self.m[i] / bc1) / ((self.v[i] / bc2).sqrt() + eps);
}
loss
}
pub fn train(&mut self, epochs: usize, lr: f64) -> f64 {
let mut l = f64::INFINITY;
for _ in 0..epochs {
l = self.train_step(lr);
}
l
}
pub fn max_error(&self) -> f64 {
let n = 100;
(0..=n)
.map(|i| {
let t = self.colloc.last().copied().unwrap_or(1.0) * i as f64 / n as f64;
(self.forward(t) - self.exact(t)).abs()
})
.fold(0.0, f64::max)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn a_pinn_solves_the_harmonic_oscillator_from_physics_alone() {
let mut pinn = Pinn::new(16, 2.0, 2.0, 30, 7);
pinn.train(3000, 6e-3);
let err = pinn.max_error();
assert!(err < 0.08, "PINN should match cos(2t) from physics alone: max error {err}");
}
#[test]
fn an_untrained_pinn_does_not_satisfy_the_equation() {
let pinn = Pinn::new(24, 2.0, 2.0, 48, 7);
assert!(pinn.max_error() > 0.1, "an untrained net should not solve the ODE");
}
}