Skip to main content

solver/models/
stationary_lq.rs

1//! Infinite-horizon scalar linear-quadratic regulator.
2//!
3//! A validating instance of
4//! [`crate::numeric::finite_difference::elliptic::EllipticControlProblem`]
5//! with a closed-form stationary solution obtained from the algebraic Riccati
6//! equation.
7//!
8//! # Dynamics
9//!
10//! ```text
11//! dx = (a x + b u) dt + c dW
12//! ```
13//!
14//! # Objective
15//!
16//! Minimize the expected discounted running quadratic cost
17//!
18//! ```text
19//! J(u) = E[ integral_0^inf e^{-rho t} (q x^2 + r u^2) dt ]
20//! ```
21//!
22//! which is maximized in negated form by
23//! [`crate::models::control::ControlProblem`]. A positive discount `rho` makes
24//! the stationary value finite and the associated elliptic operator well posed.
25//!
26//! # HJB equation
27//!
28//! ```text
29//! rho V = sup_u { -(q x^2 + r u^2) + (a x + b u) V_x + 0.5 c^2 V_xx }
30//! ```
31//!
32//! # Exact solution
33//!
34//! The stationary value is `V(x) = -P x^2 - d`, where
35//!
36//! ```text
37//! d = c^2 P / rho
38//! ```
39//!
40//! and the optimal control is
41//!
42//! ```text
43//! u*(x) = -(b / r) P x
44//! ```
45//!
46//! `P` is the positive solution of the scalar algebraic Riccati equation
47//!
48//! ```text
49//! (b^2 / r) P^2 + (rho - 2a) P - q = 0
50//! ```
51//!
52//! A stable closed loop selects the positive root. The additive constant `d`
53//! is the noise-induced correction; it vanishes when `c = 0`.
54use crate::models::control::{ControlProblem, StateDerivatives};
55use crate::numeric::finite_difference::discretization::{
56    BoundaryCondition, BoundaryConditions, DimensionKind, Transport,
57};
58use crate::numeric::finite_difference::elliptic::EllipticControlProblem;
59
60/// Infinite-horizon scalar linear-quadratic regulator with constant
61/// coefficients.
62///
63/// See the module-level documentation for the formulation and exact solution.
64#[derive(Clone, Debug)]
65pub struct StationaryLqRegulator {
66    /// State-drift coefficient `a`.
67    pub a: f64,
68    /// Control-input coefficient `b`.
69    pub b: f64,
70    /// Running state cost `q`.
71    pub q: f64,
72    /// Running control cost `r`.
73    pub r: f64,
74    /// Diffusion coefficient `c`.
75    pub c: f64,
76    /// Discount rate `rho`.
77    pub rho: f64,
78    /// Grid spacing `h` used by the finite-difference transport stencil.
79    pub h: f64,
80    /// Positive algebraic Riccati solution `P`.
81    pub p: f64,
82    /// Noise-induced additive correction `d = c^2 P / rho`.
83    pub d: f64,
84    /// Boundary conditions for the stationary solve.
85    pub boundary_conditions: BoundaryConditions<1>,
86}
87
88impl StationaryLqRegulator {
89    /// Creates a stationary LQ regulator.
90    ///
91    /// # Panics
92    ///
93    /// Panics if `r`, `q`, `rho`, or `h` is not positive, or if the scalar
94    /// algebraic Riccati equation has no positive real root.
95    ///
96    /// # Examples
97    ///
98    /// ```
99    /// use solver::models::stationary_lq::StationaryLqRegulator;
100    /// let lq = StationaryLqRegulator::new(-1.0, 1.0, 1.0, 1.0, 0.0, 0.1, 0.1);
101    /// assert!(lq.p > 0.0);
102    /// ```
103    pub fn new(a: f64, b: f64, q: f64, r: f64, c: f64, rho: f64, h: f64) -> Self {
104        assert!(r > 0.0, "control cost r must be positive");
105        assert!(q > 0.0, "state cost q must be positive");
106        assert!(rho > 0.0, "discount rate rho must be positive");
107        assert!(h > 0.0, "grid spacing h must be positive");
108
109        // (b^2/r) P^2 + (rho - 2a) P - q = 0.
110        let coeff = b * b / r;
111        let linear = rho - 2.0 * a;
112        let discriminant = linear * linear + 4.0 * coeff * q;
113        assert!(
114            discriminant >= 0.0,
115            "algebraic Riccati discriminant must be non-negative"
116        );
117        let p = (-linear + discriminant.sqrt()) / (2.0 * coeff);
118        assert!(p > 0.0, "algebraic Riccati positive root must exist");
119        let d = c * c * p / rho;
120
121        Self {
122            a,
123            b,
124            q,
125            r,
126            c,
127            rho,
128            h,
129            p,
130            d,
131            boundary_conditions: BoundaryConditions::default(),
132        }
133    }
134
135    /// Sets the grid spacing used by the transport stencil.
136    ///
137    /// # Examples
138    ///
139    /// ```
140    /// use solver::models::stationary_lq::StationaryLqRegulator;
141    /// let lq = StationaryLqRegulator::new(-1.0, 1.0, 1.0, 1.0, 0.0, 0.1, 0.1).with_h(0.05);
142    /// assert_eq!(lq.h, 0.05);
143    /// ```
144    pub fn with_h(mut self, h: f64) -> Self {
145        assert!(h > 0.0, "grid spacing h must be positive");
146        self.h = h;
147        self
148    }
149
150    /// Sets the boundary conditions used by the stationary solve.
151    ///
152    /// # Examples
153    ///
154    /// ```
155    /// use solver::models::stationary_lq::StationaryLqRegulator;
156    /// use solver::numeric::finite_difference::discretization::{
157    ///     BoundaryCondition, BoundaryConditions,
158    /// };
159    /// let lq = StationaryLqRegulator::new(-1.0, 1.0, 1.0, 1.0, 0.0, 0.1, 0.1)
160    ///     .with_boundary_conditions(BoundaryConditions::new(
161    ///         [BoundaryCondition::Dirichlet(0.0)],
162    ///         [BoundaryCondition::Dirichlet(0.0)],
163    ///     ));
164    /// assert!(matches!(lq.boundary_conditions.lower[0], BoundaryCondition::Dirichlet(_)));
165    /// ```
166    pub fn with_boundary_conditions(mut self, boundary_conditions: BoundaryConditions<1>) -> Self {
167        self.boundary_conditions = boundary_conditions;
168        self
169    }
170
171    /// Exact stationary value `V(x) = -P x^2 - d`.
172    pub fn exact_value(&self, state: &[f64; 1]) -> f64 {
173        -self.p * state[0] * state[0] - self.d
174    }
175
176    /// Exact stationary control `u*(x) = -(b / r) P x`.
177    pub fn exact_control(&self, state: &[f64; 1]) -> f64 {
178        -(self.b / self.r) * self.p * state[0]
179    }
180
181    /// Dirichlet boundary conditions matching the exact value at the grid
182    /// endpoints.
183    pub fn exact_boundary_conditions(
184        &self,
185        grid: &crate::core::grid::Grid<1>,
186    ) -> BoundaryConditions<1> {
187        let lower = self.exact_value(&[grid.min[0]]);
188        let upper =
189            self.exact_value(&[grid.min[0] + (grid.tensor_info.shape[0] - 1) as f64 * grid.dx[0]]);
190        BoundaryConditions::new(
191            [BoundaryCondition::Dirichlet(lower)],
192            [BoundaryCondition::Dirichlet(upper)],
193        )
194    }
195}
196
197impl ControlProblem<1> for StationaryLqRegulator {
198    type Control = f64;
199
200    fn optimize(&self, _t: f64, state: &[f64; 1], _derivs: &StateDerivatives<1>) -> Self::Control {
201        self.exact_control(state)
202    }
203
204    fn running_reward(&self, _t: f64, state: &[f64; 1], control: &Self::Control) -> f64 {
205        -(self.q * state[0] * state[0] + self.r * control * control)
206    }
207
208    fn generator(
209        &self,
210        _t: f64,
211        state: &[f64; 1],
212        control: &Self::Control,
213        derivs: &StateDerivatives<1>,
214    ) -> f64 {
215        let mu = self.a * state[0] + self.b * control;
216        mu * derivs.grad[0] + 0.5 * self.c * self.c * derivs.hessian[0]
217    }
218
219    fn terminal(&self, state: &[f64; 1]) -> f64 {
220        self.exact_value(state)
221    }
222
223    fn discount_rate(&self, _state: &[f64; 1]) -> f64 {
224        self.rho
225    }
226
227    fn next_step(&self, _t: f64, state: &[f64; 1], dt: f64, noise: &[f64; 1]) -> [f64; 1] {
228        let u = self.exact_control(state);
229        let drift = self.a * state[0] + self.b * u;
230        [state[0] + drift * dt + self.c * noise[0] * dt.sqrt()]
231    }
232
233    fn is_diffusion_dimension(&self, _dim: usize) -> bool {
234        true
235    }
236
237    fn gradient_step(&self, _dim: usize) -> f64 {
238        self.h
239    }
240}
241
242impl EllipticControlProblem<1> for StationaryLqRegulator {
243    fn dimension_kind(&self, _dim: usize) -> DimensionKind {
244        DimensionKind::Diffusion
245    }
246
247    fn transport(
248        &self,
249        state: &[f64; 1],
250        control: &Self::Control,
251        derivs: &StateDerivatives<1>,
252    ) -> Transport<1> {
253        let x = state[0];
254        let mu = self.a * x + self.b * control;
255        let diffusion = 0.5 * self.c * self.c;
256        let h = self.h;
257        let plus = diffusion / (h * h) + mu.max(0.0) / h;
258        let minus = diffusion / (h * h) + (-mu).max(0.0) / h;
259
260        // Discrete transport operator at the current derivative bundle.
261        let t_v = plus * h * derivs.fwd[0] - minus * h * derivs.bwd[0];
262
263        let driver =
264            self.running_reward(0.0, state, control) + self.generator(0.0, state, control, derivs);
265        Transport::new([plus], [minus], driver - t_v)
266    }
267
268    fn boundary_conditions(&self) -> BoundaryConditions<1> {
269        self.boundary_conditions.clone()
270    }
271}
272
273#[cfg(test)]
274mod tests {
275    use super::*;
276
277    fn default_lq() -> StationaryLqRegulator {
278        StationaryLqRegulator::new(-1.0, 1.0, 1.0, 1.0, 0.0, 0.1, 0.1)
279    }
280
281    #[test]
282    fn exact_value_is_negative_quadratic() {
283        let lq = default_lq();
284        let v = lq.exact_value(&[2.0]);
285        assert!(v < 0.0);
286        assert!((v - (-lq.p * 4.0 - lq.d)).abs() < 1e-12);
287    }
288
289    #[test]
290    fn exact_control_is_negative_feedback() {
291        let lq = default_lq();
292        let u = lq.exact_control(&[2.0]);
293        assert!(u < 0.0);
294        assert!((u - (-lq.p * 2.0)).abs() < 1e-12);
295    }
296
297    #[test]
298    fn riccati_positive_root_is_consistent() {
299        let lq = default_lq();
300        // (b^2/r) P^2 + (rho - 2a) P - q = 0.
301        let residual = (lq.b * lq.b / lq.r) * lq.p * lq.p + (lq.rho - 2.0 * lq.a) * lq.p - lq.q;
302        assert!(residual.abs() < 1e-12);
303    }
304
305    #[test]
306    fn transport_source_reconstructs_driver() {
307        let lq = default_lq();
308        let state = [0.5];
309        let derivs = StateDerivatives::with_directional([-0.4], [-0.8], [0.1], [-0.05]);
310        let control = lq.optimize(0.0, &state, &derivs);
311        let transport = lq.transport(&state, &control, &derivs);
312
313        let t_v =
314            transport.plus[0] * 0.1 * derivs.fwd[0] - transport.minus[0] * 0.1 * derivs.bwd[0];
315        let driver =
316            lq.running_reward(0.0, &state, &control) + lq.generator(0.0, &state, &control, &derivs);
317
318        assert!((transport.source + t_v - driver).abs() < 1e-12);
319    }
320}