use core::mem::MaybeUninit;
use core::ops::{Add, Mul};
use crate::{MaybeSendSync, Result, StridedError, StridedView, StridedViewMut};
fn same_len(actual: usize, expected: usize) -> Result<()> {
if actual == expected {
Ok(())
} else {
Err(StridedError::ShapeMismatch(vec![actual], vec![expected]))
}
}
fn product(shape: &[usize]) -> Result<usize> {
if shape.contains(&0) {
return Ok(0);
}
shape.iter().try_fold(1usize, |n, &d| {
n.checked_mul(d).ok_or(StridedError::OffsetOverflow)
})
}
fn for_each_chunk<T: MaybeSendSync, F: Fn(usize, &mut [T]) + crate::MaybeSync>(
output: &mut [T],
operation: F,
) {
#[cfg(feature = "parallel")]
{
let threads = crate::threading::parallel_threads_for_len(output.len());
if threads > 1 {
let ptr = crate::threading::SendPtr(output.as_mut_ptr());
crate::threading::parallel_for_each(0..output.len(), threads, &|range| {
let chunk = unsafe {
core::slice::from_raw_parts_mut(ptr.as_ptr().add(range.start), range.len())
};
operation(range.start, chunk);
});
return;
}
}
operation(0, output);
}
pub fn axpby_accum<T>(y: &mut [T], x: &[T], alpha: T, beta: T) -> Result<()>
where
T: Copy + Send + Sync + Add<Output = T> + Mul<Output = T>,
{
same_len(x.len(), y.len())?;
for_each_chunk(y, |start, dst| {
let len = dst.len();
for (out, &src) in dst.iter_mut().zip(&x[start..start + len]) {
*out = alpha * src + beta * *out;
}
});
Ok(())
}
pub fn triangular_mask_into_uninit<T: Copy + MaybeSendSync>(
output: &mut [MaybeUninit<T>],
input: &[T],
shape: &[usize],
k: i64,
upper: bool,
fill: T,
) -> Result<()> {
if shape.len() < 2 {
return Err(StridedError::RankMismatch(shape.len(), 2));
}
let len = product(shape)?;
same_len(input.len(), len)?;
same_len(output.len(), len)?;
if len == 0 {
return Ok(());
}
let rows = shape[0];
let cols = shape[1];
for_each_chunk(output, |start, chunk| {
let mut pos = 0;
while pos < chunk.len() {
let flat = start + pos;
let row = flat % rows;
let col = (flat / rows) % cols;
let count = (rows - row).min(chunk.len() - pos);
let boundary = col as i128 - k as i128;
let split = if upper {
boundary.saturating_add(1)
} else {
boundary
}
.clamp(row as i128, (row + count) as i128) as usize
- row;
let (low, high) = chunk[pos..pos + count].split_at_mut(split);
let (kept, source, masked) = if upper {
(low, &input[flat..flat + split], high)
} else {
(high, &input[flat + split..flat + count], low)
};
for (dst, &src) in kept.iter_mut().zip(source) {
dst.write(src);
}
masked.fill(MaybeUninit::new(fill));
pos += count;
}
});
Ok(())
}
pub fn embed_diagonal_into_uninit<T: Copy + MaybeSendSync + 'static>(
output: &mut [MaybeUninit<T>],
input: &[T],
shape: &[usize],
axis: usize,
insert_axis: usize,
zero: T,
) -> Result<()> {
if axis >= shape.len() {
return Err(StridedError::InvalidAxis {
axis,
rank: shape.len(),
});
}
if insert_axis > shape.len() {
return Err(StridedError::InvalidAxis {
axis: insert_axis,
rank: shape.len() + 1,
});
}
let len = product(shape)?;
same_len(input.len(), len)?;
let out_len = len
.checked_mul(shape[axis])
.ok_or(StridedError::OffsetOverflow)?;
same_len(output.len(), out_len)?;
if out_len == 0 {
return Ok(());
}
isize::try_from(out_len).map_err(|_| StridedError::OffsetOverflow)?;
let mut out_shape = shape.to_vec();
out_shape.insert(insert_axis, shape[axis]);
let src_strides = crate::col_major_strides(shape);
let mut dst_strides = crate::col_major_strides(&out_shape);
let inserted_stride = dst_strides.remove(insert_axis);
dst_strides[axis] = dst_strides[axis]
.checked_add(inserted_stride)
.ok_or(StridedError::OffsetOverflow)?;
let src = StridedView::<T>::new(input, shape, &src_strides, 0)?;
let mut dst = StridedViewMut::new(output, shape, &dst_strides, 0)?;
crate::map_view::validate_destination_layout_without_alloc(shape, &dst_strides)?;
for_each_chunk(dst.data_mut(), |_, chunk| {
chunk.fill(MaybeUninit::new(zero))
});
crate::map_into(&mut dst, &src, MaybeUninit::new)
}