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_pair, fuse_pair_layout, 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    pub(crate) fn execute_uninit_then<'a, T, R>(
205        &self,
206        dest: &'a mut RawStridedMut<'a, MaybeUninit<T>>,
207        src: &RawStridedRef<'_, T>,
208        f: impl for<'b> FnOnce(InitializedRawDest<'b, T>) -> R,
209    ) -> Result<R>
210    where
211        T: Copy + MaybeSendSync,
212    {
213        self.execute_uninit(dest, src)?;
214        let data = dest.data_mut();
215        let receipt = InitializedRawDest {
216            ptr: data.as_mut_ptr().cast(),
217            extent: data.len(),
218            dims: dest.dims(),
219            strides: dest.strides(),
220            offset: dest.offset(),
221            _marker: PhantomData,
222        };
223        Ok(f(receipt))
224    }
225
226    /// Compile a copy plan for the given layout pair.
227    ///
228    /// Performs the layout validation and traversal construction
229    /// (fuse + order) once:
230    ///
231    /// - `dims`, `dst_strides`, and `src_strides` must have equal length
232    ///   ([`StridedError::StrideLengthMismatch`]);
233    /// - the total element count must not overflow `usize`
234    ///   ([`StridedError::OffsetOverflow`]);
235    /// - the destination layout must be injective, i.e. map distinct logical
236    ///   indices to distinct offsets
237    ///   ([`StridedError::NonInjectiveOutputLayout`]).
238    pub fn compile(dims: &[usize], dst_strides: &[isize], src_strides: &[isize]) -> Result<Self> {
239        if dims.len() != dst_strides.len() || dims.len() != src_strides.len() {
240            return Err(StridedError::StrideLengthMismatch);
241        }
242        // A zero-sized layout has zero elements even when its other extents
243        // overflow; only a nonempty overflowing count is rejected.
244        crate::kernel::total_len(dims)?;
245        if !crate::layout_check::is_injective_layout(dims, dst_strides) {
246            return Err(StridedError::NonInjectiveOutputLayout);
247        }
248        Ok(Self {
249            dims: dims.into(),
250            dst_strides: dst_strides.into(),
251            src_strides: src_strides.into(),
252            fused: fuse_pair_layout(dims, dst_strides, src_strides),
253        })
254    }
255
256    /// Check the per-call facts: the supplied views must carry exactly the
257    /// compiled layout. Buffer bounds against that layout were already proven
258    /// by [`RawStridedRef::new`]/[`RawStridedMut::new`] (or asserted by the
259    /// caller of the `new_unchecked` constructors), so layout equality is the
260    /// complete precondition for the unsafe-free fused replay below.
261    fn check_call<D, S>(
262        &self,
263        dest: &RawStridedMut<'_, D>,
264        src: &RawStridedRef<'_, S>,
265    ) -> Result<()> {
266        if dest.dims() != &self.dims[..]
267            || src.dims() != &self.dims[..]
268            || dest.strides() != &self.dst_strides[..]
269            || src.strides() != &self.src_strides[..]
270        {
271            return Err(StridedError::PlanLayoutMismatch);
272        }
273        Ok(())
274    }
275
276    /// `dest = src` into a potentially uninitialized destination.
277    ///
278    /// On success every reachable logical destination element is initialized.
279    /// Non-reachable holes in the backing allocation are not written.
280    pub fn execute_uninit<T>(
281        &self,
282        dest: &mut RawStridedMut<'_, MaybeUninit<T>>,
283        src: &RawStridedRef<'_, T>,
284    ) -> Result<()>
285    where
286        T: Copy + MaybeSendSync,
287    {
288        self.check_call(dest, src)?;
289        match &self.fused {
290            Some(layout) => {
291                apply_fused_pair(
292                    dest,
293                    src,
294                    layout,
295                    |dst, value| {
296                        dst.write(value);
297                    },
298                    |value| value,
299                );
300                Ok(())
301            }
302            None => map_raw_into::<MaybeUninit<T>, T, Identity>(dest, src, MaybeUninit::new),
303        }
304    }
305
306    /// `dest = src`. Allocation-free for ranks at most [`RAW_FUSED_RANK_LIMIT`](crate::RAW_FUSED_RANK_LIMIT).
307    pub fn execute<T>(
308        &self,
309        dest: &mut RawStridedMut<'_, T>,
310        src: &RawStridedRef<'_, T>,
311    ) -> Result<()>
312    where
313        T: Copy + MaybeSendSync,
314    {
315        self.check_call(dest, src)?;
316        match &self.fused {
317            Some(layout) => {
318                apply_fused_pair(
319                    dest,
320                    src,
321                    layout,
322                    |dst, value| *dst = value,
323                    |value: T| value,
324                );
325                Ok(())
326            }
327            None => copy_into(&mut dest.as_view_mut(), &src.as_view()),
328        }
329    }
330
331    /// `dest = scale * src`. Allocation-free for ranks at most
332    /// [`RAW_FUSED_RANK_LIMIT`](crate::RAW_FUSED_RANK_LIMIT).
333    pub fn execute_scale<T>(
334        &self,
335        dest: &mut RawStridedMut<'_, T>,
336        src: &RawStridedRef<'_, T>,
337        scale: T,
338    ) -> Result<()>
339    where
340        T: Copy + Mul<T, Output = T> + MaybeSendSync,
341    {
342        self.check_call(dest, src)?;
343        match &self.fused {
344            Some(layout) => {
345                apply_fused_pair(
346                    dest,
347                    src,
348                    layout,
349                    |dst, value| *dst = value,
350                    |value: T| scale * value,
351                );
352                Ok(())
353            }
354            None => copy_scale(&mut dest.as_view_mut(), &src.as_view(), scale),
355        }
356    }
357
358    /// `dest = conj(src)`. Allocation-free for ranks at most
359    /// [`RAW_FUSED_RANK_LIMIT`](crate::RAW_FUSED_RANK_LIMIT).
360    pub fn execute_conj<T>(
361        &self,
362        dest: &mut RawStridedMut<'_, T>,
363        src: &RawStridedRef<'_, T>,
364    ) -> Result<()>
365    where
366        T: Copy + ElementOpApply + MaybeSendSync,
367    {
368        self.check_call(dest, src)?;
369        match &self.fused {
370            Some(layout) => {
371                apply_fused_pair(
372                    dest,
373                    src,
374                    layout,
375                    |dst, value| *dst = value,
376                    |value: T| value.conj(),
377                );
378                Ok(())
379            }
380            None => copy_conj(&mut dest.as_view_mut(), &src.as_view()),
381        }
382    }
383}
384
385#[cfg(test)]
386#[path = "copy_plan/tests/tests.rs"]
387mod tests;