use crate::error::{ADError, ADResult};
use crate::ndarray::{NDArray, NDDataBuffer, NDDataType, NDDimension};
trait BinAcc: Copy {
const ZERO: Self;
fn bin_add(self, rhs: Self) -> Self;
}
impl BinAcc for i128 {
const ZERO: Self = 0;
#[inline]
fn bin_add(self, rhs: Self) -> Self {
self.wrapping_add(rhs)
}
}
impl BinAcc for f32 {
const ZERO: Self = 0.0;
#[inline]
fn bin_add(self, rhs: Self) -> Self {
self + rhs
}
}
impl BinAcc for f64 {
const ZERO: Self = 0.0;
#[inline]
fn bin_add(self, rhs: Self) -> Self {
self + rhs
}
}
pub fn convert_type(src: &NDArray, target_type: NDDataType) -> ADResult<NDArray> {
if src.data.data_type() == target_type {
return Ok(src.clone());
}
macro_rules! cast_all {
($v:expr) => {
match target_type {
NDDataType::Int8 => NDDataBuffer::I8($v.iter().map(|&x| x as i8).collect()),
NDDataType::UInt8 => NDDataBuffer::U8($v.iter().map(|&x| x as u8).collect()),
NDDataType::Int16 => NDDataBuffer::I16($v.iter().map(|&x| x as i16).collect()),
NDDataType::UInt16 => NDDataBuffer::U16($v.iter().map(|&x| x as u16).collect()),
NDDataType::Int32 => NDDataBuffer::I32($v.iter().map(|&x| x as i32).collect()),
NDDataType::UInt32 => NDDataBuffer::U32($v.iter().map(|&x| x as u32).collect()),
NDDataType::Int64 => NDDataBuffer::I64($v.iter().map(|&x| x as i64).collect()),
NDDataType::UInt64 => NDDataBuffer::U64($v.iter().map(|&x| x as u64).collect()),
NDDataType::Float32 => NDDataBuffer::F32($v.iter().map(|&x| x as f32).collect()),
NDDataType::Float64 => NDDataBuffer::F64($v.iter().map(|&x| x as f64).collect()),
}
};
}
let out_data = match &src.data {
NDDataBuffer::I8(v) => cast_all!(v),
NDDataBuffer::U8(v) => cast_all!(v),
NDDataBuffer::I16(v) => cast_all!(v),
NDDataBuffer::U16(v) => cast_all!(v),
NDDataBuffer::I32(v) => cast_all!(v),
NDDataBuffer::U32(v) => cast_all!(v),
NDDataBuffer::I64(v) => cast_all!(v),
NDDataBuffer::U64(v) => cast_all!(v),
NDDataBuffer::F32(v) => cast_all!(v),
NDDataBuffer::F64(v) => cast_all!(v),
};
let mut arr = NDArray::new(src.dims.clone(), target_type);
arr.data = out_data;
arr.unique_id = src.unique_id;
arr.timestamp = src.timestamp;
arr.time_stamp = src.time_stamp;
arr.attributes = src.attributes.clone();
arr.codec = src.codec.clone();
Ok(arr)
}
pub fn convert_dims(
src: &NDArray,
dims_out: &[NDDimension],
target_type: NDDataType,
) -> ADResult<NDArray> {
if src.codec.is_some() {
return Err(ADError::UnsupportedConversion(
"convert: cannot convert compressed (codec) data".into(),
));
}
let ndims = src.dims.len();
if dims_out.len() != ndims {
return Err(ADError::InvalidDimensions(format!(
"convert: dims_out length {} != source ndims {}",
dims_out.len(),
ndims,
)));
}
let mut out_sizes = Vec::with_capacity(ndims);
for (i, d) in dims_out.iter().enumerate() {
let bin = d.binning.max(1);
if d.size == 0 {
return Err(ADError::InvalidDimensions(format!(
"convert: dims_out[{}].size is 0",
i,
)));
}
let out_size = d.size / bin;
if out_size == 0 {
return Err(ADError::InvalidDimensions(format!(
"convert: dims_out[{}] size {} / binning {} = 0",
i, d.size, bin,
)));
}
if d.offset + d.size > src.dims[i].size {
return Err(ADError::InvalidDimensions(format!(
"convert: dims_out[{}] offset {} + size {} > src dim size {}",
i, d.offset, d.size, src.dims[i].size,
)));
}
out_sizes.push(out_size);
}
let mut out_dims = Vec::with_capacity(ndims);
for i in 0..ndims {
let bin = dims_out[i].binning.max(1);
out_dims.push(NDDimension {
size: out_sizes[i],
offset: src.dims[i].offset + dims_out[i].offset,
binning: src.dims[i].binning * bin,
reverse: dims_out[i].reverse ^ src.dims[i].reverse,
});
}
let total_out: usize = out_sizes.iter().product();
let mut src_strides = vec![1usize; ndims];
for i in 1..ndims {
src_strides[i] = src_strides[i - 1] * src.dims[i - 1].size;
}
let mut out_strides = vec![1usize; ndims];
for i in 1..ndims {
out_strides[i] = out_strides[i - 1] * out_sizes[i - 1];
}
macro_rules! bin_loop {
($src_vec:expr, $DstT:ty, $AccT:ty, $variant:ident) => {{
let mut out = vec![0 as $DstT; total_out];
for out_idx in 0..total_out {
let mut remaining = out_idx;
let mut out_coords = [0usize; 10]; for i in (0..ndims).rev() {
out_coords[i] = remaining / out_strides[i];
remaining %= out_strides[i];
}
let mut eff_coords = [0usize; 10];
for i in 0..ndims {
eff_coords[i] = if dims_out[i].reverse {
out_sizes[i] - 1 - out_coords[i]
} else {
out_coords[i]
};
}
let mut acc = <$AccT as BinAcc>::ZERO;
let bin_total: usize = dims_out.iter().map(|d| d.binning.max(1)).product();
for bin_flat in 0..bin_total {
let mut br = bin_flat;
let mut src_flat = 0usize;
let mut valid = true;
for i in (0..ndims).rev() {
let bin = dims_out[i].binning.max(1);
let bin_off = br % bin;
br /= bin;
let src_coord = dims_out[i].offset + eff_coords[i] * bin + bin_off;
if src_coord >= src.dims[i].size {
valid = false;
break;
}
src_flat += src_coord * src_strides[i];
}
if valid {
acc = acc.bin_add($src_vec[src_flat] as $AccT);
}
}
out[out_idx] = acc as $DstT;
}
NDDataBuffer::$variant(out)
}};
}
macro_rules! bin_to_target {
($src_vec:expr) => {
match target_type {
NDDataType::Int8 => bin_loop!($src_vec, i8, i128, I8),
NDDataType::UInt8 => bin_loop!($src_vec, u8, i128, U8),
NDDataType::Int16 => bin_loop!($src_vec, i16, i128, I16),
NDDataType::UInt16 => bin_loop!($src_vec, u16, i128, U16),
NDDataType::Int32 => bin_loop!($src_vec, i32, i128, I32),
NDDataType::UInt32 => bin_loop!($src_vec, u32, i128, U32),
NDDataType::Int64 => bin_loop!($src_vec, i64, i128, I64),
NDDataType::UInt64 => bin_loop!($src_vec, u64, i128, U64),
NDDataType::Float32 => bin_loop!($src_vec, f32, f32, F32),
NDDataType::Float64 => bin_loop!($src_vec, f64, f64, F64),
}
};
}
let out_data = match &src.data {
NDDataBuffer::I8(v) => bin_to_target!(v),
NDDataBuffer::U8(v) => bin_to_target!(v),
NDDataBuffer::I16(v) => bin_to_target!(v),
NDDataBuffer::U16(v) => bin_to_target!(v),
NDDataBuffer::I32(v) => bin_to_target!(v),
NDDataBuffer::U32(v) => bin_to_target!(v),
NDDataBuffer::I64(v) => bin_to_target!(v),
NDDataBuffer::U64(v) => bin_to_target!(v),
NDDataBuffer::F32(v) => bin_to_target!(v),
NDDataBuffer::F64(v) => bin_to_target!(v),
};
let mut arr = NDArray::new(out_dims, target_type);
arr.data = out_data;
arr.timestamp = src.timestamp;
arr.time_stamp = src.time_stamp;
arr.attributes.copy_from(&src.attributes);
Ok(arr)
}