mod bool_ops;
mod linalg;
mod ops;
mod reductions;
mod search;
use crate::array::Array;
use crate::error::{NumRs2Error, Result};
use std::fmt;
use std::fmt::Debug;
#[derive(Clone)]
pub struct MaskedArray<T> {
data: Array<T>,
mask: Array<bool>,
fill_value: T,
}
impl<T: Clone> MaskedArray<T> {
pub fn new(data: Array<T>, mask: Option<Array<bool>>, fill_value: Option<T>) -> Result<Self>
where
T: Clone + Default,
{
let shape = data.shape();
let mask_array = match mask {
Some(m) => {
if m.shape() != shape {
return Err(NumRs2Error::ShapeMismatch {
expected: shape,
actual: m.shape(),
});
}
m
}
None => Array::from_vec_shape(vec![false; data.size()], &shape)?,
};
let fill_val = fill_value.unwrap_or_default();
Ok(Self {
data,
mask: mask_array,
fill_value: fill_val,
})
}
pub fn masked_values(data: Array<T>, value: T, fill_value: Option<T>) -> Result<Self>
where
T: Clone + Default + PartialEq,
{
let shape = data.shape();
let mut mask_vec = Vec::with_capacity(data.size());
for elem in data.array().iter() {
mask_vec.push(*elem == value);
}
let mask_array = Array::from_vec_shape(mask_vec, &shape)?;
let fill_val = fill_value.unwrap_or_default();
Ok(Self {
data: data.clone(),
mask: mask_array,
fill_value: fill_val,
})
}
pub fn masked_invalid(data: Array<f64>, fill_value: Option<f64>) -> Result<MaskedArray<f64>> {
let shape = data.shape();
let mut mask_vec = Vec::with_capacity(data.size());
for &elem in data.array().iter() {
mask_vec.push(elem.is_nan() || elem.is_infinite());
}
let mask_array = Array::from_vec_shape(mask_vec, &shape)?;
let fill_val = fill_value.unwrap_or(0.0);
Ok(MaskedArray {
data: data.clone(),
mask: mask_array,
fill_value: fill_val,
})
}
pub fn masked_where(
data: Array<T>,
condition: Array<bool>,
fill_value: Option<T>,
) -> Result<Self>
where
T: Clone + Default,
{
if data.shape() != condition.shape() {
return Err(NumRs2Error::ShapeMismatch {
expected: data.shape(),
actual: condition.shape(),
});
}
let fill_val = fill_value.unwrap_or_default();
Ok(Self {
data: data.clone(),
mask: condition,
fill_value: fill_val,
})
}
pub fn masked_all(data: Array<T>, fill_value: Option<T>) -> Result<Self>
where
T: Clone + Default,
{
let shape = data.shape();
let mask_array = Array::from_vec_shape(vec![true; data.size()], &shape)?;
let fill_val = fill_value.unwrap_or_default();
Ok(Self {
data: data.clone(),
mask: mask_array,
fill_value: fill_val,
})
}
pub fn get_data(&self) -> &Array<T> {
&self.data
}
pub fn get_mask(&self) -> &Array<bool> {
&self.mask
}
pub fn get_fill_value(&self) -> T {
self.fill_value.clone()
}
pub fn set_fill_value(&mut self, value: T) {
self.fill_value = value;
}
pub fn shape(&self) -> Vec<usize> {
self.data.shape()
}
pub fn ndim(&self) -> usize {
self.data.ndim()
}
pub fn size(&self) -> usize {
self.data.size()
}
pub fn count_masked(&self) -> usize {
self.mask.array().iter().filter(|&&x| x).count()
}
pub fn count_valid(&self) -> usize {
self.size() - self.count_masked()
}
pub fn filled(&self, fill_value: Option<T>) -> Array<T>
where
T: Clone,
{
let fill_val = fill_value.unwrap_or_else(|| self.fill_value.clone());
let data_op = crate::kernels::borrow::operand(&self.data);
let mask_op = crate::kernels::borrow::operand(&self.mask);
let mut filled_vec = Vec::with_capacity(self.size());
for (value, is_masked) in data_op.iter().zip(mask_op.iter()) {
if *is_masked {
filled_vec.push(fill_val.clone());
} else {
filled_vec.push(value.clone());
}
}
Array::from_vec_shape(filled_vec, &self.shape()).unwrap_or_else(|e| panic!("{e}"))
}
pub fn compressed(&self) -> Array<T>
where
T: Clone,
{
let data_op = crate::kernels::borrow::operand(&self.data);
let mask_op = crate::kernels::borrow::operand(&self.mask);
let mut compressed_vec = Vec::new();
for (value, is_masked) in data_op.iter().zip(mask_op.iter()) {
if !*is_masked {
compressed_vec.push(value.clone());
}
}
Array::from_vec(compressed_vec)
}
pub fn harden_mask(&self) -> Self
where
T: Clone,
{
self.clone()
}
pub fn soften_mask(&self) -> Self
where
T: Clone,
{
self.clone()
}
pub fn get(&self, indices: &[usize]) -> Result<T>
where
T: Clone,
{
if indices.len() != self.ndim() {
return Err(NumRs2Error::DimensionMismatch(format!(
"Expected {} indices, got {}",
self.ndim(),
indices.len()
)));
}
for (i, &idx) in indices.iter().enumerate() {
if idx >= self.shape()[i] {
return Err(NumRs2Error::IndexOutOfBounds(format!(
"Index {} out of bounds for dimension {} with size {}",
idx,
i,
self.shape()[i]
)));
}
}
let mask_array = self.mask.array();
let mask_value = mask_array.get(indices).ok_or_else(|| {
NumRs2Error::IndexOutOfBounds(format!("Failed to get mask at indices {:?}", indices))
})?;
if *mask_value {
Ok(self.fill_value.clone())
} else {
let data_array = self.data.array();
let data_value = data_array.get(indices).ok_or_else(|| {
NumRs2Error::IndexOutOfBounds(format!(
"Failed to get data at indices {:?}",
indices
))
})?;
Ok(data_value.clone())
}
}
pub fn set(&mut self, indices: &[usize], value: T, mask: Option<bool>) -> Result<()>
where
T: Clone,
{
self.data.set(indices, value)?;
if let Some(mask_value) = mask {
self.mask.set(indices, mask_value)?;
}
Ok(())
}
pub fn reshape(&self, shape: &[usize]) -> Self
where
T: Clone,
{
MaskedArray {
data: self.data.reshape(shape),
mask: self.mask.reshape(shape),
fill_value: self.fill_value.clone(),
}
}
pub fn transpose(&self) -> Self
where
T: Clone,
{
MaskedArray {
data: self.data.transpose(),
mask: self.mask.transpose(),
fill_value: self.fill_value.clone(),
}
}
}
impl<T: Clone + fmt::Display + Debug> fmt::Display for MaskedArray<T> {
fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
let data_op = crate::kernels::borrow::operand(&self.data);
let mask_op = crate::kernels::borrow::operand(&self.mask);
let shape = self.shape();
writeln!(f, "MaskedArray(")?;
if shape.len() == 1 {
write!(f, "[")?;
for (i, (val, &masked)) in data_op.iter().zip(mask_op.iter()).enumerate() {
if i > 0 {
write!(f, ", ")?;
}
if masked {
write!(f, "--")?;
} else {
write!(f, "{}", val)?;
}
}
writeln!(f, "]")?;
} else {
writeln!(f, "Shape: {:?}", shape)?;
writeln!(f, "Masked count: {}", self.count_masked())?;
}
write!(f, "Fill value: {}", self.fill_value)?;
Ok(())
}
}
impl<T: Clone + Debug> fmt::Debug for MaskedArray<T> {
fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
f.debug_struct("MaskedArray")
.field("shape", &self.shape())
.field("masked_count", &self.count_masked())
.field("fill_value", &self.fill_value)
.finish()
}
}