Skip to main content

strided_kernel/
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::ops::Mul;
13
14use crate::ops_view::{copy_conj, copy_into, copy_scale};
15use crate::raw_ops::{apply_fused_pair, fuse_pair_layout, FusedPairLayout};
16use crate::{ElementOpApply, MaybeSendSync, RawStridedMut, RawStridedRef, Result, StridedError};
17
18// Same pattern as map_view.rs / outer_product.rs: stack storage when the
19// parallel feature pulls in smallvec, plain Vec otherwise. Only `compile`
20// touches these; `execute*` never allocates either way.
21#[cfg(feature = "parallel")]
22type AxisVec<T> = smallvec::SmallVec<[T; crate::RAW_FUSED_RANK_LIMIT]>;
23#[cfg(not(feature = "parallel"))]
24type AxisVec<T> = Vec<T>;
25
26/// A compiled copy traversal for one `(dims, dst_strides, src_strides)`
27/// layout pair.
28///
29/// `compile` proves the layout facts once (rank agreement, extent overflow,
30/// destination injectivity) and fuses/orders the loop nest; each `execute*`
31/// call then only re-checks the per-call facts (that the supplied views carry
32/// exactly the compiled layout) before replaying the prepared loops.
33///
34/// Overlapping `src`/`dest` memory is not supported, matching the rest of the
35/// crate.
36///
37/// Ranks above [`RAW_FUSED_RANK_LIMIT`](crate::RAW_FUSED_RANK_LIMIT) are supported through the view-based
38/// kernels; only that fallback path may allocate.
39///
40/// # Example
41///
42/// ```rust
43/// use strided_kernel::{CopyPlan, RawStridedMut, RawStridedRef};
44///
45/// let dims = [2usize, 3];
46/// let src_strides = [3isize, 1];
47/// let dst_strides = [1isize, 2]; // transposed destination
48/// let plan = CopyPlan::compile(&dims, &dst_strides, &src_strides).unwrap();
49///
50/// let src = [0.0f64, 1.0, 2.0, 10.0, 11.0, 12.0];
51/// let mut dst = [0.0f64; 6];
52/// let src_ref = RawStridedRef::new(&src, &dims, &src_strides, 0).unwrap();
53/// let mut dst_mut = RawStridedMut::new(&mut dst, &dims, &dst_strides, 0).unwrap();
54/// plan.execute(&mut dst_mut, &src_ref).unwrap();
55/// assert_eq!(dst, [0.0, 10.0, 1.0, 11.0, 2.0, 12.0]);
56/// ```
57#[derive(Clone, Debug)]
58pub struct CopyPlan {
59    dims: AxisVec<usize>,
60    dst_strides: AxisVec<isize>,
61    src_strides: AxisVec<isize>,
62    /// `None` when rank exceeds [`RAW_FUSED_RANK_LIMIT`](crate::RAW_FUSED_RANK_LIMIT); `execute*` then
63    /// falls back to the view-based kernels.
64    fused: Option<FusedPairLayout>,
65}
66
67impl CopyPlan {
68    /// Compile a copy plan for the given layout pair.
69    ///
70    /// Performs the layout validation and traversal construction
71    /// (fuse + order) once:
72    ///
73    /// - `dims`, `dst_strides`, and `src_strides` must have equal length
74    ///   ([`StridedError::StrideLengthMismatch`]);
75    /// - the total element count must not overflow `usize`
76    ///   ([`StridedError::OffsetOverflow`]);
77    /// - the destination layout must be injective, i.e. map distinct logical
78    ///   indices to distinct offsets
79    ///   ([`StridedError::NonInjectiveOutputLayout`]).
80    pub fn compile(dims: &[usize], dst_strides: &[isize], src_strides: &[isize]) -> Result<Self> {
81        if dims.len() != dst_strides.len() || dims.len() != src_strides.len() {
82            return Err(StridedError::StrideLengthMismatch);
83        }
84        if dims
85            .iter()
86            .try_fold(1usize, |acc, &dim| acc.checked_mul(dim))
87            .is_none()
88        {
89            return Err(StridedError::OffsetOverflow);
90        }
91        if !crate::fused::is_injective_layout(dims, dst_strides) {
92            return Err(StridedError::NonInjectiveOutputLayout);
93        }
94        Ok(Self {
95            dims: dims.into(),
96            dst_strides: dst_strides.into(),
97            src_strides: src_strides.into(),
98            fused: fuse_pair_layout(dims, dst_strides, src_strides),
99        })
100    }
101
102    /// Check the per-call facts: the supplied views must carry exactly the
103    /// compiled layout. Buffer bounds against that layout were already proven
104    /// by [`RawStridedRef::new`]/[`RawStridedMut::new`] (or asserted by the
105    /// caller of the `new_unchecked` constructors), so layout equality is the
106    /// complete precondition for the unsafe-free fused replay below.
107    fn check_call<T>(&self, dest: &RawStridedMut<'_, T>, src: &RawStridedRef<'_, T>) -> Result<()> {
108        if dest.dims() != &self.dims[..]
109            || src.dims() != &self.dims[..]
110            || dest.strides() != &self.dst_strides[..]
111            || src.strides() != &self.src_strides[..]
112        {
113            return Err(StridedError::PlanLayoutMismatch);
114        }
115        Ok(())
116    }
117
118    /// `dest = src`. Allocation-free for ranks at most [`RAW_FUSED_RANK_LIMIT`](crate::RAW_FUSED_RANK_LIMIT).
119    pub fn execute<T>(
120        &self,
121        dest: &mut RawStridedMut<'_, T>,
122        src: &RawStridedRef<'_, T>,
123    ) -> Result<()>
124    where
125        T: Copy + MaybeSendSync,
126    {
127        self.check_call(dest, src)?;
128        match &self.fused {
129            Some(layout) => {
130                apply_fused_pair(
131                    dest,
132                    src,
133                    layout,
134                    |dst, value| *dst = value,
135                    |value: T| value,
136                );
137                Ok(())
138            }
139            None => copy_into(&mut dest.as_view_mut(), &src.as_view()),
140        }
141    }
142
143    /// `dest = scale * src`. Allocation-free for ranks at most
144    /// [`RAW_FUSED_RANK_LIMIT`](crate::RAW_FUSED_RANK_LIMIT).
145    pub fn execute_scale<T>(
146        &self,
147        dest: &mut RawStridedMut<'_, T>,
148        src: &RawStridedRef<'_, T>,
149        scale: T,
150    ) -> Result<()>
151    where
152        T: Copy + Mul<T, Output = T> + MaybeSendSync,
153    {
154        self.check_call(dest, src)?;
155        match &self.fused {
156            Some(layout) => {
157                apply_fused_pair(
158                    dest,
159                    src,
160                    layout,
161                    |dst, value| *dst = value,
162                    |value: T| scale * value,
163                );
164                Ok(())
165            }
166            None => copy_scale(&mut dest.as_view_mut(), &src.as_view(), scale),
167        }
168    }
169
170    /// `dest = conj(src)`. Allocation-free for ranks at most
171    /// [`RAW_FUSED_RANK_LIMIT`](crate::RAW_FUSED_RANK_LIMIT).
172    pub fn execute_conj<T>(
173        &self,
174        dest: &mut RawStridedMut<'_, T>,
175        src: &RawStridedRef<'_, T>,
176    ) -> Result<()>
177    where
178        T: Copy + ElementOpApply + MaybeSendSync,
179    {
180        self.check_call(dest, src)?;
181        match &self.fused {
182            Some(layout) => {
183                apply_fused_pair(
184                    dest,
185                    src,
186                    layout,
187                    |dst, value| *dst = value,
188                    |value: T| value.conj(),
189                );
190                Ok(())
191            }
192            None => copy_conj(&mut dest.as_view_mut(), &src.as_view()),
193        }
194    }
195}
196
197#[cfg(test)]
198mod tests {
199    use super::*;
200    use num_complex::{Complex32, Complex64};
201
202    /// Reference: the per-call raw kernel (which itself is differential-tested
203    /// against the view kernels in raw_ops.rs).
204    fn plan_matches_direct<T>(
205        dims: &[usize],
206        dst_strides: &[isize],
207        src_strides: &[isize],
208        src: &[T],
209    ) where
210        T: Copy
211            + PartialEq
212            + core::fmt::Debug
213            + Default
214            + Mul<T, Output = T>
215            + ElementOpApply
216            + MaybeSendSync
217            + num_traits::One,
218    {
219        let len = src.len();
220        let plan = CopyPlan::compile(dims, dst_strides, src_strides).unwrap();
221
222        let mut expected = vec![T::default(); len];
223        {
224            let mut dest = RawStridedMut::new(&mut expected, dims, dst_strides, 0).unwrap();
225            let source = RawStridedRef::new(src, dims, src_strides, 0).unwrap();
226            crate::copy_scale_raw(&mut dest, &source, T::one()).unwrap();
227        }
228
229        let mut actual = vec![T::default(); len];
230        {
231            let mut dest = RawStridedMut::new(&mut actual, dims, dst_strides, 0).unwrap();
232            let source = RawStridedRef::new(src, dims, src_strides, 0).unwrap();
233            plan.execute(&mut dest, &source).unwrap();
234        }
235        assert_eq!(actual, expected);
236    }
237
238    fn fill_f64(len: usize) -> Vec<f64> {
239        (0..len).map(|value| value as f64 - 2.5).collect()
240    }
241
242    #[test]
243    fn plan_copy_matches_direct_rank0() {
244        plan_matches_direct::<f64>(&[], &[], &[], &[7.0]);
245    }
246
247    #[test]
248    fn plan_copy_matches_direct_rank1() {
249        plan_matches_direct::<f64>(&[5], &[1], &[1], &fill_f64(5));
250    }
251
252    #[test]
253    fn plan_copy_matches_direct_rank2_transposed() {
254        plan_matches_direct::<f64>(&[3, 4], &[1, 3], &[4, 1], &fill_f64(12));
255    }
256
257    #[test]
258    fn plan_copy_matches_direct_rank4() {
259        plan_matches_direct::<f64>(&[2, 3, 2, 2], &[12, 4, 2, 1], &[1, 2, 6, 12], &fill_f64(24));
260    }
261
262    #[test]
263    fn plan_copy_matches_direct_rank8() {
264        let dims = [2usize; 8];
265        let dst: Vec<isize> = (0..8).map(|axis| 1isize << axis).collect();
266        let src: Vec<isize> = (0..8).rev().map(|axis| 1isize << axis).collect();
267        plan_matches_direct::<f64>(&dims, &dst, &src, &fill_f64(256));
268    }
269
270    #[test]
271    fn plan_copy_matches_direct_zero_size() {
272        plan_matches_direct::<f64>(&[2, 0, 3], &[3, 3, 1], &[1, 6, 2], &fill_f64(6));
273    }
274
275    #[test]
276    fn plan_copy_matches_direct_f32_and_complex() {
277        let dims = [2usize, 3];
278        let dst = [1isize, 2];
279        let src = [3isize, 1];
280        plan_matches_direct::<f32>(&dims, &dst, &src, &[1.0, 2.0, 3.0, 4.0, 5.0, 6.0]);
281        let complex: Vec<Complex32> = (0..6)
282            .map(|value| Complex32::new(value as f32, -(value as f32)))
283            .collect();
284        plan_matches_direct::<Complex32>(&dims, &dst, &src, &complex);
285        let complex: Vec<Complex64> = (0..6)
286            .map(|value| Complex64::new(value as f64, 1.0 - value as f64))
287            .collect();
288        plan_matches_direct::<Complex64>(&dims, &dst, &src, &complex);
289    }
290
291    #[test]
292    fn plan_copy_negative_stride_matches_view_kernel() {
293        // Negative source stride: src viewed reversed, offset at the end.
294        let dims = [4usize];
295        let src_strides = [-1isize];
296        let dst_strides = [1isize];
297        let src = [1.0f64, 2.0, 3.0, 4.0];
298        let plan = CopyPlan::compile(&dims, &dst_strides, &src_strides).unwrap();
299
300        let mut actual = [0.0f64; 4];
301        let mut dest = RawStridedMut::new(&mut actual, &dims, &dst_strides, 0).unwrap();
302        let source = RawStridedRef::new(&src, &dims, &src_strides, 3).unwrap();
303        plan.execute(&mut dest, &source).unwrap();
304        assert_eq!(actual, [4.0, 3.0, 2.0, 1.0]);
305    }
306
307    #[test]
308    fn plan_execute_scale_and_conj() {
309        let dims = [2usize, 2];
310        let strides = [2isize, 1];
311        let src = [
312            Complex64::new(1.0, 2.0),
313            Complex64::new(-3.0, 4.0),
314            Complex64::new(0.5, -1.0),
315            Complex64::new(2.0, 0.0),
316        ];
317        let plan = CopyPlan::compile(&dims, &strides, &strides).unwrap();
318
319        let mut scaled = [Complex64::default(); 4];
320        let mut dest = RawStridedMut::new(&mut scaled, &dims, &strides, 0).unwrap();
321        let source = RawStridedRef::new(&src, &dims, &strides, 0).unwrap();
322        plan.execute_scale(&mut dest, &source, Complex64::new(2.0, 0.0))
323            .unwrap();
324        assert_eq!(scaled[1], Complex64::new(-6.0, 8.0));
325
326        let mut conjugated = [Complex64::default(); 4];
327        let mut dest = RawStridedMut::new(&mut conjugated, &dims, &strides, 0).unwrap();
328        let source = RawStridedRef::new(&src, &dims, &strides, 0).unwrap();
329        plan.execute_conj(&mut dest, &source).unwrap();
330        assert_eq!(conjugated[0], Complex64::new(1.0, -2.0));
331        assert_eq!(conjugated[3], Complex64::new(2.0, 0.0));
332    }
333
334    #[test]
335    fn plan_rank_above_limit_falls_back_to_view_kernels() {
336        let dims = [2usize; 9];
337        let dst: Vec<isize> = (0..9).map(|axis| 1isize << axis).collect();
338        let src: Vec<isize> = (0..9).rev().map(|axis| 1isize << axis).collect();
339        let source_data = fill_f64(512);
340        let plan = CopyPlan::compile(&dims, &dst, &src).unwrap();
341        assert!(plan.fused.is_none());
342
343        let mut expected = vec![0.0f64; 512];
344        {
345            let mut dest = RawStridedMut::new(&mut expected, &dims, &dst, 0).unwrap();
346            let source = RawStridedRef::new(&source_data, &dims, &src, 0).unwrap();
347            crate::copy_scale_raw(&mut dest, &source, 1.0).unwrap();
348        }
349        let mut actual = vec![0.0f64; 512];
350        let mut dest = RawStridedMut::new(&mut actual, &dims, &dst, 0).unwrap();
351        let source = RawStridedRef::new(&source_data, &dims, &src, 0).unwrap();
352        plan.execute(&mut dest, &source).unwrap();
353        assert_eq!(actual, expected);
354
355        // Fallback also serves scale and conj.
356        let mut scaled = vec![0.0f64; 512];
357        let mut dest = RawStridedMut::new(&mut scaled, &dims, &dst, 0).unwrap();
358        plan.execute_scale(&mut dest, &source, 2.0).unwrap();
359        assert_eq!(scaled[0], 2.0 * actual[0]);
360        let mut conjugated = vec![0.0f64; 512];
361        let mut dest = RawStridedMut::new(&mut conjugated, &dims, &dst, 0).unwrap();
362        plan.execute_conj(&mut dest, &source).unwrap();
363        assert_eq!(conjugated, actual);
364    }
365
366    #[test]
367    fn compile_rejects_length_mismatch() {
368        let err = CopyPlan::compile(&[2, 3], &[3, 1], &[1]).unwrap_err();
369        assert!(matches!(err, StridedError::StrideLengthMismatch));
370        let err = CopyPlan::compile(&[2, 3], &[3], &[1, 2]).unwrap_err();
371        assert!(matches!(err, StridedError::StrideLengthMismatch));
372    }
373
374    #[test]
375    fn compile_rejects_extent_overflow() {
376        let err = CopyPlan::compile(&[usize::MAX, 2], &[1, 1], &[1, 1]).unwrap_err();
377        assert!(matches!(err, StridedError::OffsetOverflow));
378    }
379
380    #[test]
381    fn compile_rejects_non_injective_destination() {
382        // Two logical columns land on the same offsets: forbidden overlap in
383        // the mutable destination.
384        let err = CopyPlan::compile(&[2, 2], &[1, 0], &[2, 1]).unwrap_err();
385        assert!(matches!(err, StridedError::NonInjectiveOutputLayout));
386        // Broadcast-like (stride 0) source layouts remain allowed.
387        CopyPlan::compile(&[2, 2], &[2, 1], &[0, 1]).unwrap();
388    }
389
390    #[test]
391    fn execute_rejects_layout_drift() {
392        let dims = [2usize, 3];
393        let strides = [3isize, 1];
394        let plan = CopyPlan::compile(&dims, &strides, &strides).unwrap();
395        let src = fill_f64(6);
396        let mut dst = vec![0.0f64; 6];
397
398        // Different dims than compiled.
399        let other_dims = [3usize, 2];
400        let other_strides = [2isize, 1];
401        let mut dest = RawStridedMut::new(&mut dst, &other_dims, &other_strides, 0).unwrap();
402        let source = RawStridedRef::new(&src, &other_dims, &other_strides, 0).unwrap();
403        let err = plan.execute(&mut dest, &source).unwrap_err();
404        assert!(matches!(err, StridedError::PlanLayoutMismatch));
405
406        // Same dims, different source strides.
407        let column_major = [1isize, 2];
408        let mut dest = RawStridedMut::new(&mut dst, &dims, &strides, 0).unwrap();
409        let source = RawStridedRef::new(&src, &dims, &column_major, 0).unwrap();
410        let err = plan.execute_scale(&mut dest, &source, 1.0).unwrap_err();
411        assert!(matches!(err, StridedError::PlanLayoutMismatch));
412
413        // Same dims, different destination strides.
414        let mut dest = RawStridedMut::new(&mut dst, &dims, &column_major, 0).unwrap();
415        let source = RawStridedRef::new(&src, &dims, &strides, 0).unwrap();
416        let err = plan.execute_conj(&mut dest, &source).unwrap_err();
417        assert!(matches!(err, StridedError::PlanLayoutMismatch));
418    }
419
420    #[test]
421    fn identity_layout_uses_single_fused_axis() {
422        let plan = CopyPlan::compile(&[2, 3, 4], &[12, 4, 1], &[12, 4, 1]).unwrap();
423        let fused = plan.fused.expect("rank 3 stays on the fused path");
424        assert_eq!(fused.rank, 1);
425        assert_eq!(fused.dims[0], 24);
426    }
427}