1use 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#[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 unsafe fn data_ptr(&mut self) -> *mut T;
41 unsafe fn write_at(&mut self, offset: isize, value: T);
44}
45
46pub(crate) trait ReadModifyWrite<T>: OverwriteWriter<T> {
47 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 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 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 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 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 unsafe {
146 let ptr = self.ptr.offset(offset);
147 ptr.write(combine(ptr.read(), value));
148 }
149 }
150}
151
152#[derive(Clone, Debug)]
184pub struct CopyPlan {
185 dims: AxisVec<usize>,
186 dst_strides: AxisVec<isize>,
187 src_strides: AxisVec<isize>,
188 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 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 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 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 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 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 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 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 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 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 let err = CopyPlan::compile(&[2, 2], &[1, 0], &[2, 1]).unwrap_err();
604 assert!(matches!(err, StridedError::NonInjectiveOutputLayout));
605 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 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 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 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}