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}
229
230#[derive(Clone, Debug)]
236pub struct ErasedReducePlan {
237 dtype: KernelDType,
238 op: ReduceOp,
239 layout: ReduceLayout,
240}
241
242#[derive(Clone, Debug)]
243enum ReduceLayout {
244 Full {
245 dims: Vec<usize>,
246 src_strides: Vec<isize>,
247 },
248 Axes {
249 src_dims: Vec<usize>,
250 src_strides: Vec<isize>,
251 dest_dims: Vec<usize>,
252 dest_strides: Vec<isize>,
253 axes: Vec<usize>,
254 kept_axes: Vec<usize>,
255 outer_axes: Vec<ReduceOuterAxis>,
256 inner_axes: Vec<ReduceInnerAxis>,
257 dest_total: usize,
258 reduce_total: usize,
259 },
260}
261
262#[derive(Clone, Copy, Debug)]
263struct ReduceOuterAxis {
264 extent: usize,
265 source_step: isize,
266 source_reset: isize,
267 dest_step: isize,
268 dest_reset: isize,
269}
270#[derive(Clone, Copy, Debug)]
271struct ReduceInnerAxis {
272 extent: usize,
273 source_step: isize,
274 source_reset: isize,
275}
276impl ReduceLayout {
277 fn src_dims(&self) -> &[usize] {
278 match self {
279 Self::Full { dims, .. } => dims,
280 Self::Axes { src_dims, .. } => src_dims,
281 }
282 }
283
284 fn src_strides(&self) -> &[isize] {
285 match self {
286 Self::Full { src_strides, .. } | Self::Axes { src_strides, .. } => src_strides,
287 }
288 }
289
290 fn check_src_layout(&self, src: &ErasedRawStridedRef<'_>) -> Result<()> {
291 if src.dims() != self.src_dims() || src.strides() != self.src_strides() {
292 return Err(StridedError::PlanLayoutMismatch);
293 }
294 Ok(())
295 }
296}
297
298#[derive(Clone, Copy, Debug)]
299struct AxesLayout<'a> {
300 src_dims: &'a [usize],
301 axes: &'a [usize],
302 kept_axes: &'a [usize],
303 outer_axes: &'a [ReduceOuterAxis],
304 inner_axes: &'a [ReduceInnerAxis],
305 dest_total: usize,
306 reduce_total: usize,
307}
308
309impl ErasedReducePlan {
310 pub fn compile(
312 dtype: KernelDType,
313 op: ReduceOp,
314 dims: &[usize],
315 src_strides: &[isize],
316 ) -> Result<Self> {
317 check_reduce_op_dtype(dtype, op)?;
318 if dims.len() != src_strides.len() {
319 return Err(StridedError::StrideLengthMismatch);
320 }
321 checked_total_len(dims)?;
322 Ok(Self {
323 dtype,
324 op,
325 layout: ReduceLayout::Full {
326 dims: dims.to_vec(),
327 src_strides: src_strides.to_vec(),
328 },
329 })
330 }
331
332 #[allow(clippy::too_many_arguments)]
338 pub fn compile_axes(
339 dtype: KernelDType,
340 op: ReduceOp,
341 src_dims: &[usize],
342 src_strides: &[isize],
343 dest_dims: &[usize],
344 dest_strides: &[isize],
345 axes: &[usize],
346 ) -> Result<Self> {
347 check_reduce_op_dtype(dtype, op)?;
348 if src_dims.len() != src_strides.len() || dest_dims.len() != dest_strides.len() {
349 return Err(StridedError::StrideLengthMismatch);
350 }
351 checked_total_len(src_dims)?;
352 check_reduce_layout_offset_arithmetic(src_dims, src_strides)?;
353 let dest_total = checked_total_len(dest_dims)?;
354 check_reduce_layout_offset_arithmetic(dest_dims, dest_strides)?;
355 if !crate::layout_check::is_injective_layout(dest_dims, dest_strides) {
356 return Err(StridedError::NonInjectiveOutputLayout);
357 }
358 validate_unique_axes(axes, src_dims.len())?;
359
360 let kept_axes: Vec<usize> = (0..src_dims.len())
361 .filter(|axis| !axes.contains(axis))
362 .collect();
363 let expected_dest_dims: Vec<usize> = kept_axes.iter().map(|&axis| src_dims[axis]).collect();
364 if expected_dest_dims.is_empty() {
365 if dest_total != 1 {
366 return Err(StridedError::ShapeMismatch(
367 dest_dims.to_vec(),
368 expected_dest_dims,
369 ));
370 }
371 } else if dest_dims != expected_dest_dims.as_slice() {
372 return Err(StridedError::ShapeMismatch(
373 dest_dims.to_vec(),
374 expected_dest_dims,
375 ));
376 }
377
378 let reduce_total = axes
379 .iter()
380 .try_fold(1usize, |total, &axis| total.checked_mul(src_dims[axis]))
381 .ok_or(StridedError::OffsetOverflow)?;
382 let outer_axes = compress_reduce_outer_axes(
383 kept_axes
384 .iter()
385 .enumerate()
386 .map(|(dest_axis, &src_axis)| {
387 let extent = src_dims[src_axis];
388 Ok(ReduceOuterAxis {
389 extent,
390 source_step: src_strides[src_axis],
391 source_reset: checked_reduce_reset(extent, src_strides[src_axis])?,
392 dest_step: dest_strides[dest_axis],
393 dest_reset: checked_reduce_reset(extent, dest_strides[dest_axis])?,
394 })
395 })
396 .collect::<Result<Vec<_>>>()?,
397 )?;
398 let inner_axes = compress_reduce_inner_axes(
399 axes.iter()
400 .map(|&src_axis| {
401 let extent = src_dims[src_axis];
402 Ok(ReduceInnerAxis {
403 extent,
404 source_step: src_strides[src_axis],
405 source_reset: checked_reduce_reset(extent, src_strides[src_axis])?,
406 })
407 })
408 .collect::<Result<Vec<_>>>()?,
409 )?;
410 Ok(Self {
411 dtype,
412 op,
413 layout: ReduceLayout::Axes {
414 src_dims: src_dims.to_vec(),
415 src_strides: src_strides.to_vec(),
416 dest_dims: dest_dims.to_vec(),
417 dest_strides: dest_strides.to_vec(),
418 axes: axes.to_vec(),
419 kept_axes,
420 outer_axes,
421 inner_axes,
422 dest_total,
423 reduce_total,
424 },
425 })
426 }
427
428 #[inline]
429 pub fn dtype(&self) -> KernelDType {
430 self.dtype
431 }
432
433 #[inline]
434 pub fn op(&self) -> ReduceOp {
435 self.op
436 }
437
438 pub fn execute(
440 &self,
441 ctx: &ExecContext,
442 dest: &mut ErasedRawStridedMut<'_>,
443 src: &ErasedRawStridedRef<'_>,
444 ) -> Result<()> {
445 check_dtype(self.dtype, dest.dtype())?;
446 check_dtype(self.dtype, src.dtype())?;
447 self.layout.check_src_layout(src)?;
448 match &self.layout {
449 ReduceLayout::Full { .. } => {
450 let dest_len = checked_total_len(dest.dims())?;
451 if dest_len != 1 {
452 return Err(StridedError::RankMismatch(dest_len, 1));
453 }
454 }
455 ReduceLayout::Axes {
456 dest_dims,
457 dest_strides,
458 ..
459 } => {
460 if dest.dims() != dest_dims.as_slice() || dest.strides() != dest_strides.as_slice()
461 {
462 return Err(StridedError::PlanLayoutMismatch);
463 }
464 }
465 }
466
467 let result = match self.dtype {
468 KernelDType::F32 => {
469 let mut writer = reduce_writer::<f32>(dest)?;
470 dispatch_reduce::<f32, _>(self.op, &self.layout, ctx, &mut writer, src)
471 }
472 KernelDType::F64 => {
473 let mut writer = reduce_writer::<f64>(dest)?;
474 dispatch_reduce::<f64, _>(self.op, &self.layout, ctx, &mut writer, src)
475 }
476 KernelDType::I32 => {
477 let mut writer = reduce_writer::<i32>(dest)?;
478 dispatch_reduce::<i32, _>(self.op, &self.layout, ctx, &mut writer, src)
479 }
480 KernelDType::I64 => {
481 let mut writer = reduce_writer::<i64>(dest)?;
482 dispatch_reduce::<i64, _>(self.op, &self.layout, ctx, &mut writer, src)
483 }
484 KernelDType::C32 => {
485 let mut writer = reduce_writer::<Complex32>(dest)?;
486 dispatch_reduce::<Complex32, _>(self.op, &self.layout, ctx, &mut writer, src)
487 }
488 KernelDType::C64 => {
489 let mut writer = reduce_writer::<Complex64>(dest)?;
490 dispatch_reduce::<Complex64, _>(self.op, &self.layout, ctx, &mut writer, src)
491 }
492 _ => Err(StridedError::UnsupportedDType {
493 dtype: self.dtype.label(),
494 }),
495 };
496 result
497 }
498
499 pub fn execute_uninit(
506 &self,
507 ctx: &ExecContext,
508 dest: &mut ErasedRawStridedUninitMut<'_>,
509 src: &ErasedRawStridedPtr<'_>,
510 ) -> Result<()> {
511 check_dtype(self.dtype, dest.dtype())?;
512 check_dtype(self.dtype, src.dtype())?;
513 validate_uninit_no_overlap(dest, src, 0)?;
514 let src = unsafe { src.try_as_ref_after_no_overlap() }?;
516 self.layout.check_src_layout(&src)?;
517 match &self.layout {
518 ReduceLayout::Full { .. } => {
519 let total = checked_total_len(dest.dims())?;
520 if total != 1 {
521 return Err(StridedError::RankMismatch(total, 1));
522 }
523 }
524 ReduceLayout::Axes {
525 dest_dims,
526 dest_strides,
527 ..
528 } => {
529 if dest.dims() != dest_dims.as_slice() || dest.strides() != dest_strides.as_slice()
530 {
531 return Err(StridedError::PlanLayoutMismatch);
532 }
533 }
534 }
535 macro_rules! run {
536 ($ty:ty) => {{
537 let mut writer = reduce_uninit_writer::<$ty>(dest)?;
538 dispatch_reduce::<$ty, _>(self.op, &self.layout, ctx, &mut writer, &src)
539 }};
540 }
541 match self.dtype {
542 KernelDType::F32 => run!(f32),
543 KernelDType::F64 => run!(f64),
544 KernelDType::I32 => run!(i32),
545 KernelDType::I64 => run!(i64),
546 KernelDType::C32 => run!(Complex32),
547 KernelDType::C64 => run!(Complex64),
548 _ => Err(StridedError::UnsupportedDType {
549 dtype: self.dtype.label(),
550 }),
551 }
552 }
553}
554
555fn reduce_writer<'a, T>(dest: &'a mut ErasedRawStridedMut<'_>) -> Result<RawReduceWriter<'a, T>>
556where
557 T: KernelStorageElement,
558{
559 let offset = dest.offset();
560 let data = dest.data_as_mut::<T>()?;
561 let ptr = data.as_mut_ptr();
562 let extent = data.len();
563 Ok(RawReduceWriter {
564 ptr,
565 extent,
566 offset,
567 _marker: core::marker::PhantomData,
568 })
569}
570
571fn reduce_uninit_writer<'a, T>(
572 dest: &'a mut ErasedRawStridedUninitMut<'_>,
573) -> Result<RawReduceWriter<'a, T>>
574where
575 T: KernelStorageElement,
576{
577 let offset = dest.offset();
578 let data = dest.data_as_uninit_mut::<T>()?;
579 let ptr = data.as_mut_ptr().cast::<T>();
580 let extent = data.len();
581 Ok(RawReduceWriter {
582 ptr,
583 extent,
584 offset,
585 _marker: core::marker::PhantomData,
586 })
587}
588
589fn check_reduce_dtype(dtype: KernelDType) -> Result<()> {
590 match dtype {
591 KernelDType::F32
592 | KernelDType::F64
593 | KernelDType::I32
594 | KernelDType::I64
595 | KernelDType::C32
596 | KernelDType::C64 => Ok(()),
597 _ => Err(StridedError::UnsupportedDType {
598 dtype: dtype.label(),
599 }),
600 }
601}
602
603fn check_reduce_op_dtype(dtype: KernelDType, op: ReduceOp) -> Result<()> {
604 if op == ReduceOp::SumSquares && !matches!(dtype, KernelDType::F32 | KernelDType::F64) {
605 return Err(StridedError::UnsupportedDType {
606 dtype: dtype.label(),
607 });
608 }
609 check_reduce_dtype(dtype)
610}
611
612fn checked_total_len(dims: &[usize]) -> Result<usize> {
613 if dims.is_empty() {
614 return Ok(1);
615 }
616 dims.iter()
617 .try_fold(1usize, |acc, &dim| acc.checked_mul(dim))
618 .ok_or(StridedError::OffsetOverflow)
619}
620
621fn execute_copy<T>(
622 plan: &CopyPlan,
623 dest: &mut ErasedRawStridedMut<'_>,
624 src: &ErasedRawStridedRef<'_>,
625) -> Result<()>
626where
627 T: Copy + crate::MaybeSendSync + KernelStorageElement,
628{
629 let source_data = src.data_as::<T>()?;
630 let dest_dims = dest.dims();
631 let dest_strides = dest.strides();
632 let dest_offset = dest.offset();
633 let dest_data = dest.data_as_mut::<T>()?;
634 let source = unsafe {
635 RawStridedRef::new_unchecked(source_data, src.dims(), src.strides(), src.offset())
636 };
637 let mut dest =
638 unsafe { RawStridedMut::new_unchecked(dest_data, dest_dims, dest_strides, dest_offset) };
639 plan.execute(&mut dest, &source)
640}
641
642fn execute_concatenate<T>(
643 plan: &ConcatenatePlan,
644 dest: &mut ErasedRawStridedMut<'_>,
645 inputs: &[ErasedRawStridedRef<'_>],
646) -> Result<()>
647where
648 T: Copy + crate::MaybeSendSync + KernelStorageElement,
649{
650 if inputs.len() != plan.input_count() {
651 return Err(StridedError::RankMismatch(inputs.len(), plan.input_count()));
652 }
653 let dest_dims = dest.dims();
654 let dest_strides = dest.strides();
655 let dest_offset = dest.offset();
656 let dest_data = dest.data_as_mut::<T>()?;
657 let mut dest_ref =
658 unsafe { RawStridedMut::new_unchecked(dest_data, dest_dims, dest_strides, dest_offset) };
659 plan.check_dest_layout(&dest_ref)?;
660
661 for (position, input) in inputs.iter().enumerate() {
662 let input_data = input.data_as::<T>()?;
663 let input_ref = unsafe {
664 RawStridedRef::new_unchecked(input_data, input.dims(), input.strides(), input.offset())
665 };
666 plan.check_input_layout(position, &input_ref)?;
667 plan.segment_offset(position, dest_offset)?;
668 }
669 for (position, input) in inputs.iter().enumerate() {
670 let input_data = input.data_as::<T>()?;
671 let input_ref = unsafe {
672 RawStridedRef::new_unchecked(input_data, input.dims(), input.strides(), input.offset())
673 };
674 plan.execute_segment(position, &mut dest_ref, &input_ref)?;
675 }
676 Ok(())
677}
678
679fn execute_concatenate_uninit<T>(
680 plan: &ConcatenatePlan,
681 dest: &mut ErasedRawStridedUninitMut<'_>,
682 inputs: &[ErasedRawStridedPtr<'_>],
683) -> Result<()>
684where
685 T: Copy + crate::MaybeSendSync + KernelStorageElement,
686{
687 let dest_dims = dest.dims();
688 let dest_strides = dest.strides();
689 let dest_offset = dest.offset();
690 let dest_data = dest.data_as_uninit_mut::<T>()?;
691 let mut dest_ref =
692 unsafe { RawStridedMut::new_unchecked(dest_data, dest_dims, dest_strides, dest_offset) };
693 plan.check_dest_layout(&dest_ref)?;
694
695 for (position, input) in inputs.iter().enumerate() {
696 let input = unsafe { input.try_as_ref_after_no_overlap() }?;
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.check_input_layout(position, &input_ref)?;
703 plan.segment_offset(position, dest_offset)?;
704 }
705 for (position, input) in inputs.iter().enumerate() {
706 let input = unsafe { input.try_as_ref_after_no_overlap() }?;
708 let input_data = input.data_as::<T>()?;
709 let input_ref = unsafe {
710 RawStridedRef::new_unchecked(input_data, input.dims(), input.strides(), input.offset())
711 };
712 plan.execute_segment_uninit(position, &mut dest_ref, &input_ref)?;
713 }
714 Ok(())
715}
716
717fn execute_reduce<T, W>(
718 op: ReduceOp,
719 ctx: &ExecContext,
720 dest: &mut W,
721 src: &ErasedRawStridedRef<'_>,
722) -> Result<()>
723where
724 T: ErasedReduceScalar,
725 W: ReduceWriter<T>,
726{
727 let use_serial = ctx.is_serial()
728 || ctx
729 .max_threads_limit()
730 .is_some_and(|max_threads| max_threads.get() == 1);
731 let value = if use_serial {
732 if let Some(value) = reduce_contiguous_serial(op, src) {
733 value
734 } else {
735 let source = erased_view::<T>(src)?;
736 crate::reduce_view::reduce_serial(
737 &source,
738 |value| reduce_map_value(op, value),
739 |a, b| reduce_values(op, a, b),
740 reduce_identity(op),
741 )?
742 }
743 } else {
744 let source = erased_view::<T>(src)?;
745 ctx.run(|| {
746 crate::reduce(
747 &source,
748 |value| reduce_map_value(op, value),
749 |a, b| reduce_values(op, a, b),
750 reduce_identity(op),
751 )
752 })?
753 };
754
755 unsafe { dest.write_at(dest.offset(), value) };
757 Ok(())
758}
759
760fn reduce_contiguous_serial<T>(op: ReduceOp, src: &ErasedRawStridedRef<'_>) -> Option<T>
761where
762 T: ErasedReduceScalar,
763{
764 crate::kernel::same_contiguous_layout(src.dims(), &[src.strides()])?;
765 let len = checked_total_len(src.dims()).ok()?;
766 if len == 0 {
767 return Some(reduce_identity(op));
768 }
769
770 let source_data = src.data_as::<T>().ok()?;
771 let start = usize::try_from(src.offset()).ok()?;
772 let end = start.checked_add(len)?;
773 let values = source_data.get(start..end)?;
774 Some(match op {
775 ReduceOp::Sum => T::try_simd_sum(values)
776 .unwrap_or_else(|| reduce_contiguous_lanes(values, T::zero(), T::reduce_sum)),
777 ReduceOp::Product => T::try_simd_product(values)
778 .unwrap_or_else(|| reduce_contiguous_lanes(values, T::one(), T::reduce_product)),
779 ReduceOp::SumSquares => T::try_simd_sum_squares(values).unwrap_or_else(|| {
780 reduce_contiguous_mapped_lanes(
781 values,
782 T::zero(),
783 |value| T::reduce_product(value, value),
784 T::reduce_sum,
785 )
786 }),
787 })
788}
789
790#[inline]
791fn reduce_contiguous_lanes<T>(values: &[T], identity: T, combine: impl Fn(T, T) -> T) -> T
792where
793 T: Copy,
794{
795 reduce_contiguous_mapped_lanes(values, identity, |value| value, combine)
796}
797
798#[inline]
799fn reduce_contiguous_mapped_lanes<T>(
800 values: &[T],
801 identity: T,
802 map: impl Fn(T) -> T,
803 combine: impl Fn(T, T) -> T,
804) -> T
805where
806 T: Copy,
807{
808 let mut lanes = [identity; SERIAL_REDUCE_LANES];
809 let mut chunks = values.chunks_exact(SERIAL_REDUCE_LANES);
810 for chunk in chunks.by_ref() {
811 for lane in 0..SERIAL_REDUCE_LANES {
812 lanes[lane] = combine(lanes[lane], map(chunk[lane]));
813 }
814 }
815 for (lane, &value) in chunks.remainder().iter().enumerate() {
816 lanes[lane] = combine(lanes[lane], map(value));
817 }
818 lanes.into_iter().fold(identity, combine)
819}
820
821fn dispatch_reduce<T, W>(
822 op: ReduceOp,
823 layout: &ReduceLayout,
824 ctx: &ExecContext,
825 dest: &mut W,
826 src: &ErasedRawStridedRef<'_>,
827) -> Result<()>
828where
829 T: ErasedReduceScalar,
830 W: ReduceWriter<T>,
831{
832 match layout {
833 ReduceLayout::Full { .. } => execute_reduce::<T, W>(op, ctx, dest, src),
834 ReduceLayout::Axes {
835 src_dims,
836 axes,
837 kept_axes,
838 outer_axes,
839 inner_axes,
840 dest_total,
841 reduce_total,
842 ..
843 } => execute_reduce_axes::<T, W>(
844 op,
845 ctx,
846 dest,
847 src,
848 AxesLayout {
849 src_dims,
850 axes,
851 kept_axes,
852 outer_axes,
853 inner_axes,
854 dest_total: *dest_total,
855 reduce_total: *reduce_total,
856 },
857 ),
858 }
859}
860
861fn execute_reduce_axes<T, W>(
862 op: ReduceOp,
863 ctx: &ExecContext,
864 dest: &mut W,
865 src: &ErasedRawStridedRef<'_>,
866 layout: AxesLayout<'_>,
867) -> Result<()>
868where
869 T: ErasedReduceScalar,
870 W: ReduceWriter<T>,
871{
872 if layout.kept_axes.is_empty()
873 && layout.axes.len() == layout.src_dims.len()
874 && layout.dest_total == 1
875 {
876 return execute_reduce::<T, W>(op, ctx, dest, src);
877 }
878
879 if layout.dest_total == 0 {
880 return Ok(());
881 }
882
883 if layout.reduce_total == 0 {
884 if ctx.is_serial() {
885 execute_reduce_axes_identity_serial(op, dest, layout)
886 } else {
887 ctx.run(|| execute_reduce_axes_identity_policy(op, dest, layout))
888 }
889 } else if ctx.is_serial() {
890 execute_reduce_axes_serial::<T, W>(op, dest, src, layout)
891 } else {
892 ctx.run(|| execute_reduce_axes_policy::<T, W>(op, dest, src, layout))
893 }
894}
895
896fn execute_reduce_axes_policy<T, W>(
897 op: ReduceOp,
898 dest: &mut W,
899 src: &ErasedRawStridedRef<'_>,
900 layout: AxesLayout<'_>,
901) -> Result<()>
902where
903 T: ErasedReduceScalar,
904 W: ReduceWriter<T>,
905{
906 let source_data = src.data_as::<T>()?;
907 let dest_offset_base = dest.offset();
908 #[cfg(feature = "parallel")]
909 {
910 let nthreads = crate::threading::parallel_threads_for_len(layout.dest_total);
911 if nthreads > 1 {
912 return execute_reduce_axes_parallel(
913 op,
914 dest_offset_base,
915 dest,
916 src.offset(),
917 source_data,
918 layout,
919 nthreads,
920 );
921 }
922 }
923
924 execute_reduce_axes_serial_data(
925 op,
926 dest_offset_base,
927 dest,
928 src.offset(),
929 source_data,
930 layout,
931 )
932}
933
934fn execute_reduce_axes_serial<T, W>(
935 op: ReduceOp,
936 dest: &mut W,
937 src: &ErasedRawStridedRef<'_>,
938 layout: AxesLayout<'_>,
939) -> Result<()>
940where
941 T: ErasedReduceScalar,
942 W: ReduceWriter<T>,
943{
944 let source_data = src.data_as::<T>()?;
945 execute_reduce_axes_serial_data(op, dest.offset(), dest, src.offset(), source_data, layout)
946}
947
948fn execute_reduce_axes_serial_data<T, W>(
949 op: ReduceOp,
950 dest_offset_base: isize,
951 dest: &mut W,
952 source_offset_base: isize,
953 source_data: &[T],
954 layout: AxesLayout<'_>,
955) -> Result<()>
956where
957 T: ErasedReduceScalar,
958 W: ReduceWriter<T>,
959{
960 let mut outer =
961 ReduceOuterCursor::decode(0, source_offset_base, dest_offset_base, layout.outer_axes)?;
962 let reduce_inner = |inner: &mut ReduceInnerCursor<'_>| {
967 let mut acc = reduce_identity(op);
968 for value_index in 0..layout.reduce_total {
969 let value = unsafe { *source_data.as_ptr().offset(inner.source_offset) };
972 acc = reduce_values(op, acc, reduce_map_value(op, value));
973 if value_index + 1 < layout.reduce_total {
974 inner.advance();
975 }
976 }
977 acc
978 };
979
980 if layout.inner_axes.len() <= RAW_FUSED_RANK_LIMIT {
981 for output in 0..layout.dest_total {
982 let mut inner = ReduceInnerCursor::new(outer.source_offset, layout.inner_axes);
983 let acc = reduce_inner(&mut inner);
984 unsafe { dest.write_at(outer.dest_offset, acc) };
987 if output + 1 < layout.dest_total {
988 outer.advance();
989 }
990 }
991 } else {
992 let mut inner = ReduceInnerCursor::new(source_offset_base, layout.inner_axes);
993 for output in 0..layout.dest_total {
994 inner.reset(outer.source_offset);
995 let acc = reduce_inner(&mut inner);
996 unsafe { dest.write_at(outer.dest_offset, acc) };
999 if output + 1 < layout.dest_total {
1000 outer.advance();
1001 }
1002 }
1003 }
1004 Ok(())
1005}
1006
1007fn execute_reduce_axes_identity_serial<T, W>(
1008 op: ReduceOp,
1009 dest: &mut W,
1010 layout: AxesLayout<'_>,
1011) -> Result<()>
1012where
1013 T: ErasedReduceScalar,
1014 W: ReduceWriter<T>,
1015{
1016 let mut outer = ReduceOuterCursor::decode(0, 0, dest.offset(), layout.outer_axes)?;
1017 for output in 0..layout.dest_total {
1018 unsafe { dest.write_at(outer.dest_offset, reduce_identity(op)) };
1026 if output + 1 < layout.dest_total {
1027 outer.advance();
1028 }
1029 }
1030 Ok(())
1031}
1032fn execute_reduce_axes_identity_policy<T, W>(
1033 op: ReduceOp,
1034 dest: &mut W,
1035 layout: AxesLayout<'_>,
1036) -> Result<()>
1037where
1038 T: ErasedReduceScalar,
1039 W: ReduceWriter<T>,
1040{
1041 #[cfg(feature = "parallel")]
1042 {
1043 let nthreads = crate::threading::parallel_threads_for_len(layout.dest_total);
1044 if nthreads > 1 {
1045 return execute_reduce_axes_identity_parallel(op, dest, layout, nthreads);
1046 }
1047 }
1048 execute_reduce_axes_identity_serial(op, dest, layout)
1049}
1050#[cfg(feature = "parallel")]
1051fn execute_reduce_axes_identity_parallel<T, W>(
1052 op: ReduceOp,
1053 dest: &mut W,
1054 layout: AxesLayout<'_>,
1055 nthreads: usize,
1056) -> Result<()>
1057where
1058 T: ErasedReduceScalar,
1059 W: ReduceWriter<T>,
1060{
1061 let dest_ptr = crate::threading::SendPtr(unsafe { dest.ptr() });
1063 let dest_offset_base = dest.offset();
1064 crate::threading::parallel_map_reduce(
1065 0..layout.dest_total,
1066 nthreads,
1067 &|range| {
1068 let range_end = range.end;
1069 let mut outer =
1070 ReduceOuterCursor::decode(range.start, 0, dest_offset_base, layout.outer_axes)?;
1071 let dest_ptr = dest_ptr.as_ptr();
1072 for output in range {
1073 unsafe {
1081 dest_ptr
1082 .offset(outer.dest_offset)
1083 .write(reduce_identity(op))
1084 };
1085 if output + 1 < range_end {
1086 outer.advance();
1087 }
1088 }
1089 Ok(())
1090 },
1091 &|left, right| left.and(right),
1092 )
1093}
1094#[cfg(feature = "parallel")]
1095fn execute_reduce_axes_parallel<T, W>(
1096 op: ReduceOp,
1097 dest_offset_base: isize,
1098 dest: &mut W,
1099 source_offset_base: isize,
1100 source_data: &[T],
1101 layout: AxesLayout<'_>,
1102 nthreads: usize,
1103) -> Result<()>
1104where
1105 T: ErasedReduceScalar,
1106 W: ReduceWriter<T>,
1107{
1108 let dest_ptr = crate::threading::SendPtr(unsafe { dest.ptr() });
1110 let source_ptr = crate::threading::SendPtr(source_data.as_ptr() as *mut T);
1111 crate::threading::parallel_map_reduce(
1112 0..layout.dest_total,
1113 nthreads,
1114 &|range| {
1115 let range_end = range.end;
1116 let mut outer = ReduceOuterCursor::decode(
1117 range.start,
1118 source_offset_base,
1119 dest_offset_base,
1120 layout.outer_axes,
1121 )?;
1122 let dest_ptr = dest_ptr.as_ptr();
1123 let source_ptr = source_ptr.as_const();
1124 let reduce_inner = |inner: &mut ReduceInnerCursor<'_>| {
1129 let mut acc = reduce_identity(op);
1130 for value_index in 0..layout.reduce_total {
1131 let value = unsafe { *source_ptr.offset(inner.source_offset) };
1134 acc = reduce_values(op, acc, reduce_map_value(op, value));
1135 if value_index + 1 < layout.reduce_total {
1136 inner.advance();
1137 }
1138 }
1139 acc
1140 };
1141
1142 if layout.inner_axes.len() <= RAW_FUSED_RANK_LIMIT {
1143 for output in range {
1144 let mut inner = ReduceInnerCursor::new(outer.source_offset, layout.inner_axes);
1145 let acc = reduce_inner(&mut inner);
1146 unsafe { dest_ptr.offset(outer.dest_offset).write(acc) };
1149 if output + 1 < range_end {
1150 outer.advance();
1151 }
1152 }
1153 } else {
1154 let mut inner = ReduceInnerCursor::new(outer.source_offset, layout.inner_axes);
1155 for output in range {
1156 inner.reset(outer.source_offset);
1157 let acc = reduce_inner(&mut inner);
1158 unsafe { dest_ptr.offset(outer.dest_offset).write(acc) };
1161 if output + 1 < range_end {
1162 outer.advance();
1163 }
1164 }
1165 }
1166 Ok(())
1167 },
1168 &|left, right| left.and(right),
1169 )
1170}
1171
1172#[inline]
1173fn reduce_identity<T>(op: ReduceOp) -> T
1174where
1175 T: One + Zero,
1176{
1177 match op {
1178 ReduceOp::Sum => T::zero(),
1179 ReduceOp::Product => T::one(),
1180 ReduceOp::SumSquares => T::zero(),
1181 }
1182}
1183
1184#[inline]
1185fn reduce_values<T>(op: ReduceOp, a: T, b: T) -> T
1186where
1187 T: ErasedReduceScalar,
1188{
1189 match op {
1190 ReduceOp::Sum => T::reduce_sum(a, b),
1191 ReduceOp::Product => T::reduce_product(a, b),
1192 ReduceOp::SumSquares => T::reduce_sum(a, b),
1193 }
1194}
1195
1196#[inline]
1197fn reduce_map_value<T>(op: ReduceOp, value: T) -> T
1198where
1199 T: ErasedReduceScalar,
1200{
1201 match op {
1202 ReduceOp::Sum | ReduceOp::Product => value,
1203 ReduceOp::SumSquares => T::reduce_product(value, value),
1204 }
1205}
1206
1207trait ErasedReduceScalar:
1208 KernelStorageElement
1209 + Copy
1210 + One
1211 + Zero
1212 + crate::MaybeSendSync
1213 + crate::simd::MaybeSimdOps
1214 + crate::simd::MaybeSimdProduct
1215 + crate::simd::MaybeSimdSumSquares
1216{
1217 fn reduce_sum(lhs: Self, rhs: Self) -> Self;
1218 fn reduce_product(lhs: Self, rhs: Self) -> Self;
1219}
1220
1221macro_rules! impl_default_erased_reduce_scalar {
1222 ($($ty:ty),* $(,)?) => {
1223 $(
1224 impl ErasedReduceScalar for $ty {
1225 #[inline(always)]
1226 fn reduce_sum(lhs: Self, rhs: Self) -> Self {
1227 lhs + rhs
1228 }
1229
1230 #[inline(always)]
1231 fn reduce_product(lhs: Self, rhs: Self) -> Self {
1232 lhs * rhs
1233 }
1234 }
1235 )*
1236 };
1237}
1238
1239macro_rules! impl_wrapping_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.wrapping_add(rhs)
1246 }
1247
1248 #[inline(always)]
1249 fn reduce_product(lhs: Self, rhs: Self) -> Self {
1250 lhs.wrapping_mul(rhs)
1251 }
1252 }
1253 )*
1254 };
1255}
1256
1257impl_default_erased_reduce_scalar!(f32, f64, Complex32, Complex64);
1258
1259impl_wrapping_erased_reduce_scalar!(i32, i64);
1260
1261fn validate_unique_axes(axes: &[usize], rank: usize) -> Result<()> {
1262 let mut seen = vec![false; rank];
1263 for &axis in axes {
1264 if axis >= rank {
1265 return Err(StridedError::InvalidAxis { axis, rank });
1266 }
1267 if seen[axis] {
1268 return Err(StridedError::InvalidAxis { axis, rank });
1269 }
1270 seen[axis] = true;
1271 }
1272 Ok(())
1273}
1274
1275struct ReduceOuterCursor<'a> {
1276 axes: &'a [ReduceOuterAxis],
1277 coords: CoordScratch,
1278 source_offset: isize,
1279 dest_offset: isize,
1280}
1281impl<'a> ReduceOuterCursor<'a> {
1282 fn decode(
1283 mut linear: usize,
1284 source_base: isize,
1285 dest_base: isize,
1286 axes: &'a [ReduceOuterAxis],
1287 ) -> Result<Self> {
1288 let mut coords = CoordScratch::new(axes.len());
1289 let mut source_offset = source_base;
1290 let mut dest_offset = dest_base;
1291 for (coord, axis) in coords.as_mut_slice().iter_mut().zip(axes) {
1292 debug_assert!(axis.extent != 0);
1296 *coord = linear % axis.extent;
1297 linear /= axis.extent;
1298 source_offset = checked_offset_add(source_offset, axis.source_step, *coord)?;
1299 dest_offset = checked_offset_add(dest_offset, axis.dest_step, *coord)?;
1300 }
1301 Ok(Self {
1302 axes,
1303 coords,
1304 source_offset,
1305 dest_offset,
1306 })
1307 }
1308
1309 #[inline]
1310 fn advance(&mut self) {
1311 for (coord, axis) in self.coords.as_mut_slice().iter_mut().zip(self.axes) {
1315 let next = *coord + 1;
1316 if next < axis.extent {
1317 *coord = next;
1318 self.source_offset += axis.source_step;
1319 self.dest_offset += axis.dest_step;
1320 return;
1321 }
1322 *coord = 0;
1323 self.source_offset += axis.source_reset;
1324 self.dest_offset += axis.dest_reset;
1325 }
1326 }
1327}
1328struct ReduceInnerCursor<'a> {
1329 axes: &'a [ReduceInnerAxis],
1330 coords: CoordScratch,
1331 source_offset: isize,
1332}
1333impl<'a> ReduceInnerCursor<'a> {
1334 fn new(source_base: isize, axes: &'a [ReduceInnerAxis]) -> Self {
1335 Self {
1336 axes,
1337 coords: CoordScratch::new(axes.len()),
1338 source_offset: source_base,
1339 }
1340 }
1341
1342 #[inline]
1343 fn reset(&mut self, source_base: isize) {
1344 self.coords.as_mut_slice().fill(0);
1345 self.source_offset = source_base;
1346 }
1347
1348 #[inline]
1349 fn advance(&mut self) {
1350 for (coord, axis) in self.coords.as_mut_slice().iter_mut().zip(self.axes) {
1354 let next = *coord + 1;
1355 if next < axis.extent {
1356 *coord = next;
1357 self.source_offset += axis.source_step;
1358 return;
1359 }
1360 *coord = 0;
1361 self.source_offset += axis.source_reset;
1362 }
1363 }
1364}
1365fn check_reduce_layout_offset_arithmetic(dims: &[usize], strides: &[isize]) -> Result<()> {
1366 if dims.len() != strides.len() {
1367 return Err(StridedError::StrideLengthMismatch);
1368 }
1369 let mut min_offset = 0isize;
1370 let mut max_offset = 0isize;
1371 for (&dim, &stride) in dims.iter().zip(strides) {
1372 let last =
1373 isize::try_from(dim.saturating_sub(1)).map_err(|_| StridedError::OffsetOverflow)?;
1374 let extent = stride
1375 .checked_mul(last)
1376 .ok_or(StridedError::OffsetOverflow)?;
1377 if extent < 0 {
1378 min_offset = min_offset
1379 .checked_add(extent)
1380 .ok_or(StridedError::OffsetOverflow)?;
1381 } else {
1382 max_offset = max_offset
1383 .checked_add(extent)
1384 .ok_or(StridedError::OffsetOverflow)?;
1385 }
1386 }
1387 let _ = (min_offset, max_offset);
1388 Ok(())
1389}
1390fn compress_reduce_outer_axes(axes: Vec<ReduceOuterAxis>) -> Result<Vec<ReduceOuterAxis>> {
1391 let mut compressed: Vec<ReduceOuterAxis> = Vec::with_capacity(axes.len());
1392 for axis in axes {
1393 if let Some(previous) = compressed.last_mut() {
1394 let previous_extent =
1395 isize::try_from(previous.extent).map_err(|_| StridedError::OffsetOverflow)?;
1396 let expected_source = previous
1397 .source_step
1398 .checked_mul(previous_extent)
1399 .ok_or(StridedError::OffsetOverflow)?;
1400 let expected_dest = previous
1401 .dest_step
1402 .checked_mul(previous_extent)
1403 .ok_or(StridedError::OffsetOverflow)?;
1404 if axis.source_step == expected_source && axis.dest_step == expected_dest {
1405 let fused_extent = previous
1406 .extent
1407 .checked_mul(axis.extent)
1408 .ok_or(StridedError::OffsetOverflow)?;
1409 previous.extent = fused_extent;
1410 previous.source_reset = checked_reduce_reset(fused_extent, previous.source_step)?;
1411 previous.dest_reset = checked_reduce_reset(fused_extent, previous.dest_step)?;
1412 continue;
1413 }
1414 }
1415 compressed.push(axis);
1416 }
1417 Ok(compressed)
1418}
1419fn compress_reduce_inner_axes(axes: Vec<ReduceInnerAxis>) -> Result<Vec<ReduceInnerAxis>> {
1420 let mut compressed: Vec<ReduceInnerAxis> = Vec::with_capacity(axes.len());
1421 for axis in axes {
1422 if let Some(previous) = compressed.last_mut() {
1423 let previous_extent =
1424 isize::try_from(previous.extent).map_err(|_| StridedError::OffsetOverflow)?;
1425 let expected_source = previous
1426 .source_step
1427 .checked_mul(previous_extent)
1428 .ok_or(StridedError::OffsetOverflow)?;
1429 if axis.source_step == expected_source {
1430 let fused_extent = previous
1431 .extent
1432 .checked_mul(axis.extent)
1433 .ok_or(StridedError::OffsetOverflow)?;
1434 previous.extent = fused_extent;
1435 previous.source_reset = checked_reduce_reset(fused_extent, previous.source_step)?;
1436 continue;
1437 }
1438 }
1439 compressed.push(axis);
1440 }
1441 Ok(compressed)
1442}
1443fn checked_reduce_reset(extent: usize, stride: isize) -> Result<isize> {
1444 if extent == 0 {
1445 return Ok(0);
1446 }
1447 let last = isize::try_from(extent - 1).map_err(|_| StridedError::OffsetOverflow)?;
1448 stride
1449 .checked_mul(last)
1450 .and_then(isize::checked_neg)
1451 .ok_or(StridedError::OffsetOverflow)
1452}
1453fn checked_offset_add(base: isize, stride: isize, coord: usize) -> Result<isize> {
1454 let coord = isize::try_from(coord).map_err(|_| StridedError::OffsetOverflow)?;
1455 let scaled = stride
1456 .checked_mul(coord)
1457 .ok_or(StridedError::OffsetOverflow)?;
1458 base.checked_add(scaled).ok_or(StridedError::OffsetOverflow)
1459}
1460
1461struct CoordScratch {
1462 inline: [usize; RAW_FUSED_RANK_LIMIT],
1463 heap: Option<Vec<usize>>,
1464 len: usize,
1465}
1466
1467impl CoordScratch {
1468 fn new(len: usize) -> Self {
1469 if len <= RAW_FUSED_RANK_LIMIT {
1470 Self {
1471 inline: [0; RAW_FUSED_RANK_LIMIT],
1472 heap: None,
1473 len,
1474 }
1475 } else {
1476 Self {
1477 inline: [0; RAW_FUSED_RANK_LIMIT],
1478 heap: Some(vec![0; len]),
1479 len,
1480 }
1481 }
1482 }
1483
1484 fn as_mut_slice(&mut self) -> &mut [usize] {
1485 match &mut self.heap {
1486 Some(heap) => heap,
1487 None => &mut self.inline[..self.len],
1488 }
1489 }
1490}