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