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        if dims
243            .iter()
244            .try_fold(1usize, |acc, &dim| acc.checked_mul(dim))
245            .is_none()
246        {
247            return Err(StridedError::OffsetOverflow);
248        }
249        if !crate::layout_check::is_injective_layout(dims, dst_strides) {
250            return Err(StridedError::NonInjectiveOutputLayout);
251        }
252        Ok(Self {
253            dims: dims.into(),
254            dst_strides: dst_strides.into(),
255            src_strides: src_strides.into(),
256            fused: fuse_pair_layout(dims, dst_strides, src_strides),
257        })
258    }
259
260    /// Check the per-call facts: the supplied views must carry exactly the
261    /// compiled layout. Buffer bounds against that layout were already proven
262    /// by [`RawStridedRef::new`]/[`RawStridedMut::new`] (or asserted by the
263    /// caller of the `new_unchecked` constructors), so layout equality is the
264    /// complete precondition for the unsafe-free fused replay below.
265    fn check_call<D, S>(
266        &self,
267        dest: &RawStridedMut<'_, D>,
268        src: &RawStridedRef<'_, S>,
269    ) -> Result<()> {
270        if dest.dims() != &self.dims[..]
271            || src.dims() != &self.dims[..]
272            || dest.strides() != &self.dst_strides[..]
273            || src.strides() != &self.src_strides[..]
274        {
275            return Err(StridedError::PlanLayoutMismatch);
276        }
277        Ok(())
278    }
279
280    /// `dest = src` into a potentially uninitialized destination.
281    ///
282    /// On success every reachable logical destination element is initialized.
283    /// Non-reachable holes in the backing allocation are not written.
284    pub fn execute_uninit<T>(
285        &self,
286        dest: &mut RawStridedMut<'_, MaybeUninit<T>>,
287        src: &RawStridedRef<'_, T>,
288    ) -> Result<()>
289    where
290        T: Copy + MaybeSendSync,
291    {
292        self.check_call(dest, src)?;
293        match &self.fused {
294            Some(layout) => {
295                apply_fused_pair(
296                    dest,
297                    src,
298                    layout,
299                    |dst, value| {
300                        dst.write(value);
301                    },
302                    |value| value,
303                );
304                Ok(())
305            }
306            None => map_raw_into::<MaybeUninit<T>, T, Identity>(dest, src, MaybeUninit::new),
307        }
308    }
309
310    /// `dest = src`. Allocation-free for ranks at most [`RAW_FUSED_RANK_LIMIT`](crate::RAW_FUSED_RANK_LIMIT).
311    pub fn execute<T>(
312        &self,
313        dest: &mut RawStridedMut<'_, T>,
314        src: &RawStridedRef<'_, T>,
315    ) -> Result<()>
316    where
317        T: Copy + MaybeSendSync,
318    {
319        self.check_call(dest, src)?;
320        match &self.fused {
321            Some(layout) => {
322                apply_fused_pair(
323                    dest,
324                    src,
325                    layout,
326                    |dst, value| *dst = value,
327                    |value: T| value,
328                );
329                Ok(())
330            }
331            None => copy_into(&mut dest.as_view_mut(), &src.as_view()),
332        }
333    }
334
335    /// `dest = scale * src`. Allocation-free for ranks at most
336    /// [`RAW_FUSED_RANK_LIMIT`](crate::RAW_FUSED_RANK_LIMIT).
337    pub fn execute_scale<T>(
338        &self,
339        dest: &mut RawStridedMut<'_, T>,
340        src: &RawStridedRef<'_, T>,
341        scale: T,
342    ) -> Result<()>
343    where
344        T: Copy + Mul<T, Output = T> + MaybeSendSync,
345    {
346        self.check_call(dest, src)?;
347        match &self.fused {
348            Some(layout) => {
349                apply_fused_pair(
350                    dest,
351                    src,
352                    layout,
353                    |dst, value| *dst = value,
354                    |value: T| scale * value,
355                );
356                Ok(())
357            }
358            None => copy_scale(&mut dest.as_view_mut(), &src.as_view(), scale),
359        }
360    }
361
362    /// `dest = conj(src)`. Allocation-free for ranks at most
363    /// [`RAW_FUSED_RANK_LIMIT`](crate::RAW_FUSED_RANK_LIMIT).
364    pub fn execute_conj<T>(
365        &self,
366        dest: &mut RawStridedMut<'_, T>,
367        src: &RawStridedRef<'_, T>,
368    ) -> Result<()>
369    where
370        T: Copy + ElementOpApply + MaybeSendSync,
371    {
372        self.check_call(dest, src)?;
373        match &self.fused {
374            Some(layout) => {
375                apply_fused_pair(
376                    dest,
377                    src,
378                    layout,
379                    |dst, value| *dst = value,
380                    |value: T| value.conj(),
381                );
382                Ok(())
383            }
384            None => copy_conj(&mut dest.as_view_mut(), &src.as_view()),
385        }
386    }
387}
388
389#[cfg(test)]
390#[path = "copy_plan/tests/tests.rs"]
391mod tests;