use super::MaskedArray;
use crate::array::Array;
use crate::error::{NumRs2Error, Result};
use num_traits::{Float, Zero};
use std::ops::{Add, Div};
pub(super) fn normalize_axis(ax: isize, ndim: usize) -> Result<usize> {
let normalized = if ax < 0 { ax + ndim as isize } else { ax };
if normalized < 0 || normalized as usize >= ndim {
return Err(NumRs2Error::DimensionMismatch(format!(
"axis {ax} out of bounds for array of dimension {ndim}"
)));
}
Ok(normalized as usize)
}
pub(super) fn axis_lane_shape(shape: &[usize], ax: usize) -> (usize, usize, usize) {
let outer: usize = shape[..ax].iter().product();
let axis_size = shape[ax];
let inner: usize = shape[ax + 1..].iter().product();
(outer, axis_size, inner)
}
pub(super) fn collapsed_shape(shape: &[usize], ax: usize, keepdims: bool) -> Vec<usize> {
let mut out_shape = shape.to_vec();
if keepdims {
out_shape[ax] = 1;
} else {
out_shape.remove(ax);
if out_shape.is_empty() {
out_shape.push(1);
}
}
out_shape
}
pub(super) fn reduce_lanes<T, U>(
data: &Array<T>,
mask: &Array<bool>,
axis: Option<isize>,
keepdims: bool,
f: impl Fn(&[T], &[bool]) -> Option<U>,
) -> Result<MaskedArray<U>>
where
T: Clone,
U: Clone + Default,
{
let shape = data.shape();
let ndim = shape.len();
let data_op = crate::kernels::borrow::operand(data);
let mask_op = crate::kernels::borrow::operand(mask);
let data_slice: &[T] = &data_op;
let mask_slice: &[bool] = &mask_op;
let (out_shape, outer, axis_size, inner) = match axis {
None => {
let out_shape = if keepdims {
vec![1; ndim.max(1)]
} else {
vec![1]
};
(out_shape, 1usize, data_slice.len(), 1usize)
}
Some(ax) => {
let ax = normalize_axis(ax, ndim)?;
let (outer, axis_size, inner) = axis_lane_shape(&shape, ax);
(
collapsed_shape(&shape, ax, keepdims),
outer,
axis_size,
inner,
)
}
};
let out_len = outer * inner;
let mut out_data: Vec<U> = Vec::with_capacity(out_len);
let mut out_mask: Vec<bool> = Vec::with_capacity(out_len);
let mut lane_vals: Vec<T> = Vec::with_capacity(axis_size);
let mut lane_mask: Vec<bool> = Vec::with_capacity(axis_size);
for o in 0..outer {
for i in 0..inner {
lane_vals.clear();
lane_mask.clear();
let base = o * axis_size * inner + i;
for k in 0..axis_size {
let idx = base + k * inner;
lane_vals.push(data_slice[idx].clone());
lane_mask.push(mask_slice[idx]);
}
match f(&lane_vals, &lane_mask) {
Some(v) => {
out_data.push(v);
out_mask.push(false);
}
None => {
out_data.push(U::default());
out_mask.push(true);
}
}
}
}
Ok(MaskedArray {
data: Array::from_vec_shape(out_data, &out_shape)?,
mask: Array::from_vec_shape(out_mask, &out_shape)?,
fill_value: U::default(),
})
}
impl<T: Clone + Add<Output = T> + Div<Output = T> + Zero + From<f64> + Into<f64>> MaskedArray<T> {
pub fn mean(&self) -> Option<T> {
let data_op = crate::kernels::borrow::operand(&self.data);
let mask_op = crate::kernels::borrow::operand(&self.mask);
let mut sum = T::zero();
let mut count = 0;
for (value, is_masked) in data_op.iter().zip(mask_op.iter()) {
if !*is_masked {
sum = sum + value.clone();
count += 1;
}
}
if count == 0 {
None
} else {
let count_f64 = count as f64;
let sum_f64: f64 = sum.into();
Some(T::from(sum_f64 / count_f64))
}
}
pub fn sum(&self) -> Option<T> {
let data_op = crate::kernels::borrow::operand(&self.data);
let mask_op = crate::kernels::borrow::operand(&self.mask);
let mut sum = T::zero();
let mut count = 0;
for (value, is_masked) in data_op.iter().zip(mask_op.iter()) {
if !*is_masked {
sum = sum + value.clone();
count += 1;
}
}
if count == 0 {
None
} else {
Some(sum)
}
}
pub fn min(&self) -> Option<T>
where
T: PartialOrd,
{
let data_op = crate::kernels::borrow::operand(&self.data);
let mask_op = crate::kernels::borrow::operand(&self.mask);
let mut min_val = None;
for (value, is_masked) in data_op.iter().zip(mask_op.iter()) {
if !*is_masked {
match min_val {
None => min_val = Some(value.clone()),
Some(ref current_min) if value < current_min => min_val = Some(value.clone()),
_ => {}
}
}
}
min_val
}
pub fn max(&self) -> Option<T>
where
T: PartialOrd,
{
let data_op = crate::kernels::borrow::operand(&self.data);
let mask_op = crate::kernels::borrow::operand(&self.mask);
let mut max_val = None;
for (value, is_masked) in data_op.iter().zip(mask_op.iter()) {
if !*is_masked {
match max_val {
None => max_val = Some(value.clone()),
Some(ref current_max) if value > current_max => max_val = Some(value.clone()),
_ => {}
}
}
}
max_val
}
}
impl<T: Float + Default> MaskedArray<T> {
pub fn mean_axis(&self, axis: Option<isize>, keepdims: bool) -> Result<MaskedArray<T>> {
reduce_lanes(&self.data, &self.mask, axis, keepdims, |vals, masks| {
let mut sum = T::zero();
let mut count = 0usize;
for (v, m) in vals.iter().zip(masks) {
if !*m {
sum = sum + *v;
count += 1;
}
}
if count == 0 {
None
} else {
Some(sum / T::from(count).expect("lane length fits in T"))
}
})
}
pub fn sum_axis(&self, axis: Option<isize>, keepdims: bool) -> Result<MaskedArray<T>> {
reduce_lanes(&self.data, &self.mask, axis, keepdims, |vals, masks| {
let mut sum = T::zero();
let mut count = 0usize;
for (v, m) in vals.iter().zip(masks) {
if !*m {
sum = sum + *v;
count += 1;
}
}
if count == 0 {
None
} else {
Some(sum)
}
})
}
pub fn min_axis(&self, axis: Option<isize>, keepdims: bool) -> Result<MaskedArray<T>> {
reduce_lanes(&self.data, &self.mask, axis, keepdims, |vals, masks| {
let mut best: Option<T> = None;
for (v, m) in vals.iter().zip(masks) {
if !*m {
best = Some(match best {
None => *v,
Some(cur) if *v < cur => *v,
Some(cur) => cur,
});
}
}
best
})
}
pub fn max_axis(&self, axis: Option<isize>, keepdims: bool) -> Result<MaskedArray<T>> {
reduce_lanes(&self.data, &self.mask, axis, keepdims, |vals, masks| {
let mut best: Option<T> = None;
for (v, m) in vals.iter().zip(masks) {
if !*m {
best = Some(match best {
None => *v,
Some(cur) if *v > cur => *v,
Some(cur) => cur,
});
}
}
best
})
}
pub fn var(&self, axis: Option<isize>, ddof: usize, keepdims: bool) -> Result<MaskedArray<T>> {
reduce_lanes(&self.data, &self.mask, axis, keepdims, |vals, masks| {
let mut sum = T::zero();
let mut count = 0usize;
for (v, m) in vals.iter().zip(masks) {
if !*m {
sum = sum + *v;
count += 1;
}
}
if count == 0 {
return None;
}
let n = T::from(count).expect("lane length fits in T");
let mean = sum / n;
let mut sum_sq = T::zero();
for (v, m) in vals.iter().zip(masks) {
if !*m {
let d = *v - mean;
sum_sq = sum_sq + d * d;
}
}
let divisor_n = count.checked_sub(ddof)?;
if divisor_n == 0 {
return None;
}
let divisor = T::from(divisor_n).expect("divisor fits in T");
Some(sum_sq / divisor)
})
}
pub fn std(&self, axis: Option<isize>, ddof: usize, keepdims: bool) -> Result<MaskedArray<T>> {
let variance = self.var(axis, ddof, keepdims)?;
Ok(MaskedArray {
data: variance.data.map(|x| x.sqrt()),
mask: variance.mask,
fill_value: self.fill_value,
})
}
pub fn prod(&self, axis: Option<isize>, keepdims: bool) -> Result<MaskedArray<T>> {
reduce_lanes(&self.data, &self.mask, axis, keepdims, |vals, masks| {
let mut p = T::one();
let mut count = 0usize;
for (v, m) in vals.iter().zip(masks) {
if !*m {
p = p * *v;
count += 1;
}
}
if count == 0 {
None
} else {
Some(p)
}
})
}
pub fn median(&self, axis: Option<isize>, keepdims: bool) -> Result<MaskedArray<T>> {
reduce_lanes(&self.data, &self.mask, axis, keepdims, |vals, masks| {
let mut unmasked: Vec<T> = vals
.iter()
.zip(masks)
.filter(|(_, m)| !**m)
.map(|(v, _)| *v)
.collect();
if unmasked.is_empty() {
return None;
}
if unmasked.iter().any(|v| v.is_nan()) {
return Some(T::nan());
}
unmasked.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal));
let n = unmasked.len();
if n % 2 == 1 {
Some(unmasked[n / 2])
} else {
let two = T::one() + T::one();
Some((unmasked[n / 2 - 1] + unmasked[n / 2]) / two)
}
})
}
pub fn ptp(&self, axis: Option<isize>, keepdims: bool) -> Result<MaskedArray<T>> {
reduce_lanes(&self.data, &self.mask, axis, keepdims, |vals, masks| {
let mut min_v: Option<T> = None;
let mut max_v: Option<T> = None;
for (v, m) in vals.iter().zip(masks) {
if !*m {
min_v = Some(match min_v {
None => *v,
Some(cur) if *v < cur => *v,
Some(cur) => cur,
});
max_v = Some(match max_v {
None => *v,
Some(cur) if *v > cur => *v,
Some(cur) => cur,
});
}
}
match (min_v, max_v) {
(Some(mn), Some(mx)) => Some(mx - mn),
_ => None,
}
})
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::array::Array;
fn ma(data: Vec<f64>, mask: Vec<bool>, shape: &[usize]) -> MaskedArray<f64> {
MaskedArray {
data: Array::from_vec_shape(data, shape).expect("valid shape"),
mask: Array::from_vec_shape(mask, shape).expect("valid shape"),
fill_value: 0.0,
}
}
#[test]
fn mean_axis_none_agrees_with_scalar_mean() {
let m = ma(
vec![1.0, 2.0, 3.0, 4.0, 5.0],
vec![false, true, false, true, false],
&[5],
);
let scalar = m.mean().expect("has unmasked elements");
let axis_form = m.mean_axis(None, false).expect("reduces");
assert!(!axis_form.get_mask().to_vec()[0]);
assert_eq!(axis_form.get_data().to_vec()[0], scalar);
}
#[test]
fn sum_min_max_axis_none_agree_with_scalar_forms() {
let m = ma(
vec![1.0, 2.0, 3.0, 4.0, 5.0],
vec![false, true, false, true, false],
&[5],
);
assert_eq!(
m.sum_axis(None, false)
.expect("reduces")
.get_data()
.to_vec()[0],
m.sum().expect("has data")
);
assert_eq!(
m.min_axis(None, false)
.expect("reduces")
.get_data()
.to_vec()[0],
m.min().expect("has data")
);
assert_eq!(
m.max_axis(None, false)
.expect("reduces")
.get_data()
.to_vec()[0],
m.max().expect("has data")
);
}
#[test]
fn all_masked_axis_none_is_masked_not_none() {
let m = ma(vec![1.0, 2.0, 3.0], vec![true, true, true], &[3]);
assert!(m.mean().is_none());
let r = m.mean_axis(None, false).expect("reduces to a masked slot");
assert!(r.get_mask().to_vec()[0]);
}
#[test]
fn mean_axis_0_matches_numpy_ma() {
let m = ma(
vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0],
vec![false, true, false, false, false, true],
&[2, 3],
);
let r = m.mean_axis(Some(0), false).expect("axis 0 valid");
assert_eq!(r.get_data().to_vec(), vec![2.5, 5.0, 3.0]);
assert_eq!(r.get_mask().to_vec(), vec![false, false, false]);
}
#[test]
fn mean_axis_agrees_on_transposed_non_contiguous_input() {
let base = ma(
vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0],
vec![false, true, false, false, false, true],
&[2, 3],
);
let m = MaskedArray {
data: base.get_data().transpose_axis(0, 1),
mask: base.get_mask().transpose_axis(0, 1),
fill_value: 0.0,
};
assert!(!m.get_data().is_c_contiguous());
let r = m.mean_axis(Some(1), false).expect("axis 1 valid");
assert_eq!(r.get_data().to_vec(), vec![2.5, 5.0, 3.0]);
}
#[test]
fn mean_axis_1_matches_numpy_ma_on_a_3d_array() {
let data: Vec<f64> = (0..24).map(|i| i as f64).collect();
let mut mask = vec![false; 24];
for &(d0, d1, d2) in &[
(0usize, 1usize, 2usize),
(1, 0, 0),
(1, 2, 3),
(0, 0, 3),
(0, 1, 3),
(0, 2, 3),
] {
mask[d0 * 12 + d1 * 4 + d2] = true;
}
let m = ma(data, mask, &[2, 3, 4]);
let r = m.mean_axis(Some(1), false).expect("axis 1 valid");
assert_eq!(r.shape(), vec![2, 4]);
assert_eq!(
r.get_mask().to_vec(),
vec![false, false, false, true, false, false, false, false]
);
let got = r.get_data().to_vec();
let want = [4.0, 5.0, 6.0, 0.0, 18.0, 17.0, 18.0, 17.0];
for (g, w) in got.iter().zip(want.iter()) {
assert!((g - w).abs() < 1e-12, "got {got:?}, want {want:?}");
}
}
#[test]
fn mean_axis_masks_a_fully_masked_lane_but_not_others() {
let m = ma(
vec![1.0, 2.0, 3.0, 4.0],
vec![true, true, false, false],
&[2, 2],
);
let r = m.mean_axis(Some(1), false).expect("axis 1 valid");
assert_eq!(r.get_mask().to_vec(), vec![true, false]);
assert_eq!(r.get_data().to_vec()[1], 3.5);
}
#[test]
fn var_axis_0_matches_numpy_ma() {
let m = ma(
vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0],
vec![false, true, false, false, false, true],
&[2, 3],
);
let r = m.var(Some(0), 0, false).expect("axis 0 valid");
let got = r.get_data().to_vec();
assert!((got[0] - 2.25).abs() < 1e-12);
assert!((got[1] - 0.0).abs() < 1e-12);
assert!((got[2] - 0.0).abs() < 1e-12);
}
#[test]
fn std_axis_1_ddof_1_matches_numpy_ma() {
let m = ma(
vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0],
vec![false, true, false, false, false, true],
&[2, 3],
);
let r = m.std(Some(1), 1, false).expect("axis 1 valid");
let got = r.get_data().to_vec();
assert!((got[0] - std::f64::consts::SQRT_2).abs() < 1e-12);
assert!((got[1] - std::f64::consts::FRAC_1_SQRT_2).abs() < 1e-12);
}
#[test]
fn var_masks_a_lane_with_valid_count_at_or_below_ddof() {
let m = ma(
vec![1.0, 2.0, 3.0, 4.0],
vec![false, true, false, true],
&[2, 2],
);
let r = m.var(Some(1), 1, false).expect("axis 1 valid");
assert_eq!(r.get_mask().to_vec(), vec![true, true]);
}
#[test]
fn median_matches_numpy_ma() {
let m = ma(
vec![5.0, 1.0, 3.0, 2.0, 4.0, 6.0],
vec![false, true, false, false, false, true],
&[2, 3],
);
let r1 = m.median(Some(1), false).expect("axis 1 valid");
assert_eq!(r1.get_data().to_vec(), vec![4.0, 3.0]);
let r0 = m.median(Some(0), false).expect("axis 0 valid");
assert_eq!(r0.get_data().to_vec(), vec![3.5, 4.0, 3.0]);
}
#[test]
fn median_propagates_unmasked_nan_at_every_position() {
let nan = f64::NAN;
for data in [
vec![nan, 5.0, 3.0, 2.0, 4.0],
vec![5.0, nan, 3.0, 2.0, 4.0],
vec![5.0, 3.0, 2.0, 4.0, nan],
] {
let m = ma(data.clone(), vec![false; 5], &[5]);
let r = m.median(None, false).expect("has unmasked data");
assert!(
r.get_data().to_vec()[0].is_nan(),
"data={data:?} should have produced a NaN median"
);
assert!(!r.get_mask().to_vec()[0], "NaN median must be unmasked");
}
for data in [
vec![nan, 1.0, 3.0, 4.0],
vec![1.0, nan, 3.0, 4.0],
vec![1.0, 3.0, 4.0, nan],
] {
let m = ma(data.clone(), vec![false; 4], &[4]);
let r = m.median(None, false).expect("has unmasked data");
assert!(
r.get_data().to_vec()[0].is_nan(),
"data={data:?} should have produced a NaN median"
);
}
}
#[test]
fn median_ignores_nan_under_a_mask() {
let m = ma(vec![1.0, f64::NAN, 3.0], vec![false, true, false], &[3]);
let r = m.median(None, false).expect("has unmasked data");
assert_eq!(r.get_data().to_vec()[0], 2.0);
assert!(!r.get_mask().to_vec()[0]);
}
#[test]
fn ptp_single_valid_element_lane_is_zero_not_masked() {
let m = ma(
vec![1.0, 2.0, 3.0, 4.0],
vec![true, false, false, false],
&[2, 2],
);
let r = m.ptp(Some(1), false).expect("axis 1 valid");
assert_eq!(r.get_mask().to_vec(), vec![false, false]);
assert_eq!(r.get_data().to_vec(), vec![0.0, 1.0]);
}
#[test]
fn ptp_fully_masked_lane_is_masked() {
let m = ma(
vec![1.0, 2.0, 3.0, 4.0],
vec![true, true, false, false],
&[2, 2],
);
let r = m.ptp(Some(1), false).expect("axis 1 valid");
assert_eq!(r.get_mask().to_vec(), vec![true, false]);
assert_eq!(r.get_data().to_vec()[1], 1.0);
}
#[test]
fn prod_skips_masked_like_identity() {
let m = ma(
vec![2.0, 100.0, 3.0, 4.0],
vec![false, true, false, false],
&[4],
);
let r = m.prod(None, false).expect("reduces");
assert_eq!(r.get_data().to_vec()[0], 24.0); }
#[test]
fn negative_axis_matches_positive_equivalent() {
let m = ma(
vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0],
vec![false, true, false, false, false, true],
&[2, 3],
);
let pos = m.mean_axis(Some(1), false).expect("axis 1 valid");
let neg = m.mean_axis(Some(-1), false).expect("axis -1 valid");
assert_eq!(pos.get_data().to_vec(), neg.get_data().to_vec());
assert_eq!(pos.get_mask().to_vec(), neg.get_mask().to_vec());
}
#[test]
fn keepdims_true_keeps_reduced_axis_as_size_one() {
let m = ma(vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0], vec![false; 6], &[2, 3]);
let r = m.sum_axis(Some(1), true).expect("axis 1 valid");
assert_eq!(r.shape(), vec![2, 1]);
let r_none = m.sum_axis(None, true).expect("valid");
assert_eq!(r_none.shape(), vec![1, 1]);
}
#[test]
fn out_of_bounds_axis_is_an_error() {
let m = ma(vec![1.0, 2.0], vec![false, false], &[2]);
assert!(m.mean_axis(Some(1), false).is_err());
assert!(m.mean_axis(Some(-2), false).is_err());
}
}