use strided_kernel::reduce_axis;
use strided_view::{ElementOp, StridedArray, StridedView};
pub fn find_trace_indices<ID: PartialEq>(labels: &[ID], other: &[ID], output: &[ID]) -> Vec<usize> {
labels
.iter()
.enumerate()
.filter(|(_, id)| !other.contains(id) && !output.contains(id))
.map(|(i, _)| i)
.collect()
}
pub fn reduce_trace_axes<T, Op>(
src: &StridedView<T, Op>,
trace_axes: &[usize],
) -> strided_view::Result<StridedArray<T>>
where
T: Copy + Send + Sync + std::ops::Add<Output = T> + num_traits::Zero,
Op: ElementOp<T>,
{
if trace_axes.is_empty() {
panic!("reduce_trace_axes called with empty trace_axes");
}
let mut axes: Vec<usize> = trace_axes.to_vec();
axes.sort_unstable();
axes.reverse();
let first_reduced = reduce_axis(src, axes[0], |x| x, |a, b| a + b, T::zero())?;
let mut current = first_reduced;
for &ax in &axes[1..] {
current = reduce_axis(¤t.view(), ax, |x| x, |a, b| a + b, T::zero())?;
}
Ok(current)
}
#[cfg(test)]
mod tests {
use super::*;
use strided_view::Identity;
#[test]
#[should_panic(expected = "reduce_trace_axes called with empty trace_axes")]
fn test_reduce_no_trace() {
let a =
StridedArray::<f64>::from_fn_row_major(&[2, 3], |idx| (idx[0] * 3 + idx[1] + 1) as f64);
let _ = reduce_trace_axes::<f64, Identity>(&a.view(), &[]);
}
#[test]
fn test_reduce_single_trace() {
let a =
StridedArray::<f64>::from_fn_row_major(&[2, 3], |idx| (idx[0] * 3 + idx[1] + 1) as f64);
let result = reduce_trace_axes::<f64, Identity>(&a.view(), &[1]).unwrap();
assert_eq!(result.dims(), &[2]);
assert_eq!(result.get(&[0]), 6.0); assert_eq!(result.get(&[1]), 15.0); }
#[test]
fn test_reduce_two_traces() {
let a = StridedArray::<f64>::from_fn_row_major(&[2, 3, 4], |idx| {
(idx[0] * 12 + idx[1] * 4 + idx[2]) as f64
});
let result = reduce_trace_axes::<f64, Identity>(&a.view(), &[0, 2]).unwrap();
assert_eq!(result.dims(), &[3]);
assert_eq!(result.get(&[0]), 60.0); assert_eq!(result.get(&[1]), 92.0); assert_eq!(result.get(&[2]), 124.0); }
}