1use nalgebra::DMatrix;
4
5use crate::astro::math::portable;
6
7use super::ekf::{
8 apply_closed_loop_navigation_error, apply_closed_loop_scale_error,
9 normalized_innovation_squared, EkfCorrection, EkfCorrectionReport, InnovationGate,
10 InnovationGateReport,
11};
12use super::state::{
13 covariance_eigenvalue_tolerance, dmatrix_from_rows, invalid_input, matmul, matrix_sub,
14 reproject_covariance_psd, solve_spd, symmetrize_in_place, transpose,
15 validate_covariance_matrix, validate_finite_slice, validate_matrix_cols, validate_nonnegative,
16 validate_positive, FusionError, InsFilterState,
17};
18
19#[derive(Debug, Clone, Copy, PartialEq)]
24pub struct UnscentedTransformOptions {
25 pub alpha: f64,
27 pub beta: f64,
29 pub kappa: f64,
31}
32
33impl Default for UnscentedTransformOptions {
34 fn default() -> Self {
35 Self {
36 alpha: 0.5,
37 beta: 2.0,
38 kappa: 0.0,
39 }
40 }
41}
42
43impl UnscentedTransformOptions {
44 pub fn validate_for_dimension(&self, dimension: usize) -> Result<(), FusionError> {
46 if dimension == 0 {
47 return Err(invalid_input("dimension", "must be positive"));
48 }
49 validate_positive(self.alpha, "ukf_alpha")?;
50 validate_nonnegative(self.beta, "ukf_beta")?;
51 validate_finite_slice(&[self.kappa], "ukf_kappa")?;
52 let scale = self.scale(dimension);
53 if scale.is_finite() && scale > 0.0 {
54 Ok(())
55 } else {
56 Err(invalid_input("ukf_scale", "must be positive"))
57 }
58 }
59
60 fn lambda(self, dimension: usize) -> f64 {
61 self.alpha * self.alpha * (dimension as f64 + self.kappa) - dimension as f64
62 }
63
64 fn scale(self, dimension: usize) -> f64 {
65 dimension as f64 + self.lambda(dimension)
66 }
67}
68
69#[derive(Debug, Clone, Copy, PartialEq, Default)]
71pub struct UkfUpdateOptions {
72 pub transform: UnscentedTransformOptions,
74 pub innovation_gate: Option<InnovationGate>,
76}
77
78impl UkfUpdateOptions {
79 pub fn validate_for_dimension(&self, dimension: usize) -> Result<(), FusionError> {
81 self.transform.validate_for_dimension(dimension)?;
82 if let Some(gate) = self.innovation_gate {
83 gate.validate()?;
84 }
85 Ok(())
86 }
87}
88
89pub fn ukf_correct_closed_loop(
95 state: &mut InsFilterState,
96 correction: &EkfCorrection,
97 options: UkfUpdateOptions,
98) -> Result<EkfCorrectionReport, FusionError> {
99 state.validate()?;
100 correction.validate_for_dimension(state.dimension())?;
101 options.validate_for_dimension(state.dimension())?;
102
103 let report = ukf_measurement_update(
104 &state.covariance,
105 &correction.innovation,
106 &correction.measurement_covariance,
107 options,
108 |sigma| super::state::matvec(&correction.design, sigma),
109 )?;
110 if !report.applied {
111 return Ok(report.into_public_report());
112 }
113
114 apply_closed_loop_navigation_error(&mut state.nominal, &report.dx)?;
115 apply_closed_loop_scale_error(state, &report.dx);
116 state.covariance = report.posterior_covariance.clone();
117 state.reset_error_state();
118 state.validate()?;
119 Ok(report.into_public_report())
120}
121
122#[derive(Debug, Clone, PartialEq)]
123pub(crate) struct InternalUkfReport {
124 pub(crate) applied: bool,
125 pub(crate) normalized_innovation_squared: f64,
126 pub(crate) accepted_rows: usize,
127 pub(crate) rejected_rows: usize,
128 pub(crate) innovation_gate: Option<InnovationGateReport>,
129 pub(crate) innovation_covariance: Vec<Vec<f64>>,
130 pub(crate) kalman_gain: Vec<Vec<f64>>,
131 pub(crate) dx: Vec<f64>,
132 pub(crate) posterior_covariance: Vec<Vec<f64>>,
133}
134
135impl InternalUkfReport {
136 pub(crate) fn into_public_report(self) -> EkfCorrectionReport {
137 EkfCorrectionReport {
138 applied: self.applied,
139 normalized_innovation_squared: self.normalized_innovation_squared,
140 accepted_rows: self.accepted_rows,
141 rejected_rows: self.rejected_rows,
142 innovation_gate: self.innovation_gate,
143 innovation_covariance: self.innovation_covariance,
144 kalman_gain: self.kalman_gain,
145 dx: self.dx,
146 }
147 }
148}
149
150pub(crate) fn ukf_measurement_update<F>(
151 covariance: &[Vec<f64>],
152 innovation: &[f64],
153 measurement_covariance: &[Vec<f64>],
154 options: UkfUpdateOptions,
155 measurement_model: F,
156) -> Result<InternalUkfReport, FusionError>
157where
158 F: Fn(&[f64]) -> Result<Vec<f64>, FusionError>,
159{
160 let dimension = covariance.len();
161 validate_covariance_matrix(covariance, dimension, "covariance")?;
162 validate_finite_slice(innovation, "innovation")?;
163 validate_covariance_matrix(
164 measurement_covariance,
165 innovation.len(),
166 "measurement_covariance",
167 )?;
168 options.validate_for_dimension(dimension)?;
169
170 let sigma = sigma_points(covariance, options.transform)?;
171 let prediction = measurement_statistics(&sigma, innovation.len(), &measurement_model)?;
172 let full = predicted_update(
173 covariance,
174 innovation,
175 measurement_covariance,
176 &prediction,
177 None,
178 )?;
179
180 let Some(gate) = options.innovation_gate else {
181 return Ok(full);
182 };
183
184 let (accepted, gate_report) = screen_rows(
185 innovation,
186 &prediction.mean,
187 &full.innovation_covariance,
188 gate,
189 )?;
190 if gate_report.coasted {
191 let full_nis = normalized_innovation_squared(
192 &full.innovation_covariance,
193 &innovation_residual(innovation, &prediction.mean)?,
194 )?;
195 return Ok(InternalUkfReport {
196 applied: false,
197 normalized_innovation_squared: full_nis,
198 accepted_rows: gate_report.accepted_rows,
199 rejected_rows: gate_report.rejected_rows,
200 innovation_gate: Some(gate_report),
201 innovation_covariance: full.innovation_covariance,
202 kalman_gain: vec![vec![0.0; innovation.len()]; dimension],
203 dx: vec![0.0; dimension],
204 posterior_covariance: covariance.to_vec(),
205 });
206 }
207
208 let mut screened = predicted_update(
209 covariance,
210 innovation,
211 measurement_covariance,
212 &prediction,
213 Some(&accepted),
214 )?;
215 screened.accepted_rows = gate_report.accepted_rows;
216 screened.rejected_rows = gate_report.rejected_rows;
217 screened.innovation_gate = Some(gate_report);
218 Ok(screened)
219}
220
221#[derive(Debug, Clone, PartialEq)]
222struct SigmaSet {
223 points: Vec<Vec<f64>>,
224 mean_weights: Vec<f64>,
225 covariance_weights: Vec<f64>,
226}
227
228#[derive(Debug, Clone, PartialEq)]
229struct MeasurementPrediction {
230 values: Vec<Vec<f64>>,
231 mean: Vec<f64>,
232 cross_covariance: Vec<Vec<f64>>,
233 covariance_weights: Vec<f64>,
234}
235
236fn sigma_points(
237 covariance: &[Vec<f64>],
238 options: UnscentedTransformOptions,
239) -> Result<SigmaSet, FusionError> {
240 let dimension = covariance.len();
241 options.validate_for_dimension(dimension)?;
242 let scale = options.scale(dimension);
243 let lambda = options.lambda(dimension);
244 let gamma = scale.sqrt();
245 let sqrt = covariance_square_root(covariance)?;
246
247 let point_count = 2 * dimension + 1;
248 let mut points = Vec::with_capacity(point_count);
249 points.push(vec![0.0; dimension]);
250 for col in 0..dimension {
251 let mut point = vec![0.0; dimension];
252 for row in 0..dimension {
253 point[row] = gamma * sqrt[(row, col)];
254 }
255 points.push(point);
256 }
257 for col in 0..dimension {
258 let mut point = vec![0.0; dimension];
259 for row in 0..dimension {
260 point[row] = -gamma * sqrt[(row, col)];
261 }
262 points.push(point);
263 }
264
265 let mut mean_weights = vec![0.5 / scale; point_count];
266 let mut covariance_weights = mean_weights.clone();
267 mean_weights[0] = lambda / scale;
268 covariance_weights[0] = mean_weights[0] + (1.0 - options.alpha * options.alpha + options.beta);
269
270 Ok(SigmaSet {
271 points,
272 mean_weights,
273 covariance_weights,
274 })
275}
276
277fn covariance_square_root(covariance: &[Vec<f64>]) -> Result<DMatrix<f64>, FusionError> {
278 let dimension = covariance.len();
279 validate_covariance_matrix(covariance, dimension, "covariance")?;
280 let matrix = dmatrix_from_rows(covariance);
281 if let Some(cholesky) = portable::cholesky_lower_dynamic(&matrix) {
282 return Ok(cholesky);
283 }
284
285 let (eigenvectors, eigenvalues) = portable::symmetric_eigen_dynamic(&matrix);
286 let mut diagonal = DMatrix::<f64>::zeros(dimension, dimension);
287 for idx in 0..dimension {
288 let eigenvalue = eigenvalues[idx];
289 if eigenvalue < 0.0 {
290 let tolerance = covariance_eigenvalue_tolerance(covariance, &eigenvectors, idx);
291 if eigenvalue < -tolerance {
292 return Err(FusionError::NonPositiveSemidefinite {
293 field: "covariance",
294 });
295 }
296 diagonal[(idx, idx)] = 0.0;
297 } else {
298 diagonal[(idx, idx)] = eigenvalue.sqrt();
299 }
300 }
301 Ok(portable::product(&eigenvectors, &diagonal))
302}
303
304fn measurement_statistics<F>(
305 sigma: &SigmaSet,
306 measurement_dimension: usize,
307 measurement_model: &F,
308) -> Result<MeasurementPrediction, FusionError>
309where
310 F: Fn(&[f64]) -> Result<Vec<f64>, FusionError>,
311{
312 let mut values = Vec::with_capacity(sigma.points.len());
313 for point in &sigma.points {
314 let value = measurement_model(point)?;
315 if value.len() != measurement_dimension {
316 return Err(FusionError::DimensionMismatch {
317 field: "ukf_measurement",
318 expected: measurement_dimension,
319 actual: value.len(),
320 });
321 }
322 validate_finite_slice(&value, "ukf_measurement")?;
323 values.push(value);
324 }
325
326 let mut mean = vec![0.0; measurement_dimension];
327 for (weight, value) in sigma.mean_weights.iter().zip(values.iter()) {
328 for col in 0..measurement_dimension {
329 mean[col] += weight * value[col];
330 }
331 }
332
333 let state_dimension = sigma.points[0].len();
334 let mut cross_covariance = vec![vec![0.0; measurement_dimension]; state_dimension];
335 for (idx, point) in sigma.points.iter().enumerate() {
336 let weight = sigma.covariance_weights[idx];
337 for row in 0..state_dimension {
338 for col in 0..measurement_dimension {
339 cross_covariance[row][col] += weight * point[row] * (values[idx][col] - mean[col]);
340 }
341 }
342 }
343
344 Ok(MeasurementPrediction {
345 values,
346 mean,
347 cross_covariance,
348 covariance_weights: sigma.covariance_weights.clone(),
349 })
350}
351
352fn predicted_update(
353 covariance: &[Vec<f64>],
354 innovation: &[f64],
355 measurement_covariance: &[Vec<f64>],
356 prediction: &MeasurementPrediction,
357 accepted: Option<&[usize]>,
358) -> Result<InternalUkfReport, FusionError> {
359 let selected = accepted
360 .map(<[usize]>::to_vec)
361 .unwrap_or_else(|| (0..innovation.len()).collect());
362 let innovation = select_vector(innovation, &selected)?;
363 let mean = select_vector(&prediction.mean, &selected)?;
364 let measurement_covariance = select_matrix(measurement_covariance, &selected)?;
365 let values = prediction
366 .values
367 .iter()
368 .map(|value| select_vector(value, &selected))
369 .collect::<Result<Vec<_>, _>>()?;
370 let cross_covariance = select_columns(&prediction.cross_covariance, &selected)?;
371
372 let residual = innovation_residual(&innovation, &mean)?;
373 let mut innovation_covariance = measurement_covariance;
374 for (idx, value) in values.iter().enumerate() {
375 let weight = prediction.covariance_weights[idx];
376 for row in 0..selected.len() {
377 let dy_row = value[row] - mean[row];
378 for col in 0..selected.len() {
379 innovation_covariance[row][col] += weight * dy_row * (value[col] - mean[col]);
380 }
381 }
382 }
383 symmetrize_in_place(&mut innovation_covariance);
384 validate_covariance_matrix(
385 &innovation_covariance,
386 selected.len(),
387 "innovation_covariance",
388 )?;
389
390 let mut kalman_gain = vec![vec![0.0; selected.len()]; covariance.len()];
391 let mut scratch = crate::astro::math::linear::FlatCholeskySolveScratch::default();
392 for row in 0..covariance.len() {
393 kalman_gain[row] = solve_spd(&innovation_covariance, &cross_covariance[row], &mut scratch)?;
394 }
395
396 let dx = super::state::matvec(&kalman_gain, &residual)?;
397 let nis = normalized_innovation_squared(&innovation_covariance, &residual)?;
398 let ks = matmul(&kalman_gain, &innovation_covariance)?;
399 let k_t = transpose(&kalman_gain)?;
400 let ksk_t = matmul(&ks, &k_t)?;
401 let mut posterior_covariance = matrix_sub(covariance, &ksk_t)?;
402 symmetrize_in_place(&mut posterior_covariance);
403 reproject_covariance_psd(&mut posterior_covariance, "ukf_covariance")?;
404
405 Ok(InternalUkfReport {
406 applied: true,
407 normalized_innovation_squared: nis,
408 accepted_rows: selected.len(),
409 rejected_rows: innovation.len().saturating_sub(selected.len()),
410 innovation_gate: None,
411 innovation_covariance,
412 kalman_gain,
413 dx,
414 posterior_covariance,
415 })
416}
417
418fn innovation_residual(innovation: &[f64], mean: &[f64]) -> Result<Vec<f64>, FusionError> {
419 if innovation.len() != mean.len() {
420 return Err(FusionError::DimensionMismatch {
421 field: "innovation_mean",
422 expected: innovation.len(),
423 actual: mean.len(),
424 });
425 }
426 Ok(innovation
427 .iter()
428 .zip(mean.iter())
429 .map(|(actual, predicted)| actual - predicted)
430 .collect())
431}
432
433fn screen_rows(
434 innovation: &[f64],
435 mean: &[f64],
436 innovation_covariance: &[Vec<f64>],
437 gate: InnovationGate,
438) -> Result<(Vec<usize>, InnovationGateReport), FusionError> {
439 gate.validate()?;
440 let residual = innovation_residual(innovation, mean)?;
441 let mut accepted = Vec::with_capacity(innovation.len());
442 let mut rejected_rows = 0usize;
443 let mut max_abs_normalized_innovation = None;
444 let mut max_rejected_abs_normalized_innovation = None;
445
446 for (row, value) in residual.iter().enumerate() {
447 let variance = innovation_covariance[row][row];
448 validate_positive(variance, "innovation_covariance_diagonal")?;
449 let normalized = (value / variance.sqrt()).abs();
450 max_abs_normalized_innovation = Some(
451 max_abs_normalized_innovation
452 .map_or(normalized, |current: f64| current.max(normalized)),
453 );
454 if normalized <= gate.threshold_sigma {
455 accepted.push(row);
456 } else {
457 rejected_rows += 1;
458 max_rejected_abs_normalized_innovation = Some(
459 max_rejected_abs_normalized_innovation
460 .map_or(normalized, |current: f64| current.max(normalized)),
461 );
462 }
463 }
464
465 let coasted = accepted.len() < gate.min_rows;
466 let report = InnovationGateReport {
467 threshold_sigma: gate.threshold_sigma,
468 min_rows: gate.min_rows,
469 input_rows: innovation.len(),
470 accepted_rows: accepted.len(),
471 rejected_rows,
472 max_abs_normalized_innovation,
473 max_rejected_abs_normalized_innovation,
474 coasted,
475 };
476 Ok((accepted, report))
477}
478
479fn select_vector(values: &[f64], indices: &[usize]) -> Result<Vec<f64>, FusionError> {
480 let mut selected = Vec::with_capacity(indices.len());
481 for idx in indices {
482 let Some(value) = values.get(*idx) else {
483 return Err(FusionError::DimensionMismatch {
484 field: "selected_measurement",
485 expected: values.len(),
486 actual: *idx,
487 });
488 };
489 selected.push(*value);
490 }
491 Ok(selected)
492}
493
494fn select_matrix(matrix: &[Vec<f64>], indices: &[usize]) -> Result<Vec<Vec<f64>>, FusionError> {
495 let mut out = vec![vec![0.0; indices.len()]; indices.len()];
496 for (row_out, row_in) in indices.iter().enumerate() {
497 for (col_out, col_in) in indices.iter().enumerate() {
498 out[row_out][col_out] = matrix[*row_in][*col_in];
499 }
500 }
501 Ok(out)
502}
503
504fn select_columns(matrix: &[Vec<f64>], indices: &[usize]) -> Result<Vec<Vec<f64>>, FusionError> {
505 if matrix.is_empty() {
506 return Err(invalid_input("matrix", "must not be empty"));
507 }
508 validate_matrix_cols(matrix, matrix[0].len(), "matrix")?;
509 let mut out = vec![vec![0.0; indices.len()]; matrix.len()];
510 for (row_out, row) in matrix.iter().enumerate() {
511 for (col_out, col_in) in indices.iter().enumerate() {
512 out[row_out][col_out] = row[*col_in];
513 }
514 }
515 Ok(out)
516}
517
518#[cfg(test)]
519mod tests {
520 use super::*;
527 use crate::astro::constants::earth::WGS84_A_M;
528 use crate::fusion::ekf::{ekf_correct_closed_loop, EkfUpdateOptions};
529 use crate::fusion::state::{ErrorStateLayout, ERROR_STATE_DIMENSION_15};
530 use crate::inertial::state::mat3_identity;
531 use crate::inertial::NavState;
532
533 fn assert_close(actual: f64, expected: f64, tolerance: f64) {
534 assert!(
535 (actual - expected).abs() <= tolerance,
536 "actual {actual:.17e}, expected {expected:.17e}, tolerance {tolerance:.17e}"
537 );
538 }
539
540 fn linear_test_state() -> InsFilterState {
541 let nominal =
542 NavState::new(0.0, [WGS84_A_M, 0.0, 0.0], [0.0; 3], mat3_identity()).expect("nominal");
543 let mut covariance = vec![vec![0.0; ERROR_STATE_DIMENSION_15]; ERROR_STATE_DIMENSION_15];
544 for (idx, row) in covariance.iter_mut().enumerate() {
545 row[idx] = 1.0;
546 }
547 covariance[0][0] = 4.0;
548 covariance[0][1] = 1.0;
549 covariance[1][0] = 1.0;
550 covariance[1][1] = 9.0;
551 InsFilterState::new(nominal, ErrorStateLayout::Fifteen, covariance).expect("state")
552 }
553
554 #[test]
555 fn linear_measurement_matches_closed_form_and_ekf() {
556 let mut design = vec![vec![0.0; ERROR_STATE_DIMENSION_15]];
557 design[0][0] = 0.5;
558 design[0][1] = -2.0;
559 let correction =
560 EkfCorrection::new(vec![1.25], design, vec![vec![0.25]]).expect("correction");
561 let mut ekf_state = linear_test_state();
562 let mut ukf_state = linear_test_state();
563
564 let ekf = ekf_correct_closed_loop(&mut ekf_state, &correction, EkfUpdateOptions::default())
565 .expect("ekf");
566 let ukf = ukf_correct_closed_loop(
567 &mut ukf_state,
568 &correction,
569 UkfUpdateOptions {
570 transform: UnscentedTransformOptions {
571 alpha: 1.0,
572 beta: 2.0,
573 kappa: 0.0,
574 },
575 innovation_gate: None,
576 },
577 )
578 .expect("ukf");
579
580 let expected_s = 35.25_f64;
581 let expected_k0 = 0.0_f64;
582 let expected_k1 = -17.5 / expected_s;
583 let expected_dx1 = expected_k1 * 1.25;
584 assert_close(ukf.innovation_covariance[0][0], expected_s, 1.0e-13);
585 assert_close(ukf.kalman_gain[0][0], expected_k0, 1.0e-14);
586 assert_close(ukf.kalman_gain[1][0], expected_k1, 1.0e-14);
587 assert_close(ukf.dx[1], expected_dx1, 1.0e-14);
588
589 for row in 0..ERROR_STATE_DIMENSION_15 {
590 assert_close(ukf.kalman_gain[row][0], ekf.kalman_gain[row][0], 1.0e-15);
591 assert_close(ukf.dx[row], ekf.dx[row], 1.0e-15);
592 for col in 0..ERROR_STATE_DIMENSION_15 {
593 assert_close(
594 ukf_state.covariance[row][col],
595 ekf_state.covariance[row][col],
596 1.0e-15,
597 );
598 }
599 }
600 assert_close(
601 ukf_state.nominal.position_ecef_m[1],
602 ekf_state.nominal.position_ecef_m[1],
603 3.0e-13,
604 );
605 }
606}