Skip to main content

strided_basic/
update_view.rs

1//! In-place update operations: the destination is also an input.
2//!
3//! `map_into` and `zip_map*_into` write `dest[i] = f(src...)` and never read
4//! the old destination. These functions compute
5//!
6//! ```text
7//! dest[i] = f(OpD(dest[i]), OpA(a[i]), OpB(b[i]))
8//! ```
9//!
10//! reading the previous `dest[i]` through the destination view itself, so no
11//! second view over the destination's storage is ever formed (two live views
12//! of one allocation, one mutable, would be aliasing). Every element of an
13//! injective destination is read and written exactly once, by one call of `f`.
14//!
15//! The element operation of each input, and of the destination read, is a type
16//! parameter (`Identity`, `Conj`, ...), so a conjugating update is a
17//! monomorphized loop with no per-element flag. Traversal, blocking, contiguous
18//! inner loops and threading are those of the other map kernels (threading
19//! follows the active [`ExecContext`](crate::ExecContext)).
20
21use 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/// The per-inner-block callback: `(offsets, len, strides)`.
38#[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
43/// A raw pointer that may cross threads: dereferenced only at the validated,
44/// disjoint (for the destination) offsets the traversal generates.
45struct 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> {}
52// SAFETY: see the type documentation; every user upholds it.
53unsafe impl<T> Send for Raw<T> {}
54unsafe impl<T> Sync for Raw<T> {}
55impl<T> Raw<T> {
56    // A method, so closures capture the whole wrapper rather than its field.
57    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
70/// Walk `dims` with the strided traversal and call `body(offsets, len,
71/// strides)` per inner block; the destination is `strides_list[0]`.
72///
73/// The destination layout must already be validated injective, which makes
74/// the blocks' destination regions disjoint when they run on different
75/// threads.
76///
77/// Not generic: the traversal and threading machinery is compiled once here,
78/// and each operation passes its inner loop as a `dyn` callback. The call is
79/// per inner block (a run of elements), never per element, so it costs nothing
80/// measurable while keeping the machinery out of every caller's monomorphization.
81fn 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    // One element (rank 0, or every extent one): a single length-1 block.
92    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    // Small tensor fast path: skip compute_order and compute_block_sizes
98    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// INVARIANT: validated layouts keep all dereferences in bounds. Wrapping
150// advances permit an unused final cursor outside a reversed/gapped allocation.
151/// Unary inner loop: `d[i] = f(OpD(d[i]))`.
152///
153/// # Safety
154/// `dp` addresses `len` live, exclusively owned elements at stride `ds`.
155#[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/// Binary inner loop: `d[i] = f(OpD(d[i]), OpA(a[i]))`.
179///
180/// # Safety
181/// As [`inner_loop_update1`]; `ap` addresses `len` live elements that do not
182/// overlap the destination.
183#[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/// Ternary inner loop: `d[i] = f(OpD(d[i]), OpA(a[i]), OpB(b[i]))`.
211///
212/// # Safety
213/// As [`inner_loop_update2`], with `bp` likewise.
214#[inline(always)]
215#[allow(clippy::too_many_arguments)] // INVARIANT: three strided operands, length and body.
216unsafe 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
253/// Update in place: `dest[i] = f(OpD(dest[i]))`.
254///
255/// `OpD` is applied lazily to the old value before `f` sees it.
256///
257/// # Errors
258///
259/// [`StridedError::NonInjectiveOutputLayout`] when `dest` maps two indices to
260/// one element (checked before any write).
261///
262/// # Examples
263///
264/// ```
265/// use strided_basic::{map_update_into, Identity, StridedArray};
266/// let mut d = StridedArray::<f64>::from_fn_col_major(&[2, 2], |i| (i[0] + 2 * i[1]) as f64);
267/// map_update_into::<_, Identity>(&mut d.view_mut(), |x| 2.0 * x + 1.0).unwrap();
268/// assert_eq!(d.get(&[1, 1]), 7.0);
269/// ```
270pub 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            // SAFETY: the destination is injective and in bounds; blocks are disjoint.
286            unsafe {
287                inner_loop_update1::<D, OpD>(dp.get().offset(offsets[0]), strides[0], len, &f);
288            }
289        },
290    )
291}
292
293/// Update in place from one input: `dest[i] = f(OpD(dest[i]), OpA(a[i]))`.
294///
295/// # Errors
296///
297/// [`StridedError::ShapeMismatch`] for unequal shapes and
298/// [`StridedError::NonInjectiveOutputLayout`], both before any write.
299///
300/// # Examples
301///
302/// ```
303/// use strided_basic::{zip_update2_into, Identity, StridedArray};
304/// let a = StridedArray::<f64>::from_fn_col_major(&[3], |i| i[0] as f64);
305/// let mut d = StridedArray::<f64>::from_fn_col_major(&[3], |_| 10.0);
306/// // d = 2 * a + 0.5 * d
307/// zip_update2_into::<_, _, Identity, Identity>(&mut d.view_mut(), &a.view(), |d, a| {
308///     2.0 * a + 0.5 * d
309/// })
310/// .unwrap();
311/// assert_eq!(d.get(&[2]), 9.0);
312/// ```
313pub 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            // SAFETY: bounds are the views'; `a` is a distinct borrow from `dest`.
334            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
348/// Update in place from two inputs:
349/// `dest[i] = f(OpD(dest[i]), OpA(a[i]), OpB(b[i]))`.
350///
351/// # Errors
352///
353/// As [`zip_update2_into`].
354///
355/// # Examples
356///
357/// ```
358/// use strided_basic::{zip_update3_into, Conj, Identity, StridedArray};
359/// use num_complex::Complex64;
360/// let one = Complex64::new(1.0, 1.0);
361/// let a = StridedArray::<Complex64>::from_fn_col_major(&[2], |_| one);
362/// let b = StridedArray::<Complex64>::from_fn_col_major(&[2], |_| one);
363/// let mut d = StridedArray::<Complex64>::from_fn_col_major(&[2], |_| one);
364/// // d = conj(d) + a * b
365/// zip_update3_into::<_, _, _, Conj, Identity, Identity>(
366///     &mut d.view_mut(), &a.view(), &b.view(), |d, a, b| d + a * b,
367/// )
368/// .unwrap();
369/// assert_eq!(d.get(&[0]), Complex64::new(1.0, 1.0)); // (1 - i) + 2i
370/// ```
371pub 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            // SAFETY: bounds are the views'; `a`, `b` are distinct borrows from `dest`.
399            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;