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