1use std::sync::OnceLock;
54
55use crate::models::control::{ControlProblem, StateDerivatives};
56
57#[derive(Clone, Debug)]
65pub struct LqRegulator<const N: usize, const M: usize> {
66 pub a: Vec<f64>,
68 pub b: Vec<f64>,
70 pub c: Vec<f64>,
72 pub q: Vec<f64>,
74 pub q_terminal: Vec<f64>,
76 pub r: Vec<f64>,
78 pub horizon: f64,
80 cache: OnceLock<RiccatiCache<N, M>>,
82}
83
84#[derive(Clone, Debug)]
93struct RiccatiCache<const N: usize, const M: usize> {
94 p_traj: Vec<Vec<f64>>,
97 q_traj: Vec<f64>,
99 r_inv: Vec<f64>,
101}
102
103impl<const N: usize, const M: usize> LqRegulator<N, M> {
104 #[allow(clippy::too_many_arguments)]
129 pub fn new(
130 a: &[f64],
131 b: &[f64],
132 c: &[f64],
133 q: &[f64],
134 q_terminal: &[f64],
135 r: &[f64],
136 horizon: f64,
137 ) -> Self {
138 assert!(horizon > 0.0, "horizon must be positive");
139 assert_eq!(a.len(), N * N, "A must be n x n");
140 assert_eq!(b.len(), N * M, "B must be n x m");
141 assert_eq!(c.len(), N * N, "C must be n x n");
142 assert_eq!(q.len(), N * N, "Q must be n x n");
143 assert_eq!(q_terminal.len(), N * N, "Q_T must be n x n");
144 assert_eq!(r.len(), M * M, "R must be m x m");
145 Self {
146 a: a.to_vec(),
147 b: b.to_vec(),
148 c: c.to_vec(),
149 q: q.to_vec(),
150 q_terminal: q_terminal.to_vec(),
151 r: r.to_vec(),
152 horizon,
153 cache: OnceLock::new(),
154 }
155 }
156
157 pub fn riccati_solution(&self) -> (Vec<f64>, f64) {
176 self.riccati_solution_with_steps(self.riccati_steps())
177 }
178
179 pub fn riccati_solution_with_steps(&self, steps: usize) -> (Vec<f64>, f64) {
188 let cache = self.riccati_cache();
189 if steps < cache.p_traj.len() {
190 return (cache.p_traj[steps].clone(), cache.q_traj[steps]);
191 }
192
193 let dt = self.horizon / self.riccati_steps() as f64;
196 let mut p = cache.p_traj[cache.p_traj.len() - 1].clone();
197 let mut q_scalar = cache.q_traj[cache.q_traj.len() - 1];
198
199 for _ in cache.p_traj.len() - 1..steps {
200 let dp = riccati_derivative::<N, M>(&p, &self.a, &self.b, &self.q, &cache.r_inv);
201 let dq = -trace_of_product::<N>(&self.c, &p);
202
203 for i in 0..N * N {
204 p[i] -= dt * dp[i];
205 }
206 q_scalar -= dt * dq;
207 }
208
209 (p, q_scalar)
210 }
211
212 fn riccati_cache(&self) -> &RiccatiCache<N, M> {
214 self.cache.get_or_init(|| self.compute_riccati_cache())
215 }
216
217 fn compute_riccati_cache(&self) -> RiccatiCache<N, M> {
220 let steps = self.riccati_steps();
221 let dt = self.horizon / steps as f64;
222
223 let r_inv = inverse_matrix(&self.r, M);
224
225 let mut p = self.q_terminal.clone();
226 let mut q_scalar = 0.0;
227
228 let mut p_traj = Vec::with_capacity(steps + 1);
229 let mut q_traj = Vec::with_capacity(steps + 1);
230 p_traj.push(p.clone());
231 q_traj.push(q_scalar);
232
233 for _ in 0..steps {
234 let dp = riccati_derivative::<N, M>(&p, &self.a, &self.b, &self.q, &r_inv);
235 let dq = -trace_of_product::<N>(&self.c, &p);
236
237 for i in 0..N * N {
238 p[i] -= dt * dp[i];
239 }
240 q_scalar -= dt * dq;
241
242 p_traj.push(p.clone());
243 q_traj.push(q_scalar);
244 }
245
246 RiccatiCache {
247 p_traj,
248 q_traj,
249 r_inv,
250 }
251 }
252
253 pub fn riccati_steps(&self) -> usize {
260 1000
261 }
262
263 pub fn feedback_gain(&self, t: f64) -> Vec<f64> {
267 let steps = self.riccati_steps();
268 let remaining = self.horizon - t;
269 let backward_steps = ((remaining / self.horizon) * steps as f64).round() as usize;
270
271 let (p, _) = self.riccati_solution_with_steps(backward_steps);
272 let r_inv = &self.riccati_cache().r_inv;
273
274 let mut k = vec![0.0; M * N];
276 for i in 0..M {
277 for j in 0..N {
278 let mut btp = 0.0;
279 for l in 0..N {
280 btp += self.b[l * M + i] * p[l * N + j];
281 }
282 let mut rbtp = 0.0;
283 for l in 0..M {
284 rbtp += r_inv[i * M + l] * btp;
285 }
286 k[i * N + j] = -rbtp;
287 }
288 }
289 k
290 }
291
292 pub fn exact_value(&self, state: &[f64; N], tau: f64) -> f64 {
294 let steps = self.riccati_steps();
295 let backward_steps = ((tau / self.horizon) * steps as f64).round() as usize;
296
297 let (p, q) = self.riccati_solution_with_steps(backward_steps);
298 -quadratic_form::<N>(&p, state) - q
299 }
300
301 pub fn exact_control(&self, state: &[f64; N], tau: f64) -> [f64; M] {
303 let t = self.horizon - tau;
304 let k = self.feedback_gain(t);
305 let mut u = [0.0; M];
306 for i in 0..M {
307 let mut acc = 0.0;
308 for j in 0..N {
309 acc += k[i * N + j] * state[j];
310 }
311 u[i] = acc;
312 }
313 u
314 }
315}
316
317impl<const N: usize, const M: usize> ControlProblem<N> for LqRegulator<N, M> {
318 type Control = [f64; M];
319
320 fn optimize(&self, t: f64, state: &[f64; N], _derivs: &StateDerivatives<N>) -> Self::Control {
321 self.exact_control(state, self.horizon - t)
322 }
323
324 fn running_reward(&self, _t: f64, state: &[f64; N], control: &Self::Control) -> f64 {
325 -(quadratic_form::<N>(&self.q, state) + quadratic_form::<M>(&self.r, control))
327 }
328
329 fn generator(
330 &self,
331 _t: f64,
332 state: &[f64; N],
333 control: &Self::Control,
334 derivs: &StateDerivatives<N>,
335 ) -> f64 {
336 let mut ax = [0.0; N];
341 for (i, out) in ax.iter_mut().enumerate() {
342 for (j, &xj) in state.iter().enumerate() {
343 *out += self.a[i * N + j] * xj;
344 }
345 }
346 let mut bu = [0.0; N];
347 for (i, out) in bu.iter_mut().enumerate() {
348 for (j, &uj) in control.iter().enumerate() {
349 *out += self.b[i * M + j] * uj;
350 }
351 }
352 let drift = ax
353 .iter()
354 .zip(bu.iter())
355 .zip(derivs.grad.iter())
356 .map(|((&a, &b), &g)| (a + b) * g)
357 .sum::<f64>();
358
359 let mut diffusion = 0.0;
360 for i in 0..N {
361 for j in 0..N {
362 let mut cov = 0.0;
363 for l in 0..N {
364 cov += self.c[i * N + l] * self.c[j * N + l];
365 }
366 diffusion += cov * derivs.hessian_full[i][j];
367 }
368 }
369
370 drift + 0.5 * diffusion
371 }
372
373 fn terminal(&self, state: &[f64; N]) -> f64 {
374 -quadratic_form::<N>(&self.q_terminal, state)
375 }
376
377 fn next_step(&self, t: f64, state: &[f64; N], dt: f64, noise: &[f64; N]) -> [f64; N] {
378 let u = self.exact_control(state, self.horizon - t);
379
380 let mut ax = [0.0; N];
381 for (i, out) in ax.iter_mut().enumerate() {
382 for (j, &xj) in state.iter().enumerate() {
383 *out += self.a[i * N + j] * xj;
384 }
385 }
386 let mut bu = [0.0; N];
387 for (i, out) in bu.iter_mut().enumerate() {
388 for (j, &uj) in u.iter().enumerate() {
389 *out += self.b[i * M + j] * uj;
390 }
391 }
392
393 let mut diff = [0.0; N];
395 for (i, out) in diff.iter_mut().enumerate() {
396 for (j, &nj) in noise.iter().enumerate() {
397 *out += self.c[i * N + j] * nj;
398 }
399 }
400
401 let mut next = [0.0; N];
402 for (i, out) in next.iter_mut().enumerate() {
403 *out = state[i] + (ax[i] + bu[i]) * dt + diff[i] * dt.sqrt();
404 }
405 next
406 }
407
408 fn is_diffusion_dimension(&self, _dim: usize) -> bool {
409 true
410 }
411}
412
413fn quadratic_form<const N: usize>(m: &[f64], x: &[f64]) -> f64 {
415 let mut acc = 0.0;
416 for i in 0..N {
417 for j in 0..N {
418 acc += x[i] * m[i * N + j] * x[j];
419 }
420 }
421 acc
422}
423
424fn inverse_matrix(m: &[f64], n: usize) -> Vec<f64> {
426 let mut aug = vec![0.0; n * (2 * n)];
428 for i in 0..n {
429 for j in 0..n {
430 aug[i * (2 * n) + j] = m[i * n + j];
431 }
432 aug[i * (2 * n) + n + i] = 1.0;
433 }
434
435 for col in 0..n {
436 let mut pivot = col;
438 for row in col + 1..n {
439 if aug[row * (2 * n) + col].abs() > aug[pivot * (2 * n) + col].abs() {
440 pivot = row;
441 }
442 }
443 if pivot != col {
444 for k in 0..(2 * n) {
445 aug.swap(col * (2 * n) + k, pivot * (2 * n) + k);
446 }
447 }
448 let diag = aug[col * (2 * n) + col];
449 debug_assert!(diag.abs() > 1e-12, "Riccati matrix must be invertible");
450
451 for k in 0..(2 * n) {
453 aug[col * (2 * n) + k] /= diag;
454 }
455 for row in 0..n {
457 if row == col {
458 continue;
459 }
460 let factor = aug[row * (2 * n) + col];
461 if factor == 0.0 {
462 continue;
463 }
464 for k in 0..(2 * n) {
465 aug[row * (2 * n) + k] -= factor * aug[col * (2 * n) + k];
466 }
467 }
468 }
469
470 let mut inv = vec![0.0; n * n];
471 for i in 0..n {
472 for j in 0..n {
473 inv[i * n + j] = aug[i * (2 * n) + n + j];
474 }
475 }
476 inv
477}
478
479fn riccati_derivative<const N: usize, const M: usize>(
481 p: &[f64],
482 a: &[f64],
483 b: &[f64],
484 q: &[f64],
485 r_inv: &[f64],
486) -> Vec<f64> {
487 let mut dp = vec![0.0; N * N];
488
489 for i in 0..N {
490 for j in 0..N {
491 let mut atp = 0.0;
492 let mut pa = 0.0;
493 for l in 0..N {
494 atp += a[l * N + i] * p[l * N + j];
495 pa += p[i * N + l] * a[l * N + j];
496 }
497
498 let mut brb = 0.0;
502 for k in 0..M {
503 let mut btp_kj = 0.0;
504 for l in 0..N {
505 btp_kj += b[l * M + k] * p[l * N + j];
506 }
507 let mut pbr_ik = 0.0;
508 for m in 0..M {
509 let mut pb_im = 0.0;
510 for l in 0..N {
511 pb_im += p[i * N + l] * b[l * M + m];
512 }
513 pbr_ik += pb_im * r_inv[m * M + k];
514 }
515 brb += pbr_ik * btp_kj;
516 }
517
518 dp[i * N + j] = -atp - pa - q[i * N + j] + brb;
519 }
520 }
521
522 dp
523}
524
525fn trace_of_product<const N: usize>(c: &[f64], p: &[f64]) -> f64 {
527 let mut acc = 0.0;
530 for i in 0..N {
531 for j in 0..N {
532 let mut cp = 0.0;
533 for l in 0..N {
534 cp += c[l * N + j] * p[l * N + i];
535 }
536 acc += c[i * N + j] * cp;
537 }
538 }
539 acc
540}
541
542#[cfg(test)]
543mod tests {
544 use super::*;
545
546 fn scalar_lq() -> LqRegulator<1, 1> {
547 LqRegulator::<1, 1>::new(&[-1.0], &[1.0], &[0.0], &[1.0], &[1.0], &[1.0], 1.0)
550 }
551
552 #[test]
553 fn scalar_riccati_converges_to_finite_positive_p() {
554 let lq = scalar_lq();
555 let (p, q) = lq.riccati_solution();
556 assert!(p[0] > 0.0);
557 assert!(p[0].is_finite());
558 assert_eq!(q, 0.0, "no diffusion means q is identically zero");
559 }
560
561 #[test]
562 fn scalar_policy_is_linear_feedback() {
563 let lq = scalar_lq();
564 let u = lq.exact_control(&[2.0], 1.0);
565 assert!(u[0] < 0.0);
567 }
568
569 #[test]
570 fn exact_value_is_negative_quadratic() {
571 let lq = scalar_lq();
572 let v = lq.exact_value(&[3.0], 1.0);
573 assert!(v <= 0.0);
574 }
575
576 #[test]
577 fn terminal_matches_negative_quadratic() {
578 let lq = scalar_lq();
579 let v = lq.terminal(&[4.0]);
580 assert!((v - (-16.0)).abs() < 1e-12);
581 }
582
583 #[test]
584 fn riccati_steps_select_remaining_horizon() {
585 let lq = LqRegulator::<1, 1>::new(&[-0.5], &[1.0], &[0.0], &[1.0], &[1.0], &[1.0], 0.5);
594
595 let (p, q) = lq.riccati_solution_with_steps(1);
596 assert!((p[0] - 0.9995).abs() < 1e-12, "got {}", p[0]);
597 assert_eq!(q, 0.0);
598 }
599
600 #[test]
601 fn diffusion_adds_negative_value_correction() {
602 let lq = LqRegulator::<1, 1>::new(&[-1.0], &[1.0], &[1.0], &[1.0], &[1.0], &[1.0], 1.0);
604 let (_, q) = lq.riccati_solution();
605 assert!(q > 0.0, "noise increases the minimal cost so q(0) > 0");
606 let v = lq.exact_value(&[0.0], 1.0);
607 assert!(v < 0.0);
608 }
609
610 #[test]
611 fn riccati_p_is_symmetric_for_symmetric_q() {
612 let lq = LqRegulator::<2, 1>::new(
613 &[0.0, 1.0, 0.0, 0.0],
614 &[0.0, 1.0],
615 &[0.0, 0.0, 0.0, 0.0],
616 &[1.0, 0.0, 0.0, 0.0],
617 &[1.0, 0.0, 0.0, 0.0],
618 &[1.0],
619 1.0,
620 );
621 let (p, _) = lq.riccati_solution();
622 assert!(p[0].is_finite() && p[3].is_finite());
623 assert!((p[1] - p[2]).abs() < 1e-9);
624 }
625
626 #[test]
627 fn correlated_generator_includes_off_diagonal_covariance() {
628 let rho: f64 = 0.5;
635 let sqrt_1mr2: f64 = (1.0 - rho * rho).sqrt();
636 let lq = LqRegulator::<2, 0>::new(
637 &[0.0; 4],
638 &[],
639 &[1.0, 0.0, rho, sqrt_1mr2],
640 &[0.0; 4],
641 &[0.0; 4],
642 &[],
643 1.0,
644 );
645 let derivs = StateDerivatives::<2>::with_full_hessian(
646 [0.0; 2],
647 [[0.0, 1.0], [1.0, 0.0]],
648 [0.0; 2],
649 [0.0; 2],
650 );
651 let g = ControlProblem::generator(&lq, 0.0, &[0.0; 2], &[], &derivs);
652 assert!(
653 (g - 0.5).abs() < 1e-12,
654 "correlated generator should be 0.5, got {g}"
655 );
656 }
657}