Skip to main content

rust_robotics_control/
cgmres_nmpc.rs

1#![allow(dead_code, clippy::too_many_arguments, clippy::type_complexity)]
2
3//! Nonlinear MPC simulation with CGMRES (Continuation GMRES)
4//!
5//! author: Atsushi Sakai (@Atsushi_twi)
6//!         Ryohei Sasaki (@rsasaki0109) - Rust port
7//!
8//! Reference:
9//! - PythonRobotics: <https://github.com/AtsushiSakai/PythonRobotics>
10//! - Shunichi09/nonlinear_control: Implementing the nonlinear model predictive
11//!   control, sliding mode control <https://github.com/Shunichi09/PythonLinearNonlinearControl>
12//!
13//! Note: This implementation may exhibit numerical instability in some scenarios.
14//! The C-GMRES algorithm is known to be sensitive to parameter tuning.
15//! For production use, consider adjusting zeta, ht, threshold, and other parameters.
16
17use nalgebra::{DMatrix, DVector};
18use std::f64::consts::PI;
19
20// System parameters
21const U_A_MAX: f64 = 1.0; // Maximum acceleration input
22const U_OMEGA_MAX: f64 = PI / 4.0; // Maximum steering angle (45 degrees)
23const PHI_V: f64 = 0.01; // Penalty weight for acceleration
24const PHI_OMEGA: f64 = 0.01; // Penalty weight for steering
25const WB: f64 = 0.25; // Wheelbase [m]
26
27/// Differential model for two-wheeled robot
28/// Returns (dx, dy, d_yaw, dv) derivatives
29fn differential_model(v: f64, yaw: f64, u_1: f64, u_2: f64) -> (f64, f64, f64, f64) {
30    let dx = yaw.cos() * v;
31    let dy = yaw.sin() * v;
32    let dv = u_1;
33    // Using sin(u_2) instead of tan(u_2) for better nonlinear optimization
34    let d_yaw = v / WB * u_2.sin();
35
36    (dx, dy, d_yaw, dv)
37}
38
39/// Two-wheeled robot system state
40pub struct TwoWheeledSystem {
41    pub x: f64,
42    pub y: f64,
43    pub yaw: f64,
44    pub v: f64,
45    pub history_x: Vec<f64>,
46    pub history_y: Vec<f64>,
47    pub history_yaw: Vec<f64>,
48    pub history_v: Vec<f64>,
49}
50
51impl TwoWheeledSystem {
52    pub fn new(init_x: f64, init_y: f64, init_yaw: f64, init_v: f64) -> Self {
53        TwoWheeledSystem {
54            x: init_x,
55            y: init_y,
56            yaw: init_yaw,
57            v: init_v,
58            history_x: vec![init_x],
59            history_y: vec![init_y],
60            history_yaw: vec![init_yaw],
61            history_v: vec![init_v],
62        }
63    }
64
65    /// Update state using Euler integration
66    pub fn update_state(&mut self, u_1: f64, u_2: f64, dt: f64) {
67        let (dx, dy, d_yaw, dv) = differential_model(self.v, self.yaw, u_1, u_2);
68
69        self.x += dt * dx;
70        self.y += dt * dy;
71        self.yaw += dt * d_yaw;
72        self.v += dt * dv;
73
74        // Save history
75        self.history_x.push(self.x);
76        self.history_y.push(self.y);
77        self.history_yaw.push(self.yaw);
78        self.history_v.push(self.v);
79    }
80}
81
82/// NMPC Simulator System for state prediction and adjoint calculation
83pub struct NMPCSimulatorSystem;
84
85impl NMPCSimulatorSystem {
86    pub fn new() -> Self {
87        NMPCSimulatorSystem
88    }
89
90    /// Calculate predicted states and adjoint (costate) variables
91    pub fn calc_predict_and_adjoint_state(
92        &self,
93        x: f64,
94        y: f64,
95        yaw: f64,
96        v: f64,
97        u_1s: &[f64],
98        u_2s: &[f64],
99        n: usize,
100        dt: f64,
101    ) -> (
102        Vec<f64>,
103        Vec<f64>,
104        Vec<f64>,
105        Vec<f64>,
106        Vec<f64>,
107        Vec<f64>,
108        Vec<f64>,
109        Vec<f64>,
110    ) {
111        // Forward prediction using state equations
112        let (x_s, y_s, yaw_s, v_s) = self.calc_predict_states(x, y, yaw, v, u_1s, u_2s, n, dt);
113
114        // Backward adjoint calculation using adjoint equations
115        let (lam_1s, lam_2s, lam_3s, lam_4s) =
116            self.calc_adjoint_states(&x_s, &y_s, &yaw_s, &v_s, u_2s, n, dt);
117
118        (x_s, y_s, yaw_s, v_s, lam_1s, lam_2s, lam_3s, lam_4s)
119    }
120
121    /// Forward state prediction using Euler integration
122    fn calc_predict_states(
123        &self,
124        x: f64,
125        y: f64,
126        yaw: f64,
127        v: f64,
128        u_1s: &[f64],
129        u_2s: &[f64],
130        n: usize,
131        dt: f64,
132    ) -> (Vec<f64>, Vec<f64>, Vec<f64>, Vec<f64>) {
133        let mut x_s = vec![x];
134        let mut y_s = vec![y];
135        let mut yaw_s = vec![yaw];
136        let mut v_s = vec![v];
137
138        for i in 0..n {
139            let (dx, dy, d_yaw, dv) = differential_model(v_s[i], yaw_s[i], u_1s[i], u_2s[i]);
140
141            x_s.push(x_s[i] + dt * dx);
142            y_s.push(y_s[i] + dt * dy);
143            yaw_s.push(yaw_s[i] + dt * d_yaw);
144            v_s.push(v_s[i] + dt * dv);
145        }
146
147        (x_s, y_s, yaw_s, v_s)
148    }
149
150    /// Backward adjoint state calculation (returns N elements)
151    fn calc_adjoint_states(
152        &self,
153        x_s: &[f64],
154        y_s: &[f64],
155        yaw_s: &[f64],
156        v_s: &[f64],
157        u_2s: &[f64],
158        n: usize,
159        dt: f64,
160    ) -> (Vec<f64>, Vec<f64>, Vec<f64>, Vec<f64>) {
161        // Initialize with terminal conditions
162        let mut lam_1s = vec![x_s[n]];
163        let mut lam_2s = vec![y_s[n]];
164        let mut lam_3s = vec![yaw_s[n]];
165        let mut lam_4s = vec![v_s[n]];
166
167        for i in (1..n).rev() {
168            let yaw = yaw_s[i];
169            let v = v_s[i];
170            let u_2 = u_2s[i];
171            let lam_1 = lam_1s[0];
172            let lam_2 = lam_2s[0];
173            let lam_3 = lam_3s[0];
174            let lam_4 = lam_4s[0];
175
176            let pre_lam_1 = lam_1;
177            let pre_lam_2 = lam_2;
178
179            let tmp1 = -lam_1 * yaw.sin() * v + lam_2 * yaw.cos() * v;
180            let pre_lam_3 = lam_3 + dt * tmp1;
181
182            let tmp2 = lam_1 * yaw.cos() + lam_2 * yaw.sin() + lam_3 * u_2.sin() / WB;
183            let pre_lam_4 = lam_4 + dt * tmp2;
184
185            lam_1s.insert(0, pre_lam_1);
186            lam_2s.insert(0, pre_lam_2);
187            lam_3s.insert(0, pre_lam_3);
188            lam_4s.insert(0, pre_lam_4);
189        }
190
191        (lam_1s, lam_2s, lam_3s, lam_4s)
192    }
193}
194
195impl Default for NMPCSimulatorSystem {
196    fn default() -> Self {
197        Self::new()
198    }
199}
200
201/// NMPC Controller using C-GMRES (Continuation GMRES) algorithm
202pub struct NMPCControllerCGMRES {
203    // Parameters
204    pub zeta: f64,
205    pub ht: f64,
206    pub tf: f64,
207    pub alpha: f64,
208    pub n: usize,
209    pub threshold: f64,
210    pub input_num: usize,
211    pub max_iteration: usize,
212
213    simulator: NMPCSimulatorSystem,
214
215    // Control inputs
216    pub u_1s: Vec<f64>,
217    pub u_2s: Vec<f64>,
218    pub dummy_u_1s: Vec<f64>,
219    pub dummy_u_2s: Vec<f64>,
220    pub raw_1s: Vec<f64>,
221    pub raw_2s: Vec<f64>,
222
223    // History
224    pub history_u_1: Vec<f64>,
225    pub history_u_2: Vec<f64>,
226    pub history_f: Vec<f64>,
227}
228
229impl NMPCControllerCGMRES {
230    pub fn new() -> Self {
231        let n = 10;
232        let input_num = 6;
233
234        NMPCControllerCGMRES {
235            zeta: 100.0,
236            ht: 0.01,
237            tf: 3.0,
238            alpha: 0.5,
239            n,
240            threshold: 0.001,
241            input_num,
242            max_iteration: input_num * n,
243
244            simulator: NMPCSimulatorSystem::new(),
245
246            u_1s: vec![1.0; n],
247            u_2s: vec![1.0; n],
248            dummy_u_1s: vec![1.0; n],
249            dummy_u_2s: vec![1.0; n],
250            raw_1s: vec![0.0; n],
251            raw_2s: vec![0.0; n],
252
253            history_u_1: Vec::new(),
254            history_u_2: Vec::new(),
255            history_f: Vec::new(),
256        }
257    }
258
259    /// Calculate control input using C-GMRES algorithm
260    pub fn calc_input(
261        &mut self,
262        x: f64,
263        y: f64,
264        yaw: f64,
265        v: f64,
266        time: f64,
267    ) -> (Vec<f64>, Vec<f64>) {
268        let dt = self.tf * (1.0 - (-self.alpha * time).exp()) / (self.n as f64);
269
270        let (x_1_dot, x_2_dot, x_3_dot, x_4_dot) =
271            differential_model(v, yaw, self.u_1s[0], self.u_2s[0]);
272
273        let dx_1 = x_1_dot * self.ht;
274        let dx_2 = x_2_dot * self.ht;
275        let dx_3 = x_3_dot * self.ht;
276        let dx_4 = x_4_dot * self.ht;
277
278        let (_, _, _, v_s_fxt, _, _, lam_3s_fxt, lam_4s_fxt) =
279            self.simulator.calc_predict_and_adjoint_state(
280                x + dx_1,
281                y + dx_2,
282                yaw + dx_3,
283                v + dx_4,
284                &self.u_1s,
285                &self.u_2s,
286                self.n,
287                dt,
288            );
289
290        let fxt = Self::calc_f(
291            &v_s_fxt,
292            &lam_3s_fxt,
293            &lam_4s_fxt,
294            &self.u_1s,
295            &self.u_2s,
296            &self.dummy_u_1s,
297            &self.dummy_u_2s,
298            &self.raw_1s,
299            &self.raw_2s,
300            self.n,
301        );
302
303        let (_, _, _, v_s_f, _, _, lam_3s_f, lam_4s_f) = self
304            .simulator
305            .calc_predict_and_adjoint_state(x, y, yaw, v, &self.u_1s, &self.u_2s, self.n, dt);
306
307        let f = Self::calc_f(
308            &v_s_f,
309            &lam_3s_f,
310            &lam_4s_f,
311            &self.u_1s,
312            &self.u_2s,
313            &self.dummy_u_1s,
314            &self.dummy_u_2s,
315            &self.raw_1s,
316            &self.raw_2s,
317            self.n,
318        );
319
320        let right: Vec<f64> = f
321            .iter()
322            .zip(fxt.iter())
323            .map(|(&f_i, &fxt_i)| -self.zeta * f_i - (fxt_i - f_i) / self.ht)
324            .collect();
325
326        let du_1: Vec<f64> = self.u_1s.iter().map(|&u| u * self.ht).collect();
327        let du_2: Vec<f64> = self.u_2s.iter().map(|&u| u * self.ht).collect();
328        let ddummy_u_1: Vec<f64> = self.dummy_u_1s.iter().map(|&u| u * self.ht).collect();
329        let ddummy_u_2: Vec<f64> = self.dummy_u_2s.iter().map(|&u| u * self.ht).collect();
330        let draw_1: Vec<f64> = self.raw_1s.iter().map(|&u| u * self.ht).collect();
331        let draw_2: Vec<f64> = self.raw_2s.iter().map(|&u| u * self.ht).collect();
332
333        let u_1s_pert: Vec<f64> = self.u_1s.iter().zip(&du_1).map(|(&u, &d)| u + d).collect();
334        let u_2s_pert: Vec<f64> = self.u_2s.iter().zip(&du_2).map(|(&u, &d)| u + d).collect();
335
336        let (_, _, _, v_s_fuxt, _, _, lam_3s_fuxt, lam_4s_fuxt) =
337            self.simulator.calc_predict_and_adjoint_state(
338                x + dx_1,
339                y + dx_2,
340                yaw + dx_3,
341                v + dx_4,
342                &u_1s_pert,
343                &u_2s_pert,
344                self.n,
345                dt,
346            );
347
348        let dummy_1s_pert: Vec<f64> = self
349            .dummy_u_1s
350            .iter()
351            .zip(&ddummy_u_1)
352            .map(|(&u, &d)| u + d)
353            .collect();
354        let dummy_2s_pert: Vec<f64> = self
355            .dummy_u_2s
356            .iter()
357            .zip(&ddummy_u_2)
358            .map(|(&u, &d)| u + d)
359            .collect();
360        let raw_1s_pert: Vec<f64> = self
361            .raw_1s
362            .iter()
363            .zip(&draw_1)
364            .map(|(&u, &d)| u + d)
365            .collect();
366        let raw_2s_pert: Vec<f64> = self
367            .raw_2s
368            .iter()
369            .zip(&draw_2)
370            .map(|(&u, &d)| u + d)
371            .collect();
372
373        let fuxt = Self::calc_f(
374            &v_s_fuxt,
375            &lam_3s_fuxt,
376            &lam_4s_fuxt,
377            &u_1s_pert,
378            &u_2s_pert,
379            &dummy_1s_pert,
380            &dummy_2s_pert,
381            &raw_1s_pert,
382            &raw_2s_pert,
383            self.n,
384        );
385
386        let left: Vec<f64> = fuxt
387            .iter()
388            .zip(&fxt)
389            .map(|(&fuxt_i, &fxt_i)| (fuxt_i - fxt_i) / self.ht)
390            .collect();
391
392        let r0: Vec<f64> = right.iter().zip(&left).map(|(&r, &l)| r - l).collect();
393        let r0_norm = vec_norm(&r0);
394
395        let m = self.max_iteration;
396        let mut vs = DMatrix::zeros(m, m + 1);
397        for i in 0..m {
398            vs[(i, 0)] = r0[i] / r0_norm;
399        }
400
401        let mut hs = DMatrix::zeros(m + 1, m + 1);
402        let mut ys_pre: Option<DVector<f64>> = None;
403
404        let mut du_1_new: Option<Vec<f64>> = None;
405        let mut du_2_new: Option<Vec<f64>> = None;
406        let mut ddummy_u_1_new: Option<Vec<f64>> = None;
407        let mut ddummy_u_2_new: Option<Vec<f64>> = None;
408        let mut draw_1_new: Option<Vec<f64>> = None;
409        let mut draw_2_new: Option<Vec<f64>> = None;
410
411        for i in 0..m {
412            let mut du_1_i = vec![0.0; self.n];
413            let mut du_2_i = vec![0.0; self.n];
414            let mut ddummy_u_1_i = vec![0.0; self.n];
415            let mut ddummy_u_2_i = vec![0.0; self.n];
416            let mut draw_1_i = vec![0.0; self.n];
417            let mut draw_2_i = vec![0.0; self.n];
418
419            for k in 0..self.n {
420                du_1_i[k] = vs[(k * self.input_num, i)] * self.ht;
421                du_2_i[k] = vs[(k * self.input_num + 1, i)] * self.ht;
422                ddummy_u_1_i[k] = vs[(k * self.input_num + 2, i)] * self.ht;
423                ddummy_u_2_i[k] = vs[(k * self.input_num + 3, i)] * self.ht;
424                draw_1_i[k] = vs[(k * self.input_num + 4, i)] * self.ht;
425                draw_2_i[k] = vs[(k * self.input_num + 5, i)] * self.ht;
426            }
427
428            let u_1s_i: Vec<f64> = self
429                .u_1s
430                .iter()
431                .zip(&du_1_i)
432                .map(|(&u, &d)| u + d)
433                .collect();
434            let u_2s_i: Vec<f64> = self
435                .u_2s
436                .iter()
437                .zip(&du_2_i)
438                .map(|(&u, &d)| u + d)
439                .collect();
440
441            let (_, _, _, v_s_i, _, _, lam_3s_i, lam_4s_i) =
442                self.simulator.calc_predict_and_adjoint_state(
443                    x + dx_1,
444                    y + dx_2,
445                    yaw + dx_3,
446                    v + dx_4,
447                    &u_1s_i,
448                    &u_2s_i,
449                    self.n,
450                    dt,
451                );
452
453            let dummy_1s_i: Vec<f64> = self
454                .dummy_u_1s
455                .iter()
456                .zip(&ddummy_u_1_i)
457                .map(|(&u, &d)| u + d)
458                .collect();
459            let dummy_2s_i: Vec<f64> = self
460                .dummy_u_2s
461                .iter()
462                .zip(&ddummy_u_2_i)
463                .map(|(&u, &d)| u + d)
464                .collect();
465            let raw_1s_i: Vec<f64> = self
466                .raw_1s
467                .iter()
468                .zip(&draw_1_i)
469                .map(|(&u, &d)| u + d)
470                .collect();
471            let raw_2s_i: Vec<f64> = self
472                .raw_2s
473                .iter()
474                .zip(&draw_2_i)
475                .map(|(&u, &d)| u + d)
476                .collect();
477
478            let fuxt_i = Self::calc_f(
479                &v_s_i,
480                &lam_3s_i,
481                &lam_4s_i,
482                &u_1s_i,
483                &u_2s_i,
484                &dummy_1s_i,
485                &dummy_2s_i,
486                &raw_1s_i,
487                &raw_2s_i,
488                self.n,
489            );
490
491            let av: Vec<f64> = fuxt_i
492                .iter()
493                .zip(&fxt)
494                .map(|(&fi, &fxti)| (fi - fxti) / self.ht)
495                .collect();
496
497            let mut sum_av = vec![0.0; m];
498            for j in 0..=i {
499                let mut dot = 0.0;
500                for k in 0..m {
501                    dot += av[k] * vs[(k, j)];
502                }
503                hs[(j, i)] = dot;
504                for k in 0..m {
505                    sum_av[k] += hs[(j, i)] * vs[(k, j)];
506                }
507            }
508
509            let v_est: Vec<f64> = av.iter().zip(&sum_av).map(|(&a, &s)| a - s).collect();
510            let v_est_norm = vec_norm(&v_est);
511            hs[(i + 1, i)] = v_est_norm;
512
513            if v_est_norm > 1e-15 {
514                for k in 0..m {
515                    vs[(k, i + 1)] = v_est[k] / v_est_norm;
516                }
517            }
518
519            if i == 0 {
520                ys_pre = Some(DVector::zeros(1));
521                continue;
522            }
523
524            let hs_sub = hs.view((0, 0), (i + 1, i)).clone_owned();
525
526            let mut e_scaled = DVector::zeros(i + 1);
527            e_scaled[0] = r0_norm;
528
529            let ys = match hs_sub.clone().svd(true, true).solve(&e_scaled, 1e-15) {
530                Ok(sol) => sol,
531                Err(_) => DVector::zeros(i),
532            };
533
534            let hs_ys = &hs_sub * &ys;
535            let residual = &e_scaled - &hs_ys;
536            let judge_norm = residual.norm();
537
538            if judge_norm < self.threshold || i == m - 1 {
539                let mut du_1_tmp = du_1_i.clone();
540                let mut du_2_tmp = du_2_i.clone();
541                let mut ddummy_u_1_tmp = ddummy_u_1_i.clone();
542                let mut ddummy_u_2_tmp = ddummy_u_2_i.clone();
543                let mut draw_1_tmp = draw_1_i.clone();
544                let mut draw_2_tmp = draw_2_i.clone();
545
546                if let Some(ref ys_p) = ys_pre {
547                    let update_len = i.saturating_sub(1);
548                    let update_len = update_len.min(ys_p.len());
549                    if update_len > 0 {
550                        let vs_sub = vs.view((0, 0), (m, update_len)).clone_owned();
551                        let ys_sub = ys_p.rows(0, update_len).clone_owned();
552                        let update_val = &vs_sub * &ys_sub;
553
554                        for k in 0..self.n {
555                            du_1_tmp[k] = du_1_i[k] + update_val[k * self.input_num];
556                            du_2_tmp[k] = du_2_i[k] + update_val[k * self.input_num + 1];
557                            ddummy_u_1_tmp[k] =
558                                ddummy_u_1_i[k] + update_val[k * self.input_num + 2];
559                            ddummy_u_2_tmp[k] =
560                                ddummy_u_2_i[k] + update_val[k * self.input_num + 3];
561                            draw_1_tmp[k] = draw_1_i[k] + update_val[k * self.input_num + 4];
562                            draw_2_tmp[k] = draw_2_i[k] + update_val[k * self.input_num + 5];
563                        }
564                    }
565                }
566
567                du_1_new = Some(du_1_tmp);
568                du_2_new = Some(du_2_tmp);
569                ddummy_u_1_new = Some(ddummy_u_1_tmp);
570                ddummy_u_2_new = Some(ddummy_u_2_tmp);
571                draw_1_new = Some(draw_1_tmp);
572                draw_2_new = Some(draw_2_tmp);
573                break;
574            }
575
576            ys_pre = Some(ys);
577        }
578
579        // Update inputs (only if converged)
580        if let (Some(du1), Some(du2), Some(dd1), Some(dd2), Some(dr1), Some(dr2)) = (
581            du_1_new,
582            du_2_new,
583            ddummy_u_1_new,
584            ddummy_u_2_new,
585            draw_1_new,
586            draw_2_new,
587        ) {
588            for i in 0..self.n {
589                self.u_1s[i] += du1[i] * self.ht;
590                self.u_2s[i] += du2[i] * self.ht;
591                self.dummy_u_1s[i] += dd1[i] * self.ht;
592                self.dummy_u_2s[i] += dd2[i] * self.ht;
593                self.raw_1s[i] += dr1[i] * self.ht;
594                self.raw_2s[i] += dr2[i] * self.ht;
595
596                self.u_1s[i] = self.u_1s[i].clamp(-U_A_MAX, U_A_MAX);
597                self.u_2s[i] = self.u_2s[i].clamp(-U_OMEGA_MAX, U_OMEGA_MAX);
598                self.dummy_u_1s[i] = self.dummy_u_1s[i].max(0.0);
599                self.dummy_u_2s[i] = self.dummy_u_2s[i].max(0.0);
600            }
601        }
602
603        // Calculate final F norm
604        let (_, _, _, v_s_final, _, _, lam_3s_final, lam_4s_final) = self
605            .simulator
606            .calc_predict_and_adjoint_state(x, y, yaw, v, &self.u_1s, &self.u_2s, self.n, dt);
607
608        let f_final = Self::calc_f(
609            &v_s_final,
610            &lam_3s_final,
611            &lam_4s_final,
612            &self.u_1s,
613            &self.u_2s,
614            &self.dummy_u_1s,
615            &self.dummy_u_2s,
616            &self.raw_1s,
617            &self.raw_2s,
618            self.n,
619        );
620
621        let f_norm = vec_norm(&f_final);
622
623        self.history_f.push(f_norm);
624        self.history_u_1.push(self.u_1s[0]);
625        self.history_u_2.push(self.u_2s[0]);
626
627        (self.u_1s.clone(), self.u_2s.clone())
628    }
629
630    /// Calculate optimality condition F
631    fn calc_f(
632        v_s: &[f64],
633        lam_3s: &[f64],
634        lam_4s: &[f64],
635        u_1s: &[f64],
636        u_2s: &[f64],
637        dummy_u_1s: &[f64],
638        dummy_u_2s: &[f64],
639        raw_1s: &[f64],
640        raw_2s: &[f64],
641        n: usize,
642    ) -> Vec<f64> {
643        let mut f = Vec::with_capacity(6 * n);
644
645        for i in 0..n {
646            let lam_3_idx = i.min(lam_3s.len() - 1);
647            let lam_4_idx = i.min(lam_4s.len() - 1);
648
649            f.push(u_1s[i] + lam_4s[lam_4_idx] + 2.0 * raw_1s[i] * u_1s[i]);
650
651            f.push(
652                u_2s[i]
653                    + lam_3s[lam_3_idx] * v_s[i] / WB * u_2s[i].cos().powi(2)
654                    + 2.0 * raw_2s[i] * u_2s[i],
655            );
656
657            f.push(-PHI_V + 2.0 * raw_1s[i] * dummy_u_1s[i]);
658            f.push(-PHI_OMEGA + 2.0 * raw_2s[i] * dummy_u_2s[i]);
659            f.push(u_1s[i].powi(2) + dummy_u_1s[i].powi(2) - U_A_MAX.powi(2));
660            f.push(u_2s[i].powi(2) + dummy_u_2s[i].powi(2) - U_OMEGA_MAX.powi(2));
661        }
662
663        f
664    }
665}
666
667impl Default for NMPCControllerCGMRES {
668    fn default() -> Self {
669        Self::new()
670    }
671}
672
673fn vec_norm(v: &[f64]) -> f64 {
674    v.iter().map(|x| x * x).sum::<f64>().sqrt()
675}
676
677pub fn goal_distance(plant: &TwoWheeledSystem) -> f64 {
678    plant.x.hypot(plant.y)
679}
680
681pub fn run_demo_simulation() -> (TwoWheeledSystem, NMPCControllerCGMRES) {
682    let dt = 0.1;
683    let iteration_time = 15.0;
684
685    let init_x: f64 = -1.5;
686    let init_y: f64 = -1.0;
687    let init_yaw = (-init_y).atan2(-init_x);
688    let init_v: f64 = 0.0;
689
690    let mut plant = TwoWheeledSystem::new(init_x, init_y, init_yaw, init_v);
691    let mut controller = NMPCControllerCGMRES::new();
692    let iteration_num = (iteration_time / dt) as usize;
693
694    for i in 1..iteration_num {
695        let time = (i as f64) * dt;
696
697        let (u_1s, u_2s) = controller.calc_input(plant.x, plant.y, plant.yaw, plant.v, time);
698
699        plant.update_state(u_1s[0], u_2s[0], dt);
700
701        if goal_distance(&plant) < 0.5 && plant.v.abs() < 0.5 {
702            break;
703        }
704
705        if controller.history_f.last().copied().unwrap_or(0.0) > 1.0e5 {
706            break;
707        }
708
709        if !plant.x.is_finite()
710            || !plant.y.is_finite()
711            || !plant.yaw.is_finite()
712            || !plant.v.is_finite()
713            || plant.x.abs().max(plant.y.abs()) > 12.0
714        {
715            break;
716        }
717    }
718
719    (plant, controller)
720}
721
722#[cfg(test)]
723mod tests {
724    use super::*;
725
726    #[test]
727    fn test_cgmres_demo_reaches_goal() {
728        let (plant, controller) = run_demo_simulation();
729        assert!(
730            goal_distance(&plant) < 0.5,
731            "final distance was {}",
732            goal_distance(&plant)
733        );
734        assert!(plant.v.abs() < 0.5, "final speed was {}", plant.v.abs());
735        assert!(
736            controller
737                .history_f
738                .last()
739                .copied()
740                .unwrap_or(f64::INFINITY)
741                .is_finite(),
742            "final optimality error was not finite"
743        );
744    }
745}