1use 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#[derive(Clone, Debug)]
65pub struct StationaryLqRegulator {
66 pub a: f64,
68 pub b: f64,
70 pub q: f64,
72 pub r: f64,
74 pub c: f64,
76 pub rho: f64,
78 pub h: f64,
80 pub p: f64,
82 pub d: f64,
84 pub boundary_conditions: BoundaryConditions<1>,
86}
87
88impl StationaryLqRegulator {
89 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 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 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 pub fn with_boundary_conditions(mut self, boundary_conditions: BoundaryConditions<1>) -> Self {
167 self.boundary_conditions = boundary_conditions;
168 self
169 }
170
171 pub fn exact_value(&self, state: &[f64; 1]) -> f64 {
173 -self.p * state[0] * state[0] - self.d
174 }
175
176 pub fn exact_control(&self, state: &[f64; 1]) -> f64 {
178 -(self.b / self.r) * self.p * state[0]
179 }
180
181 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 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 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}