Skip to main content

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}