use super::reductions::reduce_lanes;
use super::MaskedArray;
use crate::error::Result;
impl MaskedArray<bool> {
pub fn any(&self, axis: Option<isize>, keepdims: bool) -> Result<MaskedArray<bool>> {
reduce_lanes(
self.get_data(),
self.get_mask(),
axis,
keepdims,
|vals, masks| {
let mut any_valid = false;
let mut result = false;
for (v, m) in vals.iter().zip(masks) {
if !*m {
any_valid = true;
result |= *v;
}
}
any_valid.then_some(result)
},
)
}
pub fn all(&self, axis: Option<isize>, keepdims: bool) -> Result<MaskedArray<bool>> {
reduce_lanes(
self.get_data(),
self.get_mask(),
axis,
keepdims,
|vals, masks| {
let mut any_valid = false;
let mut result = true;
for (v, m) in vals.iter().zip(masks) {
if !*m {
any_valid = true;
result &= *v;
}
}
any_valid.then_some(result)
},
)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::array::Array;
fn mb(data: Vec<bool>, mask: Vec<bool>, shape: &[usize]) -> MaskedArray<bool> {
MaskedArray {
data: Array::from_vec_shape(data, shape).expect("valid shape"),
mask: Array::from_vec_shape(mask, shape).expect("valid shape"),
fill_value: false,
}
}
#[test]
fn any_all_skip_masked_elements() {
let m = mb(vec![true, true, false], vec![true, false, false], &[3]);
let any = m.any(None, false).expect("reduces");
assert!(!any.get_mask().to_vec()[0]);
assert!(any.get_data().to_vec()[0]);
let all = m.all(None, false).expect("reduces");
assert!(!all.get_data().to_vec()[0]);
}
#[test]
fn any_false_when_only_unmasked_element_is_false() {
let m = mb(vec![true, true, false], vec![true, true, false], &[3]);
let any = m.any(None, false).expect("reduces");
assert!(!any.get_mask().to_vec()[0]);
assert!(!any.get_data().to_vec()[0]);
}
#[test]
fn any_all_fully_masked_is_masked() {
let m = mb(vec![true, true, false], vec![true, true, true], &[3]);
assert!(m.any(None, false).expect("reduces").get_mask().to_vec()[0]);
assert!(m.all(None, false).expect("reduces").get_mask().to_vec()[0]);
}
#[test]
fn any_all_axis_0() {
let m = mb(
vec![true, false, false, false],
vec![true, false, false, false],
&[2, 2],
);
let any = m.any(Some(0), false).expect("axis 0 valid");
assert_eq!(any.get_data().to_vec(), vec![false, false]);
assert_eq!(any.get_mask().to_vec(), vec![false, false]);
let all = m.all(Some(0), false).expect("axis 0 valid");
assert_eq!(all.get_data().to_vec(), vec![false, false]);
}
#[test]
fn keepdims_preserves_ndim() {
let m = mb(vec![true, false, true, false], vec![false; 4], &[2, 2]);
let r = m.any(Some(1), true).expect("axis 1 valid");
assert_eq!(r.shape(), vec![2, 1]);
}
}