1use nalgebra::{Matrix2, Matrix2x4, Matrix4, Matrix4x2, Vector2, Vector4};
7use std::f64::consts::PI;
8
9const WHEEL_BASE: f64 = 2.5;
10const TERMINAL_WEIGHT_SCALE: f64 = 20.0;
11const REGULARIZATION: f64 = 1.0e-6;
12const MAX_SPEED: f64 = 8.0;
13const MIN_SPEED: f64 = -5.0;
14const MAX_ACCEL: f64 = 3.0;
15const MAX_STEER: f64 = 0.6;
16
17type FeedforwardTerms = Vec<Vector2<f64>>;
18type FeedbackGains = Vec<Matrix2x4<f64>>;
19type StageCostDerivatives = (
20 Vector4<f64>,
21 Vector2<f64>,
22 Matrix4<f64>,
23 Matrix2<f64>,
24 Matrix2x4<f64>,
25);
26type SecondOrderDyn = ([Matrix4<f64>; 4], [Matrix2x4<f64>; 4], [Matrix2<f64>; 4]);
27
28#[derive(Debug, Clone, Copy)]
30pub struct DDPConfig {
31 pub horizon: usize,
33 pub dt: f64,
35 pub max_iterations: usize,
37 pub tolerance: f64,
39 pub q_pos: f64,
41 pub q_theta: f64,
43 pub q_v: f64,
45 pub r_accel: f64,
47 pub r_steer: f64,
49}
50
51impl Default for DDPConfig {
52 fn default() -> Self {
53 Self {
54 horizon: 50,
55 dt: 0.1,
56 max_iterations: 50,
57 tolerance: 1.0e-6,
58 q_pos: 1.0,
59 q_theta: 0.2,
60 q_v: 0.2,
61 r_accel: 0.05,
62 r_steer: 0.05,
63 }
64 }
65}
66
67#[derive(Debug, Clone)]
69pub struct DDPResult {
70 pub states: Vec<Vector4<f64>>,
71 pub controls: Vec<Vector2<f64>>,
72 pub cost: f64,
73 pub iterations: usize,
74 pub converged: bool,
75}
76
77pub struct DDPPlanner {
79 config: DDPConfig,
80}
81
82impl DDPPlanner {
83 pub fn new(config: DDPConfig) -> Self {
84 Self { config }
85 }
86
87 pub fn plan(&self, start: Vector4<f64>, goal: Vector4<f64>) -> DDPResult {
88 let mut controls = self.initial_controls(start, goal);
89 let mut states = self.rollout(start, &controls);
90 let mut cost = self.total_cost(&states, &controls, goal);
91 let mut converged = false;
92 let mut iterations = 0usize;
93
94 for iter in 0..self.config.max_iterations {
95 iterations = iter + 1;
96 let Some((feedforward, feedback)) = self.backward_pass(&states, &controls, goal) else {
97 break;
98 };
99
100 let mut accepted = None;
101 for alpha in [1.0, 0.5, 0.25, 0.1, 0.05, 0.01] {
102 let (candidate_states, candidate_controls) =
103 self.forward_pass(start, &states, &controls, &feedforward, &feedback, alpha);
104 let candidate_cost = self.total_cost(&candidate_states, &candidate_controls, goal);
105 if candidate_cost < cost {
106 accepted = Some((candidate_states, candidate_controls, candidate_cost));
107 break;
108 }
109 }
110
111 let Some((candidate_states, candidate_controls, candidate_cost)) = accepted else {
112 break;
113 };
114
115 let delta = (cost - candidate_cost).abs();
116 states = candidate_states;
117 controls = candidate_controls;
118 cost = candidate_cost;
119 if delta < self.config.tolerance {
120 converged = true;
121 break;
122 }
123 }
124
125 if !converged {
126 let final_state = states.last().copied().unwrap_or(start);
127 let pos_error =
128 ((final_state[0] - goal[0]).powi(2) + (final_state[1] - goal[1]).powi(2)).sqrt();
129 let heading_error = normalize_angle(final_state[2] - goal[2]).abs();
130 converged = pos_error < 0.35 && heading_error < 0.2;
131 }
132
133 DDPResult {
134 states,
135 controls,
136 cost,
137 iterations,
138 converged,
139 }
140 }
141
142 fn initial_controls(&self, start: Vector4<f64>, goal: Vector4<f64>) -> Vec<Vector2<f64>> {
143 let distance = ((goal[0] - start[0]).powi(2) + (goal[1] - start[1]).powi(2)).sqrt();
144 let total_time = (self.config.horizon as f64 * self.config.dt).max(self.config.dt);
145 let desired_v = (distance / total_time).clamp(MIN_SPEED, MAX_SPEED);
146 let accel = ((desired_v - start[3]) / total_time).clamp(-MAX_ACCEL, MAX_ACCEL);
147 let steer = (normalize_angle(goal[2] - start[2]) / total_time).clamp(-MAX_STEER, MAX_STEER);
148 vec![Vector2::new(accel, steer); self.config.horizon]
149 }
150
151 fn rollout(&self, start: Vector4<f64>, controls: &[Vector2<f64>]) -> Vec<Vector4<f64>> {
152 let mut states = Vec::with_capacity(controls.len() + 1);
153 let mut state = start;
154 states.push(state);
155 for control in controls {
156 state = dynamics(&state, control, self.config.dt);
157 states.push(state);
158 }
159 states
160 }
161
162 fn total_cost(
163 &self,
164 states: &[Vector4<f64>],
165 controls: &[Vector2<f64>],
166 goal: Vector4<f64>,
167 ) -> f64 {
168 let stage_cost: f64 = states
169 .iter()
170 .take(controls.len())
171 .zip(controls.iter())
172 .map(|(state, control)| self.stage_cost(state, control, goal))
173 .sum();
174 stage_cost + self.terminal_cost(states.last().expect("trajectory"), goal)
175 }
176
177 fn stage_cost(&self, state: &Vector4<f64>, control: &Vector2<f64>, goal: Vector4<f64>) -> f64 {
178 let error = state_error(state, goal);
179 self.config.q_pos * (error[0].powi(2) + error[1].powi(2))
180 + self.config.q_theta * error[2].powi(2)
181 + self.config.q_v * error[3].powi(2)
182 + self.config.r_accel * control[0].powi(2)
183 + self.config.r_steer * control[1].powi(2)
184 }
185
186 fn terminal_cost(&self, state: &Vector4<f64>, goal: Vector4<f64>) -> f64 {
187 let error = state_error(state, goal);
188 TERMINAL_WEIGHT_SCALE
189 * (self.config.q_pos * (error[0].powi(2) + error[1].powi(2))
190 + self.config.q_theta * error[2].powi(2)
191 + self.config.q_v * error[3].powi(2))
192 }
193
194 fn stage_cost_derivatives(
195 &self,
196 state: &Vector4<f64>,
197 control: &Vector2<f64>,
198 goal: Vector4<f64>,
199 ) -> StageCostDerivatives {
200 let error = state_error(state, goal);
201 let l_x = Vector4::new(
202 2.0 * self.config.q_pos * error[0],
203 2.0 * self.config.q_pos * error[1],
204 2.0 * self.config.q_theta * error[2],
205 2.0 * self.config.q_v * error[3],
206 );
207 let l_u = Vector2::new(
208 2.0 * self.config.r_accel * control[0],
209 2.0 * self.config.r_steer * control[1],
210 );
211 let l_xx = Matrix4::from_diagonal(&Vector4::new(
212 2.0 * self.config.q_pos,
213 2.0 * self.config.q_pos,
214 2.0 * self.config.q_theta,
215 2.0 * self.config.q_v,
216 ));
217 let l_uu = Matrix2::from_diagonal(&Vector2::new(
218 2.0 * self.config.r_accel,
219 2.0 * self.config.r_steer,
220 ));
221 let l_ux = Matrix2x4::zeros();
222 (l_x, l_u, l_xx, l_uu, l_ux)
223 }
224
225 fn terminal_cost_derivatives(
226 &self,
227 state: &Vector4<f64>,
228 goal: Vector4<f64>,
229 ) -> (Vector4<f64>, Matrix4<f64>) {
230 let error = state_error(state, goal);
231 let v_x = Vector4::new(
232 2.0 * TERMINAL_WEIGHT_SCALE * self.config.q_pos * error[0],
233 2.0 * TERMINAL_WEIGHT_SCALE * self.config.q_pos * error[1],
234 2.0 * TERMINAL_WEIGHT_SCALE * self.config.q_theta * error[2],
235 2.0 * TERMINAL_WEIGHT_SCALE * self.config.q_v * error[3],
236 );
237 let v_xx = Matrix4::from_diagonal(&Vector4::new(
238 2.0 * TERMINAL_WEIGHT_SCALE * self.config.q_pos,
239 2.0 * TERMINAL_WEIGHT_SCALE * self.config.q_pos,
240 2.0 * TERMINAL_WEIGHT_SCALE * self.config.q_theta,
241 2.0 * TERMINAL_WEIGHT_SCALE * self.config.q_v,
242 ));
243 (v_x, v_xx)
244 }
245
246 fn backward_pass(
247 &self,
248 states: &[Vector4<f64>],
249 controls: &[Vector2<f64>],
250 goal: Vector4<f64>,
251 ) -> Option<(FeedforwardTerms, FeedbackGains)> {
252 let horizon = controls.len();
253 let mut feedforward = vec![Vector2::zeros(); horizon];
254 let mut feedback = vec![Matrix2x4::zeros(); horizon];
255
256 let (mut value_x, mut value_xx) =
257 self.terminal_cost_derivatives(states.last().expect("trajectory"), goal);
258
259 for t in (0..horizon).rev() {
260 let (a, b) = linearize_dynamics(&states[t], &controls[t], self.config.dt);
261 let (f_xx, f_ux, f_uu) =
262 second_order_dynamics(&states[t], &controls[t], self.config.dt);
263 let (l_x, l_u, l_xx, l_uu, l_ux) =
264 self.stage_cost_derivatives(&states[t], &controls[t], goal);
265
266 let mut q_x = l_x + a.transpose() * value_x;
267 let q_u = l_u + b.transpose() * value_x;
268 let mut q_xx = l_xx + a.transpose() * value_xx * a;
269 let mut q_ux = l_ux + b.transpose() * value_xx * a;
270 let mut q_uu = l_uu + b.transpose() * value_xx * b;
271
272 for i in 0..4 {
273 q_xx += value_x[i] * f_xx[i];
274 q_ux += value_x[i] * f_ux[i];
275 q_uu += value_x[i] * f_uu[i];
276 }
277
278 q_xx = symmetrize_4x4(q_xx);
279 q_uu = symmetrize_2x2(q_uu) + REGULARIZATION * Matrix2::identity();
280 let q_uu_inv = q_uu.try_inverse()?;
281
282 let k = -q_uu_inv * q_u;
283 let k_feedback = -q_uu_inv * q_ux;
284 feedforward[t] = k;
285 feedback[t] = k_feedback;
286
287 q_x += k_feedback.transpose() * q_uu * k
288 + k_feedback.transpose() * q_u
289 + q_ux.transpose() * k;
290 value_x = q_x;
291 value_xx = q_xx
292 + k_feedback.transpose() * q_uu * k_feedback
293 + k_feedback.transpose() * q_ux
294 + q_ux.transpose() * k_feedback;
295 value_xx = symmetrize_4x4(value_xx);
296 }
297
298 Some((feedforward, feedback))
299 }
300
301 fn forward_pass(
302 &self,
303 start: Vector4<f64>,
304 nominal_states: &[Vector4<f64>],
305 nominal_controls: &[Vector2<f64>],
306 feedforward: &[Vector2<f64>],
307 feedback: &[Matrix2x4<f64>],
308 alpha: f64,
309 ) -> (Vec<Vector4<f64>>, Vec<Vector2<f64>>) {
310 let mut states = Vec::with_capacity(nominal_controls.len() + 1);
311 let mut controls = Vec::with_capacity(nominal_controls.len());
312 states.push(start);
313
314 for t in 0..nominal_controls.len() {
315 let delta_x = state_error(&states[t], nominal_states[t]);
316 let delta_u = alpha * feedforward[t] + feedback[t] * delta_x;
317 let mut control = nominal_controls[t] + delta_u;
318 control[0] = control[0].clamp(-MAX_ACCEL, MAX_ACCEL);
319 control[1] = control[1].clamp(-MAX_STEER, MAX_STEER);
320 controls.push(control);
321 states.push(dynamics(&states[t], &control, self.config.dt));
322 }
323 (states, controls)
324 }
325}
326
327fn state_error(state: &Vector4<f64>, goal: Vector4<f64>) -> Vector4<f64> {
328 Vector4::new(
329 state[0] - goal[0],
330 state[1] - goal[1],
331 normalize_angle(state[2] - goal[2]),
332 state[3] - goal[3],
333 )
334}
335
336fn normalize_angle(mut angle: f64) -> f64 {
337 while angle > PI {
338 angle -= 2.0 * PI;
339 }
340 while angle < -PI {
341 angle += 2.0 * PI;
342 }
343 angle
344}
345
346fn dynamics(state: &Vector4<f64>, control: &Vector2<f64>, dt: f64) -> Vector4<f64> {
347 let accel = control[0].clamp(-MAX_ACCEL, MAX_ACCEL);
348 let steer = control[1].clamp(-MAX_STEER, MAX_STEER);
349 let v = (state[3] + accel * dt).clamp(MIN_SPEED, MAX_SPEED);
350 let x = state[0] + state[3] * state[2].cos() * dt;
351 let y = state[1] + state[3] * state[2].sin() * dt;
352 let theta = normalize_angle(state[2] + (state[3] / WHEEL_BASE) * steer.tan() * dt);
353 Vector4::new(x, y, theta, v)
354}
355
356fn linearize_dynamics(
357 state: &Vector4<f64>,
358 control: &Vector2<f64>,
359 dt: f64,
360) -> (Matrix4<f64>, Matrix4x2<f64>) {
361 let theta = state[2];
362 let v = state[3];
363 let steer = control[1].clamp(-MAX_STEER, MAX_STEER);
364 let sec2 = 1.0 / steer.cos().powi(2);
365
366 let a = Matrix4::new(
367 1.0,
368 0.0,
369 -v * theta.sin() * dt,
370 theta.cos() * dt,
371 0.0,
372 1.0,
373 v * theta.cos() * dt,
374 theta.sin() * dt,
375 0.0,
376 0.0,
377 1.0,
378 (steer.tan() / WHEEL_BASE) * dt,
379 0.0,
380 0.0,
381 0.0,
382 1.0,
383 );
384
385 let b = Matrix4x2::new(
386 0.0,
387 0.0,
388 0.0,
389 0.0,
390 0.0,
391 (v / WHEEL_BASE) * sec2 * dt,
392 dt,
393 0.0,
394 );
395
396 (a, b)
397}
398
399fn second_order_dynamics(state: &Vector4<f64>, control: &Vector2<f64>, dt: f64) -> SecondOrderDyn {
400 let theta = state[2];
401 let v = state[3];
402 let steer = control[1].clamp(-MAX_STEER, MAX_STEER);
403 let sec2 = 1.0 / steer.cos().powi(2);
404 let tan = steer.tan();
405
406 let mut f_xx = [
407 Matrix4::zeros(),
408 Matrix4::zeros(),
409 Matrix4::zeros(),
410 Matrix4::zeros(),
411 ];
412 let mut f_ux = [
413 Matrix2x4::zeros(),
414 Matrix2x4::zeros(),
415 Matrix2x4::zeros(),
416 Matrix2x4::zeros(),
417 ];
418 let mut f_uu = [
419 Matrix2::zeros(),
420 Matrix2::zeros(),
421 Matrix2::zeros(),
422 Matrix2::zeros(),
423 ];
424
425 f_xx[0][(2, 2)] = -v * theta.cos() * dt;
427 f_xx[0][(2, 3)] = -theta.sin() * dt;
428 f_xx[0][(3, 2)] = -theta.sin() * dt;
429
430 f_xx[1][(2, 2)] = -v * theta.sin() * dt;
432 f_xx[1][(2, 3)] = theta.cos() * dt;
433 f_xx[1][(3, 2)] = theta.cos() * dt;
434
435 f_ux[2][(1, 3)] = (sec2 / WHEEL_BASE) * dt;
437 f_uu[2][(1, 1)] = (v / WHEEL_BASE) * 2.0 * sec2 * tan * dt;
438
439 (f_xx, f_ux, f_uu)
440}
441
442fn symmetrize_2x2(matrix: Matrix2<f64>) -> Matrix2<f64> {
443 0.5 * (matrix + matrix.transpose())
444}
445
446fn symmetrize_4x4(matrix: Matrix4<f64>) -> Matrix4<f64> {
447 0.5 * (matrix + matrix.transpose())
448}
449
450#[cfg(test)]
451mod tests {
452 use super::*;
453
454 #[test]
455 fn test_ddp_config_defaults() {
456 let config = DDPConfig::default();
457 assert_eq!(config.horizon, 50);
458 assert_eq!(config.dt, 0.1);
459 assert_eq!(config.max_iterations, 50);
460 assert_eq!(config.tolerance, 1.0e-6);
461 }
462
463 #[test]
464 fn test_ddp_cost_decreases_from_initial_rollout() {
465 let planner = DDPPlanner::new(DDPConfig::default());
466 let start = Vector4::new(0.0, 0.0, 0.0, 0.0);
467 let goal = Vector4::new(5.0, 1.0, 0.2, 1.0);
468
469 let initial_controls = planner.initial_controls(start, goal);
470 let initial_states = planner.rollout(start, &initial_controls);
471 let initial_cost = planner.total_cost(&initial_states, &initial_controls, goal);
472
473 let result = planner.plan(start, goal);
474 assert!(result.cost < initial_cost);
475 }
476
477 #[test]
478 fn test_ddp_reaches_goal_with_turn() {
479 let planner = DDPPlanner::new(DDPConfig {
480 horizon: 100,
481 ..DDPConfig::default()
482 });
483 let start = Vector4::new(0.0, 0.0, 0.0, 0.0);
484 let goal = Vector4::new(2.0, 0.5, 0.2, 0.8);
485 let result = planner.plan(start, goal);
486 let final_state = result.states.last().copied().expect("trajectory");
487 let pos_error =
488 ((final_state[0] - goal[0]).powi(2) + (final_state[1] - goal[1]).powi(2)).sqrt();
489
490 assert!(pos_error < 0.6);
491 assert!(normalize_angle(final_state[2] - goal[2]).abs() < 0.5);
492 }
493
494 #[test]
495 fn test_ddp_converges_on_straight_line_case() {
496 let planner = DDPPlanner::new(DDPConfig::default());
497 let start = Vector4::new(0.0, 0.0, 0.0, 0.0);
498 let goal = Vector4::new(5.0, 0.0, 0.0, 1.0);
499 let result = planner.plan(start, goal);
500 let final_state = result.states.last().copied().expect("trajectory");
501
502 assert!(result.converged);
503 assert!((final_state[0] - goal[0]).abs() < 0.4);
504 assert!(final_state[1].abs() < 0.2);
505 }
506}