use crate::autodiff::Tape;
use nalgebra::{DMatrix, DVector};
pub struct ModelBasedControl {
true_mass: f64,
true_damping: f64,
true_disturbance: f64,
setpoint: f64,
dt: f64,
steps: usize,
learned: [f64; 4],
identified: bool,
gains: [f64; 3],
m: [f64; 3],
v: [f64; 3],
t: u64,
}
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)
}
impl ModelBasedControl {
pub fn new(mass: f64, damping: f64, disturbance: f64, setpoint: f64) -> Self {
ModelBasedControl {
true_mass: mass,
true_damping: damping,
true_disturbance: disturbance,
setpoint,
dt: 0.05,
steps: 90,
learned: [0.0; 4],
identified: false,
gains: [1.0, 0.0, 0.0],
m: [0.0; 3],
v: [0.0; 3],
t: 0,
}
}
fn true_accel(&self, x: f64, v: f64, u: f64) -> f64 {
let _ = x; (u - self.true_damping * v - self.true_disturbance) / self.true_mass
}
pub fn identify(&mut self) {
let mut s = 20240101u64;
let mut rnd = |lo: f64, hi: f64| lo + (splitmix64(&mut s) as f64 / u64::MAX as f64) * (hi - lo);
let n = 200;
let mut a = DMatrix::zeros(n, 4);
let mut y = DVector::zeros(n);
for i in 0..n {
let (x, v, u) = (rnd(-2.0, 2.0), rnd(-2.0, 2.0), rnd(-5.0, 5.0));
a[(i, 0)] = x;
a[(i, 1)] = v;
a[(i, 2)] = u;
a[(i, 3)] = 1.0;
y[i] = self.true_accel(x, v, u);
}
let sol = a.svd(true, true).solve(&y, 1e-12).unwrap_or_else(|_| DVector::zeros(4));
self.learned = [sol[0], sol[1], sol[2], sol[3]];
self.identified = true;
}
pub fn learned_coeffs(&self) -> [f64; 4] {
self.learned
}
pub fn true_coeffs(&self) -> [f64; 4] {
[0.0, -self.true_damping / self.true_mass, 1.0 / self.true_mass, -self.true_disturbance / self.true_mass]
}
pub fn gains(&self) -> [f64; 3] {
self.gains
}
fn simulate(&self, use_true: bool) -> Vec<f64> {
let (mut x, mut v, mut integ) = (0.0, 0.0, 0.0);
let mut out = vec![x];
for _ in 0..self.steps {
let e = self.setpoint - x;
integ += e * self.dt;
let u = self.gains[0] * e + self.gains[1] * integ + self.gains[2] * (-v);
let acc = if use_true {
self.true_accel(x, v, u)
} else {
self.learned[0] * x + self.learned[1] * v + self.learned[2] * u + self.learned[3]
};
v += acc * self.dt;
x += v * self.dt;
out.push(x);
}
out
}
pub fn deploy(&self) -> Vec<f64> {
self.simulate(true)
}
pub fn imagine(&self) -> Vec<f64> {
self.simulate(false)
}
pub fn deploy_error(&self) -> f64 {
(self.deploy().last().copied().unwrap_or(0.0) - self.setpoint).abs()
}
pub fn train_step(&mut self, lr: f64) -> f64 {
let tape = Tape::new();
let kp = tape.var(self.gains[0]);
let ki = tape.var(self.gains[1]);
let kd = tape.var(self.gains[2]);
let mut x = tape.constant(0.0);
let mut v = tape.constant(0.0);
let mut integ = tape.constant(0.0);
let mut cost = tape.constant(0.0);
let (sp, dt) = (self.setpoint, self.dt);
let [a, b, c, d] = self.learned;
for step in 0..self.steps {
let e = tape.constant(sp) - x;
integ = integ + e * dt;
let u = kp * e + ki * integ + kd * (v * -1.0);
let acc = x * a + v * b + u * c + tape.constant(d);
v = v + acc * dt;
x = x + v * dt;
let w = 0.5 + step as f64 / self.steps as f64;
cost = cost + e * e * (w * dt) + u * u * (1e-4 * dt);
}
let g = cost.backward();
let grad = [g.wrt(kp), g.wrt(ki), g.wrt(kd)];
self.t += 1;
let (b1, b2, eps) = (0.9_f64, 0.999_f64, 1e-8);
let (bc1, bc2) = (1.0 - b1.powi(self.t as i32), 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.gains[i] -= lr * (self.m[i] / bc1) / ((self.v[i] / bc2).sqrt() + eps);
}
cost.value()
}
pub fn train(&mut self, epochs: usize, lr: f64) -> f64 {
let mut c = f64::INFINITY;
for _ in 0..epochs {
c = self.train_step(lr);
}
c
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn learn_a_model_then_control_the_real_plant_through_it() {
let mut mbc = ModelBasedControl::new(1.0, 0.6, 2.0, 1.0);
mbc.identify();
let (lc, tc) = (mbc.learned_coeffs(), mbc.true_coeffs());
for i in 0..4 {
assert!((lc[i] - tc[i]).abs() < 1e-6, "coeff {i}: learned {} vs true {}", lc[i], tc[i]);
}
let before = mbc.deploy_error();
assert!(before > 0.3, "proportional-only should miss on the true plant: {before}");
mbc.train(800, 0.05);
let after = mbc.deploy_error();
assert!(after < 0.06, "controller trained in imagination should settle the real plant: {after}");
assert!(mbc.gains()[1] > 0.1, "it should discover integral action: Ki={}", mbc.gains()[1]);
}
}