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 merged = isize::try_from(layout.dims[fused])
.ok()
.filter(|&extent| {
layout.dst_strides[fused].checked_mul(extent) == Some(layout.dst_strides[axis])
&& layout.src_strides[fused].checked_mul(extent)
== Some(layout.src_strides[axis])
})
.and_then(|_| layout.dims[fused].checked_mul(layout.dims[axis]));
if let Some(merged) = merged {
layout.dims[fused] = merged;
} 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)
}
#[inline]
pub(crate) fn fused_total(layout: &FusedPairLayout) -> usize {
layout.dims[..layout.rank].iter().product()
}
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,
{
let total = fused_total(layout);
if total == 0 {
return;
}
let src_ptr = src.data().as_ptr();
let src_base = src.offset();
let dst_base = dst.offset();
let dst_ptr = dst.data_mut().as_mut_ptr();
unsafe {
apply_fused_range(
dst_ptr, dst_base, src_ptr, src_base, layout, 0, total, &apply, &op,
);
}
}
#[allow(clippy::too_many_arguments)]
pub(crate) unsafe fn apply_fused_range<D, S, Apply, Op>(
dst_ptr: *mut D,
dst_base: isize,
src_ptr: *const S,
src_base: isize,
layout: &FusedPairLayout,
start: usize,
len: usize,
apply: &Apply,
op: &Op,
) where
D: Copy,
S: Copy,
Apply: Fn(&mut D, S),
Op: Fn(S) -> S,
{
if len == 0 {
return;
}
let rank = layout.rank;
let inner_len = layout.dims[0];
let inner_dst = layout.dst_strides[0];
let inner_src = layout.src_strides[0];
let mut index = [0usize; RAW_FUSED_RANK_LIMIT];
let mut rest = start;
let mut dst_outer = dst_base;
let mut src_outer = src_base;
for axis in 0..rank {
let dim = layout.dims[axis];
index[axis] = rest % dim;
rest /= dim;
if axis > 0 {
dst_outer += index[axis] as isize * layout.dst_strides[axis];
src_outer += index[axis] as isize * layout.src_strides[axis];
}
}
let mut inner_start = index[0];
let mut remaining = len;
loop {
let run = (inner_len - inner_start).min(remaining);
let dst_run = dst_outer + inner_start as isize * inner_dst;
let src_run = src_outer + inner_start as isize * inner_src;
unsafe {
apply_run(
dst_ptr, dst_run, inner_dst, src_ptr, src_run, inner_src, run, apply, op,
)
};
remaining -= run;
if remaining == 0 {
return;
}
inner_start = 0;
let mut axis = 1;
while axis < rank {
if index[axis] + 1 < layout.dims[axis] {
index[axis] += 1;
dst_outer += layout.dst_strides[axis];
src_outer += layout.src_strides[axis];
break;
}
let last = (layout.dims[axis] - 1) as isize;
dst_outer -= last * layout.dst_strides[axis];
src_outer -= last * layout.src_strides[axis];
index[axis] = 0;
axis += 1;
}
}
}
#[allow(clippy::too_many_arguments)]
#[inline(always)]
unsafe fn apply_run<D, S, Apply, Op>(
dst_ptr: *mut D,
dst_offset: isize,
dst_stride: isize,
src_ptr: *const S,
src_offset: isize,
src_stride: isize,
len: usize,
apply: &Apply,
op: &Op,
) where
D: Copy,
S: Copy,
Apply: Fn(&mut D, S),
Op: Fn(S) -> S,
{
if dst_stride == 1 {
let dst_run = unsafe { core::slice::from_raw_parts_mut(dst_ptr.offset(dst_offset), len) };
match src_stride {
1 => {
let src_run =
unsafe { core::slice::from_raw_parts(src_ptr.offset(src_offset), len) };
for (dst, &value) in dst_run.iter_mut().zip(src_run) {
apply(dst, op(value));
}
}
-1 => {
let src_run = unsafe {
core::slice::from_raw_parts(
src_ptr.offset(src_offset - (len as isize - 1)),
len,
)
};
for (dst, &value) in dst_run.iter_mut().zip(src_run.iter().rev()) {
apply(dst, op(value));
}
}
_ => {
let src_start = unsafe { src_ptr.offset(src_offset) };
for (position, dst) in dst_run.iter_mut().enumerate() {
let value = unsafe { *src_start.offset(position as isize * src_stride) };
apply(dst, op(value));
}
}
}
return;
}
let dst_start = unsafe { dst_ptr.offset(dst_offset) };
let src_start = unsafe { src_ptr.offset(src_offset) };
for position in 0..len as isize {
unsafe {
let value = *src_start.offset(position * src_stride);
apply(&mut *dst_start.offset(position * dst_stride), op(value));
}
}
}
fn ensure_same_dims(dst: &[usize], src: &[usize]) -> Result<()> {
if dst != src {
return Err(crate::StridedError::ShapeMismatch(
dst.to_vec(),
src.to_vec(),
));
}
crate::kernel::total_len(dst)?;
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;