1use 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#[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#[derive(Clone, Debug)]
58pub struct CopyPlan {
59 dims: AxisVec<usize>,
60 dst_strides: AxisVec<isize>,
61 src_strides: AxisVec<isize>,
62 fused: Option<FusedPairLayout>,
65}
66
67impl CopyPlan {
68 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 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 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 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 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 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 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 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 let err = CopyPlan::compile(&[2, 2], &[1, 0], &[2, 1]).unwrap_err();
385 assert!(matches!(err, StridedError::NonInjectiveOutputLayout));
386 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 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 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 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}