use crate::error::FdarError;
use crate::helpers::{linear_interp, simpsons_weights};
use crate::iter_maybe_parallel;
use crate::matrix::FdMatrix;
#[cfg(feature = "parallel")]
use rayon::iter::ParallelIterator;
pub const DEFAULT_MAX_SHIFT_FRACTION: f64 = 0.25;
const GS_TOL: f64 = 1e-6;
const GS_MAX_ITER: usize = 100;
#[derive(Debug, Clone, PartialEq)]
#[non_exhaustive]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub struct ShiftRegistrationResult {
pub registered_data: FdMatrix,
pub shifts: Vec<f64>,
}
fn golden_section_search<F>(f: F, mut lo: f64, mut hi: f64, tol: f64, max_iter: usize) -> f64
where
F: Fn(f64) -> f64,
{
const PHI: f64 = 1.618_033_988_749_895;
let mut x1 = hi - (hi - lo) / PHI;
let mut x2 = lo + (hi - lo) / PHI;
let mut f1 = f(x1);
let mut f2 = f(x2);
for _ in 0..max_iter {
if (hi - lo) < tol {
break;
}
if f1 < f2 {
hi = x2;
x2 = x1;
f2 = f1;
x1 = hi - (hi - lo) / PHI;
f1 = f(x1);
} else {
lo = x1;
x1 = x2;
f1 = f2;
x2 = lo + (hi - lo) / PHI;
f2 = f(x2);
}
}
(lo + hi) / 2.0
}
fn l2_shift_objective(
row: &[f64],
argvals: &[f64],
mean: &[f64],
weights: &[f64],
delta: f64,
) -> f64 {
argvals
.iter()
.zip(mean.iter())
.zip(weights.iter())
.map(|((&t, &m_j), &w)| {
let fi_shifted = linear_interp(argvals, row, t - delta);
let diff = fi_shifted - m_j;
diff * diff * w
})
.sum::<f64>()
}
pub fn least_squares_shift_registration(
data: &FdMatrix,
argvals: &[f64],
max_shift: f64,
) -> Result<ShiftRegistrationResult, FdarError> {
let (n, m) = data.shape();
if n == 0 || m == 0 {
return Err(FdarError::InvalidDimension {
parameter: "data",
expected: "non-empty matrix".to_string(),
actual: format!("{}x{}", n, m),
});
}
if argvals.len() != m {
return Err(FdarError::InvalidDimension {
parameter: "argvals",
expected: m.to_string(),
actual: argvals.len().to_string(),
});
}
if argvals.len() < 2 {
return Err(FdarError::InvalidParameter {
parameter: "argvals",
message: "must have at least 2 evaluation points".to_string(),
});
}
if max_shift <= 0.0 {
return Err(FdarError::InvalidParameter {
parameter: "max_shift",
message: format!("must be positive, got {max_shift}"),
});
}
let weights = simpsons_weights(argvals);
let mean = crate::fdata::mean_1d(data);
let results: Vec<(f64, Vec<f64>)> = iter_maybe_parallel!(0..n)
.map(|i| {
let row = data.row(i);
let delta = golden_section_search(
|d| l2_shift_objective(&row, argvals, &mean, &weights, d),
-max_shift,
max_shift,
GS_TOL,
GS_MAX_ITER,
);
let shifted: Vec<f64> = argvals
.iter()
.map(|&t| linear_interp(argvals, &row, t - delta))
.collect();
(delta, shifted)
})
.collect();
let mut registered_data = FdMatrix::zeros(n, m);
let mut shifts = Vec::with_capacity(n);
for (i, (delta, shifted_row)) in results.into_iter().enumerate() {
for j in 0..m {
registered_data[(i, j)] = shifted_row[j];
}
shifts.push(delta);
}
Ok(ShiftRegistrationResult {
registered_data,
shifts,
})
}
#[cfg(test)]
mod tests {
use super::*;
use crate::matrix::FdMatrix;
fn uniform_grid(n: usize) -> Vec<f64> {
(0..n).map(|i| i as f64 / (n - 1) as f64).collect()
}
fn gaussian_bump(argvals: &[f64], mu: f64, sigma: f64) -> Vec<f64> {
argvals
.iter()
.map(|&t| (-(t - mu).powi(2) / (2.0 * sigma * sigma)).exp())
.collect()
}
#[test]
fn test_shift_already_aligned() {
let m = 51;
let argvals = uniform_grid(m);
let n = 4;
let mut data = FdMatrix::zeros(n, m);
let bump = gaussian_bump(&argvals, 0.5, 0.08);
for i in 0..n {
for j in 0..m {
data[(i, j)] = bump[j];
}
}
let max_shift = 0.25 * (argvals[m - 1] - argvals[0]);
let result = least_squares_shift_registration(&data, &argvals, max_shift).unwrap();
for (i, &delta) in result.shifts.iter().enumerate() {
assert!(
delta.abs() < 1e-3,
"curve {i}: expected shift ≈ 0, got {delta}"
);
}
}
#[test]
fn test_shift_recovers_injected_offset() {
let m = 101;
let argvals = uniform_grid(m);
let sigma = 0.05_f64;
let centres = [0.5_f64, 0.4, 0.6];
let n = centres.len();
let mut data = FdMatrix::zeros(n, m);
for (i, &mu) in centres.iter().enumerate() {
let row = gaussian_bump(&argvals, mu, sigma);
for j in 0..m {
data[(i, j)] = row[j];
}
}
let true_shifts = [0.0_f64, 0.1, -0.1];
let max_shift = 0.25 * (argvals[m - 1] - argvals[0]);
let result = least_squares_shift_registration(&data, &argvals, max_shift).unwrap();
for (i, (&recovered, &expected)) in result.shifts.iter().zip(true_shifts.iter()).enumerate()
{
assert!(
(recovered - expected).abs() < 0.05,
"curve {i}: expected shift ≈ {expected}, got {recovered}"
);
}
}
#[test]
fn test_shift_registration_curve_values() {
let m = 5;
let argvals = uniform_grid(m);
let n = 2;
let mut data = FdMatrix::zeros(n, m);
for (i, &mu) in [0.3_f64, 0.7].iter().enumerate() {
let row = gaussian_bump(&argvals, mu, 0.15);
for j in 0..m {
data[(i, j)] = row[j];
}
}
let max_shift = 0.25;
let result = least_squares_shift_registration(&data, &argvals, max_shift).unwrap();
for i in 0..n {
let row = data.row(i);
let delta = result.shifts[i];
for j in 0..m {
let expected = linear_interp(&argvals, &row, argvals[j] - delta);
let actual = result.registered_data[(i, j)];
assert!(
(actual - expected).abs() < 1e-9,
"registered_data[({i},{j})] = {actual}, expected {expected} (shift={delta})"
);
}
}
}
#[test]
fn test_shift_registration_empty_data() {
let argvals = uniform_grid(5);
let data_n0 = FdMatrix::zeros(0, 5);
let result = least_squares_shift_registration(&data_n0, &argvals, 0.1);
assert!(
matches!(result, Err(FdarError::InvalidDimension { .. })),
"expected Err(InvalidDimension) for n=0, got {result:?}"
);
let data_m0 = FdMatrix::zeros(3, 0);
let result_m0 = least_squares_shift_registration(&data_m0, &[], 0.1);
assert!(
matches!(result_m0, Err(FdarError::InvalidDimension { .. })),
"expected Err(InvalidDimension) for m=0, got {result_m0:?}"
);
}
#[test]
fn test_shift_registration_argvals_mismatch() {
let m = 5;
let data = FdMatrix::zeros(2, m);
let wrong_argvals = uniform_grid(m + 1);
let result = least_squares_shift_registration(&data, &wrong_argvals, 0.1);
assert!(
matches!(result, Err(FdarError::InvalidDimension { .. })),
"expected Err(InvalidDimension) for argvals length mismatch, got {result:?}"
);
let short_argvals = uniform_grid(m - 1);
let result2 = least_squares_shift_registration(&data, &short_argvals, 0.1);
assert!(
matches!(result2, Err(FdarError::InvalidDimension { .. })),
"expected Err(InvalidDimension) for argvals too short, got {result2:?}"
);
}
}