use crate::ops_view::{axpy, copy_scale};
use crate::{ElementOpApply, RawStridedMut, RawStridedRef, Result};
use core::ops::{Add, Mul};
use crate::maybe_sync::MaybeSendSync;
pub const RAW_FUSED_RANK_LIMIT: usize = 8;
#[derive(Clone, Copy, Debug)]
pub(crate) struct FusedPairLayout {
pub(crate) rank: usize,
pub(crate) dims: [usize; RAW_FUSED_RANK_LIMIT],
pub(crate) dst_strides: [isize; RAW_FUSED_RANK_LIMIT],
pub(crate) src_strides: [isize; RAW_FUSED_RANK_LIMIT],
}
pub(crate) fn fuse_pair_layout(
dims: &[usize],
dst_strides: &[isize],
src_strides: &[isize],
) -> Option<FusedPairLayout> {
if dims.len() > RAW_FUSED_RANK_LIMIT {
return None;
}
let mut layout = FusedPairLayout {
rank: 0,
dims: [1; RAW_FUSED_RANK_LIMIT],
dst_strides: [0; RAW_FUSED_RANK_LIMIT],
src_strides: [0; RAW_FUSED_RANK_LIMIT],
};
for axis in 0..dims.len() {
if dims[axis] == 1 {
continue;
}
if dims[axis] == 0 {
return Some(FusedPairLayout {
rank: 1,
dims: [0; RAW_FUSED_RANK_LIMIT],
dst_strides: [0; RAW_FUSED_RANK_LIMIT],
src_strides: [0; RAW_FUSED_RANK_LIMIT],
});
}
let mut position = layout.rank;
while position > 0 && layout.dst_strides[position - 1] > dst_strides[axis] {
layout.dims[position] = layout.dims[position - 1];
layout.dst_strides[position] = layout.dst_strides[position - 1];
layout.src_strides[position] = layout.src_strides[position - 1];
position -= 1;
}
layout.dims[position] = dims[axis];
layout.dst_strides[position] = dst_strides[axis];
layout.src_strides[position] = src_strides[axis];
layout.rank += 1;
}
if layout.rank == 0 {
layout.rank = 1;
layout.dims[0] = 1;
}
let mut fused = 0usize;
for axis in 1..layout.rank {
let extent = layout.dims[fused] as isize;
if layout.dst_strides[fused].checked_mul(extent) == Some(layout.dst_strides[axis])
&& layout.src_strides[fused].checked_mul(extent) == Some(layout.src_strides[axis])
{
layout.dims[fused] *= layout.dims[axis];
} else {
fused += 1;
layout.dims[fused] = layout.dims[axis];
layout.dst_strides[fused] = layout.dst_strides[axis];
layout.src_strides[fused] = layout.src_strides[axis];
}
}
layout.rank = fused + 1;
Some(layout)
}
pub(crate) fn apply_fused_pair<D, S, Apply, Op>(
dst: &mut RawStridedMut<'_, D>,
src: &RawStridedRef<'_, S>,
layout: &FusedPairLayout,
apply: Apply,
op: Op,
) where
D: Copy,
S: Copy,
Apply: Fn(&mut D, S),
Op: Fn(S) -> S,
{
if layout.dims[..layout.rank].iter().any(|&dim| dim == 0) {
return;
}
let inner_len = layout.dims[0];
let inner_dst = layout.dst_strides[0];
let inner_src = layout.src_strides[0];
let src_data = src.data();
let src_offset = src.offset();
let dst_offset = dst.offset();
let dst_data = dst.data_mut();
let mut index = [0usize; RAW_FUSED_RANK_LIMIT];
let mut dst_base = dst_offset;
let mut src_base = src_offset;
loop {
if inner_dst == 1 && inner_src == 1 {
let dst_start = dst_base as usize;
let src_start = src_base as usize;
let dst_run = &mut dst_data[dst_start..dst_start + inner_len];
let src_run = &src_data[src_start..src_start + inner_len];
for position in 0..inner_len {
apply(&mut dst_run[position], op(src_run[position]));
}
} else {
for position in 0..inner_len {
let dst_position = (dst_base + position as isize * inner_dst) as usize;
let src_position = (src_base + position as isize * inner_src) as usize;
apply(&mut dst_data[dst_position], op(src_data[src_position]));
}
}
let mut axis = 1;
loop {
if axis >= layout.rank {
return;
}
index[axis] += 1;
dst_base += layout.dst_strides[axis];
src_base += layout.src_strides[axis];
if index[axis] < layout.dims[axis] {
break;
}
dst_base -= layout.dims[axis] as isize * layout.dst_strides[axis];
src_base -= layout.dims[axis] as isize * layout.src_strides[axis];
index[axis] = 0;
axis += 1;
}
}
}
fn ensure_same_dims(dst: &[usize], src: &[usize]) -> Result<()> {
if dst != src {
return Err(crate::StridedError::ShapeMismatch(
dst.to_vec(),
src.to_vec(),
));
}
Ok(())
}
pub fn copy_scale_raw<T>(
dest: &mut RawStridedMut<'_, T>,
src: &RawStridedRef<'_, T>,
scale: T,
) -> Result<()>
where
T: Copy + Mul<T, Output = T> + ElementOpApply + MaybeSendSync,
{
ensure_same_dims(dest.dims(), src.dims())?;
match fuse_pair_layout(dest.dims(), dest.strides(), src.strides()) {
Some(layout) => {
apply_fused_pair(
dest,
src,
&layout,
|dst, value| *dst = value,
|value: T| scale * value,
);
Ok(())
}
None => copy_scale(&mut dest.as_view_mut(), &src.as_view(), scale),
}
}
pub fn copy_scale_conj_raw<T>(
dest: &mut RawStridedMut<'_, T>,
src: &RawStridedRef<'_, T>,
scale: T,
) -> Result<()>
where
T: Copy + Mul<T, Output = T> + ElementOpApply + MaybeSendSync,
{
ensure_same_dims(dest.dims(), src.dims())?;
match fuse_pair_layout(dest.dims(), dest.strides(), src.strides()) {
Some(layout) => {
apply_fused_pair(
dest,
src,
&layout,
|dst, value| *dst = value,
|value: T| scale * value.conj(),
);
Ok(())
}
None => copy_scale(&mut dest.as_view_mut(), &src.as_view().conj(), scale),
}
}
pub fn axpy_raw<T>(
dest: &mut RawStridedMut<'_, T>,
src: &RawStridedRef<'_, T>,
alpha: T,
) -> Result<()>
where
T: Copy + Add<T, Output = T> + Mul<T, Output = T> + ElementOpApply + MaybeSendSync,
{
ensure_same_dims(dest.dims(), src.dims())?;
match fuse_pair_layout(dest.dims(), dest.strides(), src.strides()) {
Some(layout) => {
apply_fused_pair(
dest,
src,
&layout,
|dst, value| *dst = *dst + value,
|value: T| alpha * value,
);
Ok(())
}
None => axpy(&mut dest.as_view_mut(), &src.as_view(), alpha),
}
}
pub fn axpy_conj_raw<T>(
dest: &mut RawStridedMut<'_, T>,
src: &RawStridedRef<'_, T>,
alpha: T,
) -> Result<()>
where
T: Copy + Add<T, Output = T> + Mul<T, Output = T> + ElementOpApply + MaybeSendSync,
{
ensure_same_dims(dest.dims(), src.dims())?;
match fuse_pair_layout(dest.dims(), dest.strides(), src.strides()) {
Some(layout) => {
apply_fused_pair(
dest,
src,
&layout,
|dst, value| *dst = *dst + value,
|value: T| alpha * value.conj(),
);
Ok(())
}
None => axpy(&mut dest.as_view_mut(), &src.as_view().conj(), alpha),
}
}
#[cfg(test)]
#[path = "raw_ops/tests/tests.rs"]
mod tests;