Skip to main content

ferromotion_learn/
diff_control.rs

1//! **Differentiable control** — tune a controller by backpropagating a trajectory cost *through* the closed
2//! loop. Because the plant is differentiable, you can roll out the controlled system, measure how far it is
3//! from the goal, and take the gradient of that cost w.r.t. the controller's parameters — then step them
4//! downhill. No trial-and-error over thousands of episodes as in reinforcement learning; one gradient per
5//! rollout. This is the payoff of every learned/differentiable model in the course: a model you can
6//! differentiate is a model you can control by gradient descent.
7//!
8//! Here the controller is a **PID** (three interpretable gains) and the plant is a damped point mass under a
9//! constant disturbance, so proportional-derivative action alone leaves a steady-state offset and the loop
10//! must *discover* integral action to null it. The gains are tuned by rolling the closed loop out on the
11//! reverse-mode tape and backpropagating a setpoint-tracking cost. Verified: from a proportional-only start
12//! that never reaches the setpoint, training drives the steady-state error to near zero and the integral gain
13//! becomes positive — the loop learned it needs an integrator, from the gradient alone.
14
15use crate::autodiff::Tape;
16
17/// A PID controller (gains `[Kp, Ki, Kd]`) tuned by differentiating through the plant rollout.
18pub struct PidController {
19    gains: [f64; 3],
20    m: [f64; 3],
21    v: [f64; 3],
22    t: u64,
23    // plant + task
24    mass: f64,
25    damping: f64,
26    disturbance: f64,
27    setpoint: f64,
28    dt: f64,
29    steps: usize,
30}
31
32impl PidController {
33    /// A controller for a damped point mass `mass·ẍ = u − damping·ẋ − disturbance`, tasked to hold
34    /// `setpoint`. Starts proportional-only (`Kp=1, Ki=0, Kd=0`).
35    pub fn new(mass: f64, damping: f64, disturbance: f64, setpoint: f64) -> Self {
36        PidController {
37            gains: [1.0, 0.0, 0.0],
38            m: [0.0; 3],
39            v: [0.0; 3],
40            t: 0,
41            mass,
42            damping,
43            disturbance,
44            setpoint,
45            dt: 0.05,
46            steps: 90,
47        }
48    }
49
50    pub fn gains(&self) -> [f64; 3] {
51        self.gains
52    }
53
54    /// Roll out the closed loop in plain `f64` and return the position trajectory (for plotting).
55    pub fn simulate(&self) -> Vec<f64> {
56        let (mut x, mut v, mut integ) = (0.0, 0.0, 0.0);
57        let mut out = vec![x];
58        for _ in 0..self.steps {
59            let e = self.setpoint - x;
60            integ += e * self.dt;
61            let de = -v; // ė = −ẋ (setpoint constant)
62            let u = self.gains[0] * e + self.gains[1] * integ + self.gains[2] * de;
63            let acc = (u - self.damping * v - self.disturbance) / self.mass;
64            v += acc * self.dt;
65            x += v * self.dt;
66            out.push(x);
67        }
68        out
69    }
70
71    /// Steady-state tracking error (distance of the final position from the setpoint).
72    pub fn final_error(&self) -> f64 {
73        (self.simulate().last().copied().unwrap_or(0.0) - self.setpoint).abs()
74    }
75
76    /// One Adam step of the tracking cost, differentiating through the closed-loop rollout. Returns the cost.
77    pub fn train_step(&mut self, lr: f64) -> f64 {
78        let tape = Tape::new();
79        let kp = tape.var(self.gains[0]);
80        let ki = tape.var(self.gains[1]);
81        let kd = tape.var(self.gains[2]);
82        let mut x = tape.constant(0.0);
83        let mut v = tape.constant(0.0);
84        let mut integ = tape.constant(0.0);
85        let mut cost = tape.constant(0.0);
86        let (m, c, d, sp, dt) = (self.mass, self.damping, self.disturbance, self.setpoint, self.dt);
87        for step in 0..self.steps {
88            let e = tape.constant(sp) - x;
89            integ = integ + e * dt;
90            let de = v * (-1.0);
91            let u = kp * e + ki * integ + kd * de;
92            let acc = (u - v * c - tape.constant(d)) * (1.0 / m);
93            v = v + acc * dt;
94            x = x + v * dt;
95            // weight later steps more (settle to the setpoint), plus a small control-effort penalty
96            let w = 0.5 + step as f64 / self.steps as f64;
97            cost = cost + e * e * (w * dt) + u * u * (1e-4 * dt);
98        }
99        let g = cost.backward();
100        let grad = [g.wrt(kp), g.wrt(ki), g.wrt(kd)];
101
102        self.t += 1;
103        let (b1, b2, eps) = (0.9_f64, 0.999_f64, 1e-8);
104        let bc1 = 1.0 - b1.powi(self.t as i32);
105        let bc2 = 1.0 - b2.powi(self.t as i32);
106        for (i, &gi) in grad.iter().enumerate() {
107            self.m[i] = b1 * self.m[i] + (1.0 - b1) * gi;
108            self.v[i] = b2 * self.v[i] + (1.0 - b2) * gi * gi;
109            self.gains[i] -= lr * (self.m[i] / bc1) / ((self.v[i] / bc2).sqrt() + eps);
110        }
111        cost.value()
112    }
113
114    /// Train for `epochs` steps; returns the final cost.
115    pub fn train(&mut self, epochs: usize, lr: f64) -> f64 {
116        let mut c = f64::INFINITY;
117        for _ in 0..epochs {
118            c = self.train_step(lr);
119        }
120        c
121    }
122}
123
124#[cfg(test)]
125mod tests {
126    use super::*;
127
128    #[test]
129    fn differentiable_tuning_nulls_the_steady_state_error_by_discovering_integral_action() {
130        // Plant: mass 1, damping 0.6, constant disturbance 2.0, setpoint 1.0. Proportional-only leaves a large
131        // offset; differentiating the tracking cost through the rollout must find gains (incl. Ki > 0) that
132        // reach the setpoint.
133        let mut pid = PidController::new(1.0, 0.6, 2.0, 1.0);
134        let e_before = pid.final_error();
135        assert!(e_before > 0.3, "proportional-only should miss the setpoint: {e_before}");
136        pid.train(800, 0.05);
137        let e_after = pid.final_error();
138        assert!(e_after < 0.05, "differentiable tuning should reach the setpoint: error {e_after}");
139        assert!(pid.gains()[1] > 0.1, "it should discover integral action: Ki = {}", pid.gains()[1]);
140    }
141}