Skip to main content

strided_basic/erased/
scan.rs

1//! Dtype-erased cumulative scans along one axis.
2
3use super::line::{for_each_unit, LineLayout, UnitKernel, UnitPtr, PANEL};
4use super::{
5    check_reduce_layout_offset_arithmetic, checked_total_len, reduce_uninit_writer, reduce_writer,
6    ReduceWriter,
7};
8use crate::erased_common::{check_dtype, validate_uninit_no_overlap};
9use crate::*;
10use num_complex::{Complex32, Complex64};
11
12/// Runtime operation of an [`ErasedScanPlan`].
13///
14/// # Examples
15///
16/// ```
17/// use strided_basic::ScanOp;
18/// assert_ne!(ScanOp::Sum, ScanOp::Product);
19/// ```
20#[non_exhaustive]
21#[derive(Clone, Copy, Debug, Eq, PartialEq)]
22pub enum ScanOp {
23    /// Cumulative sum (`cumsum`). Integers wrap on overflow.
24    Sum,
25    /// Cumulative product (`cumprod`). Integers wrap on overflow.
26    Product,
27}
28
29/// Direction and inclusivity of an [`ErasedScanPlan`].
30///
31/// The default is an inclusive forward scan: output `k` combines inputs
32/// `0..=k`. `exclusive` combines inputs `0..k` instead, so the first output is
33/// the identity (`0` for sums, `1` for products). `reverse` scans from the end
34/// of the axis: output `k` combines inputs `k..n` (inclusive) or `k+1..n`
35/// (exclusive).
36///
37/// # Examples
38///
39/// ```
40/// use strided_basic::ScanOptions;
41/// let options = ScanOptions::new().exclusive(true).reverse(true);
42/// assert!(options.is_exclusive() && options.is_reverse());
43/// assert_eq!(ScanOptions::default(), ScanOptions::new());
44/// ```
45#[non_exhaustive]
46#[derive(Clone, Copy, Debug, Default, Eq, PartialEq)]
47pub struct ScanOptions {
48    exclusive: bool,
49    reverse: bool,
50}
51
52impl ScanOptions {
53    /// An inclusive forward scan.
54    ///
55    /// # Examples
56    ///
57    /// ```
58    /// use strided_basic::ScanOptions;
59    /// assert!(!ScanOptions::new().is_exclusive());
60    /// ```
61    #[inline]
62    pub const fn new() -> Self {
63        Self {
64            exclusive: false,
65            reverse: false,
66        }
67    }
68
69    /// Select an exclusive (`true`) or inclusive (`false`) scan.
70    ///
71    /// # Examples
72    ///
73    /// ```
74    /// use strided_basic::ScanOptions;
75    /// assert!(ScanOptions::new().exclusive(true).is_exclusive());
76    /// ```
77    #[inline]
78    pub const fn exclusive(mut self, exclusive: bool) -> Self {
79        self.exclusive = exclusive;
80        self
81    }
82
83    /// Select a reverse (`true`) or forward (`false`) scan.
84    ///
85    /// # Examples
86    ///
87    /// ```
88    /// use strided_basic::ScanOptions;
89    /// assert!(ScanOptions::new().reverse(true).is_reverse());
90    /// ```
91    #[inline]
92    pub const fn reverse(mut self, reverse: bool) -> Self {
93        self.reverse = reverse;
94        self
95    }
96
97    /// Whether the scan is exclusive.
98    ///
99    /// # Examples
100    ///
101    /// ```
102    /// use strided_basic::ScanOptions;
103    /// assert!(!ScanOptions::default().is_exclusive());
104    /// ```
105    #[inline]
106    pub const fn is_exclusive(&self) -> bool {
107        self.exclusive
108    }
109
110    /// Whether the scan runs from the end of the axis.
111    ///
112    /// # Examples
113    ///
114    /// ```
115    /// use strided_basic::ScanOptions;
116    /// assert!(!ScanOptions::default().is_reverse());
117    /// ```
118    #[inline]
119    pub const fn is_reverse(&self) -> bool {
120        self.reverse
121    }
122}
123
124/// Dtype-erased cumulative scan (`cumsum` / `cumprod`) along one axis.
125///
126/// The destination has the source dimensions; source and destination strides
127/// are independent, and negative strides and offsets are accepted. Supported
128/// dtypes are `f32`, `f64`, `i32`, `i64`, `c32` and `c64`; `bool` is rejected.
129///
130/// # Evaluation order
131///
132/// Every output is the sequential left fold of its inputs in scan order (from
133/// the start of the axis, or from its end for a reverse scan), with no
134/// reassociation. The result is therefore bitwise identical for every layout,
135/// execution context and thread count. Integer scans wrap on overflow.
136///
137/// An empty axis or an empty set of lines writes nothing.
138///
139/// # Examples
140///
141/// ```
142/// use strided_basic::{
143///     ErasedRawStridedMut, ErasedRawStridedRef, ErasedScanPlan, ExecContext, KernelDType,
144///     ScanOp, ScanOptions,
145/// };
146///
147/// // A column-major 3 x 2 matrix, scanned along axis 0.
148/// let src = [1.0_f64, 2.0, 3.0, 10.0, 20.0, 30.0];
149/// let mut out = [0.0_f64; 6];
150/// let plan = ErasedScanPlan::compile(
151///     KernelDType::F64,
152///     ScanOp::Sum,
153///     &[3, 2],
154///     &[1, 3],
155///     &[1, 3],
156///     0,
157///     ScanOptions::new(),
158/// )
159/// .unwrap();
160/// let src_ref = ErasedRawStridedRef::from_slice(&src, &[3, 2], &[1, 3], 0).unwrap();
161/// let mut dest = ErasedRawStridedMut::from_slice_mut(&mut out, &[3, 2], &[1, 3], 0).unwrap();
162/// plan.execute(&ExecContext::serial(), &mut dest, &src_ref).unwrap();
163/// assert_eq!(out, [1.0, 3.0, 6.0, 10.0, 30.0, 60.0]);
164/// ```
165#[derive(Clone, Debug)]
166pub struct ErasedScanPlan {
167    dtype: KernelDType,
168    op: ScanOp,
169    options: ScanOptions,
170    dims: Vec<usize>,
171    src_strides: Vec<isize>,
172    dest_strides: Vec<isize>,
173    layout: LineLayout,
174}
175
176impl ErasedScanPlan {
177    /// Validate and store a scan plan for one dtype and fixed layouts.
178    ///
179    /// `dims` are the shared source and destination dimensions; `axis` is the
180    /// scanned axis.
181    ///
182    /// # Examples
183    ///
184    /// ```
185    /// use strided_basic::{ErasedScanPlan, KernelDType, ScanOp, ScanOptions};
186    /// let plan = ErasedScanPlan::compile(
187    ///     KernelDType::I64, ScanOp::Product, &[4], &[1], &[-1], 0, ScanOptions::new(),
188    /// )
189    /// .unwrap();
190    /// assert_eq!(plan.op(), ScanOp::Product);
191    /// assert!(ErasedScanPlan::compile(
192    ///     KernelDType::Bool, ScanOp::Sum, &[4], &[1], &[1], 0, ScanOptions::new(),
193    /// )
194    /// .is_err());
195    /// ```
196    ///
197    /// # Errors
198    ///
199    /// Returns `UnsupportedDType` for `bool`, `InvalidAxis` for an axis out of
200    /// range, `StrideLengthMismatch` for inconsistent ranks,
201    /// `NonInjectiveOutputLayout` for an aliasing destination layout, and
202    /// `OffsetOverflow` when a layout's offsets are not representable.
203    pub fn compile(
204        dtype: KernelDType,
205        op: ScanOp,
206        dims: &[usize],
207        src_strides: &[isize],
208        dest_strides: &[isize],
209        axis: usize,
210        options: ScanOptions,
211    ) -> Result<Self> {
212        check_scan_dtype(dtype)?;
213        if dims.len() != src_strides.len() || dims.len() != dest_strides.len() {
214            return Err(StridedError::StrideLengthMismatch);
215        }
216        if axis >= dims.len() {
217            return Err(StridedError::InvalidAxis {
218                axis,
219                rank: dims.len(),
220            });
221        }
222        checked_total_len(dims)?;
223        check_reduce_layout_offset_arithmetic(dims, src_strides)?;
224        check_reduce_layout_offset_arithmetic(dims, dest_strides)?;
225        if !crate::layout_check::is_injective_layout(dims, dest_strides) {
226            return Err(StridedError::NonInjectiveOutputLayout);
227        }
228        let dest_outer: Vec<isize> = (0..dims.len())
229            .filter(|&a| a != axis)
230            .map(|a| dest_strides[a])
231            .collect();
232        let layout = LineLayout::compile(
233            dims,
234            src_strides,
235            &dest_outer,
236            dest_strides[axis],
237            axis,
238            true,
239        )?;
240        Ok(Self {
241            dtype,
242            op,
243            options,
244            dims: dims.to_vec(),
245            src_strides: src_strides.to_vec(),
246            dest_strides: dest_strides.to_vec(),
247            layout,
248        })
249    }
250
251    /// Element dtype of the source and destination.
252    ///
253    /// # Examples
254    ///
255    /// ```
256    /// use strided_basic::{ErasedScanPlan, KernelDType, ScanOp, ScanOptions};
257    /// let plan = ErasedScanPlan::compile(
258    ///     KernelDType::F32, ScanOp::Sum, &[2], &[1], &[1], 0, ScanOptions::new(),
259    /// )
260    /// .unwrap();
261    /// assert_eq!(plan.dtype(), KernelDType::F32);
262    /// ```
263    #[inline]
264    pub fn dtype(&self) -> KernelDType {
265        self.dtype
266    }
267
268    /// Scan operation.
269    ///
270    /// # Examples
271    ///
272    /// ```
273    /// use strided_basic::{ErasedScanPlan, KernelDType, ScanOp, ScanOptions};
274    /// let plan = ErasedScanPlan::compile(
275    ///     KernelDType::F32, ScanOp::Sum, &[2], &[1], &[1], 0, ScanOptions::new(),
276    /// )
277    /// .unwrap();
278    /// assert_eq!(plan.op(), ScanOp::Sum);
279    /// ```
280    #[inline]
281    pub fn op(&self) -> ScanOp {
282        self.op
283    }
284
285    /// Direction and inclusivity.
286    ///
287    /// # Examples
288    ///
289    /// ```
290    /// use strided_basic::{ErasedScanPlan, KernelDType, ScanOp, ScanOptions};
291    /// let options = ScanOptions::new().reverse(true);
292    /// let plan =
293    ///     ErasedScanPlan::compile(KernelDType::F32, ScanOp::Sum, &[2], &[1], &[1], 0, options)
294    ///         .unwrap();
295    /// assert_eq!(plan.options(), options);
296    /// ```
297    #[inline]
298    pub fn options(&self) -> ScanOptions {
299        self.options
300    }
301
302    fn check_layouts(
303        &self,
304        dest_dims: &[usize],
305        dest_strides: &[isize],
306        src: &ErasedRawStridedRef<'_>,
307    ) -> Result<()> {
308        if src.dims() != self.dims.as_slice()
309            || src.strides() != self.src_strides.as_slice()
310            || dest_dims != self.dims.as_slice()
311            || dest_strides != self.dest_strides.as_slice()
312        {
313            return Err(StridedError::PlanLayoutMismatch);
314        }
315        Ok(())
316    }
317
318    /// Execute the scan into an initialized destination.
319    ///
320    /// # Examples
321    ///
322    /// ```
323    /// use strided_basic::{
324    ///     ErasedRawStridedMut, ErasedRawStridedRef, ErasedScanPlan, ExecContext, KernelDType,
325    ///     ScanOp, ScanOptions,
326    /// };
327    /// let src = [1_i32, 2, 3, 4];
328    /// let mut out = [0_i32; 4];
329    /// let options = ScanOptions::new().exclusive(true).reverse(true);
330    /// let plan =
331    ///     ErasedScanPlan::compile(KernelDType::I32, ScanOp::Sum, &[4], &[1], &[1], 0, options)
332    ///         .unwrap();
333    /// let src_ref = ErasedRawStridedRef::from_slice(&src, &[4], &[1], 0).unwrap();
334    /// let mut dest = ErasedRawStridedMut::from_slice_mut(&mut out, &[4], &[1], 0).unwrap();
335    /// plan.execute(&ExecContext::serial(), &mut dest, &src_ref).unwrap();
336    /// assert_eq!(out, [9, 7, 4, 0]);
337    /// ```
338    ///
339    /// # Errors
340    ///
341    /// Returns `DTypeMismatch` or `PlanLayoutMismatch` when a descriptor does
342    /// not match the plan, before any destination write.
343    pub fn execute(
344        &self,
345        ctx: &ExecContext,
346        dest: &mut ErasedRawStridedMut<'_>,
347        src: &ErasedRawStridedRef<'_>,
348    ) -> Result<()> {
349        check_dtype(self.dtype, dest.dtype())?;
350        check_dtype(self.dtype, src.dtype())?;
351        self.check_layouts(dest.dims(), dest.strides(), src)?;
352        macro_rules! run {
353            ($ty:ty) => {{
354                let mut writer = reduce_writer::<$ty>(dest)?;
355                self.dispatch::<$ty, _>(ctx, &mut writer, src)
356            }};
357        }
358        match self.dtype {
359            KernelDType::F32 => run!(f32),
360            KernelDType::F64 => run!(f64),
361            KernelDType::I32 => run!(i32),
362            KernelDType::I64 => run!(i64),
363            KernelDType::C32 => run!(Complex32),
364            KernelDType::C64 => run!(Complex64),
365            _ => Err(unsupported(self.dtype)),
366        }
367    }
368
369    /// Execute the scan into an uninitialized destination.
370    ///
371    /// On success every reachable destination element is written; validation
372    /// errors are returned before any write.
373    ///
374    /// # Examples
375    ///
376    /// ```
377    /// use core::mem::MaybeUninit;
378    /// use strided_basic::{
379    ///     ErasedRawStridedPtr, ErasedRawStridedRef, ErasedRawStridedUninitMut, ErasedScanPlan,
380    ///     ExecContext, KernelDType, ScanOp, ScanOptions,
381    /// };
382    /// let src = [2.0_f32, 3.0, 4.0];
383    /// let mut out = [MaybeUninit::<f32>::uninit(); 3];
384    /// let plan = ErasedScanPlan::compile(
385    ///     KernelDType::F32, ScanOp::Product, &[3], &[1], &[1], 0, ScanOptions::new(),
386    /// )
387    /// .unwrap();
388    /// let src_ref = ErasedRawStridedRef::from_slice(&src, &[3], &[1], 0).unwrap();
389    /// let src_ptr = ErasedRawStridedPtr::from_ref(&src_ref);
390    /// let mut dest =
391    ///     ErasedRawStridedUninitMut::from_uninit_slice(&mut out, &[3], &[1], 0).unwrap();
392    /// plan.execute_uninit(&ExecContext::serial(), &mut dest, &src_ptr).unwrap();
393    /// let out: Vec<f32> = out.iter().map(|v| unsafe { v.assume_init() }).collect();
394    /// assert_eq!(out, [2.0, 6.0, 24.0]);
395    /// ```
396    ///
397    /// # Errors
398    ///
399    /// As [`Self::execute`], plus `OverlappingInputOutput` when the source
400    /// overlaps the destination allocation.
401    pub fn execute_uninit(
402        &self,
403        ctx: &ExecContext,
404        dest: &mut ErasedRawStridedUninitMut<'_>,
405        src: &ErasedRawStridedPtr<'_>,
406    ) -> Result<()> {
407        check_dtype(self.dtype, dest.dtype())?;
408        check_dtype(self.dtype, src.dtype())?;
409        validate_uninit_no_overlap(dest, src, 0)?;
410        // SAFETY: the owning erased entry rejected all input/output overlap before conversion.
411        let src = unsafe { src.try_as_ref_after_no_overlap() }?;
412        self.check_layouts(dest.dims(), dest.strides(), &src)?;
413        macro_rules! run {
414            ($ty:ty) => {{
415                let mut writer = reduce_uninit_writer::<$ty>(dest)?;
416                self.dispatch::<$ty, _>(ctx, &mut writer, &src)
417            }};
418        }
419        match self.dtype {
420            KernelDType::F32 => run!(f32),
421            KernelDType::F64 => run!(f64),
422            KernelDType::I32 => run!(i32),
423            KernelDType::I64 => run!(i64),
424            KernelDType::C32 => run!(Complex32),
425            KernelDType::C64 => run!(Complex64),
426            _ => Err(unsupported(self.dtype)),
427        }
428    }
429
430    fn dispatch<T, W>(
431        &self,
432        ctx: &ExecContext,
433        dest: &mut W,
434        src: &ErasedRawStridedRef<'_>,
435    ) -> Result<()>
436    where
437        T: ScanScalar,
438        W: ReduceWriter<T>,
439    {
440        // The op and options are matched once per execution; the loops below
441        // are monomorphized for each combination.
442        match (self.op, self.options.exclusive, self.options.reverse) {
443            (ScanOp::Sum, false, false) => self.run::<T, W, SumScan, false, false>(ctx, dest, src),
444            (ScanOp::Sum, false, true) => self.run::<T, W, SumScan, false, true>(ctx, dest, src),
445            (ScanOp::Sum, true, false) => self.run::<T, W, SumScan, true, false>(ctx, dest, src),
446            (ScanOp::Sum, true, true) => self.run::<T, W, SumScan, true, true>(ctx, dest, src),
447            (ScanOp::Product, false, false) => {
448                self.run::<T, W, ProductScan, false, false>(ctx, dest, src)
449            }
450            (ScanOp::Product, false, true) => {
451                self.run::<T, W, ProductScan, false, true>(ctx, dest, src)
452            }
453            (ScanOp::Product, true, false) => {
454                self.run::<T, W, ProductScan, true, false>(ctx, dest, src)
455            }
456            (ScanOp::Product, true, true) => {
457                self.run::<T, W, ProductScan, true, true>(ctx, dest, src)
458            }
459        }
460    }
461
462    fn run<T, W, K, const EXCLUSIVE: bool, const REVERSE: bool>(
463        &self,
464        ctx: &ExecContext,
465        dest: &mut W,
466        src: &ErasedRawStridedRef<'_>,
467    ) -> Result<()>
468    where
469        T: ScanScalar,
470        W: ReduceWriter<T>,
471        K: ScanKernel<T>,
472    {
473        let layout = &self.layout;
474        let source = UnitPtr(src.data_as::<T>()?.as_ptr() as *mut T);
475        // SAFETY: the validated writer owns the destination allocation.
476        let target = UnitPtr(unsafe { dest.ptr() });
477        let n = layout.axis_len;
478        let ss = layout.src_axis_stride;
479        let ds = layout.dest_axis_stride;
480        let kernel = ScanUnit::<T, K, EXCLUSIVE, REVERSE> {
481            source,
482            target,
483            n,
484            ss,
485            ds,
486            _kernel: core::marker::PhantomData,
487        };
488        // INVARIANT: (1) compile checked the signed source/destination spans
489        // and every cursor step/reset; (2) the raw descriptors validated every
490        // reachable offset; (3) execute checked exact plan-layout equality.
491        // Units cover disjoint lines, so their destination writes are disjoint.
492        // SAFETY: the three-link invariant above.
493        unsafe { for_each_unit(ctx, layout, src.offset(), dest.offset(), kernel) }
494    }
495}
496
497/// Per-unit scan kernel of one execution.
498struct ScanUnit<T, K, const EXCLUSIVE: bool, const REVERSE: bool> {
499    source: UnitPtr<T>,
500    target: UnitPtr<T>,
501    n: usize,
502    ss: isize,
503    ds: isize,
504    _kernel: core::marker::PhantomData<fn() -> K>,
505}
506
507impl<T, K, const EXCLUSIVE: bool, const REVERSE: bool> Clone
508    for ScanUnit<T, K, EXCLUSIVE, REVERSE>
509{
510    fn clone(&self) -> Self {
511        *self
512    }
513}
514impl<T, K, const EXCLUSIVE: bool, const REVERSE: bool> Copy for ScanUnit<T, K, EXCLUSIVE, REVERSE> {}
515
516impl<T, K, const EXCLUSIVE: bool, const REVERSE: bool> UnitKernel
517    for ScanUnit<T, K, EXCLUSIVE, REVERSE>
518where
519    T: ScanScalar,
520    K: ScanKernel<T>,
521{
522    #[inline(always)]
523    unsafe fn unit(self, so: isize, d_o: isize, width: usize) {
524        let Self {
525            source,
526            target,
527            n,
528            ss,
529            ds,
530            ..
531        } = self;
532        // SAFETY: the caller passes offsets of the validated layout.
533        unsafe {
534            if width == 1 {
535                scan_line::<T, K, EXCLUSIVE, REVERSE>(
536                    source.get(),
537                    so,
538                    ss,
539                    target.get(),
540                    d_o,
541                    ds,
542                    n,
543                )
544            } else {
545                scan_panel::<T, K, EXCLUSIVE, REVERSE>(
546                    source.get(),
547                    so,
548                    ss,
549                    target.get(),
550                    d_o,
551                    ds,
552                    n,
553                    width,
554                )
555            }
556        }
557    }
558}
559
560fn unsupported(dtype: KernelDType) -> StridedError {
561    StridedError::UnsupportedDType {
562        dtype: dtype.label(),
563    }
564}
565
566fn check_scan_dtype(dtype: KernelDType) -> Result<()> {
567    match dtype {
568        KernelDType::F32
569        | KernelDType::F64
570        | KernelDType::I32
571        | KernelDType::I64
572        | KernelDType::C32
573        | KernelDType::C64 => Ok(()),
574        _ => Err(unsupported(dtype)),
575    }
576}
577
578/// Element type supported by scans.
579pub(super) trait ScanScalar: KernelStorageElement + MaybeSendSync {
580    fn zero() -> Self;
581    fn one() -> Self;
582    fn scan_add(lhs: Self, rhs: Self) -> Self;
583    fn scan_mul(lhs: Self, rhs: Self) -> Self;
584}
585
586macro_rules! impl_scan_scalar {
587    ($add:ident, $mul:ident; $($ty:ty => $zero:expr, $one:expr),* $(,)?) => {$(
588        impl ScanScalar for $ty {
589            #[inline(always)]
590            fn zero() -> Self { $zero }
591            #[inline(always)]
592            fn one() -> Self { $one }
593            #[inline(always)]
594            fn scan_add(lhs: Self, rhs: Self) -> Self { $add(lhs, rhs) }
595            #[inline(always)]
596            fn scan_mul(lhs: Self, rhs: Self) -> Self { $mul(lhs, rhs) }
597        }
598    )*};
599}
600
601#[inline(always)]
602fn plain_add<T: core::ops::Add<Output = T>>(lhs: T, rhs: T) -> T {
603    lhs + rhs
604}
605#[inline(always)]
606fn plain_mul<T: core::ops::Mul<Output = T>>(lhs: T, rhs: T) -> T {
607    lhs * rhs
608}
609trait WrappingArith {
610    fn wadd(self, rhs: Self) -> Self;
611    fn wmul(self, rhs: Self) -> Self;
612}
613impl WrappingArith for i32 {
614    #[inline(always)]
615    fn wadd(self, rhs: Self) -> Self {
616        self.wrapping_add(rhs)
617    }
618    #[inline(always)]
619    fn wmul(self, rhs: Self) -> Self {
620        self.wrapping_mul(rhs)
621    }
622}
623impl WrappingArith for i64 {
624    #[inline(always)]
625    fn wadd(self, rhs: Self) -> Self {
626        self.wrapping_add(rhs)
627    }
628    #[inline(always)]
629    fn wmul(self, rhs: Self) -> Self {
630        self.wrapping_mul(rhs)
631    }
632}
633#[inline(always)]
634fn wrapping_add<T: WrappingArith>(lhs: T, rhs: T) -> T {
635    lhs.wadd(rhs)
636}
637#[inline(always)]
638fn wrapping_mul<T: WrappingArith>(lhs: T, rhs: T) -> T {
639    lhs.wmul(rhs)
640}
641
642impl_scan_scalar!(plain_add, plain_mul;
643    f32 => 0.0, 1.0,
644    f64 => 0.0, 1.0,
645    Complex32 => Complex32::new(0.0, 0.0), Complex32::new(1.0, 0.0),
646    Complex64 => Complex64::new(0.0, 0.0), Complex64::new(1.0, 0.0),
647);
648impl_scan_scalar!(wrapping_add, wrapping_mul;
649    i32 => 0, 1,
650    i64 => 0, 1,
651);
652
653/// One scan operation, fixed at compile time.
654pub(super) trait ScanKernel<T>: 'static {
655    fn identity() -> T;
656    fn combine(acc: T, value: T) -> T;
657}
658
659pub(super) struct SumScan;
660pub(super) struct ProductScan;
661
662impl<T: ScanScalar> ScanKernel<T> for SumScan {
663    #[inline(always)]
664    fn identity() -> T {
665        T::zero()
666    }
667    #[inline(always)]
668    fn combine(acc: T, value: T) -> T {
669        T::scan_add(acc, value)
670    }
671}
672
673impl<T: ScanScalar> ScanKernel<T> for ProductScan {
674    #[inline(always)]
675    fn identity() -> T {
676        T::one()
677    }
678    #[inline(always)]
679    fn combine(acc: T, value: T) -> T {
680        T::scan_mul(acc, value)
681    }
682}
683
684/// Start offset and signed step of a line visited in scan order.
685#[inline(always)]
686fn scan_order(base: isize, stride: isize, n: usize, reverse: bool) -> (isize, isize) {
687    if reverse {
688        // INVARIANT: compile checked (n - 1) * stride for this axis.
689        (base + (n as isize - 1) * stride, -stride)
690    } else {
691        (base, stride)
692    }
693}
694
695/// Scans one line.
696///
697/// # Safety
698///
699/// Every `src.offset(so + k * ss)` and `dst.offset(d_o + k * ds)` for
700/// `k < n` must be valid, and `n > 0`.
701#[inline(always)]
702unsafe fn scan_line<T, K, const EXCLUSIVE: bool, const REVERSE: bool>(
703    src: *const T,
704    so: isize,
705    ss: isize,
706    dst: *mut T,
707    d_o: isize,
708    ds: isize,
709    n: usize,
710) where
711    T: ScanScalar,
712    K: ScanKernel<T>,
713{
714    let (mut s, ss) = scan_order(so, ss, n, REVERSE);
715    let (mut d, ds) = scan_order(d_o, ds, n, REVERSE);
716    let mut acc = K::identity();
717    for _ in 0..n {
718        // SAFETY: the caller guarantees every visited offset.
719        unsafe {
720            let value = src.offset(s).read();
721            if EXCLUSIVE {
722                dst.offset(d).write(acc);
723                acc = K::combine(acc, value);
724            } else {
725                acc = K::combine(acc, value);
726                dst.offset(d).write(acc);
727            }
728        }
729        s += ss;
730        d += ds;
731    }
732}
733
734/// Scans `width` adjacent lines whose elements are contiguous across the
735/// lines in both source and destination.
736///
737/// # Safety
738///
739/// As [`scan_line`] for each line `j < width`, with source base `so + j` and
740/// destination base `d_o + j`; `width <= PANEL`.
741#[allow(clippy::too_many_arguments)]
742#[inline(always)]
743unsafe fn scan_panel<T, K, const EXCLUSIVE: bool, const REVERSE: bool>(
744    src: *const T,
745    so: isize,
746    ss: isize,
747    dst: *mut T,
748    d_o: isize,
749    ds: isize,
750    n: usize,
751    width: usize,
752) where
753    T: ScanScalar,
754    K: ScanKernel<T>,
755{
756    debug_assert!(width <= PANEL);
757    let (mut s, ss) = scan_order(so, ss, n, REVERSE);
758    let (mut d, ds) = scan_order(d_o, ds, n, REVERSE);
759    let mut acc = [K::identity(); PANEL];
760    let acc = &mut acc[..width];
761    for _ in 0..n {
762        // SAFETY: the caller guarantees `width` contiguous elements at both
763        // offsets for every visited axis position.
764        unsafe {
765            let input = core::slice::from_raw_parts(src.offset(s), width);
766            // Raw writes: the destination may be uninitialized.
767            let output = dst.offset(d);
768            for (lane, (acc, &value)) in acc.iter_mut().zip(input).enumerate() {
769                if EXCLUSIVE {
770                    output.add(lane).write(*acc);
771                    *acc = K::combine(*acc, value);
772                } else {
773                    *acc = K::combine(*acc, value);
774                    output.add(lane).write(*acc);
775                }
776            }
777        }
778        s += ss;
779        d += ds;
780    }
781}