Skip to main content

ErasedArgReducePlan

Struct ErasedArgReducePlan 

Source
pub struct ErasedArgReducePlan { /* private fields */ }
Expand description

Dtype-erased argmax / argmin along one axis.

The destination holds one index per line: its dimensions are the source dimensions with axis removed (batch axes are preserved in order), and its dtype is the index dtype, i32 or i64. Indices count from the start of the logical axis, whatever the sign of its stride.

Supported source dtypes are f32, f64, i32 and i64 for every ArgReduceOp, and c32 / c64 for the magnitude variants (ArgReduceOp::MaxAbs, ArgReduceOp::MinAbs) only.

§Ties, NaN and magnitudes

  • Ties resolve to the lowest index. -0.0 and +0.0 compare equal, so they tie.
  • NaN propagates like ReduceOp::Max: if a line contains a NaN, the result is the index of its first NaN, for both the max and the min variants. A complex element is NaN when either component is NaN. This matches NumPy, PyTorch and JAX.
  • Infinities order normally; +inf beats every finite value for Max.
  • The complex magnitude is the overflow-safe modulus (hypot of the components), so values near the top of the exponent range still order correctly. A component of ±inf makes the magnitude +inf unless the other component is NaN.
  • The integer magnitude is the unsigned absolute value, so i32::MIN is larger than i32::MAX.

The result does not depend on the layout, execution context or thread count.

§Examples

use strided_basic::{
    ArgReduceOp, ErasedArgReducePlan, ErasedRawStridedMut, ErasedRawStridedRef, ExecContext,
    KernelDType,
};

// Column-major 3 x 2 matrix; argmax down each column.
let src = [1.0_f64, 5.0, 5.0, -2.0, -7.0, 3.0];
let mut out = [0_i64; 2];
let plan = ErasedArgReducePlan::compile(
    KernelDType::F64,
    KernelDType::I64,
    ArgReduceOp::Max,
    &[3, 2],
    &[1, 3],
    &[2],
    &[1],
    0,
)
.unwrap();
let src_ref = ErasedRawStridedRef::from_slice(&src, &[3, 2], &[1, 3], 0).unwrap();
let mut dest = ErasedRawStridedMut::from_slice_mut(&mut out, &[2], &[1], 0).unwrap();
plan.execute(&ExecContext::serial(), &mut dest, &src_ref).unwrap();
assert_eq!(out, [1, 2]); // lowest index wins the 5.0 tie

Implementations§

Source§

impl ErasedArgReducePlan

Source

pub fn compile( dtype: KernelDType, index_dtype: KernelDType, op: ArgReduceOp, src_dims: &[usize], src_strides: &[isize], dest_dims: &[usize], dest_strides: &[isize], axis: usize, ) -> Result<ErasedArgReducePlan, StridedError>

Validate and store an arg-reduction plan for fixed layouts.

dest_dims must equal src_dims with axis removed.

§Examples
use strided_basic::{ArgReduceOp, ErasedArgReducePlan, KernelDType};
let plan = ErasedArgReducePlan::compile(
    KernelDType::C64, KernelDType::I32, ArgReduceOp::MaxAbs, &[4, 3], &[3, 1], &[3], &[1], 0,
)
.unwrap();
assert_eq!(plan.index_dtype(), KernelDType::I32);
// Complex input has no order without the magnitude.
assert!(ErasedArgReducePlan::compile(
    KernelDType::C64, KernelDType::I64, ArgReduceOp::Max, &[4], &[1], &[], &[], 0,
)
.is_err());
// An empty axis has no index to return.
assert!(ErasedArgReducePlan::compile(
    KernelDType::F64, KernelDType::I64, ArgReduceOp::Max, &[0, 2], &[1, 1], &[2], &[1], 0,
)
.is_err());
§Errors
  • UnsupportedDType for a bool source or an index dtype other than i32 / i64;
  • UnsupportedOp for Max / Min on complex input, and for a zero-length axis, which has no index to return;
  • InvalidAxis, StrideLengthMismatch, ShapeMismatch and NonInjectiveOutputLayout for inconsistent layouts;
  • OffsetOverflow when a layout’s offsets are not representable, or when the axis is too long for an i32 index.
Source

pub fn dtype(&self) -> KernelDType

Source element dtype.

§Examples
use strided_basic::{ArgReduceOp, ErasedArgReducePlan, KernelDType};
let plan = ErasedArgReducePlan::compile(
    KernelDType::F32, KernelDType::I64, ArgReduceOp::Min, &[2], &[1], &[], &[], 0,
)
.unwrap();
assert_eq!(plan.dtype(), KernelDType::F32);
Source

pub fn index_dtype(&self) -> KernelDType

Destination index dtype (i32 or i64).

§Examples
use strided_basic::{ArgReduceOp, ErasedArgReducePlan, KernelDType};
let plan = ErasedArgReducePlan::compile(
    KernelDType::F32, KernelDType::I64, ArgReduceOp::Min, &[2], &[1], &[], &[], 0,
)
.unwrap();
assert_eq!(plan.index_dtype(), KernelDType::I64);
Source

pub fn op(&self) -> ArgReduceOp

Arg-reduction operation.

§Examples
use strided_basic::{ArgReduceOp, ErasedArgReducePlan, KernelDType};
let plan = ErasedArgReducePlan::compile(
    KernelDType::F32, KernelDType::I64, ArgReduceOp::Min, &[2], &[1], &[], &[], 0,
)
.unwrap();
assert_eq!(plan.op(), ArgReduceOp::Min);
Source

pub fn execute( &self, ctx: &ExecContext, dest: &mut ErasedRawStridedMut<'_>, src: &ErasedRawStridedRef<'_>, ) -> Result<(), StridedError>

Execute into an initialized index destination.

§Examples
use strided_basic::{
    ArgReduceOp, ErasedArgReducePlan, ErasedRawStridedMut, ErasedRawStridedRef,
    ExecContext, KernelDType,
};
let src = [3_i32, -9, 9, 1];
let mut out = [0_i32; 1];
let plan = ErasedArgReducePlan::compile(
    KernelDType::I32, KernelDType::I32, ArgReduceOp::MaxAbs, &[4], &[1], &[], &[], 0,
)
.unwrap();
let src_ref = ErasedRawStridedRef::from_slice(&src, &[4], &[1], 0).unwrap();
let mut dest = ErasedRawStridedMut::from_slice_mut(&mut out, &[], &[], 0).unwrap();
plan.execute(&ExecContext::serial(), &mut dest, &src_ref).unwrap();
assert_eq!(out, [1]);
§Errors

Returns DTypeMismatch or PlanLayoutMismatch when a descriptor does not match the plan, before any destination write.

Source

pub fn execute_uninit( &self, ctx: &ExecContext, dest: &mut ErasedRawStridedUninitMut<'_>, src: &ErasedRawStridedPtr<'_>, ) -> Result<(), StridedError>

Execute into an uninitialized index destination.

On success every reachable destination element is written; validation errors are returned before any write.

§Examples
use core::mem::MaybeUninit;
use num_complex::Complex64;
use strided_basic::{
    ArgReduceOp, ErasedArgReducePlan, ErasedRawStridedPtr, ErasedRawStridedRef,
    ErasedRawStridedUninitMut, ExecContext, KernelDType,
};
let src = [Complex64::new(3.0, 4.0), Complex64::new(0.0, 6.0)];
let mut out = [MaybeUninit::<i64>::uninit()];
let plan = ErasedArgReducePlan::compile(
    KernelDType::C64, KernelDType::I64, ArgReduceOp::MinAbs, &[2], &[1], &[], &[], 0,
)
.unwrap();
let src_ref = ErasedRawStridedRef::from_slice(&src, &[2], &[1], 0).unwrap();
let src_ptr = ErasedRawStridedPtr::from_ref(&src_ref);
let mut dest = ErasedRawStridedUninitMut::from_uninit_slice(&mut out, &[], &[], 0).unwrap();
plan.execute_uninit(&ExecContext::serial(), &mut dest, &src_ptr).unwrap();
assert_eq!(unsafe { out[0].assume_init() }, 0);
§Errors

As Self::execute, plus OverlappingInputOutput when the source overlaps the destination allocation.

Trait Implementations§

Source§

impl Clone for ErasedArgReducePlan

Source§

fn clone(&self) -> ErasedArgReducePlan

Returns a duplicate of the value. Read more
1.0.0 (const: unstable) · Source§

fn clone_from(&mut self, source: &Self)

Performs copy-assignment from source. Read more
Source§

impl Debug for ErasedArgReducePlan

Source§

fn fmt(&self, f: &mut Formatter<'_>) -> Result<(), Error>

Formats the value using the given formatter. Read more

Auto Trait Implementations§

Blanket Implementations§

Source§

impl<T> Any for T
where T: 'static + ?Sized,

Source§

fn type_id(&self) -> TypeId

Gets the TypeId of self. Read more
Source§

impl<T> Borrow<T> for T
where T: ?Sized,

Source§

fn borrow(&self) -> &T

Immutably borrows from an owned value. Read more
Source§

impl<T> BorrowMut<T> for T
where T: ?Sized,

Source§

fn borrow_mut(&mut self) -> &mut T

Mutably borrows from an owned value. Read more
Source§

impl<T> CloneToUninit for T
where T: Clone,

Source§

unsafe fn clone_to_uninit(&self, dest: *mut u8)

🔬This is a nightly-only experimental API. (clone_to_uninit)
Performs copy-assignment from self to dest. Read more
Source§

impl<T> From<T> for T

Source§

fn from(t: T) -> T

Returns the argument unchanged.

Source§

impl<T, U> Into<U> for T
where U: From<T>,

Source§

fn into(self) -> U

Calls U::from(self).

That is, this conversion is whatever the implementation of From<T> for U chooses to do.

Source§

impl<T> MaybeSend for T

Source§

impl<T> MaybeSendSync for T

Source§

impl<T> MaybeSync for T

Source§

impl<T> ToOwned for T
where T: Clone,

Source§

type Owned = T

The resulting type after obtaining ownership.
Source§

fn to_owned(&self) -> T

Creates owned data from borrowed data, usually by cloning. Read more
Source§

fn clone_into(&self, target: &mut T)

Uses borrowed data to replace owned data, usually by cloning. Read more
Source§

impl<T, U> TryFrom<U> for T
where U: Into<T>,

Source§

type Error = !

The type returned in the event of a conversion error.
Source§

fn try_from(value: U) -> Result<T, !>

Performs the conversion.
Source§

impl<T, U> TryInto<U> for T
where U: TryFrom<T>,

Source§

type Error = <U as TryFrom<T>>::Error

The type returned in the event of a conversion error.
Source§

fn try_into(self) -> Result<U, <U as TryFrom<T>>::Error>

Performs the conversion.