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}