1use crate::kernel::{
22 build_plan_fused, build_plan_fused_small, ensure_same_shape, for_each_inner_block_preordered,
23 total_len, SMALL_TENSOR_THRESHOLD,
24};
25use crate::layout_check::is_injective_layout;
26use crate::maybe_sync::{MaybeSendSync, MaybeSync};
27use crate::simd;
28use crate::view::{StridedView, StridedViewMut};
29use crate::{Result, StridedError};
30use strided_view::ElementOp;
31
32#[cfg(feature = "parallel")]
33use crate::fuse::compute_costs;
34#[cfg(feature = "parallel")]
35use crate::threading::{for_each_inner_block_with_offsets, mapreduce_threaded, MINTHREADLENGTH};
36
37#[cfg(feature = "parallel")]
39type Body<'a> = &'a (dyn Fn(&[isize], usize, &[isize]) + Sync);
40#[cfg(not(feature = "parallel"))]
41type Body<'a> = &'a dyn Fn(&[isize], usize, &[isize]);
42
43struct Raw<T>(*mut T);
46impl<T> Clone for Raw<T> {
47 fn clone(&self) -> Self {
48 *self
49 }
50}
51impl<T> Copy for Raw<T> {}
52unsafe impl<T> Send for Raw<T> {}
54unsafe impl<T> Sync for Raw<T> {}
55impl<T> Raw<T> {
56 fn get(self) -> *mut T {
58 self.0
59 }
60}
61
62fn validate_destination(dims: &[usize], strides: &[isize]) -> Result<()> {
63 if is_injective_layout(dims, strides) {
64 Ok(())
65 } else {
66 Err(StridedError::NonInjectiveOutputLayout)
67 }
68}
69
70fn run_update(
82 dims: &[usize],
83 strides_list: &[&[isize]],
84 elem_size: usize,
85 body: Body<'_>,
86) -> Result<()> {
87 let total = total_len(dims)?;
88 if total == 0 {
89 return Ok(());
90 }
91 if dims.iter().all(|&d| d == 1) {
93 let zeros = vec![0isize; strides_list.len()];
94 body(&zeros, 1, &zeros);
95 return Ok(());
96 }
97 let (fused_dims, ordered_strides, plan) = if total <= SMALL_TENSOR_THRESHOLD {
99 build_plan_fused_small(dims, strides_list)
100 } else {
101 build_plan_fused(dims, strides_list, Some(0), elem_size)
102 };
103
104 #[cfg(feature = "parallel")]
105 {
106 let total = total_len(&fused_dims)?;
107 let nthreads = crate::execution_policy::rayon_threads();
108 if total > MINTHREADLENGTH && nthreads > 1 {
109 let costs = compute_costs(&ordered_strides);
110 let initial_offsets = vec![0isize; strides_list.len()];
111 return mapreduce_threaded(
112 &fused_dims,
113 &plan.block,
114 &ordered_strides,
115 &initial_offsets,
116 &costs,
117 nthreads,
118 0,
119 1,
120 &|dims, blocks, strides_list, offsets| {
121 for_each_inner_block_with_offsets(
122 dims,
123 blocks,
124 strides_list,
125 offsets,
126 |offsets, len, strides| {
127 body(offsets, len, strides);
128 Ok(())
129 },
130 )
131 },
132 );
133 }
134 }
135
136 let initial_offsets = vec![0isize; ordered_strides.len()];
137 for_each_inner_block_preordered(
138 &fused_dims,
139 &plan.block,
140 &ordered_strides,
141 &initial_offsets,
142 |offsets, len, strides| {
143 body(offsets, len, strides);
144 Ok(())
145 },
146 )
147}
148
149#[inline(always)]
154unsafe fn inner_loop_update1<D: Copy, OpD: ElementOp<D>>(
155 dp: *mut D,
156 ds: isize,
157 len: usize,
158 f: &impl Fn(D) -> D,
159) {
160 if ds == 1 {
161 let dst = std::slice::from_raw_parts_mut(dp, len);
162 simd::dispatch_if_large(len, || {
163 for d in dst.iter_mut() {
164 *d = f(OpD::apply(*d));
165 }
166 });
167 } else {
168 let mut dp = dp;
169 for _ in 0..len {
170 *dp = f(OpD::apply(*dp));
171 dp = dp.offset(ds);
172 }
173 }
174}
175
176#[inline(always)]
182unsafe fn inner_loop_update2<D: Copy, A: Copy, OpD: ElementOp<D>, OpA: ElementOp<A>>(
183 dp: *mut D,
184 ds: isize,
185 ap: *const A,
186 a_s: isize,
187 len: usize,
188 f: &impl Fn(D, A) -> D,
189) {
190 if ds == 1 && a_s == 1 {
191 let src_a = std::slice::from_raw_parts(ap, len);
192 let dst = std::slice::from_raw_parts_mut(dp, len);
193 simd::dispatch_if_large(len, || {
194 for (d, &a) in dst.iter_mut().zip(src_a) {
195 *d = f(OpD::apply(*d), OpA::apply(a));
196 }
197 });
198 } else {
199 let (mut dp, mut ap) = (dp, ap);
200 for _ in 0..len {
201 *dp = f(OpD::apply(*dp), OpA::apply(*ap));
202 dp = dp.offset(ds);
203 ap = ap.offset(a_s);
204 }
205 }
206}
207
208#[inline(always)]
213#[allow(clippy::too_many_arguments)] unsafe fn inner_loop_update3<
215 D: Copy,
216 A: Copy,
217 B: Copy,
218 OpD: ElementOp<D>,
219 OpA: ElementOp<A>,
220 OpB: ElementOp<B>,
221>(
222 dp: *mut D,
223 ds: isize,
224 ap: *const A,
225 a_s: isize,
226 bp: *const B,
227 b_s: isize,
228 len: usize,
229 f: &impl Fn(D, A, B) -> D,
230) {
231 if ds == 1 && a_s == 1 && b_s == 1 {
232 let src_a = std::slice::from_raw_parts(ap, len);
233 let src_b = std::slice::from_raw_parts(bp, len);
234 let dst = std::slice::from_raw_parts_mut(dp, len);
235 simd::dispatch_if_large(len, || {
236 for ((d, &a), &b) in dst.iter_mut().zip(src_a).zip(src_b) {
237 *d = f(OpD::apply(*d), OpA::apply(a), OpB::apply(b));
238 }
239 });
240 } else {
241 let (mut dp, mut ap, mut bp) = (dp, ap, bp);
242 for _ in 0..len {
243 *dp = f(OpD::apply(*dp), OpA::apply(*ap), OpB::apply(*bp));
244 dp = dp.offset(ds);
245 ap = ap.offset(a_s);
246 bp = bp.offset(b_s);
247 }
248 }
249}
250
251pub fn map_update_into<D, OpD>(
269 dest: &mut StridedViewMut<D>,
270 f: impl Fn(D) -> D + MaybeSync,
271) -> Result<()>
272where
273 D: Copy + MaybeSendSync,
274 OpD: ElementOp<D>,
275{
276 validate_destination(dest.dims(), dest.strides())?;
277 let dp = Raw(dest.as_mut_ptr());
278 run_update(
279 dest.dims(),
280 &[dest.strides()],
281 std::mem::size_of::<D>(),
282 &|offsets, len, strides| {
283 unsafe {
285 inner_loop_update1::<D, OpD>(dp.get().offset(offsets[0]), strides[0], len, &f);
286 }
287 },
288 )
289}
290
291pub fn zip_update2_into<D, A, OpD, OpA>(
312 dest: &mut StridedViewMut<D>,
313 a: &StridedView<A, OpA>,
314 f: impl Fn(D, A) -> D + MaybeSync,
315) -> Result<()>
316where
317 D: Copy + MaybeSendSync,
318 A: Copy + MaybeSendSync,
319 OpD: ElementOp<D>,
320 OpA: ElementOp<A>,
321{
322 ensure_same_shape(dest.dims(), a.dims())?;
323 validate_destination(dest.dims(), dest.strides())?;
324 let dp = Raw(dest.as_mut_ptr());
325 let ap = Raw(a.ptr() as *mut A);
326 run_update(
327 dest.dims(),
328 &[dest.strides(), a.strides()],
329 std::mem::size_of::<D>().max(std::mem::size_of::<A>()),
330 &|offsets, len, strides| {
331 unsafe {
333 inner_loop_update2::<D, A, OpD, OpA>(
334 dp.get().offset(offsets[0]),
335 strides[0],
336 ap.get().offset(offsets[1]).cast_const(),
337 strides[1],
338 len,
339 &f,
340 );
341 }
342 },
343 )
344}
345
346pub fn zip_update3_into<D, A, B, OpD, OpA, OpB>(
370 dest: &mut StridedViewMut<D>,
371 a: &StridedView<A, OpA>,
372 b: &StridedView<B, OpB>,
373 f: impl Fn(D, A, B) -> D + MaybeSync,
374) -> Result<()>
375where
376 D: Copy + MaybeSendSync,
377 A: Copy + MaybeSendSync,
378 B: Copy + MaybeSendSync,
379 OpD: ElementOp<D>,
380 OpA: ElementOp<A>,
381 OpB: ElementOp<B>,
382{
383 ensure_same_shape(dest.dims(), a.dims())?;
384 ensure_same_shape(dest.dims(), b.dims())?;
385 validate_destination(dest.dims(), dest.strides())?;
386 let dp = Raw(dest.as_mut_ptr());
387 let ap = Raw(a.ptr() as *mut A);
388 let bp = Raw(b.ptr() as *mut B);
389 run_update(
390 dest.dims(),
391 &[dest.strides(), a.strides(), b.strides()],
392 std::mem::size_of::<D>()
393 .max(std::mem::size_of::<A>())
394 .max(std::mem::size_of::<B>()),
395 &|offsets, len, strides| {
396 unsafe {
398 inner_loop_update3::<D, A, B, OpD, OpA, OpB>(
399 dp.get().offset(offsets[0]),
400 strides[0],
401 ap.get().offset(offsets[1]).cast_const(),
402 strides[1],
403 bp.get().offset(offsets[2]).cast_const(),
404 strides[2],
405 len,
406 &f,
407 );
408 }
409 },
410 )
411}
412
413#[cfg(test)]
414#[path = "update_view/tests/tests.rs"]
415mod tests;