Skip to main content

strided_basic/
copy_plan.rs

1//! Prepared (compile-once, execute-many) copy plans over raw strided layouts.
2//!
3//! [`copy_scale_raw`](crate::copy_scale_raw) and friends rebuild the fused
4//! loop nest on every call; for prepared-replay consumers that issue many
5//! small copies with a fixed layout, that per-call planning dominates
6//! (see issue #139). [`CopyPlan`] splits the work: [`CopyPlan::compile`]
7//! validates the layout pair and builds the fused traversal once,
8//! [`CopyPlan::execute`]/[`CopyPlan::execute_scale`]/[`CopyPlan::execute_conj`]
9//! replay it with no planning and no heap allocation for ranks at most
10//! [`RAW_FUSED_RANK_LIMIT`](crate::RAW_FUSED_RANK_LIMIT).
11
12use core::{
13    marker::PhantomData,
14    mem::MaybeUninit,
15    ops::{Add, Mul},
16};
17
18use crate::map_view::map_raw_into;
19use crate::ops_view::{copy_conj, copy_into, copy_scale};
20use crate::raw_ops::{apply_fused_range, fuse_pair_layout, fused_total, FusedPairLayout};
21use crate::{
22    ElementOpApply, Identity, MaybeSendSync, RawStridedMut, RawStridedRef, Result, StridedError,
23};
24
25// Same pattern as map_view.rs / outer_product.rs: stack storage when the
26// parallel feature pulls in smallvec, plain Vec otherwise. Only `compile`
27// touches these; `execute*` never allocates either way.
28#[cfg(feature = "parallel")]
29type AxisVec<T> = smallvec::SmallVec<[T; crate::RAW_FUSED_RANK_LIMIT]>;
30#[cfg(not(feature = "parallel"))]
31type AxisVec<T> = Vec<T>;
32
33pub(crate) trait OverwriteWriter<T> {
34    fn dims(&self) -> &[usize];
35    fn strides(&self) -> &[isize];
36    fn offset(&self) -> isize;
37    /// # Safety
38    /// The caller must use the pointer only for the validated allocation and
39    /// logical layout represented by this writer.
40    unsafe fn data_ptr(&mut self) -> *mut T;
41    /// # Safety
42    /// The offset must be an in-bounds logical destination proven by layout.
43    unsafe fn write_at(&mut self, offset: isize, value: T);
44}
45
46pub(crate) trait ReadModifyWrite<T>: OverwriteWriter<T> {
47    /// # Safety
48    /// The offset must be an in-bounds initialized slot covered by the
49    /// traversal's copy and disjointness proof.
50    unsafe fn add_at<F>(&mut self, offset: isize, value: T, combine: F)
51    where
52        F: FnOnce(T, T) -> T;
53}
54
55impl<'a, T> OverwriteWriter<T> for RawStridedMut<'a, T> {
56    fn dims(&self) -> &[usize] {
57        self.dims()
58    }
59    fn strides(&self) -> &[isize] {
60        self.strides()
61    }
62    fn offset(&self) -> isize {
63        self.offset()
64    }
65    unsafe fn data_ptr(&mut self) -> *mut T {
66        self.data_mut().as_mut_ptr()
67    }
68    unsafe fn write_at(&mut self, offset: isize, value: T) {
69        // SAFETY: the prepared layout validates every logical destination.
70        unsafe { self.data_mut().as_mut_ptr().offset(offset).write(value) }
71    }
72}
73
74impl<'a, T> ReadModifyWrite<T> for RawStridedMut<'a, T>
75where
76    T: Add<Output = T>,
77{
78    #[inline(always)]
79    unsafe fn add_at<F>(&mut self, offset: isize, value: T, combine: F)
80    where
81        F: FnOnce(T, T) -> T,
82    {
83        // SAFETY: the copy or initialized caller proves this logical slot.
84        unsafe {
85            let ptr = self.data_mut().as_mut_ptr().offset(offset);
86            ptr.write(combine(ptr.read(), value));
87        }
88    }
89}
90
91impl<'a, T> OverwriteWriter<T> for RawStridedMut<'a, MaybeUninit<T>> {
92    fn dims(&self) -> &[usize] {
93        self.dims()
94    }
95    fn strides(&self) -> &[isize] {
96        self.strides()
97    }
98    fn offset(&self) -> isize {
99        self.offset()
100    }
101    unsafe fn data_ptr(&mut self) -> *mut T {
102        self.data_mut().as_mut_ptr().cast()
103    }
104    unsafe fn write_at(&mut self, offset: isize, value: T) {
105        // SAFETY: the prepared layout validates every logical destination.
106        unsafe {
107            self.data_mut()
108                .as_mut_ptr()
109                .offset(offset)
110                .write(MaybeUninit::new(value))
111        }
112    }
113}
114
115pub(crate) struct InitializedRawDest<'a, T> {
116    ptr: *mut T,
117    extent: usize,
118    dims: &'a [usize],
119    strides: &'a [isize],
120    offset: isize,
121    _marker: PhantomData<&'a mut [MaybeUninit<T>]>,
122}
123
124impl<'a, T> OverwriteWriter<T> for InitializedRawDest<'a, T> {
125    fn dims(&self) -> &[usize] {
126        self.dims
127    }
128    fn strides(&self) -> &[isize] {
129        self.strides
130    }
131    fn offset(&self) -> isize {
132        self.offset
133    }
134    unsafe fn data_ptr(&mut self) -> *mut T {
135        self.ptr
136    }
137    unsafe fn write_at(&mut self, offset: isize, value: T) {
138        debug_assert!(offset >= 0 && (offset as usize) < self.extent);
139        // SAFETY: the copy proof and extent check cover this logical slot.
140        unsafe { self.ptr.offset(offset).write(value) }
141    }
142}
143
144impl<'a, T> ReadModifyWrite<T> for InitializedRawDest<'a, T>
145where
146    T: Add<Output = T>,
147{
148    #[inline(always)]
149    unsafe fn add_at<F>(&mut self, offset: isize, value: T, combine: F)
150    where
151        F: FnOnce(T, T) -> T,
152    {
153        debug_assert!(offset >= 0 && (offset as usize) < self.extent);
154        // SAFETY: the copy proof and extent check cover this logical slot.
155        unsafe {
156            let ptr = self.ptr.offset(offset);
157            ptr.write(combine(ptr.read(), value));
158        }
159    }
160}
161
162/// A compiled copy traversal for one `(dims, dst_strides, src_strides)`
163/// layout pair.
164///
165/// `compile` proves the layout facts once (rank agreement, extent overflow,
166/// destination injectivity) and fuses/orders the loop nest; each `execute*`
167/// call then only re-checks the per-call facts (that the supplied views carry
168/// exactly the compiled layout) before replaying the prepared loops.
169///
170/// Overlapping `src`/`dest` memory is not supported, matching the rest of the
171/// crate.
172///
173/// Ranks above [`RAW_FUSED_RANK_LIMIT`](crate::RAW_FUSED_RANK_LIMIT) are supported through the view-based
174/// kernels; only that fallback path may allocate.
175///
176/// # Example
177///
178/// ```rust
179/// use strided_basic::{CopyPlan, RawStridedMut, RawStridedRef};
180///
181/// let dims = [2usize, 3];
182/// let src_strides = [3isize, 1];
183/// let dst_strides = [1isize, 2]; // transposed destination
184/// let plan = CopyPlan::compile(&dims, &dst_strides, &src_strides).unwrap();
185///
186/// let src = [0.0f64, 1.0, 2.0, 10.0, 11.0, 12.0];
187/// let mut dst = [0.0f64; 6];
188/// let src_ref = RawStridedRef::new(&src, &dims, &src_strides, 0).unwrap();
189/// let mut dst_mut = RawStridedMut::new(&mut dst, &dims, &dst_strides, 0).unwrap();
190/// plan.execute(&mut dst_mut, &src_ref).unwrap();
191/// assert_eq!(dst, [0.0, 10.0, 1.0, 11.0, 2.0, 12.0]);
192/// ```
193#[derive(Clone, Debug)]
194pub struct CopyPlan {
195    dims: AxisVec<usize>,
196    dst_strides: AxisVec<isize>,
197    src_strides: AxisVec<isize>,
198    /// `None` when rank exceeds [`RAW_FUSED_RANK_LIMIT`](crate::RAW_FUSED_RANK_LIMIT); `execute*` then
199    /// falls back to the view-based kernels.
200    fused: Option<FusedPairLayout>,
201}
202
203impl CopyPlan {
204    /// The fused traversal, or `None` above
205    /// [`RAW_FUSED_RANK_LIMIT`](crate::RAW_FUSED_RANK_LIMIT).
206    #[cfg(feature = "parallel")]
207    #[inline]
208    pub(crate) fn fused_layout(&self) -> Option<&FusedPairLayout> {
209        self.fused.as_ref()
210    }
211
212    pub(crate) fn execute_uninit_then<'a, T, R>(
213        &self,
214        dest: &'a mut RawStridedMut<'a, MaybeUninit<T>>,
215        src: &RawStridedRef<'_, T>,
216        f: impl for<'b> FnOnce(InitializedRawDest<'b, T>) -> R,
217    ) -> Result<R>
218    where
219        T: Copy + MaybeSendSync,
220    {
221        self.execute_uninit(dest, src)?;
222        let data = dest.data_mut();
223        let receipt = InitializedRawDest {
224            ptr: data.as_mut_ptr().cast(),
225            extent: data.len(),
226            dims: dest.dims(),
227            strides: dest.strides(),
228            offset: dest.offset(),
229            _marker: PhantomData,
230        };
231        Ok(f(receipt))
232    }
233
234    /// Compile a copy plan for the given layout pair.
235    ///
236    /// Performs the layout validation and traversal construction
237    /// (fuse + order) once:
238    ///
239    /// - `dims`, `dst_strides`, and `src_strides` must have equal length
240    ///   ([`StridedError::StrideLengthMismatch`]);
241    /// - the total element count must not overflow `usize`
242    ///   ([`StridedError::OffsetOverflow`]);
243    /// - the destination layout must be injective, i.e. map distinct logical
244    ///   indices to distinct offsets
245    ///   ([`StridedError::NonInjectiveOutputLayout`]).
246    pub fn compile(dims: &[usize], dst_strides: &[isize], src_strides: &[isize]) -> Result<Self> {
247        if dims.len() != dst_strides.len() || dims.len() != src_strides.len() {
248            return Err(StridedError::StrideLengthMismatch);
249        }
250        // A zero-sized layout has zero elements even when its other extents
251        // overflow; only a nonempty overflowing count is rejected.
252        crate::kernel::total_len(dims)?;
253        if !crate::layout_check::is_injective_layout(dims, dst_strides) {
254            return Err(StridedError::NonInjectiveOutputLayout);
255        }
256        Ok(Self {
257            dims: dims.into(),
258            dst_strides: dst_strides.into(),
259            src_strides: src_strides.into(),
260            fused: fuse_pair_layout(dims, dst_strides, src_strides),
261        })
262    }
263
264    /// Check the per-call facts: the supplied views must carry exactly the
265    /// compiled layout. Buffer bounds against that layout were already proven
266    /// by [`RawStridedRef::new`]/[`RawStridedMut::new`] (or asserted by the
267    /// caller of the `new_unchecked` constructors), so layout equality is the
268    /// complete precondition for the pointer-based fused replay below.
269    fn check_call<D, S>(
270        &self,
271        dest: &RawStridedMut<'_, D>,
272        src: &RawStridedRef<'_, S>,
273    ) -> Result<()> {
274        if dest.dims() != &self.dims[..]
275            || src.dims() != &self.dims[..]
276            || dest.strides() != &self.dst_strides[..]
277            || src.strides() != &self.src_strides[..]
278        {
279            return Err(StridedError::PlanLayoutMismatch);
280        }
281        Ok(())
282    }
283
284    /// `dest = src` into a potentially uninitialized destination.
285    ///
286    /// On success every reachable logical destination element is initialized.
287    /// Non-reachable holes in the backing allocation are not written.
288    pub fn execute_uninit<T>(
289        &self,
290        dest: &mut RawStridedMut<'_, MaybeUninit<T>>,
291        src: &RawStridedRef<'_, T>,
292    ) -> Result<()>
293    where
294        T: Copy + MaybeSendSync,
295    {
296        self.check_call(dest, src)?;
297        match &self.fused {
298            Some(layout) => {
299                replay_fused(
300                    dest,
301                    src,
302                    layout,
303                    |dst, value| {
304                        dst.write(value);
305                    },
306                    |value| value,
307                );
308                Ok(())
309            }
310            None => map_raw_into::<MaybeUninit<T>, T, Identity>(dest, src, MaybeUninit::new),
311        }
312    }
313
314    /// `dest = src`. Allocation-free for ranks at most [`RAW_FUSED_RANK_LIMIT`](crate::RAW_FUSED_RANK_LIMIT).
315    pub fn execute<T>(
316        &self,
317        dest: &mut RawStridedMut<'_, T>,
318        src: &RawStridedRef<'_, T>,
319    ) -> Result<()>
320    where
321        T: Copy + MaybeSendSync,
322    {
323        self.check_call(dest, src)?;
324        match &self.fused {
325            Some(layout) => {
326                replay_fused(
327                    dest,
328                    src,
329                    layout,
330                    |dst, value| *dst = value,
331                    |value: T| value,
332                );
333                Ok(())
334            }
335            None => copy_into(&mut dest.as_view_mut(), &src.as_view()),
336        }
337    }
338
339    /// `dest = scale * src`. Allocation-free for ranks at most
340    /// [`RAW_FUSED_RANK_LIMIT`](crate::RAW_FUSED_RANK_LIMIT).
341    pub fn execute_scale<T>(
342        &self,
343        dest: &mut RawStridedMut<'_, T>,
344        src: &RawStridedRef<'_, T>,
345        scale: T,
346    ) -> Result<()>
347    where
348        T: Copy + Mul<T, Output = T> + MaybeSendSync,
349    {
350        self.check_call(dest, src)?;
351        match &self.fused {
352            Some(layout) => {
353                replay_fused(
354                    dest,
355                    src,
356                    layout,
357                    |dst, value| *dst = value,
358                    |value: T| scale * value,
359                );
360                Ok(())
361            }
362            None => copy_scale(&mut dest.as_view_mut(), &src.as_view(), scale),
363        }
364    }
365
366    /// `dest = conj(src)`. Allocation-free for ranks at most
367    /// [`RAW_FUSED_RANK_LIMIT`](crate::RAW_FUSED_RANK_LIMIT).
368    pub fn execute_conj<T>(
369        &self,
370        dest: &mut RawStridedMut<'_, T>,
371        src: &RawStridedRef<'_, T>,
372    ) -> Result<()>
373    where
374        T: Copy + ElementOpApply + MaybeSendSync,
375    {
376        self.check_call(dest, src)?;
377        match &self.fused {
378            Some(layout) => {
379                replay_fused(
380                    dest,
381                    src,
382                    layout,
383                    |dst, value| *dst = value,
384                    |value: T| value.conj(),
385                );
386                Ok(())
387            }
388            None => copy_conj(&mut dest.as_view_mut(), &src.as_view()),
389        }
390    }
391}
392
393/// Replay a compiled fused layout, splitting it across worker threads when the
394/// repository threshold and the active execution policy allow it.
395///
396/// Parallel replay is sound only because [`CopyPlan::compile`] proved the
397/// destination layout injective: disjoint logical ranges then write disjoint
398/// destination slots. The logical index space `0..total` (column-major over
399/// the fused axes) is split into contiguous worker ranges, which divides the
400/// outer fused axes and, when there are fewer outer positions than workers,
401/// also chunks long inner runs (for example a rank-1 contiguous copy).
402fn replay_fused<D, S, Apply, Op>(
403    dest: &mut RawStridedMut<'_, D>,
404    src: &RawStridedRef<'_, S>,
405    layout: &FusedPairLayout,
406    apply: Apply,
407    op: Op,
408) where
409    D: Copy + MaybeSendSync,
410    S: Copy + MaybeSendSync,
411    Apply: Fn(&mut D, S) + MaybeSendSync,
412    Op: Fn(S) -> S + MaybeSendSync,
413{
414    let src_ptr = src.data().as_ptr();
415    let src_base = src.offset();
416    let dst_base = dest.offset();
417    let dst_ptr = dest.data_mut().as_mut_ptr();
418    // SAFETY: `check_call` proved both views carry exactly the compiled
419    // layout, and `RawStridedRef`/`RawStridedMut` guarantee every offset
420    // reachable through that layout lies inside their data. `layout` is the
421    // fusion of that layout, so every logical index in `0..total` is an
422    // in-bounds slot of both. The compiled destination layout is injective
423    // and the exclusive destination borrow cannot overlap the shared source
424    // borrow.
425    unsafe { replay_fused_raw(dst_ptr, dst_base, src_ptr, src_base, layout, apply, op) }
426}
427
428/// Pointer-level fused replay shared by [`CopyPlan`] and the dynamic slice
429/// plans, which replay a fused window layout at a runtime base offset.
430///
431/// Splits `0..total` across workers when the repository threshold and the
432/// active execution policy allow more than one thread, otherwise runs the
433/// serial kernel directly.
434///
435/// # Safety
436///
437/// Every logical index of `layout`, mapped through its destination strides
438/// from `dst_base` and its source strides from `src_base`, must be an
439/// in-bounds slot of the destination and source allocations. The
440/// destination strides must be injective over `layout`, the destination
441/// must not overlap the source, and nothing else may access the destination
442/// slots during the call.
443pub(crate) unsafe fn replay_fused_raw<D, S, Apply, Op>(
444    dst_ptr: *mut D,
445    dst_base: isize,
446    src_ptr: *const S,
447    src_base: isize,
448    layout: &FusedPairLayout,
449    apply: Apply,
450    op: Op,
451) where
452    D: Copy + MaybeSendSync,
453    S: Copy + MaybeSendSync,
454    Apply: Fn(&mut D, S) + MaybeSendSync,
455    Op: Fn(S) -> S + MaybeSendSync,
456{
457    let total = fused_total(layout);
458    if total == 0 {
459        return;
460    }
461    #[cfg(feature = "parallel")]
462    {
463        let nthreads = crate::threading::parallel_threads_for_len(total);
464        if nthreads > 1 {
465            let dst_ptr = crate::threading::SendPtr(dst_ptr);
466            let src_ptr = crate::threading::SendPtr(src_ptr as *mut S);
467            crate::threading::parallel_for_each(0..total, nthreads, &|range| {
468                // SAFETY: the caller contract makes every index in bounds;
469                // the worker ranges are disjoint and the destination strides
470                // injective, so no two workers write the same slot, and the
471                // source is only read.
472                unsafe {
473                    apply_fused_range(
474                        dst_ptr.as_ptr(),
475                        dst_base,
476                        src_ptr.as_const(),
477                        src_base,
478                        layout,
479                        range.start,
480                        range.len(),
481                        &apply,
482                        &op,
483                    );
484                }
485            });
486            return;
487        }
488    }
489    // SAFETY: forwarded caller contract over the full range `0..total`.
490    unsafe {
491        apply_fused_range(
492            dst_ptr, dst_base, src_ptr, src_base, layout, 0, total, &apply, &op,
493        );
494    }
495}
496
497#[cfg(test)]
498#[path = "copy_plan/tests/tests.rs"]
499mod tests;
500
501#[cfg(all(test, feature = "parallel"))]
502#[path = "copy_plan/tests/parallel_tests.rs"]
503mod parallel_tests;