use crate::kernel::{
build_plan_fused, build_plan_fused_small, ensure_same_shape, for_each_inner_block_preordered,
total_len, SMALL_TENSOR_THRESHOLD,
};
use crate::layout_check::is_injective_layout;
use crate::maybe_sync::{MaybeSendSync, MaybeSync};
use crate::simd;
use crate::view::{StridedView, StridedViewMut};
use crate::{Result, StridedError};
use strided_view::ElementOp;
#[cfg(feature = "parallel")]
use crate::fuse::compute_costs;
#[cfg(feature = "parallel")]
use crate::threading::{for_each_inner_block_with_offsets, mapreduce_threaded, MINTHREADLENGTH};
#[cfg(feature = "parallel")]
type Body<'a> = &'a (dyn Fn(&[isize], usize, &[isize]) + Sync);
#[cfg(not(feature = "parallel"))]
type Body<'a> = &'a dyn Fn(&[isize], usize, &[isize]);
struct Raw<T>(*mut T);
impl<T> Clone for Raw<T> {
fn clone(&self) -> Self {
*self
}
}
impl<T> Copy for Raw<T> {}
unsafe impl<T> Send for Raw<T> {}
unsafe impl<T> Sync for Raw<T> {}
impl<T> Raw<T> {
fn get(self) -> *mut T {
self.0
}
}
fn validate_destination(dims: &[usize], strides: &[isize]) -> Result<()> {
if is_injective_layout(dims, strides) {
Ok(())
} else {
Err(StridedError::NonInjectiveOutputLayout)
}
}
fn run_update(
dims: &[usize],
strides_list: &[&[isize]],
elem_size: usize,
body: Body<'_>,
) -> Result<()> {
let total = total_len(dims)?;
if total == 0 {
return Ok(());
}
if dims.iter().all(|&d| d == 1) {
let zeros = vec![0isize; strides_list.len()];
body(&zeros, 1, &zeros);
return Ok(());
}
let (fused_dims, ordered_strides, plan) = if total <= SMALL_TENSOR_THRESHOLD {
build_plan_fused_small(dims, strides_list)
} else {
build_plan_fused(dims, strides_list, Some(0), elem_size)
};
#[cfg(feature = "parallel")]
{
let total = total_len(&fused_dims)?;
let nthreads = crate::execution_policy::rayon_threads();
if total > MINTHREADLENGTH && nthreads > 1 {
let costs = compute_costs(&ordered_strides);
let initial_offsets = vec![0isize; strides_list.len()];
return mapreduce_threaded(
&fused_dims,
&plan.block,
&ordered_strides,
&initial_offsets,
&costs,
nthreads,
0,
1,
&|dims, blocks, strides_list, offsets| {
for_each_inner_block_with_offsets(
dims,
blocks,
strides_list,
offsets,
|offsets, len, strides| {
body(offsets, len, strides);
Ok(())
},
)
},
);
}
}
let initial_offsets = vec![0isize; ordered_strides.len()];
for_each_inner_block_preordered(
&fused_dims,
&plan.block,
&ordered_strides,
&initial_offsets,
|offsets, len, strides| {
body(offsets, len, strides);
Ok(())
},
)
}
#[inline(always)]
unsafe fn inner_loop_update1<D: Copy, OpD: ElementOp<D>>(
dp: *mut D,
ds: isize,
len: usize,
f: &impl Fn(D) -> D,
) {
if ds == 1 {
let dst = std::slice::from_raw_parts_mut(dp, len);
simd::dispatch_if_large(len, || {
for d in dst.iter_mut() {
*d = f(OpD::apply(*d));
}
});
} else {
let mut dp = dp;
for _ in 0..len {
*dp = f(OpD::apply(*dp));
dp = dp.wrapping_offset(ds);
}
}
}
#[inline(always)]
unsafe fn inner_loop_update2<D: Copy, A: Copy, OpD: ElementOp<D>, OpA: ElementOp<A>>(
dp: *mut D,
ds: isize,
ap: *const A,
a_s: isize,
len: usize,
f: &impl Fn(D, A) -> D,
) {
if ds == 1 && a_s == 1 {
let src_a = std::slice::from_raw_parts(ap, len);
let dst = std::slice::from_raw_parts_mut(dp, len);
simd::dispatch_if_large(len, || {
for (d, &a) in dst.iter_mut().zip(src_a) {
*d = f(OpD::apply(*d), OpA::apply(a));
}
});
} else {
let (mut dp, mut ap) = (dp, ap);
for _ in 0..len {
*dp = f(OpD::apply(*dp), OpA::apply(*ap));
dp = dp.wrapping_offset(ds);
ap = ap.wrapping_offset(a_s);
}
}
}
#[inline(always)]
#[allow(clippy::too_many_arguments)] unsafe fn inner_loop_update3<
D: Copy,
A: Copy,
B: Copy,
OpD: ElementOp<D>,
OpA: ElementOp<A>,
OpB: ElementOp<B>,
>(
dp: *mut D,
ds: isize,
ap: *const A,
a_s: isize,
bp: *const B,
b_s: isize,
len: usize,
f: &impl Fn(D, A, B) -> D,
) {
if ds == 1 && a_s == 1 && b_s == 1 {
let src_a = std::slice::from_raw_parts(ap, len);
let src_b = std::slice::from_raw_parts(bp, len);
let dst = std::slice::from_raw_parts_mut(dp, len);
simd::dispatch_if_large(len, || {
for ((d, &a), &b) in dst.iter_mut().zip(src_a).zip(src_b) {
*d = f(OpD::apply(*d), OpA::apply(a), OpB::apply(b));
}
});
} else {
let (mut dp, mut ap, mut bp) = (dp, ap, bp);
for _ in 0..len {
*dp = f(OpD::apply(*dp), OpA::apply(*ap), OpB::apply(*bp));
dp = dp.wrapping_offset(ds);
ap = ap.wrapping_offset(a_s);
bp = bp.wrapping_offset(b_s);
}
}
}
pub fn map_update_into<D, OpD>(
dest: &mut StridedViewMut<D>,
f: impl Fn(D) -> D + MaybeSync,
) -> Result<()>
where
D: Copy + MaybeSendSync,
OpD: ElementOp<D>,
{
validate_destination(dest.dims(), dest.strides())?;
let dp = Raw(dest.as_mut_ptr());
run_update(
dest.dims(),
&[dest.strides()],
std::mem::size_of::<D>(),
&|offsets, len, strides| {
unsafe {
inner_loop_update1::<D, OpD>(dp.get().offset(offsets[0]), strides[0], len, &f);
}
},
)
}
pub fn zip_update2_into<D, A, OpD, OpA>(
dest: &mut StridedViewMut<D>,
a: &StridedView<A, OpA>,
f: impl Fn(D, A) -> D + MaybeSync,
) -> Result<()>
where
D: Copy + MaybeSendSync,
A: Copy + MaybeSendSync,
OpD: ElementOp<D>,
OpA: ElementOp<A>,
{
ensure_same_shape(dest.dims(), a.dims())?;
validate_destination(dest.dims(), dest.strides())?;
let dp = Raw(dest.as_mut_ptr());
let ap = Raw(a.ptr() as *mut A);
run_update(
dest.dims(),
&[dest.strides(), a.strides()],
std::mem::size_of::<D>().max(std::mem::size_of::<A>()),
&|offsets, len, strides| {
unsafe {
inner_loop_update2::<D, A, OpD, OpA>(
dp.get().offset(offsets[0]),
strides[0],
ap.get().offset(offsets[1]).cast_const(),
strides[1],
len,
&f,
);
}
},
)
}
pub fn zip_update3_into<D, A, B, OpD, OpA, OpB>(
dest: &mut StridedViewMut<D>,
a: &StridedView<A, OpA>,
b: &StridedView<B, OpB>,
f: impl Fn(D, A, B) -> D + MaybeSync,
) -> Result<()>
where
D: Copy + MaybeSendSync,
A: Copy + MaybeSendSync,
B: Copy + MaybeSendSync,
OpD: ElementOp<D>,
OpA: ElementOp<A>,
OpB: ElementOp<B>,
{
ensure_same_shape(dest.dims(), a.dims())?;
ensure_same_shape(dest.dims(), b.dims())?;
validate_destination(dest.dims(), dest.strides())?;
let dp = Raw(dest.as_mut_ptr());
let ap = Raw(a.ptr() as *mut A);
let bp = Raw(b.ptr() as *mut B);
run_update(
dest.dims(),
&[dest.strides(), a.strides(), b.strides()],
std::mem::size_of::<D>()
.max(std::mem::size_of::<A>())
.max(std::mem::size_of::<B>()),
&|offsets, len, strides| {
unsafe {
inner_loop_update3::<D, A, B, OpD, OpA, OpB>(
dp.get().offset(offsets[0]),
strides[0],
ap.get().offset(offsets[1]).cast_const(),
strides[1],
bp.get().offset(offsets[2]).cast_const(),
strides[2],
len,
&f,
);
}
},
)
}
#[cfg(test)]
#[path = "update_view/tests/tests.rs"]
mod tests;