use crate::diag::PathError;
use crate::path::OutOfRangeMode;
use nalgebra::{DMatrix, DMatrixView, DVector};
use rayon::prelude::*;
#[derive(Clone, Copy, Debug)]
pub enum Parametrization {
Uniform,
}
pub struct SplineConfig {
pub order: usize,
pub parametrization: Parametrization,
pub s_min: f64,
pub s_max: f64,
pub out_of_range_mode: OutOfRangeMode,
pub start_state: Option<DMatrix<f64>>,
pub end_state: Option<DMatrix<f64>>,
}
impl Default for SplineConfig {
fn default() -> Self {
Self {
order: 5,
parametrization: Parametrization::Uniform,
s_min: 0.0,
s_max: 1.0,
out_of_range_mode: OutOfRangeMode::Error,
start_state: None,
end_state: None,
}
}
}
pub struct SplinePath {
pub order: usize,
pub s_min: f64,
pub s_max: f64,
pub out_of_range_mode: OutOfRangeMode,
n_coef: usize,
n_segments: usize,
inv_ds: f64,
coeffs: Vec<f64>,
}
impl SplinePath {
pub fn from_waypoints(waypoints: &DMatrix<f64>, cfg: &SplineConfig) -> Result<Self, PathError> {
Self::from_waypoints_view(waypoints.as_view(), cfg)
}
pub fn from_waypoints_view(
waypoints: DMatrixView<'_, f64>,
cfg: &SplineConfig,
) -> Result<Self, PathError> {
let order = cfg.order;
if order < 3 || order.is_multiple_of(2) {
return Err(PathError::InvalidOrder { order });
}
let dim = waypoints.nrows();
let n_points = waypoints.ncols();
if n_points < 2 {
return Err(PathError::NotEnoughWaypoints { n: n_points });
}
let n_segments = n_points - 1;
let n_coef = order + 1;
let m = (order - 1) / 2;
let start = extract_boundary(cfg.start_state.as_ref(), dim, m)?;
let end = extract_boundary(cfg.end_state.as_ref(), dim, m)?;
let coeffs = solve_general_thomas(waypoints, order, m, &start, &end)?;
let range = cfg.s_max - cfg.s_min;
let inv_ds = n_segments as f64 / range;
Ok(Self {
order,
s_min: cfg.s_min,
s_max: cfg.s_max,
out_of_range_mode: cfg.out_of_range_mode,
n_coef,
n_segments,
inv_ds,
coeffs,
})
}
#[inline(always)]
fn segment_tau(&self, s: f64) -> (usize, f64) {
let scaled = (s - self.s_min) * self.inv_ds;
if scaled <= 0.0 {
return (0, 0.0);
}
let max_scaled = self.n_segments as f64;
if scaled >= max_scaled {
return (self.n_segments - 1, 1.0);
}
let seg = scaled.floor() as usize;
(seg, scaled - seg as f64)
}
#[inline(always)]
fn eval_poly<const ORDER: u8>(coeffs: &[f64], t: f64) -> (f64, f64, f64, f64) {
if coeffs.len() == 6 {
let [a0, a1, a2, a3, a4, a5] = [
coeffs[0], coeffs[1], coeffs[2], coeffs[3], coeffs[4], coeffs[5],
];
let q = a5
.mul_add(t, a4)
.mul_add(t, a3)
.mul_add(t, a2)
.mul_add(t, a1)
.mul_add(t, a0);
if ORDER == 0 {
return (q, 0.0, 0.0, 0.0);
}
let dq = (5.0 * a5)
.mul_add(t, 4.0 * a4)
.mul_add(t, 3.0 * a3)
.mul_add(t, 2.0 * a2)
.mul_add(t, a1);
let ddq = (20.0 * a5)
.mul_add(t, 12.0 * a4)
.mul_add(t, 6.0 * a3)
.mul_add(t, 2.0 * a2);
if ORDER == 2 {
return (q, dq, ddq, 0.0);
}
let dddq = (60.0 * a5).mul_add(t, 24.0 * a4).mul_add(t, 6.0 * a3);
return (q, dq, ddq, dddq);
}
if coeffs.len() == 4 {
let [a0, a1, a2, a3] = [coeffs[0], coeffs[1], coeffs[2], coeffs[3]];
let q = a3.mul_add(t, a2).mul_add(t, a1).mul_add(t, a0);
if ORDER == 0 {
return (q, 0.0, 0.0, 0.0);
}
let dq = (3.0 * a3).mul_add(t, 2.0 * a2).mul_add(t, a1);
let ddq = (6.0 * a3).mul_add(t, 2.0 * a2);
if ORDER == 2 {
return (q, dq, ddq, 0.0);
}
return (q, dq, ddq, 6.0 * a3);
}
let n = coeffs.len() - 1;
let mut b0 = coeffs[n];
let mut b1 = 0.0_f64;
let mut b2 = 0.0_f64;
let mut b3 = 0.0_f64;
for k in (0..n).rev() {
if ORDER >= 3 {
b3 = b3 * t + 3.0 * b2;
}
if ORDER >= 2 {
b2 = b2 * t + 2.0 * b1;
}
if ORDER >= 1 {
b1 = b1 * t + b0;
}
b0 = b0 * t + coeffs[k];
}
(b0, b1, b2, b3)
}
#[inline(always)]
pub fn eval_at<const ORDER: u8>(
&self,
s: f64,
_dim: usize,
q_col: &mut [f64],
dq_col: &mut [f64],
ddq_col: &mut [f64],
dddq_col: &mut [f64],
) {
let (seg, tau) = self.segment_tau(s);
let seg_offset = seg * self.n_coef;
let inv_ds2 = self.inv_ds * self.inv_ds;
let inv_ds3 = inv_ds2 * self.inv_ds;
let stride = self.n_segments * self.n_coef;
for (i, q_v) in q_col.iter_mut().enumerate() {
let start = i * stride + seg_offset;
let (q_val, dq_val, ddq_val, dddq_val) =
Self::eval_poly::<ORDER>(&self.coeffs[start..start + self.n_coef], tau);
*q_v = q_val;
if ORDER >= 2 {
dq_col[i] = dq_val * self.inv_ds;
ddq_col[i] = ddq_val * inv_ds2;
}
if ORDER >= 3 {
dddq_col[i] = dddq_val * inv_ds3;
}
}
}
}
fn extract_boundary(
state: Option<&DMatrix<f64>>,
dim: usize,
m: usize,
) -> Result<Vec<f64>, PathError> {
match state {
None => Ok(vec![0.0; dim * m]),
Some(mat) => {
if mat.nrows() != dim || mat.ncols() != m {
return Err(PathError::DimensionMismatch);
}
let mut out = vec![0.0; dim * m];
for (d, row) in out.chunks_mut(m).enumerate() {
for (r, val) in row.iter_mut().enumerate() {
*val = mat[(d, r)];
}
}
Ok(out)
}
}
}
struct BlockMatrices {
m: usize,
a: DMatrix<f64>,
b: DMatrix<f64>,
c: DMatrix<f64>,
r: DMatrix<f64>,
h_coeff: DMatrix<f64>,
}
impl BlockMatrices {
fn for_order(m: usize) -> Self {
match m {
1 => Self::m1(),
2 => Self::m2(),
3 => Self::m3(),
_ => Self::general(m),
}
}
fn m1() -> Self {
BlockMatrices {
m: 1,
a: DMatrix::from_row_slice(1, 1, &[1.0]),
b: DMatrix::from_row_slice(1, 1, &[4.0]),
c: DMatrix::from_row_slice(1, 1, &[1.0]),
r: DMatrix::from_row_slice(1, 2, &[3.0, 3.0]),
h_coeff: DMatrix::from_row_slice(2, 3, &[3.0, -2.0, -1.0, -2.0, 1.0, 1.0]),
}
}
fn m2() -> Self {
BlockMatrices {
m: 2,
a: DMatrix::from_row_slice(2, 2, &[-4.0, -1.0, -7.0, -2.0]),
b: DMatrix::from_row_slice(2, 2, &[0.0, 6.0, -16.0, 0.0]),
c: DMatrix::from_row_slice(2, 2, &[4.0, -1.0, -7.0, 2.0]),
r: DMatrix::from_row_slice(2, 2, &[-10.0, 10.0, -15.0, -15.0]),
h_coeff: DMatrix::from_row_slice(
3,
5,
&[
10.0, -6.0, -3.0, -4.0, 1.0, -15.0, 8.0, 3.0, 7.0, -2.0, 6.0, -3.0, -1.0, -3.0,
1.0,
],
),
}
}
fn m3() -> Self {
BlockMatrices {
m: 3,
a: DMatrix::from_row_slice(3, 3, &[15.0, 5.0, 1.0, 39.0, 14.0, 3.0, 34.0, 13.0, 3.0]),
b: DMatrix::from_row_slice(3, 3, &[40.0, 0.0, 8.0, 0.0, -40.0, 0.0, 72.0, 0.0, 8.0]),
c: DMatrix::from_row_slice(
3,
3,
&[15.0, -5.0, 1.0, -39.0, 14.0, -3.0, 34.0, -13.0, 3.0],
),
r: DMatrix::from_row_slice(3, 2, &[35.0, 35.0, 84.0, -84.0, 70.0, 70.0]),
h_coeff: DMatrix::from_row_slice(
4,
7,
&[
35.0, -20.0, -10.0, -4.0, -15.0, 5.0, -1.0, -84.0, 45.0, 20.0, 6.0, 39.0,
-14.0, 3.0, 70.0, -36.0, -15.0, -4.0, -34.0, 13.0, -3.0, -20.0, 10.0, 4.0, 1.0,
10.0, -4.0, 1.0,
],
),
}
}
fn general(m: usize) -> Self {
let p = 2 * m + 1;
let n = m + 1;
let binom = |nn: usize, k: usize| -> f64 {
if k > nn {
return 0.0;
}
(0..k).fold(1.0f64, |acc, i| acc * (nn - i) as f64 / (i + 1) as f64)
};
let mut aug = DMatrix::<f64>::zeros(n, 2 * n);
for r in 0..n {
for (col_idx, k) in (m + 1..=p).enumerate() {
aug[(r, col_idx)] = binom(k, r);
}
aug[(r, n + r)] = 1.0;
}
for col in 0..n {
let pivot_row = (col..n)
.max_by(|&a, &b| {
aug[(a, col)]
.abs()
.partial_cmp(&aug[(b, col)].abs())
.unwrap()
})
.unwrap();
if pivot_row != col {
aug.swap_rows(col, pivot_row);
}
let piv_inv = 1.0 / aug[(col, col)];
for j in 0..2 * n {
aug[(col, j)] *= piv_inv;
}
for r in 0..n {
if r == col {
continue;
}
let factor = aug[(r, col)];
if factor == 0.0 {
continue;
}
for j in 0..2 * n {
let sub = factor * aug[(col, j)];
aug[(r, j)] -= sub;
}
}
}
let minv = aug.columns(n, n).into_owned();
let basis = 1 + 2 * m;
let mut rhs_mat = DMatrix::<f64>::zeros(n, basis);
rhs_mat[(0, 0)] = 1.0;
for b in 1..=m {
rhs_mat[(0, b)] = -1.0;
}
for r in 1..n {
rhs_mat[(r, m + r)] = 1.0;
for k in r..n {
rhs_mat[(r, k)] -= binom(k, r);
}
}
let h_coeff = &minv * &rhs_mat;
let bs = 2 + 3 * m;
let get_coeff_vec = |k: usize, is_left: bool| -> DVector<f64> {
let mut v = DVector::<f64>::zeros(bs);
if k == 0 {
return v;
}
if k <= m {
let offset = if is_left { 2 } else { 2 + m };
v[offset + k - 1] = 1.0;
return v;
}
let ic = k - (m + 1);
for (b, &cv) in h_coeff.row(ic).iter().enumerate() {
if cv == 0.0 {
continue;
}
if b == 0 {
v[if is_left { 0 } else { 1 }] += cv;
} else if b <= m {
let offset = if is_left { 2 } else { 2 + m };
v[offset + b - 1] += cv;
} else {
let offset = if is_left { 2 + m } else { 2 + 2 * m };
v[offset + b - m - 1] += cv;
}
}
v
};
let mut a_mat = DMatrix::<f64>::zeros(m, m);
let mut b_mat = DMatrix::<f64>::zeros(m, m);
let mut c_mat = DMatrix::<f64>::zeros(m, m);
let mut r_mat = DMatrix::<f64>::zeros(m, 2);
for (eq_idx, r) in (m + 1..=2 * m).enumerate() {
let lhs = (r..=p).fold(DVector::<f64>::zeros(bs), |acc, k| {
acc + get_coeff_vec(k, true) * binom(k, r)
});
let rhs_v = get_coeff_vec(r, false);
let eq = lhs - rhs_v;
for j in 0..m {
a_mat[(eq_idx, j)] = eq[2 + j];
b_mat[(eq_idx, j)] = eq[2 + m + j];
c_mat[(eq_idx, j)] = eq[2 + 2 * m + j];
}
r_mat[(eq_idx, 0)] = -eq[0];
r_mat[(eq_idx, 1)] = -eq[1];
}
BlockMatrices {
m,
a: a_mat,
b: b_mat,
c: c_mat,
r: r_mat,
h_coeff,
}
}
#[inline(always)]
fn compute_rhs(&self, dy_prev: f64, dy_next: f64, out: &mut [f64]) {
let r0 = self.r.column(0);
let r1 = self.r.column(1);
for (o, (&c0, &c1)) in out.iter_mut().zip(r0.iter().zip(r1.iter())) {
*o = c0 * dy_prev + c1 * dy_next;
}
}
#[inline(always)]
fn upper_coeffs(&self, dy: f64, u0: &[f64], u1: &[f64], out: &mut [f64]) {
let m = self.m;
for (o, row) in out.iter_mut().zip(self.h_coeff.row_iter()) {
let dot_u0: f64 = row
.iter()
.skip(1)
.take(m)
.zip(u0.iter())
.map(|(&c, &v)| c * v)
.sum();
let dot_u1: f64 = row
.iter()
.skip(1 + m)
.zip(u1.iter())
.map(|(&c, &v)| c * v)
.sum();
*o = row[0] * dy + dot_u0 + dot_u1;
}
}
}
fn solve_general_thomas(
waypoints: DMatrixView<'_, f64>,
order: usize,
m: usize,
start_bd: &[f64],
end_bd: &[f64],
) -> Result<Vec<f64>, PathError> {
let dim = waypoints.nrows();
let n = waypoints.ncols(); let ns = n - 1; let n_coef = order + 1;
let bm = BlockMatrices::for_order(m);
let chunk_len = ns * n_coef;
let mut coeffs = vec![0.0f64; dim * chunk_len];
let waypoints = &waypoints;
coeffs.par_chunks_mut(chunk_len).enumerate().try_for_each(
|(d, chunk)| -> Result<(), PathError> {
let u_start = &start_bd[d * m..(d + 1) * m];
let u_end = &end_bd[d * m..(d + 1) * m];
thomas_row(
waypoints,
(d, n, ns, m, n_coef),
&bm,
(u_start, u_end),
chunk,
)
},
)?;
Ok(coeffs)
}
fn thomas_row(
waypoints: &DMatrixView<'_, f64>,
(d, n, ns, m, n_coef): (usize, usize, usize, usize, usize),
bm: &BlockMatrices,
(u_start, u_end): (&[f64], &[f64]),
chunk: &mut [f64],
) -> Result<(), PathError> {
let row_d = waypoints.row(d);
let dy: Vec<f64> = row_d
.iter()
.zip(row_d.iter().skip(1))
.map(|(&a, &b)| b - a)
.collect();
let mut rhs = vec![0.0f64; m];
let mut upper_buf = vec![0.0f64; m + 1];
if n == 2 {
bm.upper_coeffs(dy[0], u_start, u_end, &mut upper_buf);
write_segment_coeffs(chunk, 0, n_coef, m, waypoints[(d, 0)], u_start, &upper_buf);
return Ok(());
}
let ni = n - 2;
let mut b_inv_c_list: Vec<DMatrix<f64>> = Vec::with_capacity(ni);
let mut b_inv_r_list: Vec<DVector<f64>> = Vec::with_capacity(ni);
let u_start_vec = DVector::from_column_slice(u_start);
let u_end_vec = DVector::from_column_slice(u_end);
for i in 0..ni {
let node = i + 1;
bm.compute_rhs(dy[node - 1], dy[node], &mut rhs);
let mut rhs_curr = DVector::from_column_slice(&rhs);
if i == 0 {
rhs_curr -= &bm.a * &u_start_vec;
}
if i == ni - 1 {
rhs_curr -= &bm.c * &u_end_vec;
}
let b_cur = if i == 0 {
bm.b.clone()
} else {
&bm.b - &bm.a * &b_inv_c_list[i - 1]
};
if i > 0 {
rhs_curr -= &bm.a * &b_inv_r_list[i - 1];
}
let lu = b_cur.lu();
b_inv_c_list.push(lu.solve(&bm.c).ok_or(PathError::SingularSystem)?);
b_inv_r_list.push(lu.solve(&rhs_curr).ok_or(PathError::SingularSystem)?);
}
let mut u_list: Vec<DVector<f64>> = vec![DVector::zeros(m); ni];
u_list[ni - 1] = b_inv_r_list[ni - 1].clone();
for i in (0..ni - 1).rev() {
u_list[i] = &b_inv_r_list[i] - &b_inv_c_list[i] * &u_list[i + 1];
}
let u_at = |k: usize| -> &[f64] {
if k == 0 {
u_start
} else if k >= n - 1 {
u_end
} else {
u_list[k - 1].as_slice()
}
};
for seg in 0..ns {
let ul = u_at(seg);
let ur = u_at(seg + 1);
bm.upper_coeffs(dy[seg], ul, ur, &mut upper_buf);
write_segment_coeffs(chunk, seg, n_coef, m, waypoints[(d, seg)], ul, &upper_buf);
}
Ok(())
}
#[inline(always)]
fn write_segment_coeffs(
chunk: &mut [f64],
seg: usize,
n_coef: usize,
m: usize,
y_left: f64,
u_left: &[f64],
upper: &[f64],
) {
let base = seg * n_coef;
chunk[base] = y_left;
chunk[base + 1..base + 1 + m].copy_from_slice(u_left);
chunk[base + m + 1..base + m + 1 + (m + 1)].copy_from_slice(upper);
}