use ndarray::Array1;
use num_traits::{Float, FromPrimitive, One, Zero};
pub fn correlate1d<T>(
input: &Vec<T>,
weights: &Array1<f64>,
axis: isize,
mode: &str,
cval: f64,
origin: isize,
) -> Result<Vec<T>, String> where
T: Float + FromPrimitive + Zero + Clone + One,
{
let ndim = 1;
let _axis = normalize_axis_index(axis, ndim)?;
if weights.ndim() != 1 || weights.len() < 1 {
return Err("no filter weights given".into());
}
if invalid_origin(origin, weights.len()) {
return Err(format!(
"Invalid origin; must satisfy -(len(weights)//2) <= origin <= (len(weights)-1)//2"
));
}
let w_len = weights.len();
let origin_offset = origin;
let half = (w_len / 2) as isize;
let input_len = input.len();
let mut output = vec![T::zero(); input_len];
for i in 0..input_len {
let mut acc = T::zero();
for j in 0..w_len {
let offset = j as isize - half + origin_offset;
let idx = i as isize + offset;
let value = match mode {
"constant" => get_constant(input, idx, T::from_f64(cval).unwrap()),
"reflect" => get_reflect(input, idx),
_ => return Err(format!("Unsupported mode: {}", mode)),
};
acc = acc + T::from_f64(weights[j]).unwrap() * value;
}
output[i] = acc;
}
Ok(output)
}
fn normalize_axis_index(axis: isize, ndim: usize) -> Result<usize, String> {
if axis < -(ndim as isize) || axis >= ndim as isize {
Err(format!("axis {} is out of bounds for array of dimension {}", axis, ndim))
} else if axis < 0 {
Ok((axis + ndim as isize) as usize)
} else {
Ok(axis as usize)
}
}
fn invalid_origin(origin: isize, lenw: usize) -> bool {
origin < -((lenw as isize) / 2) || origin > ((lenw as isize - 1) / 2)
}
fn get_reflect<T>(input: &Vec<T>, idx: isize) -> T where
T: Float + FromPrimitive + Zero + Clone + One,
{
let len = input.len() as isize;
if idx < 0 {
input[(-idx - 1) as usize]
} else if idx >= len {
input[(2 * len - idx - 1) as usize]
} else {
input[idx as usize]
}
}
fn get_constant<T>(input: &Vec<T>, idx: isize, cval: T) -> T where
T: Float + FromPrimitive + Zero + Clone + One,
{
if idx < 0 || idx >= input.len() as isize {
cval
} else {
input[idx as usize]
}
}