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::{
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(&mut self, offset: isize, value: T, combine: fn(T, T) -> T);
51}
52
53impl<'a, T> OverwriteWriter<T> for RawStridedMut<'a, T> {
54    fn dims(&self) -> &[usize] {
55        self.dims()
56    }
57    fn strides(&self) -> &[isize] {
58        self.strides()
59    }
60    fn offset(&self) -> isize {
61        self.offset()
62    }
63    unsafe fn data_ptr(&mut self) -> *mut T {
64        self.data_mut().as_mut_ptr()
65    }
66    unsafe fn write_at(&mut self, offset: isize, value: T) {
67        // SAFETY: the prepared layout validates every logical destination.
68        unsafe { self.data_mut().as_mut_ptr().offset(offset).write(value) }
69    }
70}
71
72impl<'a, T> ReadModifyWrite<T> for RawStridedMut<'a, T>
73where
74    T: Add<Output = T>,
75{
76    unsafe fn add_at(&mut self, offset: isize, value: T, combine: fn(T, T) -> T) {
77        // SAFETY: the copy or initialized caller proves this logical slot.
78        unsafe {
79            let ptr = self.data_mut().as_mut_ptr().offset(offset);
80            ptr.write(combine(ptr.read(), value));
81        }
82    }
83}
84
85impl<'a, T> OverwriteWriter<T> for RawStridedMut<'a, MaybeUninit<T>> {
86    fn dims(&self) -> &[usize] {
87        self.dims()
88    }
89    fn strides(&self) -> &[isize] {
90        self.strides()
91    }
92    fn offset(&self) -> isize {
93        self.offset()
94    }
95    unsafe fn data_ptr(&mut self) -> *mut T {
96        self.data_mut().as_mut_ptr().cast()
97    }
98    unsafe fn write_at(&mut self, offset: isize, value: T) {
99        // SAFETY: the prepared layout validates every logical destination.
100        unsafe {
101            self.data_mut()
102                .as_mut_ptr()
103                .offset(offset)
104                .write(MaybeUninit::new(value))
105        }
106    }
107}
108
109pub(crate) struct InitializedRawDest<'a, T> {
110    ptr: *mut T,
111    extent: usize,
112    dims: &'a [usize],
113    strides: &'a [isize],
114    offset: isize,
115    _marker: PhantomData<&'a mut [MaybeUninit<T>]>,
116}
117
118impl<'a, T> OverwriteWriter<T> for InitializedRawDest<'a, T> {
119    fn dims(&self) -> &[usize] {
120        self.dims
121    }
122    fn strides(&self) -> &[isize] {
123        self.strides
124    }
125    fn offset(&self) -> isize {
126        self.offset
127    }
128    unsafe fn data_ptr(&mut self) -> *mut T {
129        self.ptr
130    }
131    unsafe fn write_at(&mut self, offset: isize, value: T) {
132        debug_assert!(offset >= 0 && (offset as usize) < self.extent);
133        // SAFETY: the copy proof and extent check cover this logical slot.
134        unsafe { self.ptr.offset(offset).write(value) }
135    }
136}
137
138impl<'a, T> ReadModifyWrite<T> for InitializedRawDest<'a, T>
139where
140    T: Add<Output = T>,
141{
142    unsafe fn add_at(&mut self, offset: isize, value: T, combine: fn(T, T) -> T) {
143        debug_assert!(offset >= 0 && (offset as usize) < self.extent);
144        // SAFETY: the copy proof and extent check cover this logical slot.
145        unsafe {
146            let ptr = self.ptr.offset(offset);
147            ptr.write(combine(ptr.read(), value));
148        }
149    }
150}
151
152/// A compiled copy traversal for one `(dims, dst_strides, src_strides)`
153/// layout pair.
154///
155/// `compile` proves the layout facts once (rank agreement, extent overflow,
156/// destination injectivity) and fuses/orders the loop nest; each `execute*`
157/// call then only re-checks the per-call facts (that the supplied views carry
158/// exactly the compiled layout) before replaying the prepared loops.
159///
160/// Overlapping `src`/`dest` memory is not supported, matching the rest of the
161/// crate.
162///
163/// Ranks above [`RAW_FUSED_RANK_LIMIT`](crate::RAW_FUSED_RANK_LIMIT) are supported through the view-based
164/// kernels; only that fallback path may allocate.
165///
166/// # Example
167///
168/// ```rust
169/// use strided_kernel::{CopyPlan, RawStridedMut, RawStridedRef};
170///
171/// let dims = [2usize, 3];
172/// let src_strides = [3isize, 1];
173/// let dst_strides = [1isize, 2]; // transposed destination
174/// let plan = CopyPlan::compile(&dims, &dst_strides, &src_strides).unwrap();
175///
176/// let src = [0.0f64, 1.0, 2.0, 10.0, 11.0, 12.0];
177/// let mut dst = [0.0f64; 6];
178/// let src_ref = RawStridedRef::new(&src, &dims, &src_strides, 0).unwrap();
179/// let mut dst_mut = RawStridedMut::new(&mut dst, &dims, &dst_strides, 0).unwrap();
180/// plan.execute(&mut dst_mut, &src_ref).unwrap();
181/// assert_eq!(dst, [0.0, 10.0, 1.0, 11.0, 2.0, 12.0]);
182/// ```
183#[derive(Clone, Debug)]
184pub struct CopyPlan {
185    dims: AxisVec<usize>,
186    dst_strides: AxisVec<isize>,
187    src_strides: AxisVec<isize>,
188    /// `None` when rank exceeds [`RAW_FUSED_RANK_LIMIT`](crate::RAW_FUSED_RANK_LIMIT); `execute*` then
189    /// falls back to the view-based kernels.
190    fused: Option<FusedPairLayout>,
191}
192
193impl CopyPlan {
194    pub(crate) fn execute_uninit_then<'a, T, R>(
195        &self,
196        dest: &'a mut RawStridedMut<'a, MaybeUninit<T>>,
197        src: &RawStridedRef<'_, T>,
198        f: impl for<'b> FnOnce(InitializedRawDest<'b, T>) -> R,
199    ) -> Result<R>
200    where
201        T: Copy + MaybeSendSync,
202    {
203        self.execute_uninit(dest, src)?;
204        let data = dest.data_mut();
205        let receipt = InitializedRawDest {
206            ptr: data.as_mut_ptr().cast(),
207            extent: data.len(),
208            dims: dest.dims(),
209            strides: dest.strides(),
210            offset: dest.offset(),
211            _marker: PhantomData,
212        };
213        Ok(f(receipt))
214    }
215
216    /// Compile a copy plan for the given layout pair.
217    ///
218    /// Performs the layout validation and traversal construction
219    /// (fuse + order) once:
220    ///
221    /// - `dims`, `dst_strides`, and `src_strides` must have equal length
222    ///   ([`StridedError::StrideLengthMismatch`]);
223    /// - the total element count must not overflow `usize`
224    ///   ([`StridedError::OffsetOverflow`]);
225    /// - the destination layout must be injective, i.e. map distinct logical
226    ///   indices to distinct offsets
227    ///   ([`StridedError::NonInjectiveOutputLayout`]).
228    pub fn compile(dims: &[usize], dst_strides: &[isize], src_strides: &[isize]) -> Result<Self> {
229        if dims.len() != dst_strides.len() || dims.len() != src_strides.len() {
230            return Err(StridedError::StrideLengthMismatch);
231        }
232        if dims
233            .iter()
234            .try_fold(1usize, |acc, &dim| acc.checked_mul(dim))
235            .is_none()
236        {
237            return Err(StridedError::OffsetOverflow);
238        }
239        if !crate::fused::is_injective_layout(dims, dst_strides) {
240            return Err(StridedError::NonInjectiveOutputLayout);
241        }
242        Ok(Self {
243            dims: dims.into(),
244            dst_strides: dst_strides.into(),
245            src_strides: src_strides.into(),
246            fused: fuse_pair_layout(dims, dst_strides, src_strides),
247        })
248    }
249
250    /// Check the per-call facts: the supplied views must carry exactly the
251    /// compiled layout. Buffer bounds against that layout were already proven
252    /// by [`RawStridedRef::new`]/[`RawStridedMut::new`] (or asserted by the
253    /// caller of the `new_unchecked` constructors), so layout equality is the
254    /// complete precondition for the unsafe-free fused replay below.
255    fn check_call<D, S>(
256        &self,
257        dest: &RawStridedMut<'_, D>,
258        src: &RawStridedRef<'_, S>,
259    ) -> Result<()> {
260        if dest.dims() != &self.dims[..]
261            || src.dims() != &self.dims[..]
262            || dest.strides() != &self.dst_strides[..]
263            || src.strides() != &self.src_strides[..]
264        {
265            return Err(StridedError::PlanLayoutMismatch);
266        }
267        Ok(())
268    }
269
270    /// `dest = src` into a potentially uninitialized destination.
271    ///
272    /// On success every reachable logical destination element is initialized.
273    /// Non-reachable holes in the backing allocation are not written.
274    pub fn execute_uninit<T>(
275        &self,
276        dest: &mut RawStridedMut<'_, MaybeUninit<T>>,
277        src: &RawStridedRef<'_, T>,
278    ) -> Result<()>
279    where
280        T: Copy + MaybeSendSync,
281    {
282        self.check_call(dest, src)?;
283        match &self.fused {
284            Some(layout) => {
285                apply_fused_pair(
286                    dest,
287                    src,
288                    layout,
289                    |dst, value| {
290                        dst.write(value);
291                    },
292                    |value| value,
293                );
294                Ok(())
295            }
296            None => map_raw_into::<MaybeUninit<T>, T, Identity>(dest, src, MaybeUninit::new),
297        }
298    }
299
300    /// `dest = src`. Allocation-free for ranks at most [`RAW_FUSED_RANK_LIMIT`](crate::RAW_FUSED_RANK_LIMIT).
301    pub fn execute<T>(
302        &self,
303        dest: &mut RawStridedMut<'_, T>,
304        src: &RawStridedRef<'_, T>,
305    ) -> Result<()>
306    where
307        T: Copy + MaybeSendSync,
308    {
309        self.check_call(dest, src)?;
310        match &self.fused {
311            Some(layout) => {
312                apply_fused_pair(
313                    dest,
314                    src,
315                    layout,
316                    |dst, value| *dst = value,
317                    |value: T| value,
318                );
319                Ok(())
320            }
321            None => copy_into(&mut dest.as_view_mut(), &src.as_view()),
322        }
323    }
324
325    /// `dest = scale * src`. Allocation-free for ranks at most
326    /// [`RAW_FUSED_RANK_LIMIT`](crate::RAW_FUSED_RANK_LIMIT).
327    pub fn execute_scale<T>(
328        &self,
329        dest: &mut RawStridedMut<'_, T>,
330        src: &RawStridedRef<'_, T>,
331        scale: T,
332    ) -> Result<()>
333    where
334        T: Copy + Mul<T, Output = T> + MaybeSendSync,
335    {
336        self.check_call(dest, src)?;
337        match &self.fused {
338            Some(layout) => {
339                apply_fused_pair(
340                    dest,
341                    src,
342                    layout,
343                    |dst, value| *dst = value,
344                    |value: T| scale * value,
345                );
346                Ok(())
347            }
348            None => copy_scale(&mut dest.as_view_mut(), &src.as_view(), scale),
349        }
350    }
351
352    /// `dest = conj(src)`. Allocation-free for ranks at most
353    /// [`RAW_FUSED_RANK_LIMIT`](crate::RAW_FUSED_RANK_LIMIT).
354    pub fn execute_conj<T>(
355        &self,
356        dest: &mut RawStridedMut<'_, T>,
357        src: &RawStridedRef<'_, T>,
358    ) -> Result<()>
359    where
360        T: Copy + ElementOpApply + MaybeSendSync,
361    {
362        self.check_call(dest, src)?;
363        match &self.fused {
364            Some(layout) => {
365                apply_fused_pair(
366                    dest,
367                    src,
368                    layout,
369                    |dst, value| *dst = value,
370                    |value: T| value.conj(),
371                );
372                Ok(())
373            }
374            None => copy_conj(&mut dest.as_view_mut(), &src.as_view()),
375        }
376    }
377}
378
379#[cfg(test)]
380mod tests {
381    use super::*;
382    use num_complex::{Complex32, Complex64};
383
384    #[test]
385    fn uninit_then_receipt_drops_after_panic() {
386        use std::panic::{catch_unwind, AssertUnwindSafe};
387        let plan = CopyPlan::compile(&[2], &[1], &[1]).unwrap();
388        let source_data = [3i32, 5];
389        let source = RawStridedRef::new(&source_data, &[2], &[1], 0).unwrap();
390        let result = catch_unwind(AssertUnwindSafe(|| {
391            let mut storage = vec![MaybeUninit::<i32>::uninit(); 3];
392            let mut dest = RawStridedMut::new(&mut storage, &[2], &[1], 0).unwrap();
393            let _: () = plan
394                .execute_uninit_then(&mut dest, &source, |_receipt| {
395                    panic!("post-copy update failure");
396                })
397                .unwrap();
398        }));
399        assert!(result.is_err());
400    }
401
402    /// Reference: the per-call raw kernel (which itself is differential-tested
403    /// against the view kernels in raw_ops.rs).
404    fn plan_matches_direct<T>(
405        dims: &[usize],
406        dst_strides: &[isize],
407        src_strides: &[isize],
408        src: &[T],
409    ) where
410        T: Copy
411            + PartialEq
412            + core::fmt::Debug
413            + Default
414            + Mul<T, Output = T>
415            + ElementOpApply
416            + MaybeSendSync
417            + num_traits::One,
418    {
419        let len = src.len();
420        let plan = CopyPlan::compile(dims, dst_strides, src_strides).unwrap();
421
422        let mut expected = vec![T::default(); len];
423        {
424            let mut dest = RawStridedMut::new(&mut expected, dims, dst_strides, 0).unwrap();
425            let source = RawStridedRef::new(src, dims, src_strides, 0).unwrap();
426            crate::copy_scale_raw(&mut dest, &source, T::one()).unwrap();
427        }
428
429        let mut actual = vec![T::default(); len];
430        {
431            let mut dest = RawStridedMut::new(&mut actual, dims, dst_strides, 0).unwrap();
432            let source = RawStridedRef::new(src, dims, src_strides, 0).unwrap();
433            plan.execute(&mut dest, &source).unwrap();
434        }
435        assert_eq!(actual, expected);
436    }
437
438    fn fill_f64(len: usize) -> Vec<f64> {
439        (0..len).map(|value| value as f64 - 2.5).collect()
440    }
441
442    #[test]
443    fn plan_copy_matches_direct_rank0() {
444        plan_matches_direct::<f64>(&[], &[], &[], &[7.0]);
445    }
446
447    #[test]
448    fn plan_copy_matches_direct_rank1() {
449        plan_matches_direct::<f64>(&[5], &[1], &[1], &fill_f64(5));
450    }
451
452    #[test]
453    fn plan_copy_matches_direct_rank2_transposed() {
454        plan_matches_direct::<f64>(&[3, 4], &[1, 3], &[4, 1], &fill_f64(12));
455    }
456
457    #[test]
458    fn plan_copy_matches_direct_rank4() {
459        plan_matches_direct::<f64>(&[2, 3, 2, 2], &[12, 4, 2, 1], &[1, 2, 6, 12], &fill_f64(24));
460    }
461
462    #[test]
463    fn plan_copy_matches_direct_rank8() {
464        let dims = [2usize; 8];
465        let dst: Vec<isize> = (0..8).map(|axis| 1isize << axis).collect();
466        let src: Vec<isize> = (0..8).rev().map(|axis| 1isize << axis).collect();
467        plan_matches_direct::<f64>(&dims, &dst, &src, &fill_f64(256));
468    }
469
470    #[test]
471    fn plan_copy_matches_direct_zero_size() {
472        plan_matches_direct::<f64>(&[2, 0, 3], &[3, 3, 1], &[1, 6, 2], &fill_f64(6));
473    }
474
475    #[test]
476    fn plan_copy_matches_direct_f32_and_complex() {
477        let dims = [2usize, 3];
478        let dst = [1isize, 2];
479        let src = [3isize, 1];
480        plan_matches_direct::<f32>(&dims, &dst, &src, &[1.0, 2.0, 3.0, 4.0, 5.0, 6.0]);
481        let complex: Vec<Complex32> = (0..6)
482            .map(|value| Complex32::new(value as f32, -(value as f32)))
483            .collect();
484        plan_matches_direct::<Complex32>(&dims, &dst, &src, &complex);
485        let complex: Vec<Complex64> = (0..6)
486            .map(|value| Complex64::new(value as f64, 1.0 - value as f64))
487            .collect();
488        plan_matches_direct::<Complex64>(&dims, &dst, &src, &complex);
489    }
490
491    #[test]
492    fn plan_copy_negative_stride_matches_view_kernel() {
493        // Negative source stride: src viewed reversed, offset at the end.
494        let dims = [4usize];
495        let src_strides = [-1isize];
496        let dst_strides = [1isize];
497        let src = [1.0f64, 2.0, 3.0, 4.0];
498        let plan = CopyPlan::compile(&dims, &dst_strides, &src_strides).unwrap();
499
500        let mut actual = [0.0f64; 4];
501        let mut dest = RawStridedMut::new(&mut actual, &dims, &dst_strides, 0).unwrap();
502        let source = RawStridedRef::new(&src, &dims, &src_strides, 3).unwrap();
503        plan.execute(&mut dest, &source).unwrap();
504        assert_eq!(actual, [4.0, 3.0, 2.0, 1.0]);
505    }
506
507    #[test]
508    fn plan_execute_scale_and_conj() {
509        let dims = [2usize, 2];
510        let strides = [2isize, 1];
511        let src = [
512            Complex64::new(1.0, 2.0),
513            Complex64::new(-3.0, 4.0),
514            Complex64::new(0.5, -1.0),
515            Complex64::new(2.0, 0.0),
516        ];
517        let plan = CopyPlan::compile(&dims, &strides, &strides).unwrap();
518
519        let mut scaled = [Complex64::default(); 4];
520        let mut dest = RawStridedMut::new(&mut scaled, &dims, &strides, 0).unwrap();
521        let source = RawStridedRef::new(&src, &dims, &strides, 0).unwrap();
522        plan.execute_scale(&mut dest, &source, Complex64::new(2.0, 0.0))
523            .unwrap();
524        assert_eq!(scaled[1], Complex64::new(-6.0, 8.0));
525
526        let mut conjugated = [Complex64::default(); 4];
527        let mut dest = RawStridedMut::new(&mut conjugated, &dims, &strides, 0).unwrap();
528        let source = RawStridedRef::new(&src, &dims, &strides, 0).unwrap();
529        plan.execute_conj(&mut dest, &source).unwrap();
530        assert_eq!(conjugated[0], Complex64::new(1.0, -2.0));
531        assert_eq!(conjugated[3], Complex64::new(2.0, 0.0));
532    }
533
534    #[test]
535    fn plan_rank_above_limit_falls_back_to_view_kernels() {
536        let dims = [2usize; 9];
537        let dst: Vec<isize> = (0..9).map(|axis| 1isize << axis).collect();
538        let src: Vec<isize> = (0..9).rev().map(|axis| 1isize << axis).collect();
539        let source_data = fill_f64(512);
540        let plan = CopyPlan::compile(&dims, &dst, &src).unwrap();
541        assert!(plan.fused.is_none());
542
543        let mut expected = vec![0.0f64; 512];
544        {
545            let mut dest = RawStridedMut::new(&mut expected, &dims, &dst, 0).unwrap();
546            let source = RawStridedRef::new(&source_data, &dims, &src, 0).unwrap();
547            crate::copy_scale_raw(&mut dest, &source, 1.0).unwrap();
548        }
549        let mut actual = vec![0.0f64; 512];
550        let mut dest = RawStridedMut::new(&mut actual, &dims, &dst, 0).unwrap();
551        let source = RawStridedRef::new(&source_data, &dims, &src, 0).unwrap();
552        plan.execute(&mut dest, &source).unwrap();
553        assert_eq!(actual, expected);
554
555        // Fallback also serves scale and conj.
556        let mut scaled = vec![0.0f64; 512];
557        let mut dest = RawStridedMut::new(&mut scaled, &dims, &dst, 0).unwrap();
558        plan.execute_scale(&mut dest, &source, 2.0).unwrap();
559        assert_eq!(scaled[0], 2.0 * actual[0]);
560        let mut conjugated = vec![0.0f64; 512];
561        let mut dest = RawStridedMut::new(&mut conjugated, &dims, &dst, 0).unwrap();
562        plan.execute_conj(&mut dest, &source).unwrap();
563        assert_eq!(conjugated, actual);
564    }
565
566    #[test]
567    fn compile_rejects_length_mismatch() {
568        let err = CopyPlan::compile(&[2, 3], &[3, 1], &[1]).unwrap_err();
569        assert!(matches!(err, StridedError::StrideLengthMismatch));
570        let err = CopyPlan::compile(&[2, 3], &[3], &[1, 2]).unwrap_err();
571        assert!(matches!(err, StridedError::StrideLengthMismatch));
572    }
573
574    #[test]
575    fn compile_rejects_extent_overflow() {
576        let err = CopyPlan::compile(&[usize::MAX, 2], &[1, 1], &[1, 1]).unwrap_err();
577        assert!(matches!(err, StridedError::OffsetOverflow));
578    }
579
580    #[test]
581    fn compile_rejects_unrepresentable_positive_and_negative_offset_spans() {
582        for strides in [
583            [isize::MAX / 2 + 1, isize::MAX],
584            [isize::MIN / 2 - 1, isize::MIN],
585        ] {
586            let err = CopyPlan::compile(&[2, 2], &strides, &strides).unwrap_err();
587            assert!(matches!(err, StridedError::NonInjectiveOutputLayout));
588        }
589    }
590
591    #[test]
592    fn compile_accepts_representable_mixed_sign_span_without_fusion_overflow() {
593        let positive = isize::MAX / 4;
594        let negative = -(isize::MAX - positive);
595        let strides = [positive, negative];
596        CopyPlan::compile(&[2, 2], &strides, &strides).unwrap();
597    }
598
599    #[test]
600    fn compile_rejects_non_injective_destination() {
601        // Two logical columns land on the same offsets: forbidden overlap in
602        // the mutable destination.
603        let err = CopyPlan::compile(&[2, 2], &[1, 0], &[2, 1]).unwrap_err();
604        assert!(matches!(err, StridedError::NonInjectiveOutputLayout));
605        // Broadcast-like (stride 0) source layouts remain allowed.
606        CopyPlan::compile(&[2, 2], &[2, 1], &[0, 1]).unwrap();
607    }
608
609    #[test]
610    fn execute_rejects_layout_drift() {
611        let dims = [2usize, 3];
612        let strides = [3isize, 1];
613        let plan = CopyPlan::compile(&dims, &strides, &strides).unwrap();
614        let src = fill_f64(6);
615        let mut dst = vec![0.0f64; 6];
616
617        // Different dims than compiled.
618        let other_dims = [3usize, 2];
619        let other_strides = [2isize, 1];
620        let mut dest = RawStridedMut::new(&mut dst, &other_dims, &other_strides, 0).unwrap();
621        let source = RawStridedRef::new(&src, &other_dims, &other_strides, 0).unwrap();
622        let err = plan.execute(&mut dest, &source).unwrap_err();
623        assert!(matches!(err, StridedError::PlanLayoutMismatch));
624
625        // Same dims, different source strides.
626        let column_major = [1isize, 2];
627        let mut dest = RawStridedMut::new(&mut dst, &dims, &strides, 0).unwrap();
628        let source = RawStridedRef::new(&src, &dims, &column_major, 0).unwrap();
629        let err = plan.execute_scale(&mut dest, &source, 1.0).unwrap_err();
630        assert!(matches!(err, StridedError::PlanLayoutMismatch));
631
632        // Same dims, different destination strides.
633        let mut dest = RawStridedMut::new(&mut dst, &dims, &column_major, 0).unwrap();
634        let source = RawStridedRef::new(&src, &dims, &strides, 0).unwrap();
635        let err = plan.execute_conj(&mut dest, &source).unwrap_err();
636        assert!(matches!(err, StridedError::PlanLayoutMismatch));
637    }
638
639    #[test]
640    fn identity_layout_uses_single_fused_axis() {
641        let plan = CopyPlan::compile(&[2, 3, 4], &[12, 4, 1], &[12, 4, 1]).unwrap();
642        let fused = plan.fused.expect("rank 3 stays on the fused path");
643        assert_eq!(fused.rank, 1);
644        assert_eq!(fused.dims[0], 24);
645    }
646}