use anyhow::Result;
use faer::{Mat, MatRef};
use std::collections::HashMap;
pub fn invert_spillover(spillover: MatRef<'_, f32>) -> Result<Mat<f32>> {
use faer::linalg::solvers::{DenseSolveCore, PartialPivLu};
let lu = PartialPivLu::new(spillover);
let u = lu.U();
for i in 0..u.nrows().min(u.ncols()) {
if u[(i, i)].abs() < f32::EPSILON {
anyhow::bail!(
"spillover matrix is singular or ill-conditioned at diagonal index {i}"
);
}
}
Ok(lu.inverse())
}
pub fn apply_compensation_inv(
raw_channels: &[(&str, &[f32])],
comp_inv: MatRef<'_, f32>,
matrix_channel_names: &[&str],
) -> Result<HashMap<String, Vec<f32>>> {
use rayon::prelude::*;
let n = matrix_channel_names.len();
let raw_map: HashMap<&str, &[f32]> = raw_channels.iter().copied().collect();
let channel_data: Vec<Option<&[f32]>> = matrix_channel_names
.iter()
.map(|&name| raw_map.get(name).copied())
.collect();
let n_events = channel_data
.iter()
.find_map(|c| c.map(|s| s.len()))
.unwrap_or(0);
for (i, opt) in channel_data.iter().enumerate() {
if let Some(raw) = opt {
anyhow::ensure!(
raw.len() == n_events,
"channel '{}' has {} events but expected {n_events}",
matrix_channel_names[i],
raw.len()
);
}
}
if n_events == 0 {
return Ok(HashMap::new());
}
let compensated: Vec<Vec<f32>> = (0..n)
.into_par_iter()
.map(|i| {
let mut result = vec![0.0f32; n_events];
for (event_idx, val) in result.iter_mut().enumerate() {
let mut sum = 0.0f32;
for j in 0..n {
if let Some(raw) = channel_data[j] {
sum += comp_inv[(i, j)] * raw[event_idx];
}
}
*val = sum;
}
result
})
.collect();
let mut result = HashMap::new();
for (i, name) in matrix_channel_names.iter().enumerate() {
if raw_map.contains_key(name) {
result.insert(name.to_string(), compensated[i].clone());
}
}
Ok(result)
}
pub fn compensate_channels(
raw_channels: &[(&str, &[f32])],
spillover: MatRef<'_, f32>,
matrix_channel_names: &[&str],
channels_needed: &[&str],
) -> Result<HashMap<String, Vec<f32>>> {
let comp_inv = invert_spillover(spillover)?;
let all_compensated =
apply_compensation_inv(raw_channels, comp_inv.as_ref(), matrix_channel_names)?;
let needed_set: std::collections::HashSet<&str> = channels_needed.iter().copied().collect();
Ok(all_compensated
.into_iter()
.filter(|(k, _)| needed_set.contains(k.as_str()))
.collect())
}
#[cfg(test)]
mod tests {
use super::*;
use faer::mat;
fn identity_2x2() -> Mat<f32> {
mat![[1.0f32, 0.0], [0.0, 1.0]]
}
fn known_spillover_2x2() -> Mat<f32> {
mat![[1.0f32, 0.0], [0.2, 1.0]]
}
#[test]
fn test_invert_identity_is_identity() {
let m = identity_2x2();
let inv = invert_spillover(m.as_ref()).unwrap();
for i in 0..2 {
for j in 0..2 {
let expected = if i == j { 1.0 } else { 0.0 };
assert!(
(inv[(i, j)] - expected).abs() < 1e-5,
"inv[{i},{j}] = {} expected {expected}",
inv[(i, j)]
);
}
}
}
#[test]
fn test_compensate_identity_returns_raw() {
let m = identity_2x2();
let ch_a: Vec<f32> = vec![1.0, 2.0, 3.0];
let ch_b: Vec<f32> = vec![4.0, 5.0, 6.0];
let raw = [("A", ch_a.as_slice()), ("B", ch_b.as_slice())];
let names = ["A", "B"];
let result = compensate_channels(&raw, m.as_ref(), &names, &names).unwrap();
for (i, &v) in result["A"].iter().enumerate() {
assert!((v - ch_a[i]).abs() < 1e-5);
}
for (i, &v) in result["B"].iter().enumerate() {
assert!((v - ch_b[i]).abs() < 1e-5);
}
}
#[test]
fn test_compensate_known_spillover_removes_spillover() {
let spillover = known_spillover_2x2();
let true_a: Vec<f32> = vec![100.0, 200.0];
let true_b: Vec<f32> = vec![50.0, 80.0];
let measured_a = true_a.clone();
let measured_b: Vec<f32> = true_b
.iter()
.zip(true_a.iter())
.map(|(b, a)| b + 0.2 * a)
.collect();
let raw = [("A", measured_a.as_slice()), ("B", measured_b.as_slice())];
let names = ["A", "B"];
let result = compensate_channels(&raw, spillover.as_ref(), &names, &names).unwrap();
for (i, &v) in result["B"].iter().enumerate() {
assert!(
(v - true_b[i]).abs() < 1e-3,
"compensated_b[{i}] = {v}, expected {}",
true_b[i]
);
}
}
#[test]
fn test_channels_needed_filters_result() {
let m = identity_2x2();
let ch_a: Vec<f32> = vec![1.0];
let ch_b: Vec<f32> = vec![2.0];
let raw = [("A", ch_a.as_slice()), ("B", ch_b.as_slice())];
let names = ["A", "B"];
let result = compensate_channels(&raw, m.as_ref(), &names, &["A"]).unwrap();
assert!(result.contains_key("A"));
assert!(!result.contains_key("B"));
}
}