strided_basic/erased/arg_reduce.rs
1//! Dtype-erased `argmax` / `argmin` 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 [`ErasedArgReducePlan`].
13///
14/// # Examples
15///
16/// ```
17/// use strided_basic::ArgReduceOp;
18/// assert_ne!(ArgReduceOp::Max, ArgReduceOp::MaxAbs);
19/// ```
20#[non_exhaustive]
21#[derive(Clone, Copy, Debug, Eq, PartialEq)]
22pub enum ArgReduceOp {
23 /// Index of the largest value. Real dtypes only.
24 Max,
25 /// Index of the smallest value. Real dtypes only.
26 Min,
27 /// Index of the largest magnitude (`|x|`, the modulus for complex input).
28 MaxAbs,
29 /// Index of the smallest magnitude (`|x|`, the modulus for complex input).
30 MinAbs,
31}
32
33impl ArgReduceOp {
34 #[inline]
35 fn label(self) -> &'static str {
36 match self {
37 Self::Max => "argmax",
38 Self::Min => "argmin",
39 Self::MaxAbs => "argmax_abs",
40 Self::MinAbs => "argmin_abs",
41 }
42 }
43}
44
45/// Dtype-erased `argmax` / `argmin` along one axis.
46///
47/// The destination holds one index per line: its dimensions are the source
48/// dimensions with `axis` removed (batch axes are preserved in order), and its
49/// dtype is the index dtype, `i32` or `i64`. Indices count from the start of
50/// the logical axis, whatever the sign of its stride.
51///
52/// Supported source dtypes are `f32`, `f64`, `i32` and `i64` for every
53/// [`ArgReduceOp`], and `c32` / `c64` for the magnitude variants
54/// ([`ArgReduceOp::MaxAbs`], [`ArgReduceOp::MinAbs`]) only.
55///
56/// # Ties, NaN and magnitudes
57///
58/// * Ties resolve to the **lowest index**. `-0.0` and `+0.0` compare equal,
59/// so they tie.
60/// * NaN propagates like [`ReduceOp::Max`]: if a line contains a NaN, the
61/// result is the index of its **first** NaN, for both the max and the min
62/// variants. A complex element is NaN when either component is NaN. This
63/// matches NumPy, PyTorch and JAX.
64/// * Infinities order normally; `+inf` beats every finite value for `Max`.
65/// * The complex magnitude is the overflow-safe modulus (`hypot` of the
66/// components), so values near the top of the exponent range still order
67/// correctly. A component of `±inf` makes the magnitude `+inf` unless the
68/// other component is NaN.
69/// * The integer magnitude is the unsigned absolute value, so `i32::MIN` is
70/// larger than `i32::MAX`.
71///
72/// The result does not depend on the layout, execution context or thread
73/// count.
74///
75/// # Examples
76///
77/// ```
78/// use strided_basic::{
79/// ArgReduceOp, ErasedArgReducePlan, ErasedRawStridedMut, ErasedRawStridedRef, ExecContext,
80/// KernelDType,
81/// };
82///
83/// // Column-major 3 x 2 matrix; argmax down each column.
84/// let src = [1.0_f64, 5.0, 5.0, -2.0, -7.0, 3.0];
85/// let mut out = [0_i64; 2];
86/// let plan = ErasedArgReducePlan::compile(
87/// KernelDType::F64,
88/// KernelDType::I64,
89/// ArgReduceOp::Max,
90/// &[3, 2],
91/// &[1, 3],
92/// &[2],
93/// &[1],
94/// 0,
95/// )
96/// .unwrap();
97/// let src_ref = ErasedRawStridedRef::from_slice(&src, &[3, 2], &[1, 3], 0).unwrap();
98/// let mut dest = ErasedRawStridedMut::from_slice_mut(&mut out, &[2], &[1], 0).unwrap();
99/// plan.execute(&ExecContext::serial(), &mut dest, &src_ref).unwrap();
100/// assert_eq!(out, [1, 2]); // lowest index wins the 5.0 tie
101/// ```
102#[derive(Clone, Debug)]
103pub struct ErasedArgReducePlan {
104 dtype: KernelDType,
105 index_dtype: KernelDType,
106 op: ArgReduceOp,
107 src_dims: Vec<usize>,
108 src_strides: Vec<isize>,
109 dest_dims: Vec<usize>,
110 dest_strides: Vec<isize>,
111 layout: LineLayout,
112}
113
114impl ErasedArgReducePlan {
115 /// Validate and store an arg-reduction plan for fixed layouts.
116 ///
117 /// `dest_dims` must equal `src_dims` with `axis` removed.
118 ///
119 /// # Examples
120 ///
121 /// ```
122 /// use strided_basic::{ArgReduceOp, ErasedArgReducePlan, KernelDType};
123 /// let plan = ErasedArgReducePlan::compile(
124 /// KernelDType::C64, KernelDType::I32, ArgReduceOp::MaxAbs, &[4, 3], &[3, 1], &[3], &[1], 0,
125 /// )
126 /// .unwrap();
127 /// assert_eq!(plan.index_dtype(), KernelDType::I32);
128 /// // Complex input has no order without the magnitude.
129 /// assert!(ErasedArgReducePlan::compile(
130 /// KernelDType::C64, KernelDType::I64, ArgReduceOp::Max, &[4], &[1], &[], &[], 0,
131 /// )
132 /// .is_err());
133 /// // An empty axis has no index to return.
134 /// assert!(ErasedArgReducePlan::compile(
135 /// KernelDType::F64, KernelDType::I64, ArgReduceOp::Max, &[0, 2], &[1, 1], &[2], &[1], 0,
136 /// )
137 /// .is_err());
138 /// ```
139 ///
140 /// # Errors
141 ///
142 /// * `UnsupportedDType` for a `bool` source or an index dtype other than
143 /// `i32` / `i64`;
144 /// * `UnsupportedOp` for `Max` / `Min` on complex input, and for a
145 /// zero-length `axis`, which has no index to return;
146 /// * `InvalidAxis`, `StrideLengthMismatch`, `ShapeMismatch` and
147 /// `NonInjectiveOutputLayout` for inconsistent layouts;
148 /// * `OffsetOverflow` when a layout's offsets are not representable, or
149 /// when the axis is too long for an `i32` index.
150 #[allow(clippy::too_many_arguments)]
151 pub fn compile(
152 dtype: KernelDType,
153 index_dtype: KernelDType,
154 op: ArgReduceOp,
155 src_dims: &[usize],
156 src_strides: &[isize],
157 dest_dims: &[usize],
158 dest_strides: &[isize],
159 axis: usize,
160 ) -> Result<Self> {
161 check_arg_dtype(dtype, op)?;
162 if !matches!(index_dtype, KernelDType::I32 | KernelDType::I64) {
163 return Err(StridedError::UnsupportedDType {
164 dtype: index_dtype.label(),
165 });
166 }
167 if src_dims.len() != src_strides.len() || dest_dims.len() != dest_strides.len() {
168 return Err(StridedError::StrideLengthMismatch);
169 }
170 let rank = src_dims.len();
171 if axis >= rank {
172 return Err(StridedError::InvalidAxis { axis, rank });
173 }
174 let expected: Vec<usize> = (0..rank)
175 .filter(|&a| a != axis)
176 .map(|a| src_dims[a])
177 .collect();
178 if dest_dims != expected.as_slice() {
179 return Err(StridedError::ShapeMismatch(dest_dims.to_vec(), expected));
180 }
181 if src_dims[axis] == 0 {
182 return Err(StridedError::UnsupportedOp {
183 op: "argmax/argmin over a zero-length axis",
184 dtype: dtype.label(),
185 });
186 }
187 if index_dtype == KernelDType::I32 && i32::try_from(src_dims[axis] - 1).is_err() {
188 return Err(StridedError::OffsetOverflow);
189 }
190 checked_total_len(src_dims)?;
191 check_reduce_layout_offset_arithmetic(src_dims, src_strides)?;
192 check_reduce_layout_offset_arithmetic(dest_dims, dest_strides)?;
193 if !crate::layout_check::is_injective_layout(dest_dims, dest_strides) {
194 return Err(StridedError::NonInjectiveOutputLayout);
195 }
196 let layout = LineLayout::compile(src_dims, src_strides, dest_strides, 0, axis, false)?;
197 Ok(Self {
198 dtype,
199 index_dtype,
200 op,
201 src_dims: src_dims.to_vec(),
202 src_strides: src_strides.to_vec(),
203 dest_dims: dest_dims.to_vec(),
204 dest_strides: dest_strides.to_vec(),
205 layout,
206 })
207 }
208
209 /// Source element dtype.
210 ///
211 /// # Examples
212 ///
213 /// ```
214 /// use strided_basic::{ArgReduceOp, ErasedArgReducePlan, KernelDType};
215 /// let plan = ErasedArgReducePlan::compile(
216 /// KernelDType::F32, KernelDType::I64, ArgReduceOp::Min, &[2], &[1], &[], &[], 0,
217 /// )
218 /// .unwrap();
219 /// assert_eq!(plan.dtype(), KernelDType::F32);
220 /// ```
221 #[inline]
222 pub fn dtype(&self) -> KernelDType {
223 self.dtype
224 }
225
226 /// Destination index dtype (`i32` or `i64`).
227 ///
228 /// # Examples
229 ///
230 /// ```
231 /// use strided_basic::{ArgReduceOp, ErasedArgReducePlan, KernelDType};
232 /// let plan = ErasedArgReducePlan::compile(
233 /// KernelDType::F32, KernelDType::I64, ArgReduceOp::Min, &[2], &[1], &[], &[], 0,
234 /// )
235 /// .unwrap();
236 /// assert_eq!(plan.index_dtype(), KernelDType::I64);
237 /// ```
238 #[inline]
239 pub fn index_dtype(&self) -> KernelDType {
240 self.index_dtype
241 }
242
243 /// Arg-reduction operation.
244 ///
245 /// # Examples
246 ///
247 /// ```
248 /// use strided_basic::{ArgReduceOp, ErasedArgReducePlan, KernelDType};
249 /// let plan = ErasedArgReducePlan::compile(
250 /// KernelDType::F32, KernelDType::I64, ArgReduceOp::Min, &[2], &[1], &[], &[], 0,
251 /// )
252 /// .unwrap();
253 /// assert_eq!(plan.op(), ArgReduceOp::Min);
254 /// ```
255 #[inline]
256 pub fn op(&self) -> ArgReduceOp {
257 self.op
258 }
259
260 fn check_layouts(
261 &self,
262 dest_dims: &[usize],
263 dest_strides: &[isize],
264 src: &ErasedRawStridedRef<'_>,
265 ) -> Result<()> {
266 if src.dims() != self.src_dims.as_slice()
267 || src.strides() != self.src_strides.as_slice()
268 || dest_dims != self.dest_dims.as_slice()
269 || dest_strides != self.dest_strides.as_slice()
270 {
271 return Err(StridedError::PlanLayoutMismatch);
272 }
273 Ok(())
274 }
275
276 /// Execute into an initialized index destination.
277 ///
278 /// # Examples
279 ///
280 /// ```
281 /// use strided_basic::{
282 /// ArgReduceOp, ErasedArgReducePlan, ErasedRawStridedMut, ErasedRawStridedRef,
283 /// ExecContext, KernelDType,
284 /// };
285 /// let src = [3_i32, -9, 9, 1];
286 /// let mut out = [0_i32; 1];
287 /// let plan = ErasedArgReducePlan::compile(
288 /// KernelDType::I32, KernelDType::I32, ArgReduceOp::MaxAbs, &[4], &[1], &[], &[], 0,
289 /// )
290 /// .unwrap();
291 /// let src_ref = ErasedRawStridedRef::from_slice(&src, &[4], &[1], 0).unwrap();
292 /// let mut dest = ErasedRawStridedMut::from_slice_mut(&mut out, &[], &[], 0).unwrap();
293 /// plan.execute(&ExecContext::serial(), &mut dest, &src_ref).unwrap();
294 /// assert_eq!(out, [1]);
295 /// ```
296 ///
297 /// # Errors
298 ///
299 /// Returns `DTypeMismatch` or `PlanLayoutMismatch` when a descriptor does
300 /// not match the plan, before any destination write.
301 pub fn execute(
302 &self,
303 ctx: &ExecContext,
304 dest: &mut ErasedRawStridedMut<'_>,
305 src: &ErasedRawStridedRef<'_>,
306 ) -> Result<()> {
307 check_dtype(self.index_dtype, dest.dtype())?;
308 check_dtype(self.dtype, src.dtype())?;
309 self.check_layouts(dest.dims(), dest.strides(), src)?;
310 match self.index_dtype {
311 KernelDType::I32 => {
312 let mut writer = reduce_writer::<i32>(dest)?;
313 self.dispatch_dtype(ctx, &mut writer, src)
314 }
315 _ => {
316 let mut writer = reduce_writer::<i64>(dest)?;
317 self.dispatch_dtype(ctx, &mut writer, src)
318 }
319 }
320 }
321
322 /// Execute into an uninitialized index destination.
323 ///
324 /// On success every reachable destination element is written; validation
325 /// errors are returned before any write.
326 ///
327 /// # Examples
328 ///
329 /// ```
330 /// use core::mem::MaybeUninit;
331 /// use num_complex::Complex64;
332 /// use strided_basic::{
333 /// ArgReduceOp, ErasedArgReducePlan, ErasedRawStridedPtr, ErasedRawStridedRef,
334 /// ErasedRawStridedUninitMut, ExecContext, KernelDType,
335 /// };
336 /// let src = [Complex64::new(3.0, 4.0), Complex64::new(0.0, 6.0)];
337 /// let mut out = [MaybeUninit::<i64>::uninit()];
338 /// let plan = ErasedArgReducePlan::compile(
339 /// KernelDType::C64, KernelDType::I64, ArgReduceOp::MinAbs, &[2], &[1], &[], &[], 0,
340 /// )
341 /// .unwrap();
342 /// let src_ref = ErasedRawStridedRef::from_slice(&src, &[2], &[1], 0).unwrap();
343 /// let src_ptr = ErasedRawStridedPtr::from_ref(&src_ref);
344 /// let mut dest = ErasedRawStridedUninitMut::from_uninit_slice(&mut out, &[], &[], 0).unwrap();
345 /// plan.execute_uninit(&ExecContext::serial(), &mut dest, &src_ptr).unwrap();
346 /// assert_eq!(unsafe { out[0].assume_init() }, 0);
347 /// ```
348 ///
349 /// # Errors
350 ///
351 /// As [`Self::execute`], plus `OverlappingInputOutput` when the source
352 /// overlaps the destination allocation.
353 pub fn execute_uninit(
354 &self,
355 ctx: &ExecContext,
356 dest: &mut ErasedRawStridedUninitMut<'_>,
357 src: &ErasedRawStridedPtr<'_>,
358 ) -> Result<()> {
359 check_dtype(self.index_dtype, dest.dtype())?;
360 check_dtype(self.dtype, src.dtype())?;
361 validate_uninit_no_overlap(dest, src, 0)?;
362 // SAFETY: the owning erased entry rejected all input/output overlap before conversion.
363 let src = unsafe { src.try_as_ref_after_no_overlap() }?;
364 self.check_layouts(dest.dims(), dest.strides(), &src)?;
365 match self.index_dtype {
366 KernelDType::I32 => {
367 let mut writer = reduce_uninit_writer::<i32>(dest)?;
368 self.dispatch_dtype(ctx, &mut writer, &src)
369 }
370 _ => {
371 let mut writer = reduce_uninit_writer::<i64>(dest)?;
372 self.dispatch_dtype(ctx, &mut writer, &src)
373 }
374 }
375 }
376
377 fn dispatch_dtype<I, W>(
378 &self,
379 ctx: &ExecContext,
380 dest: &mut W,
381 src: &ErasedRawStridedRef<'_>,
382 ) -> Result<()>
383 where
384 I: ArgIndex,
385 W: ReduceWriter<I>,
386 {
387 // The dtype and op are matched once per execution; each loop is
388 // monomorphized for one (dtype, key, direction, index) combination.
389 macro_rules! real {
390 ($ty:ty) => {
391 match self.op {
392 ArgReduceOp::Max => self.run::<$ty, I, W, Plain, Greater>(ctx, dest, src),
393 ArgReduceOp::Min => self.run::<$ty, I, W, Plain, Less>(ctx, dest, src),
394 ArgReduceOp::MaxAbs => {
395 self.run::<$ty, I, W, Magnitude, Greater>(ctx, dest, src)
396 }
397 ArgReduceOp::MinAbs => self.run::<$ty, I, W, Magnitude, Less>(ctx, dest, src),
398 }
399 };
400 }
401 macro_rules! complex {
402 ($ty:ty) => {
403 match self.op {
404 ArgReduceOp::MaxAbs => {
405 self.run::<$ty, I, W, Magnitude, Greater>(ctx, dest, src)
406 }
407 ArgReduceOp::MinAbs => self.run::<$ty, I, W, Magnitude, Less>(ctx, dest, src),
408 // INVARIANT: check_arg_dtype rejects ordered complex ops at compile.
409 ArgReduceOp::Max | ArgReduceOp::Min => Err(StridedError::UnsupportedOp {
410 op: self.op.label(),
411 dtype: self.dtype.label(),
412 }),
413 }
414 };
415 }
416 match self.dtype {
417 KernelDType::F32 => real!(f32),
418 KernelDType::F64 => real!(f64),
419 KernelDType::I32 => real!(i32),
420 KernelDType::I64 => real!(i64),
421 KernelDType::C32 => complex!(Complex32),
422 KernelDType::C64 => complex!(Complex64),
423 _ => Err(StridedError::UnsupportedDType {
424 dtype: self.dtype.label(),
425 }),
426 }
427 }
428
429 fn run<T, I, W, F, D>(
430 &self,
431 ctx: &ExecContext,
432 dest: &mut W,
433 src: &ErasedRawStridedRef<'_>,
434 ) -> Result<()>
435 where
436 T: KernelStorageElement + MaybeSendSync,
437 I: ArgIndex,
438 W: ReduceWriter<I>,
439 F: KeyFn<T>,
440 D: Direction,
441 {
442 let layout = &self.layout;
443 let source = UnitPtr(src.data_as::<T>()?.as_ptr() as *mut T);
444 // SAFETY: the validated writer owns the destination allocation.
445 let target = UnitPtr(unsafe { dest.ptr() });
446 let n = layout.axis_len;
447 let ss = layout.src_axis_stride;
448 let lane = layout.dest_lane_stride;
449 let kernel = ArgUnit::<T, I, F, D> {
450 source,
451 target,
452 n,
453 ss,
454 lane,
455 _ops: core::marker::PhantomData,
456 };
457 // INVARIANT: (1) compile checked the signed source/destination spans
458 // and every cursor step/reset, and rejected an empty axis; (2) the raw
459 // descriptors validated every reachable offset; (3) execute checked
460 // exact plan-layout equality. Units own disjoint outputs.
461 // SAFETY: the three-link invariant above.
462 unsafe { for_each_unit(ctx, layout, src.offset(), dest.offset(), kernel) }
463 }
464}
465
466/// Per-unit arg-reduction kernel of one execution.
467struct ArgUnit<T, I, F, D> {
468 source: UnitPtr<T>,
469 target: UnitPtr<I>,
470 n: usize,
471 ss: isize,
472 lane: isize,
473 _ops: core::marker::PhantomData<fn() -> (F, D)>,
474}
475
476impl<T, I, F, D> Clone for ArgUnit<T, I, F, D> {
477 fn clone(&self) -> Self {
478 *self
479 }
480}
481impl<T, I, F, D> Copy for ArgUnit<T, I, F, D> {}
482
483impl<T, I, F, D> UnitKernel for ArgUnit<T, I, F, D>
484where
485 T: KernelStorageElement + MaybeSendSync,
486 I: ArgIndex,
487 F: KeyFn<T>,
488 D: Direction,
489{
490 #[inline(always)]
491 unsafe fn unit(self, so: isize, d_o: isize, width: usize) {
492 let Self {
493 source,
494 target,
495 n,
496 ss,
497 lane,
498 ..
499 } = self;
500 // SAFETY: the caller passes offsets of the validated layout.
501 unsafe {
502 if width == 1 {
503 let index = arg_line::<T, F, D>(source.get(), so, ss, n);
504 target.get().offset(d_o).write(I::from_index(index));
505 } else {
506 arg_panel::<T, I, F, D>(source.get(), so, ss, n, target.get(), d_o, lane, width);
507 }
508 }
509 }
510}
511
512fn check_arg_dtype(dtype: KernelDType, op: ArgReduceOp) -> Result<()> {
513 match dtype {
514 KernelDType::F32 | KernelDType::F64 | KernelDType::I32 | KernelDType::I64 => Ok(()),
515 KernelDType::C32 | KernelDType::C64 => match op {
516 ArgReduceOp::MaxAbs | ArgReduceOp::MinAbs => Ok(()),
517 ArgReduceOp::Max | ArgReduceOp::Min => Err(StridedError::UnsupportedOp {
518 op: op.label(),
519 dtype: dtype.label(),
520 }),
521 },
522 _ => Err(StridedError::UnsupportedDType {
523 dtype: dtype.label(),
524 }),
525 }
526}
527
528/// Index element written to the destination.
529pub(super) trait ArgIndex: KernelStorageElement + MaybeSendSync {
530 fn from_index(index: usize) -> Self;
531}
532
533impl ArgIndex for i32 {
534 #[inline(always)]
535 fn from_index(index: usize) -> Self {
536 // INVARIANT: compile rejected axes whose last index exceeds i32::MAX.
537 index as i32
538 }
539}
540
541impl ArgIndex for i64 {
542 #[inline(always)]
543 fn from_index(index: usize) -> Self {
544 // INVARIANT: an addressable axis length fits in i64.
545 index as i64
546 }
547}
548
549/// Comparable key of an element.
550pub(super) trait OrderKey: Copy + PartialEq {
551 /// Whether the key is NaN (never for integer keys).
552 fn key_is_nan(self) -> bool;
553 /// Whether `candidate` replaces `best` for a maximum: strictly greater,
554 /// or the first NaN.
555 fn beats_max(candidate: Self, best: Self) -> bool;
556 /// Whether `candidate` replaces `best` for a minimum: strictly less, or
557 /// the first NaN.
558 fn beats_min(candidate: Self, best: Self) -> bool;
559}
560
561macro_rules! impl_float_key {
562 ($($ty:ty),*) => {$(
563 impl OrderKey for $ty {
564 #[inline(always)]
565 fn key_is_nan(self) -> bool {
566 self.is_nan()
567 }
568 #[inline(always)]
569 fn beats_max(candidate: Self, best: Self) -> bool {
570 // Select form without a data dependent branch; a NaN `best`
571 // is never replaced, so the first NaN wins.
572 (candidate > best) | (candidate.is_nan() & !best.is_nan())
573 }
574 #[inline(always)]
575 fn beats_min(candidate: Self, best: Self) -> bool {
576 (candidate < best) | (candidate.is_nan() & !best.is_nan())
577 }
578 }
579 )*};
580}
581impl_float_key!(f32, f64);
582
583macro_rules! impl_int_key {
584 ($($ty:ty),*) => {$(
585 impl OrderKey for $ty {
586 #[inline(always)]
587 fn key_is_nan(self) -> bool {
588 false
589 }
590 #[inline(always)]
591 fn beats_max(candidate: Self, best: Self) -> bool {
592 candidate > best
593 }
594 #[inline(always)]
595 fn beats_min(candidate: Self, best: Self) -> bool {
596 candidate < best
597 }
598 }
599 )*};
600}
601impl_int_key!(i32, i64, u32, u64);
602
603/// Maps an element to its comparison key.
604pub(super) trait KeyFn<T> {
605 type Key: OrderKey;
606 fn key(value: T) -> Self::Key;
607}
608
609/// The value itself (real dtypes).
610pub(super) struct Plain;
611/// The magnitude `|x|`.
612pub(super) struct Magnitude;
613
614macro_rules! impl_plain_key {
615 ($($ty:ty),*) => {$(
616 impl KeyFn<$ty> for Plain {
617 type Key = $ty;
618 #[inline(always)]
619 fn key(value: $ty) -> $ty {
620 value
621 }
622 }
623 )*};
624}
625impl_plain_key!(f32, f64, i32, i64);
626
627impl KeyFn<f32> for Magnitude {
628 type Key = f32;
629 #[inline(always)]
630 fn key(value: f32) -> f32 {
631 value.abs()
632 }
633}
634impl KeyFn<f64> for Magnitude {
635 type Key = f64;
636 #[inline(always)]
637 fn key(value: f64) -> f64 {
638 value.abs()
639 }
640}
641impl KeyFn<i32> for Magnitude {
642 type Key = u32;
643 #[inline(always)]
644 fn key(value: i32) -> u32 {
645 value.unsigned_abs()
646 }
647}
648impl KeyFn<i64> for Magnitude {
649 type Key = u64;
650 #[inline(always)]
651 fn key(value: i64) -> u64 {
652 value.unsigned_abs()
653 }
654}
655impl KeyFn<Complex32> for Magnitude {
656 type Key = f32;
657 #[inline(always)]
658 fn key(value: Complex32) -> f32 {
659 // `hypot` returns +inf for an infinite component even when the other
660 // is NaN; the NaN test keeps any NaN component a NaN key.
661 if value.re.is_nan() || value.im.is_nan() {
662 f32::NAN
663 } else {
664 value.re.hypot(value.im)
665 }
666 }
667}
668impl KeyFn<Complex64> for Magnitude {
669 type Key = f64;
670 #[inline(always)]
671 fn key(value: Complex64) -> f64 {
672 if value.re.is_nan() || value.im.is_nan() {
673 f64::NAN
674 } else {
675 value.re.hypot(value.im)
676 }
677 }
678}
679
680/// Maximum or minimum, fixed at compile time.
681pub(super) trait Direction {
682 fn beats<K: OrderKey>(candidate: K, best: K) -> bool;
683}
684pub(super) struct Greater;
685pub(super) struct Less;
686impl Direction for Greater {
687 #[inline(always)]
688 fn beats<K: OrderKey>(candidate: K, best: K) -> bool {
689 K::beats_max(candidate, best)
690 }
691}
692impl Direction for Less {
693 #[inline(always)]
694 fn beats<K: OrderKey>(candidate: K, best: K) -> bool {
695 K::beats_min(candidate, best)
696 }
697}
698
699/// Independent lanes of the contiguous winner search.
700const ARG_LANES: usize = 8;
701
702/// Index of the winning element of one line.
703///
704/// A unit-stride line runs two passes: a lane-parallel search for the winning
705/// key (NaN if the line has one), then a scan for the first element whose key
706/// equals it (or is NaN). Equal keys tie, so the second pass returns the
707/// lowest index of the winner, exactly as the sequential scan used for
708/// strided lines.
709///
710/// # Safety
711///
712/// Every `src.offset(so + k * ss)` for `k < n` must be readable, and `n > 0`.
713#[inline(always)]
714unsafe fn arg_line<T, F, D>(src: *const T, so: isize, ss: isize, n: usize) -> usize
715where
716 T: Copy,
717 F: KeyFn<T>,
718 D: Direction,
719{
720 // SAFETY: the caller guarantees every visited offset and `n > 0`; a unit
721 // stride line is `n` contiguous elements.
722 unsafe {
723 if ss == 1 {
724 let values = core::slice::from_raw_parts(src.offset(so), n);
725 let mut lanes = [F::key(values[0]); ARG_LANES];
726 let mut chunks = values.chunks_exact(ARG_LANES);
727 for chunk in &mut chunks {
728 for (lane, &value) in lanes.iter_mut().zip(chunk) {
729 let key = F::key(value);
730 *lane = if D::beats(key, *lane) { key } else { *lane };
731 }
732 }
733 // Which NaN or which signed zero wins here does not matter: the
734 // second pass matches by NaN-ness or by equality.
735 let mut winner = lanes[0];
736 let tail = chunks.remainder().iter().map(|&value| F::key(value));
737 for key in lanes[1..].iter().copied().chain(tail) {
738 if D::beats(key, winner) {
739 winner = key;
740 }
741 }
742 let found = if winner.key_is_nan() {
743 values.iter().position(|&value| F::key(value).key_is_nan())
744 } else {
745 values.iter().position(|&value| F::key(value) == winner)
746 };
747 // INVARIANT: the winner is the key of some element of the line.
748 return found.unwrap_or(0);
749 }
750 let mut best = F::key(src.offset(so).read());
751 let mut best_index = 0;
752 let mut offset = so;
753 for index in 1..n {
754 offset += ss;
755 let key = F::key(src.offset(offset).read());
756 if D::beats(key, best) {
757 best = key;
758 best_index = index;
759 }
760 }
761 best_index
762 }
763}
764
765/// Arg-reduces `width` adjacent lines whose source elements are contiguous
766/// across the lines; output `j` is written at `d_o + j * lane`.
767///
768/// # Safety
769///
770/// As [`arg_line`] for each line `j < width` with source base `so + j`, every
771/// destination offset must be writable, and `width <= PANEL`.
772#[allow(clippy::too_many_arguments)]
773#[inline(always)]
774unsafe fn arg_panel<T, I, F, D>(
775 src: *const T,
776 so: isize,
777 ss: isize,
778 n: usize,
779 dst: *mut I,
780 d_o: isize,
781 lane: isize,
782 width: usize,
783) where
784 T: Copy,
785 I: ArgIndex,
786 F: KeyFn<T>,
787 D: Direction,
788{
789 debug_assert!(width <= PANEL);
790 // SAFETY: the caller guarantees `width` contiguous source elements at
791 // every axis position and the destination offsets.
792 unsafe {
793 let first = core::slice::from_raw_parts(src.offset(so), width);
794 let mut best = [F::key(first[0]); PANEL];
795 let mut best_index = [0usize; PANEL];
796 for (best, &value) in best.iter_mut().zip(first) {
797 *best = F::key(value);
798 }
799 let best = &mut best[..width];
800 let best_index = &mut best_index[..width];
801 let mut offset = so;
802 for index in 1..n {
803 offset += ss;
804 let row = core::slice::from_raw_parts(src.offset(offset), width);
805 for ((best, best_index), &value) in best.iter_mut().zip(best_index.iter_mut()).zip(row)
806 {
807 let key = F::key(value);
808 let take = D::beats(key, *best);
809 *best = if take { key } else { *best };
810 *best_index = if take { index } else { *best_index };
811 }
812 }
813 for (lane_index, &index) in best_index.iter().enumerate() {
814 dst.offset(d_o + lane_index as isize * lane)
815 .write(I::from_index(index));
816 }
817 }
818}