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/// Unary inner loop: `d[i] = f(OpD(d[i]))`.
150///
151/// # Safety
152/// `dp` addresses `len` live, exclusively owned elements at stride `ds`.
153#[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/// Binary inner loop: `d[i] = f(OpD(d[i]), OpA(a[i]))`.
177///
178/// # Safety
179/// As [`inner_loop_update1`]; `ap` addresses `len` live elements that do not
180/// overlap the destination.
181#[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/// Ternary inner loop: `d[i] = f(OpD(d[i]), OpA(a[i]), OpB(b[i]))`.
209///
210/// # Safety
211/// As [`inner_loop_update2`], with `bp` likewise.
212#[inline(always)]
213#[allow(clippy::too_many_arguments)] // INVARIANT: three strided operands, length and body.
214unsafe 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
251/// Update in place: `dest[i] = f(OpD(dest[i]))`.
252///
253/// `OpD` is applied lazily to the old value before `f` sees it.
254///
255/// # Errors
256///
257/// [`StridedError::NonInjectiveOutputLayout`] when `dest` maps two indices to
258/// one element (checked before any write).
259///
260/// # Examples
261///
262/// ```
263/// use strided_basic::{map_update_into, Identity, StridedArray};
264/// let mut d = StridedArray::<f64>::from_fn_col_major(&[2, 2], |i| (i[0] + 2 * i[1]) as f64);
265/// map_update_into::<_, Identity>(&mut d.view_mut(), |x| 2.0 * x + 1.0).unwrap();
266/// assert_eq!(d.get(&[1, 1]), 7.0);
267/// ```
268pub 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            // SAFETY: the destination is injective and in bounds; blocks are disjoint.
284            unsafe {
285                inner_loop_update1::<D, OpD>(dp.get().offset(offsets[0]), strides[0], len, &f);
286            }
287        },
288    )
289}
290
291/// Update in place from one input: `dest[i] = f(OpD(dest[i]), OpA(a[i]))`.
292///
293/// # Errors
294///
295/// [`StridedError::ShapeMismatch`] for unequal shapes and
296/// [`StridedError::NonInjectiveOutputLayout`], both before any write.
297///
298/// # Examples
299///
300/// ```
301/// use strided_basic::{zip_update2_into, Identity, StridedArray};
302/// let a = StridedArray::<f64>::from_fn_col_major(&[3], |i| i[0] as f64);
303/// let mut d = StridedArray::<f64>::from_fn_col_major(&[3], |_| 10.0);
304/// // d = 2 * a + 0.5 * d
305/// zip_update2_into::<_, _, Identity, Identity>(&mut d.view_mut(), &a.view(), |d, a| {
306///     2.0 * a + 0.5 * d
307/// })
308/// .unwrap();
309/// assert_eq!(d.get(&[2]), 9.0);
310/// ```
311pub 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            // SAFETY: bounds are the views'; `a` is a distinct borrow from `dest`.
332            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
346/// Update in place from two inputs:
347/// `dest[i] = f(OpD(dest[i]), OpA(a[i]), OpB(b[i]))`.
348///
349/// # Errors
350///
351/// As [`zip_update2_into`].
352///
353/// # Examples
354///
355/// ```
356/// use strided_basic::{zip_update3_into, Conj, Identity, StridedArray};
357/// use num_complex::Complex64;
358/// let one = Complex64::new(1.0, 1.0);
359/// let a = StridedArray::<Complex64>::from_fn_col_major(&[2], |_| one);
360/// let b = StridedArray::<Complex64>::from_fn_col_major(&[2], |_| one);
361/// let mut d = StridedArray::<Complex64>::from_fn_col_major(&[2], |_| one);
362/// // d = conj(d) + a * b
363/// zip_update3_into::<_, _, _, Conj, Identity, Identity>(
364///     &mut d.view_mut(), &a.view(), &b.view(), |d, a, b| d + a * b,
365/// )
366/// .unwrap();
367/// assert_eq!(d.get(&[0]), Complex64::new(1.0, 1.0)); // (1 - i) + 2i
368/// ```
369pub 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            // SAFETY: bounds are the views'; `a`, `b` are distinct borrows from `dest`.
397            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;