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)]
156unsafe fn inner_loop_update1<D: Copy, OpD: ElementOp<D>>(
157 dp: *mut D,
158 ds: isize,
159 len: usize,
160 f: &impl Fn(D) -> D,
161) {
162 if ds == 1 {
163 let dst = std::slice::from_raw_parts_mut(dp, len);
164 simd::dispatch_if_large(len, || {
165 for d in dst.iter_mut() {
166 *d = f(OpD::apply(*d));
167 }
168 });
169 } else {
170 let mut dp = dp;
171 for _ in 0..len {
172 *dp = f(OpD::apply(*dp));
173 dp = dp.wrapping_offset(ds);
174 }
175 }
176}
177
178#[inline(always)]
184unsafe fn inner_loop_update2<D: Copy, A: Copy, OpD: ElementOp<D>, OpA: ElementOp<A>>(
185 dp: *mut D,
186 ds: isize,
187 ap: *const A,
188 a_s: isize,
189 len: usize,
190 f: &impl Fn(D, A) -> D,
191) {
192 if ds == 1 && a_s == 1 {
193 let src_a = std::slice::from_raw_parts(ap, len);
194 let dst = std::slice::from_raw_parts_mut(dp, len);
195 simd::dispatch_if_large(len, || {
196 for (d, &a) in dst.iter_mut().zip(src_a) {
197 *d = f(OpD::apply(*d), OpA::apply(a));
198 }
199 });
200 } else {
201 let (mut dp, mut ap) = (dp, ap);
202 for _ in 0..len {
203 *dp = f(OpD::apply(*dp), OpA::apply(*ap));
204 dp = dp.wrapping_offset(ds);
205 ap = ap.wrapping_offset(a_s);
206 }
207 }
208}
209
210#[inline(always)]
215#[allow(clippy::too_many_arguments)] unsafe fn inner_loop_update3<
217 D: Copy,
218 A: Copy,
219 B: Copy,
220 OpD: ElementOp<D>,
221 OpA: ElementOp<A>,
222 OpB: ElementOp<B>,
223>(
224 dp: *mut D,
225 ds: isize,
226 ap: *const A,
227 a_s: isize,
228 bp: *const B,
229 b_s: isize,
230 len: usize,
231 f: &impl Fn(D, A, B) -> D,
232) {
233 if ds == 1 && a_s == 1 && b_s == 1 {
234 let src_a = std::slice::from_raw_parts(ap, len);
235 let src_b = std::slice::from_raw_parts(bp, len);
236 let dst = std::slice::from_raw_parts_mut(dp, len);
237 simd::dispatch_if_large(len, || {
238 for ((d, &a), &b) in dst.iter_mut().zip(src_a).zip(src_b) {
239 *d = f(OpD::apply(*d), OpA::apply(a), OpB::apply(b));
240 }
241 });
242 } else {
243 let (mut dp, mut ap, mut bp) = (dp, ap, bp);
244 for _ in 0..len {
245 *dp = f(OpD::apply(*dp), OpA::apply(*ap), OpB::apply(*bp));
246 dp = dp.wrapping_offset(ds);
247 ap = ap.wrapping_offset(a_s);
248 bp = bp.wrapping_offset(b_s);
249 }
250 }
251}
252
253pub fn map_update_into<D, OpD>(
271 dest: &mut StridedViewMut<D>,
272 f: impl Fn(D) -> D + MaybeSync,
273) -> Result<()>
274where
275 D: Copy + MaybeSendSync,
276 OpD: ElementOp<D>,
277{
278 validate_destination(dest.dims(), dest.strides())?;
279 let dp = Raw(dest.as_mut_ptr());
280 run_update(
281 dest.dims(),
282 &[dest.strides()],
283 std::mem::size_of::<D>(),
284 &|offsets, len, strides| {
285 unsafe {
287 inner_loop_update1::<D, OpD>(dp.get().offset(offsets[0]), strides[0], len, &f);
288 }
289 },
290 )
291}
292
293pub fn zip_update2_into<D, A, OpD, OpA>(
314 dest: &mut StridedViewMut<D>,
315 a: &StridedView<A, OpA>,
316 f: impl Fn(D, A) -> D + MaybeSync,
317) -> Result<()>
318where
319 D: Copy + MaybeSendSync,
320 A: Copy + MaybeSendSync,
321 OpD: ElementOp<D>,
322 OpA: ElementOp<A>,
323{
324 ensure_same_shape(dest.dims(), a.dims())?;
325 validate_destination(dest.dims(), dest.strides())?;
326 let dp = Raw(dest.as_mut_ptr());
327 let ap = Raw(a.ptr() as *mut A);
328 run_update(
329 dest.dims(),
330 &[dest.strides(), a.strides()],
331 std::mem::size_of::<D>().max(std::mem::size_of::<A>()),
332 &|offsets, len, strides| {
333 unsafe {
335 inner_loop_update2::<D, A, OpD, OpA>(
336 dp.get().offset(offsets[0]),
337 strides[0],
338 ap.get().offset(offsets[1]).cast_const(),
339 strides[1],
340 len,
341 &f,
342 );
343 }
344 },
345 )
346}
347
348pub fn zip_update3_into<D, A, B, OpD, OpA, OpB>(
372 dest: &mut StridedViewMut<D>,
373 a: &StridedView<A, OpA>,
374 b: &StridedView<B, OpB>,
375 f: impl Fn(D, A, B) -> D + MaybeSync,
376) -> Result<()>
377where
378 D: Copy + MaybeSendSync,
379 A: Copy + MaybeSendSync,
380 B: Copy + MaybeSendSync,
381 OpD: ElementOp<D>,
382 OpA: ElementOp<A>,
383 OpB: ElementOp<B>,
384{
385 ensure_same_shape(dest.dims(), a.dims())?;
386 ensure_same_shape(dest.dims(), b.dims())?;
387 validate_destination(dest.dims(), dest.strides())?;
388 let dp = Raw(dest.as_mut_ptr());
389 let ap = Raw(a.ptr() as *mut A);
390 let bp = Raw(b.ptr() as *mut B);
391 run_update(
392 dest.dims(),
393 &[dest.strides(), a.strides(), b.strides()],
394 std::mem::size_of::<D>()
395 .max(std::mem::size_of::<A>())
396 .max(std::mem::size_of::<B>()),
397 &|offsets, len, strides| {
398 unsafe {
400 inner_loop_update3::<D, A, B, OpD, OpA, OpB>(
401 dp.get().offset(offsets[0]),
402 strides[0],
403 ap.get().offset(offsets[1]).cast_const(),
404 strides[1],
405 bp.get().offset(offsets[2]).cast_const(),
406 strides[2],
407 len,
408 &f,
409 );
410 }
411 },
412 )
413}
414
415#[cfg(test)]
416#[path = "update_view/tests/tests.rs"]
417mod tests;