1use crate::erased_common::*;
2use crate::*;
3use core::mem::MaybeUninit;
4use num_complex::{Complex32, Complex64};
5use num_traits::{One, Zero};
6mod arg_reduce;
7mod line;
8mod norm;
9mod reduce_kernel;
10mod scan;
11
12pub use arg_reduce::{ArgReduceOp, ErasedArgReducePlan};
13pub use norm::{ErasedNormPlan, NormKind, NormSpec};
14pub use scan::{ErasedScanPlan, ScanOp, ScanOptions};
15
16use reduce_kernel::{
17 reduce_axes_range, reduce_full_range, AxesRange, FullTraversal, MaxKernel, MinKernel,
18 ProductKernel, ReduceKernel, SumKernel, SumSquaresKernel,
19};
20
21trait ReduceWriter<T> {
22 fn offset(&self) -> isize;
23 unsafe fn ptr(&mut self) -> *mut T;
26 fn extent(&self) -> usize;
27 unsafe fn write_at(&mut self, offset: isize, value: T) {
30 debug_assert!(offset >= 0 && (offset as usize) < self.extent());
31 unsafe { self.ptr().offset(offset).write(value) }
33 }
34}
35
36struct RawReduceWriter<'a, T> {
37 ptr: *mut T,
38 extent: usize,
39 offset: isize,
40 _marker: core::marker::PhantomData<&'a mut [MaybeUninit<T>]>,
41}
42
43impl<'a, T> ReduceWriter<T> for RawReduceWriter<'a, T> {
44 fn offset(&self) -> isize {
45 self.offset
46 }
47 unsafe fn ptr(&mut self) -> *mut T {
48 self.ptr
49 }
50 fn extent(&self) -> usize {
51 self.extent
52 }
53}
54
55#[derive(Clone, Debug)]
57pub struct ErasedCopyPlan {
58 dtype: KernelDType,
59 plan: CopyPlan,
60}
61
62#[derive(Clone, Debug)]
64pub struct ErasedConcatenatePlan {
65 dtype: KernelDType,
66 plan: ConcatenatePlan,
67}
68
69impl ErasedCopyPlan {
70 pub fn compile(
72 dtype: KernelDType,
73 dims: &[usize],
74 dst_strides: &[isize],
75 src_strides: &[isize],
76 ) -> Result<Self> {
77 Ok(Self {
78 dtype,
79 plan: CopyPlan::compile(dims, dst_strides, src_strides)?,
80 })
81 }
82
83 #[inline]
84 pub fn dtype(&self) -> KernelDType {
85 self.dtype
86 }
87
88 pub fn execute(
90 &self,
91 ctx: &ExecContext,
92 dest: &mut ErasedRawStridedMut<'_>,
93 src: &ErasedRawStridedRef<'_>,
94 ) -> Result<()> {
95 self.check_dtype(dest.dtype())?;
96 self.check_dtype(src.dtype())?;
97
98 let result = ctx.run(|| match self.dtype {
99 KernelDType::F32 => execute_copy::<f32>(&self.plan, dest, src),
100 KernelDType::F64 => execute_copy::<f64>(&self.plan, dest, src),
101 KernelDType::I32 => execute_copy::<i32>(&self.plan, dest, src),
102 KernelDType::I64 => execute_copy::<i64>(&self.plan, dest, src),
103 KernelDType::Bool => execute_copy::<bool>(&self.plan, dest, src),
104 KernelDType::C32 => execute_copy::<Complex32>(&self.plan, dest, src),
105 KernelDType::C64 => execute_copy::<Complex64>(&self.plan, dest, src),
106 _ => Err(StridedError::UnsupportedDType {
107 dtype: self.dtype.label(),
108 }),
109 });
110 result
111 }
112
113 fn check_dtype(&self, actual: KernelDType) -> Result<()> {
114 if actual != self.dtype {
115 return Err(StridedError::DTypeMismatch {
116 expected: self.dtype.label(),
117 actual: actual.label(),
118 });
119 }
120 Ok(())
121 }
122}
123
124impl ErasedConcatenatePlan {
125 pub fn compile(
127 dtype: KernelDType,
128 input_dims: &[&[usize]],
129 input_strides: &[&[isize]],
130 dest_dims: &[usize],
131 dest_strides: &[isize],
132 axis: usize,
133 ) -> Result<Self> {
134 check_static_indexing_dtype(dtype)?;
135 Ok(Self {
136 dtype,
137 plan: ConcatenatePlan::compile(
138 input_dims,
139 input_strides,
140 dest_dims,
141 dest_strides,
142 axis,
143 )?,
144 })
145 }
146
147 #[inline]
148 pub fn dtype(&self) -> KernelDType {
149 self.dtype
150 }
151
152 #[inline]
153 pub fn plan(&self) -> &ConcatenatePlan {
154 &self.plan
155 }
156
157 pub fn execute(
159 &self,
160 ctx: &ExecContext,
161 dest: &mut ErasedRawStridedMut<'_>,
162 inputs: &[ErasedRawStridedRef<'_>],
163 ) -> Result<()> {
164 check_dtype(self.dtype, dest.dtype())?;
165 for input in inputs {
166 check_dtype(self.dtype, input.dtype())?;
167 }
168
169 let result = ctx.run(|| match self.dtype {
170 KernelDType::F32 => execute_concatenate::<f32>(&self.plan, dest, inputs),
171 KernelDType::F64 => execute_concatenate::<f64>(&self.plan, dest, inputs),
172 KernelDType::I32 => execute_concatenate::<i32>(&self.plan, dest, inputs),
173 KernelDType::I64 => execute_concatenate::<i64>(&self.plan, dest, inputs),
174 KernelDType::Bool => execute_concatenate::<bool>(&self.plan, dest, inputs),
175 KernelDType::C32 => execute_concatenate::<Complex32>(&self.plan, dest, inputs),
176 KernelDType::C64 => execute_concatenate::<Complex64>(&self.plan, dest, inputs),
177 _ => Err(StridedError::UnsupportedDType {
178 dtype: self.dtype.label(),
179 }),
180 });
181 result
182 }
183
184 pub fn execute_uninit(
192 &self,
193 ctx: &ExecContext,
194 dest: &mut ErasedRawStridedUninitMut<'_>,
195 inputs: &[ErasedRawStridedPtr<'_>],
196 ) -> Result<()> {
197 check_dtype(self.dtype, dest.dtype())?;
198 if inputs.len() != self.plan.input_count() {
199 return Err(StridedError::RankMismatch(
200 inputs.len(),
201 self.plan.input_count(),
202 ));
203 }
204 for input in inputs {
205 check_dtype(self.dtype, input.dtype())?;
206 }
207 for (position, input) in inputs.iter().enumerate() {
208 validate_uninit_no_overlap(dest, input, position)?;
209 }
210 for input in inputs {
211 unsafe { input.try_as_ref_after_no_overlap() }?;
213 }
214
215 ctx.run(|| match self.dtype {
216 KernelDType::F32 => execute_concatenate_uninit::<f32>(&self.plan, dest, inputs),
217 KernelDType::F64 => execute_concatenate_uninit::<f64>(&self.plan, dest, inputs),
218 KernelDType::I32 => execute_concatenate_uninit::<i32>(&self.plan, dest, inputs),
219 KernelDType::I64 => execute_concatenate_uninit::<i64>(&self.plan, dest, inputs),
220 KernelDType::Bool => execute_concatenate_uninit::<bool>(&self.plan, dest, inputs),
221 KernelDType::C32 => execute_concatenate_uninit::<Complex32>(&self.plan, dest, inputs),
222 KernelDType::C64 => execute_concatenate_uninit::<Complex64>(&self.plan, dest, inputs),
223 _ => Err(StridedError::UnsupportedDType {
224 dtype: self.dtype.label(),
225 }),
226 })
227 }
228}
229
230#[non_exhaustive]
232#[derive(Clone, Copy, Debug, Eq, PartialEq)]
233pub enum ReduceOp {
234 Sum,
235 Product,
236 SumSquares,
241 Max,
250 Min,
259}
260
261#[derive(Clone, Debug)]
303pub struct ErasedReducePlan {
304 dtype: KernelDType,
305 op: ReduceOp,
306 layout: ReduceLayout,
307}
308
309#[derive(Clone, Debug)]
310enum ReduceLayout {
311 Full {
312 dims: Vec<usize>,
313 src_strides: Vec<isize>,
314 traversal: FullTraversal,
315 },
316 Axes {
317 src_dims: Vec<usize>,
318 src_strides: Vec<isize>,
319 dest_dims: Vec<usize>,
320 dest_strides: Vec<isize>,
321 outer_axes: Vec<ReduceOuterAxis>,
322 inner_axes: Vec<ReduceInnerAxis>,
323 dest_total: usize,
324 reduce_total: usize,
325 full: Option<FullTraversal>,
328 },
329}
330
331#[derive(Clone, Copy, Debug)]
332struct ReduceOuterAxis {
333 extent: usize,
334 source_step: isize,
335 source_reset: isize,
336 dest_step: isize,
337 dest_reset: isize,
338}
339#[derive(Clone, Copy, Debug)]
340struct ReduceInnerAxis {
341 extent: usize,
342 source_step: isize,
343 source_reset: isize,
344}
345impl ReduceLayout {
346 fn src_dims(&self) -> &[usize] {
347 match self {
348 Self::Full { dims, .. } => dims,
349 Self::Axes { src_dims, .. } => src_dims,
350 }
351 }
352
353 fn src_strides(&self) -> &[isize] {
354 match self {
355 Self::Full { src_strides, .. } | Self::Axes { src_strides, .. } => src_strides,
356 }
357 }
358
359 fn check_src_layout(&self, src: &ErasedRawStridedRef<'_>) -> Result<()> {
360 if src.dims() != self.src_dims() || src.strides() != self.src_strides() {
361 return Err(StridedError::PlanLayoutMismatch);
362 }
363 Ok(())
364 }
365}
366
367#[derive(Clone, Copy, Debug)]
368struct AxesLayout<'a> {
369 outer_axes: &'a [ReduceOuterAxis],
370 inner_axes: &'a [ReduceInnerAxis],
371 dest_total: usize,
372 reduce_total: usize,
373}
374
375impl ErasedReducePlan {
376 pub fn compile(
378 dtype: KernelDType,
379 op: ReduceOp,
380 dims: &[usize],
381 src_strides: &[isize],
382 ) -> Result<Self> {
383 check_reduce_op_dtype(dtype, op)?;
384 if dims.len() != src_strides.len() {
385 return Err(StridedError::StrideLengthMismatch);
386 }
387 checked_total_len(dims)?;
388 let traversal = FullTraversal::compile(dims, src_strides)?;
389 Ok(Self {
390 dtype,
391 op,
392 layout: ReduceLayout::Full {
393 dims: dims.to_vec(),
394 src_strides: src_strides.to_vec(),
395 traversal,
396 },
397 })
398 }
399
400 #[allow(clippy::too_many_arguments)]
406 pub fn compile_axes(
407 dtype: KernelDType,
408 op: ReduceOp,
409 src_dims: &[usize],
410 src_strides: &[isize],
411 dest_dims: &[usize],
412 dest_strides: &[isize],
413 axes: &[usize],
414 ) -> Result<Self> {
415 check_reduce_op_dtype(dtype, op)?;
416 if src_dims.len() != src_strides.len() || dest_dims.len() != dest_strides.len() {
417 return Err(StridedError::StrideLengthMismatch);
418 }
419 checked_total_len(src_dims)?;
420 check_reduce_layout_offset_arithmetic(src_dims, src_strides)?;
421 let dest_total = checked_total_len(dest_dims)?;
422 check_reduce_layout_offset_arithmetic(dest_dims, dest_strides)?;
423 if !crate::layout_check::is_injective_layout(dest_dims, dest_strides) {
424 return Err(StridedError::NonInjectiveOutputLayout);
425 }
426 validate_unique_axes(axes, src_dims.len())?;
427
428 let kept_axes: Vec<usize> = (0..src_dims.len())
429 .filter(|axis| !axes.contains(axis))
430 .collect();
431 let expected_dest_dims: Vec<usize> = kept_axes.iter().map(|&axis| src_dims[axis]).collect();
432 if expected_dest_dims.is_empty() {
433 if dest_total != 1 {
434 return Err(StridedError::ShapeMismatch(
435 dest_dims.to_vec(),
436 expected_dest_dims,
437 ));
438 }
439 } else if dest_dims != expected_dest_dims.as_slice() {
440 return Err(StridedError::ShapeMismatch(
441 dest_dims.to_vec(),
442 expected_dest_dims,
443 ));
444 }
445
446 let reduce_total = axes
447 .iter()
448 .try_fold(1usize, |total, &axis| total.checked_mul(src_dims[axis]))
449 .ok_or(StridedError::OffsetOverflow)?;
450 let outer_axes = compress_reduce_outer_axes(
451 kept_axes
452 .iter()
453 .enumerate()
454 .map(|(dest_axis, &src_axis)| {
455 let extent = src_dims[src_axis];
456 Ok(ReduceOuterAxis {
457 extent,
458 source_step: src_strides[src_axis],
459 source_reset: checked_reduce_reset(extent, src_strides[src_axis])?,
460 dest_step: dest_strides[dest_axis],
461 dest_reset: checked_reduce_reset(extent, dest_strides[dest_axis])?,
462 })
463 })
464 .collect::<Result<Vec<_>>>()?,
465 )?;
466 let inner_axes = compress_reduce_inner_axes(
467 axes.iter()
468 .map(|&src_axis| {
469 let extent = src_dims[src_axis];
470 Ok(ReduceInnerAxis {
471 extent,
472 source_step: src_strides[src_axis],
473 source_reset: checked_reduce_reset(extent, src_strides[src_axis])?,
474 })
475 })
476 .collect::<Result<Vec<_>>>()?,
477 )?;
478 let full = if kept_axes.is_empty() {
479 Some(FullTraversal::compile(src_dims, src_strides)?)
480 } else {
481 None
482 };
483 Ok(Self {
484 dtype,
485 op,
486 layout: ReduceLayout::Axes {
487 src_dims: src_dims.to_vec(),
488 src_strides: src_strides.to_vec(),
489 dest_dims: dest_dims.to_vec(),
490 dest_strides: dest_strides.to_vec(),
491 outer_axes,
492 inner_axes,
493 dest_total,
494 reduce_total,
495 full,
496 },
497 })
498 }
499
500 #[inline]
501 pub fn dtype(&self) -> KernelDType {
502 self.dtype
503 }
504
505 #[inline]
506 pub fn op(&self) -> ReduceOp {
507 self.op
508 }
509
510 pub fn execute(
512 &self,
513 ctx: &ExecContext,
514 dest: &mut ErasedRawStridedMut<'_>,
515 src: &ErasedRawStridedRef<'_>,
516 ) -> Result<()> {
517 check_dtype(self.dtype, dest.dtype())?;
518 check_dtype(self.dtype, src.dtype())?;
519 self.layout.check_src_layout(src)?;
520 match &self.layout {
521 ReduceLayout::Full { .. } => {
522 let dest_len = checked_total_len(dest.dims())?;
523 if dest_len != 1 {
524 return Err(StridedError::RankMismatch(dest_len, 1));
525 }
526 }
527 ReduceLayout::Axes {
528 dest_dims,
529 dest_strides,
530 ..
531 } => {
532 if dest.dims() != dest_dims.as_slice() || dest.strides() != dest_strides.as_slice()
533 {
534 return Err(StridedError::PlanLayoutMismatch);
535 }
536 }
537 }
538
539 let result = match self.dtype {
540 KernelDType::F32 => {
541 let mut writer = reduce_writer::<f32>(dest)?;
542 dispatch_reduce::<f32, _>(self.op, &self.layout, ctx, &mut writer, src)
543 }
544 KernelDType::F64 => {
545 let mut writer = reduce_writer::<f64>(dest)?;
546 dispatch_reduce::<f64, _>(self.op, &self.layout, ctx, &mut writer, src)
547 }
548 KernelDType::I32 => {
549 let mut writer = reduce_writer::<i32>(dest)?;
550 dispatch_reduce::<i32, _>(self.op, &self.layout, ctx, &mut writer, src)
551 }
552 KernelDType::I64 => {
553 let mut writer = reduce_writer::<i64>(dest)?;
554 dispatch_reduce::<i64, _>(self.op, &self.layout, ctx, &mut writer, src)
555 }
556 KernelDType::C32 => {
557 let mut writer = reduce_writer::<Complex32>(dest)?;
558 dispatch_reduce::<Complex32, _>(self.op, &self.layout, ctx, &mut writer, src)
559 }
560 KernelDType::C64 => {
561 let mut writer = reduce_writer::<Complex64>(dest)?;
562 dispatch_reduce::<Complex64, _>(self.op, &self.layout, ctx, &mut writer, src)
563 }
564 _ => Err(StridedError::UnsupportedDType {
565 dtype: self.dtype.label(),
566 }),
567 };
568 result
569 }
570
571 pub fn execute_uninit(
578 &self,
579 ctx: &ExecContext,
580 dest: &mut ErasedRawStridedUninitMut<'_>,
581 src: &ErasedRawStridedPtr<'_>,
582 ) -> Result<()> {
583 check_dtype(self.dtype, dest.dtype())?;
584 check_dtype(self.dtype, src.dtype())?;
585 validate_uninit_no_overlap(dest, src, 0)?;
586 let src = unsafe { src.try_as_ref_after_no_overlap() }?;
588 self.layout.check_src_layout(&src)?;
589 match &self.layout {
590 ReduceLayout::Full { .. } => {
591 let total = checked_total_len(dest.dims())?;
592 if total != 1 {
593 return Err(StridedError::RankMismatch(total, 1));
594 }
595 }
596 ReduceLayout::Axes {
597 dest_dims,
598 dest_strides,
599 ..
600 } => {
601 if dest.dims() != dest_dims.as_slice() || dest.strides() != dest_strides.as_slice()
602 {
603 return Err(StridedError::PlanLayoutMismatch);
604 }
605 }
606 }
607 macro_rules! run {
608 ($ty:ty) => {{
609 let mut writer = reduce_uninit_writer::<$ty>(dest)?;
610 dispatch_reduce::<$ty, _>(self.op, &self.layout, ctx, &mut writer, &src)
611 }};
612 }
613 match self.dtype {
614 KernelDType::F32 => run!(f32),
615 KernelDType::F64 => run!(f64),
616 KernelDType::I32 => run!(i32),
617 KernelDType::I64 => run!(i64),
618 KernelDType::C32 => run!(Complex32),
619 KernelDType::C64 => run!(Complex64),
620 _ => Err(StridedError::UnsupportedDType {
621 dtype: self.dtype.label(),
622 }),
623 }
624 }
625}
626
627fn reduce_writer<'a, T>(dest: &'a mut ErasedRawStridedMut<'_>) -> Result<RawReduceWriter<'a, T>>
628where
629 T: KernelStorageElement,
630{
631 let offset = dest.offset();
632 let data = dest.data_as_mut::<T>()?;
633 let ptr = data.as_mut_ptr();
634 let extent = data.len();
635 Ok(RawReduceWriter {
636 ptr,
637 extent,
638 offset,
639 _marker: core::marker::PhantomData,
640 })
641}
642
643fn reduce_uninit_writer<'a, T>(
644 dest: &'a mut ErasedRawStridedUninitMut<'_>,
645) -> Result<RawReduceWriter<'a, T>>
646where
647 T: KernelStorageElement,
648{
649 let offset = dest.offset();
650 let data = dest.data_as_uninit_mut::<T>()?;
651 let ptr = data.as_mut_ptr().cast::<T>();
652 let extent = data.len();
653 Ok(RawReduceWriter {
654 ptr,
655 extent,
656 offset,
657 _marker: core::marker::PhantomData,
658 })
659}
660
661fn check_reduce_dtype(dtype: KernelDType) -> Result<()> {
662 match dtype {
663 KernelDType::F32
664 | KernelDType::F64
665 | KernelDType::I32
666 | KernelDType::I64
667 | KernelDType::C32
668 | KernelDType::C64 => Ok(()),
669 _ => Err(StridedError::UnsupportedDType {
670 dtype: dtype.label(),
671 }),
672 }
673}
674
675fn check_reduce_op_dtype(dtype: KernelDType, op: ReduceOp) -> Result<()> {
676 if op == ReduceOp::SumSquares && !matches!(dtype, KernelDType::F32 | KernelDType::F64) {
677 return Err(StridedError::UnsupportedDType {
678 dtype: dtype.label(),
679 });
680 }
681 if matches!(op, ReduceOp::Max | ReduceOp::Min)
682 && !matches!(
683 dtype,
684 KernelDType::F32 | KernelDType::F64 | KernelDType::I32 | KernelDType::I64
685 )
686 {
687 return Err(StridedError::UnsupportedDType {
688 dtype: dtype.label(),
689 });
690 }
691 check_reduce_dtype(dtype)
692}
693
694fn checked_total_len(dims: &[usize]) -> Result<usize> {
695 if dims.is_empty() {
696 return Ok(1);
697 }
698 dims.iter()
699 .try_fold(1usize, |acc, &dim| acc.checked_mul(dim))
700 .ok_or(StridedError::OffsetOverflow)
701}
702
703fn execute_copy<T>(
704 plan: &CopyPlan,
705 dest: &mut ErasedRawStridedMut<'_>,
706 src: &ErasedRawStridedRef<'_>,
707) -> Result<()>
708where
709 T: Copy + crate::MaybeSendSync + KernelStorageElement,
710{
711 let source_data = src.data_as::<T>()?;
712 let dest_dims = dest.dims();
713 let dest_strides = dest.strides();
714 let dest_offset = dest.offset();
715 let dest_data = dest.data_as_mut::<T>()?;
716 let source = unsafe {
717 RawStridedRef::new_unchecked(source_data, src.dims(), src.strides(), src.offset())
718 };
719 let mut dest =
720 unsafe { RawStridedMut::new_unchecked(dest_data, dest_dims, dest_strides, dest_offset) };
721 plan.execute(&mut dest, &source)
722}
723
724fn execute_concatenate<T>(
725 plan: &ConcatenatePlan,
726 dest: &mut ErasedRawStridedMut<'_>,
727 inputs: &[ErasedRawStridedRef<'_>],
728) -> Result<()>
729where
730 T: Copy + crate::MaybeSendSync + KernelStorageElement,
731{
732 if inputs.len() != plan.input_count() {
733 return Err(StridedError::RankMismatch(inputs.len(), plan.input_count()));
734 }
735 let dest_dims = dest.dims();
736 let dest_strides = dest.strides();
737 let dest_offset = dest.offset();
738 let dest_data = dest.data_as_mut::<T>()?;
739 let mut dest_ref =
740 unsafe { RawStridedMut::new_unchecked(dest_data, dest_dims, dest_strides, dest_offset) };
741 plan.check_dest_layout(&dest_ref)?;
742
743 for (position, input) in inputs.iter().enumerate() {
744 let input_data = input.data_as::<T>()?;
745 let input_ref = unsafe {
746 RawStridedRef::new_unchecked(input_data, input.dims(), input.strides(), input.offset())
747 };
748 plan.check_input_layout(position, &input_ref)?;
749 plan.segment_offset(position, dest_offset)?;
750 }
751 if plan.prefers_whole_plan() {
752 let input_refs = inputs
753 .iter()
754 .map(|input| {
755 let input_data = input.data_as::<T>()?;
756 Ok(unsafe {
757 RawStridedRef::new_unchecked(
758 input_data,
759 input.dims(),
760 input.strides(),
761 input.offset(),
762 )
763 })
764 })
765 .collect::<Result<Vec<_>>>()?;
766 return plan.execute(&mut dest_ref, &input_refs);
768 }
769 for (position, input) in inputs.iter().enumerate() {
770 let input_data = input.data_as::<T>()?;
771 let input_ref = unsafe {
772 RawStridedRef::new_unchecked(input_data, input.dims(), input.strides(), input.offset())
773 };
774 plan.execute_segment(position, &mut dest_ref, &input_ref)?;
775 }
776 Ok(())
777}
778
779fn execute_concatenate_uninit<T>(
780 plan: &ConcatenatePlan,
781 dest: &mut ErasedRawStridedUninitMut<'_>,
782 inputs: &[ErasedRawStridedPtr<'_>],
783) -> Result<()>
784where
785 T: Copy + crate::MaybeSendSync + KernelStorageElement,
786{
787 let dest_dims = dest.dims();
788 let dest_strides = dest.strides();
789 let dest_offset = dest.offset();
790 let dest_data = dest.data_as_uninit_mut::<T>()?;
791 let mut dest_ref =
792 unsafe { RawStridedMut::new_unchecked(dest_data, dest_dims, dest_strides, dest_offset) };
793 plan.check_dest_layout(&dest_ref)?;
794
795 for (position, input) in inputs.iter().enumerate() {
796 let input = unsafe { input.try_as_ref_after_no_overlap() }?;
798 let input_data = input.data_as::<T>()?;
799 let input_ref = unsafe {
800 RawStridedRef::new_unchecked(input_data, input.dims(), input.strides(), input.offset())
801 };
802 plan.check_input_layout(position, &input_ref)?;
803 plan.segment_offset(position, dest_offset)?;
804 }
805 if plan.prefers_whole_plan() {
806 let inputs = inputs
807 .iter()
808 .map(|input| unsafe { input.try_as_ref_after_no_overlap() })
810 .collect::<Result<Vec<_>>>()?;
811 let input_refs = inputs
812 .iter()
813 .map(|input| {
814 let input_data = input.data_as::<T>()?;
815 Ok(unsafe {
816 RawStridedRef::new_unchecked(
817 input_data,
818 input.dims(),
819 input.strides(),
820 input.offset(),
821 )
822 })
823 })
824 .collect::<Result<Vec<_>>>()?;
825 return plan.execute_uninit(&mut dest_ref, &input_refs);
827 }
828 for (position, input) in inputs.iter().enumerate() {
829 let input = unsafe { input.try_as_ref_after_no_overlap() }?;
831 let input_data = input.data_as::<T>()?;
832 let input_ref = unsafe {
833 RawStridedRef::new_unchecked(input_data, input.dims(), input.strides(), input.offset())
834 };
835 plan.execute_segment_uninit(position, &mut dest_ref, &input_ref)?;
836 }
837 Ok(())
838}
839
840fn reduce_context_is_serial(ctx: &ExecContext) -> bool {
843 ctx.is_serial()
844 || ctx
845 .max_threads_limit()
846 .is_some_and(|max_threads| max_threads.get() == 1)
847}
848
849fn execute_reduce_full<T, W, K>(
850 ctx: &ExecContext,
851 dest: &mut W,
852 src: &ErasedRawStridedRef<'_>,
853 traversal: &FullTraversal,
854) -> Result<()>
855where
856 T: ErasedReduceScalar,
857 W: ReduceWriter<T>,
858 K: ReduceKernel<T>,
859{
860 let source = src.data_as::<T>()?;
861 let total = traversal.total();
862 let value = if total == 0 {
863 K::identity()
864 } else if reduce_context_is_serial(ctx) {
865 unsafe { reduce_full_range::<T, K>(source.as_ptr(), src.offset(), traversal, 0..total)? }
871 } else {
872 ctx.run(|| reduce_full_policy::<T, K>(source, src.offset(), traversal))?
873 };
874
875 unsafe { dest.write_at(dest.offset(), value) };
877 Ok(())
878}
879
880fn reduce_full_policy<T, K>(
881 source: &[T],
882 source_base: isize,
883 traversal: &FullTraversal,
884) -> Result<T>
885where
886 T: ErasedReduceScalar,
887 K: ReduceKernel<T>,
888{
889 let total = traversal.total();
890 #[cfg(feature = "parallel")]
891 {
892 let nthreads = crate::threading::parallel_threads_for_len(total);
893 if nthreads > 1 {
894 let source_ptr = crate::threading::SendPtr(source.as_ptr() as *mut T);
895 return crate::threading::parallel_map_reduce(
896 0..total,
897 nthreads,
898 &|range| {
899 unsafe {
902 reduce_full_range::<T, K>(
903 source_ptr.as_const(),
904 source_base,
905 traversal,
906 range,
907 )
908 }
909 },
910 &|left, right| Ok(K::combine(left?, right?)),
911 );
912 }
913 }
914 unsafe { reduce_full_range::<T, K>(source.as_ptr(), source_base, traversal, 0..total) }
917}
918
919fn dispatch_reduce<T, W>(
920 op: ReduceOp,
921 layout: &ReduceLayout,
922 ctx: &ExecContext,
923 dest: &mut W,
924 src: &ErasedRawStridedRef<'_>,
925) -> Result<()>
926where
927 T: ErasedReduceScalar,
928 W: ReduceWriter<T>,
929{
930 match op {
933 ReduceOp::Sum => dispatch_reduce_kernel::<T, W, SumKernel>(layout, ctx, dest, src),
934 ReduceOp::Product => dispatch_reduce_kernel::<T, W, ProductKernel>(layout, ctx, dest, src),
935 ReduceOp::SumSquares => {
936 dispatch_reduce_kernel::<T, W, SumSquaresKernel>(layout, ctx, dest, src)
937 }
938 ReduceOp::Max => dispatch_reduce_kernel::<T, W, MaxKernel>(layout, ctx, dest, src),
939 ReduceOp::Min => dispatch_reduce_kernel::<T, W, MinKernel>(layout, ctx, dest, src),
940 }
941}
942
943fn dispatch_reduce_kernel<T, W, K>(
944 layout: &ReduceLayout,
945 ctx: &ExecContext,
946 dest: &mut W,
947 src: &ErasedRawStridedRef<'_>,
948) -> Result<()>
949where
950 T: ErasedReduceScalar,
951 W: ReduceWriter<T>,
952 K: ReduceKernel<T>,
953{
954 match layout {
955 ReduceLayout::Full { traversal, .. } => {
956 execute_reduce_full::<T, W, K>(ctx, dest, src, traversal)
957 }
958 ReduceLayout::Axes {
959 outer_axes,
960 inner_axes,
961 dest_total,
962 reduce_total,
963 full,
964 ..
965 } => {
966 if let Some(traversal) = full {
967 return execute_reduce_full::<T, W, K>(ctx, dest, src, traversal);
968 }
969 execute_reduce_axes::<T, W, K>(
970 ctx,
971 dest,
972 src,
973 AxesLayout {
974 outer_axes,
975 inner_axes,
976 dest_total: *dest_total,
977 reduce_total: *reduce_total,
978 },
979 )
980 }
981 }
982}
983
984fn execute_reduce_axes<T, W, K>(
985 ctx: &ExecContext,
986 dest: &mut W,
987 src: &ErasedRawStridedRef<'_>,
988 layout: AxesLayout<'_>,
989) -> Result<()>
990where
991 T: ErasedReduceScalar,
992 W: ReduceWriter<T>,
993 K: ReduceKernel<T>,
994{
995 if layout.dest_total == 0 {
996 return Ok(());
997 }
998
999 if layout.reduce_total == 0 {
1000 if ctx.is_serial() {
1001 execute_reduce_axes_identity_serial::<T, W, K>(dest, layout)
1002 } else {
1003 ctx.run(|| execute_reduce_axes_identity_policy::<T, W, K>(dest, layout))
1004 }
1005 } else if reduce_context_is_serial(ctx) {
1006 execute_reduce_axes_policy::<T, W, K>(dest, src, layout, false)
1007 } else {
1008 ctx.run(|| execute_reduce_axes_policy::<T, W, K>(dest, src, layout, true))
1009 }
1010}
1011
1012fn execute_reduce_axes_policy<T, W, K>(
1013 dest: &mut W,
1014 src: &ErasedRawStridedRef<'_>,
1015 layout: AxesLayout<'_>,
1016 allow_parallel: bool,
1017) -> Result<()>
1018where
1019 T: ErasedReduceScalar,
1020 W: ReduceWriter<T>,
1021 K: ReduceKernel<T>,
1022{
1023 let source_data = src.data_as::<T>()?;
1024 let source_base = src.offset();
1025 let dest_base = dest.offset();
1026 let dest_ptr = unsafe { dest.ptr() };
1028 #[cfg(feature = "parallel")]
1029 if allow_parallel {
1030 let work = layout.dest_total.saturating_mul(layout.reduce_total);
1036 let nthreads = crate::threading::parallel_threads_for_len(work).min(layout.dest_total);
1037 if nthreads > 1 {
1038 let dest_ptr = crate::threading::SendPtr(dest_ptr);
1039 let source_ptr = crate::threading::SendPtr(source_data.as_ptr() as *mut T);
1040 return crate::threading::parallel_map_reduce(
1041 0..layout.dest_total,
1042 nthreads,
1043 &|range| {
1044 let parts = AxesRange {
1045 source: source_ptr.as_const(),
1046 source_base,
1047 dest: dest_ptr.as_ptr(),
1048 dest_base,
1049 outer_axes: layout.outer_axes,
1050 inner_axes: layout.inner_axes,
1051 reduce_total: layout.reduce_total,
1052 };
1053 unsafe { reduce_axes_range::<T, K>(parts, range) }
1055 },
1056 &|left, right| left.and(right),
1057 );
1058 }
1059 }
1060 #[cfg(not(feature = "parallel"))]
1061 let _ = allow_parallel;
1062
1063 let parts = AxesRange {
1064 source: source_data.as_ptr(),
1065 source_base,
1066 dest: dest_ptr,
1067 dest_base,
1068 outer_axes: layout.outer_axes,
1069 inner_axes: layout.inner_axes,
1070 reduce_total: layout.reduce_total,
1071 };
1072 unsafe { reduce_axes_range::<T, K>(parts, 0..layout.dest_total) }
1079}
1080
1081fn execute_reduce_axes_identity_serial<T, W, K>(dest: &mut W, layout: AxesLayout<'_>) -> Result<()>
1082where
1083 T: ErasedReduceScalar,
1084 W: ReduceWriter<T>,
1085 K: ReduceKernel<T>,
1086{
1087 let mut outer = ReduceOuterCursor::decode(0, 0, dest.offset(), layout.outer_axes)?;
1088 for output in 0..layout.dest_total {
1089 unsafe { dest.write_at(outer.dest_offset, K::identity()) };
1097 if output + 1 < layout.dest_total {
1098 outer.advance();
1099 }
1100 }
1101 Ok(())
1102}
1103fn execute_reduce_axes_identity_policy<T, W, K>(dest: &mut W, layout: AxesLayout<'_>) -> Result<()>
1104where
1105 T: ErasedReduceScalar,
1106 W: ReduceWriter<T>,
1107 K: ReduceKernel<T>,
1108{
1109 #[cfg(feature = "parallel")]
1110 {
1111 let nthreads = crate::threading::parallel_threads_for_len(layout.dest_total);
1112 if nthreads > 1 {
1113 return execute_reduce_axes_identity_parallel::<T, W, K>(dest, layout, nthreads);
1114 }
1115 }
1116 execute_reduce_axes_identity_serial::<T, W, K>(dest, layout)
1117}
1118#[cfg(feature = "parallel")]
1119fn execute_reduce_axes_identity_parallel<T, W, K>(
1120 dest: &mut W,
1121 layout: AxesLayout<'_>,
1122 nthreads: usize,
1123) -> Result<()>
1124where
1125 T: ErasedReduceScalar,
1126 W: ReduceWriter<T>,
1127 K: ReduceKernel<T>,
1128{
1129 let dest_ptr = crate::threading::SendPtr(unsafe { dest.ptr() });
1131 let dest_offset_base = dest.offset();
1132 crate::threading::parallel_map_reduce(
1133 0..layout.dest_total,
1134 nthreads,
1135 &|range| {
1136 let range_end = range.end;
1137 let mut outer =
1138 ReduceOuterCursor::decode(range.start, 0, dest_offset_base, layout.outer_axes)?;
1139 let dest_ptr = dest_ptr.as_ptr();
1140 for output in range {
1141 unsafe { dest_ptr.offset(outer.dest_offset).write(K::identity()) };
1149 if output + 1 < range_end {
1150 outer.advance();
1151 }
1152 }
1153 Ok(())
1154 },
1155 &|left, right| left.and(right),
1156 )
1157}
1158
1159trait ErasedReduceScalar:
1160 KernelStorageElement
1161 + Copy
1162 + One
1163 + Zero
1164 + crate::MaybeSendSync
1165 + crate::simd::MaybeSimdOps
1166 + crate::simd::MaybeSimdProduct
1167 + crate::simd::MaybeSimdSumSquares
1168{
1169 fn reduce_sum(lhs: Self, rhs: Self) -> Self;
1170 fn reduce_product(lhs: Self, rhs: Self) -> Self;
1171 fn max_identity() -> Self;
1173 fn min_identity() -> Self;
1175 fn reduce_max(lhs: Self, rhs: Self) -> Self;
1176 fn reduce_min(lhs: Self, rhs: Self) -> Self;
1177}
1178
1179macro_rules! impl_float_erased_reduce_scalar {
1180 ($($ty:ty),* $(,)?) => {
1181 $(
1182 impl ErasedReduceScalar for $ty {
1183 #[inline(always)]
1184 fn reduce_sum(lhs: Self, rhs: Self) -> Self {
1185 lhs + rhs
1186 }
1187
1188 #[inline(always)]
1189 fn reduce_product(lhs: Self, rhs: Self) -> Self {
1190 lhs * rhs
1191 }
1192
1193 #[inline(always)]
1194 fn max_identity() -> Self {
1195 <$ty>::NEG_INFINITY
1196 }
1197
1198 #[inline(always)]
1199 fn min_identity() -> Self {
1200 <$ty>::INFINITY
1201 }
1202
1203 #[inline(always)]
1204 fn reduce_max(lhs: Self, rhs: Self) -> Self {
1205 if (lhs > rhs) | lhs.is_nan() {
1213 lhs
1214 } else {
1215 rhs
1216 }
1217 }
1218
1219 #[inline(always)]
1220 fn reduce_min(lhs: Self, rhs: Self) -> Self {
1221 if (lhs < rhs) | lhs.is_nan() {
1229 lhs
1230 } else {
1231 rhs
1232 }
1233 }
1234 }
1235 )*
1236 };
1237}
1238
1239macro_rules! impl_complex_erased_reduce_scalar {
1240 ($($ty:ty),* $(,)?) => {
1241 $(
1242 impl ErasedReduceScalar for $ty {
1243 #[inline(always)]
1244 fn reduce_sum(lhs: Self, rhs: Self) -> Self {
1245 lhs + rhs
1246 }
1247
1248 #[inline(always)]
1249 fn reduce_product(lhs: Self, rhs: Self) -> Self {
1250 lhs * rhs
1251 }
1252
1253 fn max_identity() -> Self {
1254 unreachable!("complex max reduction is rejected at plan compile")
1256 }
1257
1258 fn min_identity() -> Self {
1259 unreachable!("complex min reduction is rejected at plan compile")
1261 }
1262
1263 fn reduce_max(_lhs: Self, _rhs: Self) -> Self {
1264 unreachable!("complex max reduction is rejected at plan compile")
1266 }
1267
1268 fn reduce_min(_lhs: Self, _rhs: Self) -> Self {
1269 unreachable!("complex min reduction is rejected at plan compile")
1271 }
1272 }
1273 )*
1274 };
1275}
1276
1277macro_rules! impl_wrapping_erased_reduce_scalar {
1278 ($($ty:ty),* $(,)?) => {
1279 $(
1280 impl ErasedReduceScalar for $ty {
1281 #[inline(always)]
1282 fn reduce_sum(lhs: Self, rhs: Self) -> Self {
1283 lhs.wrapping_add(rhs)
1284 }
1285
1286 #[inline(always)]
1287 fn reduce_product(lhs: Self, rhs: Self) -> Self {
1288 lhs.wrapping_mul(rhs)
1289 }
1290
1291 #[inline(always)]
1292 fn max_identity() -> Self {
1293 <$ty>::MIN
1294 }
1295
1296 #[inline(always)]
1297 fn min_identity() -> Self {
1298 <$ty>::MAX
1299 }
1300
1301 #[inline(always)]
1302 fn reduce_max(lhs: Self, rhs: Self) -> Self {
1303 lhs.max(rhs)
1304 }
1305
1306 #[inline(always)]
1307 fn reduce_min(lhs: Self, rhs: Self) -> Self {
1308 lhs.min(rhs)
1309 }
1310 }
1311 )*
1312 };
1313}
1314
1315impl_float_erased_reduce_scalar!(f32, f64);
1316
1317impl_complex_erased_reduce_scalar!(Complex32, Complex64);
1318
1319impl_wrapping_erased_reduce_scalar!(i32, i64);
1320
1321fn validate_unique_axes(axes: &[usize], rank: usize) -> Result<()> {
1322 let mut seen = vec![false; rank];
1323 for &axis in axes {
1324 if axis >= rank {
1325 return Err(StridedError::InvalidAxis { axis, rank });
1326 }
1327 if seen[axis] {
1328 return Err(StridedError::InvalidAxis { axis, rank });
1329 }
1330 seen[axis] = true;
1331 }
1332 Ok(())
1333}
1334
1335struct ReduceOuterCursor<'a> {
1336 axes: &'a [ReduceOuterAxis],
1337 coords: CoordScratch,
1338 source_offset: isize,
1339 dest_offset: isize,
1340}
1341impl<'a> ReduceOuterCursor<'a> {
1342 fn decode(
1343 mut linear: usize,
1344 source_base: isize,
1345 dest_base: isize,
1346 axes: &'a [ReduceOuterAxis],
1347 ) -> Result<Self> {
1348 let mut coords = CoordScratch::new(axes.len());
1349 let mut source_offset = source_base;
1350 let mut dest_offset = dest_base;
1351 for (coord, axis) in coords.as_mut_slice().iter_mut().zip(axes) {
1352 debug_assert!(axis.extent != 0);
1356 *coord = linear % axis.extent;
1357 linear /= axis.extent;
1358 source_offset = checked_offset_add(source_offset, axis.source_step, *coord)?;
1359 dest_offset = checked_offset_add(dest_offset, axis.dest_step, *coord)?;
1360 }
1361 Ok(Self {
1362 axes,
1363 coords,
1364 source_offset,
1365 dest_offset,
1366 })
1367 }
1368
1369 #[inline]
1371 fn leading_coord(&self) -> usize {
1372 self.coords.first()
1373 }
1374
1375 #[inline]
1376 fn advance(&mut self) {
1377 for (coord, axis) in self.coords.as_mut_slice().iter_mut().zip(self.axes) {
1381 let next = *coord + 1;
1382 if next < axis.extent {
1383 *coord = next;
1384 self.source_offset += axis.source_step;
1385 self.dest_offset += axis.dest_step;
1386 return;
1387 }
1388 *coord = 0;
1389 self.source_offset += axis.source_reset;
1390 self.dest_offset += axis.dest_reset;
1391 }
1392 }
1393}
1394struct ReduceInnerCursor<'a> {
1395 axes: &'a [ReduceInnerAxis],
1396 coords: CoordScratch,
1397 source_offset: isize,
1398}
1399impl<'a> ReduceInnerCursor<'a> {
1400 fn new(source_base: isize, axes: &'a [ReduceInnerAxis]) -> Self {
1401 Self {
1402 axes,
1403 coords: CoordScratch::new(axes.len()),
1404 source_offset: source_base,
1405 }
1406 }
1407
1408 #[inline]
1409 fn reset(&mut self, source_base: isize) {
1410 self.coords.as_mut_slice().fill(0);
1411 self.source_offset = source_base;
1412 }
1413
1414 #[inline]
1415 fn advance(&mut self) {
1416 for (coord, axis) in self.coords.as_mut_slice().iter_mut().zip(self.axes) {
1420 let next = *coord + 1;
1421 if next < axis.extent {
1422 *coord = next;
1423 self.source_offset += axis.source_step;
1424 return;
1425 }
1426 *coord = 0;
1427 self.source_offset += axis.source_reset;
1428 }
1429 }
1430}
1431fn check_reduce_layout_offset_arithmetic(dims: &[usize], strides: &[isize]) -> Result<()> {
1432 if dims.len() != strides.len() {
1433 return Err(StridedError::StrideLengthMismatch);
1434 }
1435 let mut min_offset = 0isize;
1436 let mut max_offset = 0isize;
1437 for (&dim, &stride) in dims.iter().zip(strides) {
1438 let last =
1439 isize::try_from(dim.saturating_sub(1)).map_err(|_| StridedError::OffsetOverflow)?;
1440 let extent = stride
1441 .checked_mul(last)
1442 .ok_or(StridedError::OffsetOverflow)?;
1443 if extent < 0 {
1444 min_offset = min_offset
1445 .checked_add(extent)
1446 .ok_or(StridedError::OffsetOverflow)?;
1447 } else {
1448 max_offset = max_offset
1449 .checked_add(extent)
1450 .ok_or(StridedError::OffsetOverflow)?;
1451 }
1452 }
1453 let _ = (min_offset, max_offset);
1454 Ok(())
1455}
1456fn compress_reduce_outer_axes(axes: Vec<ReduceOuterAxis>) -> Result<Vec<ReduceOuterAxis>> {
1457 let mut compressed: Vec<ReduceOuterAxis> = Vec::with_capacity(axes.len());
1458 for axis in axes {
1459 if let Some(previous) = compressed.last_mut() {
1460 let previous_extent =
1461 isize::try_from(previous.extent).map_err(|_| StridedError::OffsetOverflow)?;
1462 let expected_source = previous
1463 .source_step
1464 .checked_mul(previous_extent)
1465 .ok_or(StridedError::OffsetOverflow)?;
1466 let expected_dest = previous
1467 .dest_step
1468 .checked_mul(previous_extent)
1469 .ok_or(StridedError::OffsetOverflow)?;
1470 if axis.source_step == expected_source && axis.dest_step == expected_dest {
1471 let fused_extent = previous
1472 .extent
1473 .checked_mul(axis.extent)
1474 .ok_or(StridedError::OffsetOverflow)?;
1475 previous.extent = fused_extent;
1476 previous.source_reset = checked_reduce_reset(fused_extent, previous.source_step)?;
1477 previous.dest_reset = checked_reduce_reset(fused_extent, previous.dest_step)?;
1478 continue;
1479 }
1480 }
1481 compressed.push(axis);
1482 }
1483 Ok(compressed)
1484}
1485fn compress_reduce_inner_axes(axes: Vec<ReduceInnerAxis>) -> Result<Vec<ReduceInnerAxis>> {
1486 let mut compressed: Vec<ReduceInnerAxis> = Vec::with_capacity(axes.len());
1487 for axis in axes {
1488 if let Some(previous) = compressed.last_mut() {
1489 let previous_extent =
1490 isize::try_from(previous.extent).map_err(|_| StridedError::OffsetOverflow)?;
1491 let expected_source = previous
1492 .source_step
1493 .checked_mul(previous_extent)
1494 .ok_or(StridedError::OffsetOverflow)?;
1495 if axis.source_step == expected_source {
1496 let fused_extent = previous
1497 .extent
1498 .checked_mul(axis.extent)
1499 .ok_or(StridedError::OffsetOverflow)?;
1500 previous.extent = fused_extent;
1501 previous.source_reset = checked_reduce_reset(fused_extent, previous.source_step)?;
1502 continue;
1503 }
1504 }
1505 compressed.push(axis);
1506 }
1507 Ok(compressed)
1508}
1509fn checked_reduce_reset(extent: usize, stride: isize) -> Result<isize> {
1510 if extent == 0 {
1511 return Ok(0);
1512 }
1513 let last = isize::try_from(extent - 1).map_err(|_| StridedError::OffsetOverflow)?;
1514 stride
1515 .checked_mul(last)
1516 .and_then(isize::checked_neg)
1517 .ok_or(StridedError::OffsetOverflow)
1518}
1519fn checked_offset_add(base: isize, stride: isize, coord: usize) -> Result<isize> {
1520 let coord = isize::try_from(coord).map_err(|_| StridedError::OffsetOverflow)?;
1521 let scaled = stride
1522 .checked_mul(coord)
1523 .ok_or(StridedError::OffsetOverflow)?;
1524 base.checked_add(scaled).ok_or(StridedError::OffsetOverflow)
1525}
1526
1527struct CoordScratch {
1528 inline: [usize; RAW_FUSED_RANK_LIMIT],
1529 heap: Option<Vec<usize>>,
1530 len: usize,
1531}
1532
1533impl CoordScratch {
1534 fn new(len: usize) -> Self {
1535 if len <= RAW_FUSED_RANK_LIMIT {
1536 Self {
1537 inline: [0; RAW_FUSED_RANK_LIMIT],
1538 heap: None,
1539 len,
1540 }
1541 } else {
1542 Self {
1543 inline: [0; RAW_FUSED_RANK_LIMIT],
1544 heap: Some(vec![0; len]),
1545 len,
1546 }
1547 }
1548 }
1549
1550 #[inline]
1551 fn first(&self) -> usize {
1552 match &self.heap {
1553 Some(heap) => heap.first().copied().unwrap_or(0),
1554 None => self.inline[0],
1555 }
1556 }
1557
1558 fn as_mut_slice(&mut self) -> &mut [usize] {
1559 match &mut self.heap {
1560 Some(heap) => heap,
1561 None => &mut self.inline[..self.len],
1562 }
1563 }
1564}