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.0and+0.0compare 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;
+infbeats every finite value forMax. - The complex magnitude is the overflow-safe modulus (
hypotof the components), so values near the top of the exponent range still order correctly. A component of±infmakes the magnitude+infunless the other component is NaN. - The integer magnitude is the unsigned absolute value, so
i32::MINis larger thani32::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 tieImplementations§
Source§impl ErasedArgReducePlan
impl ErasedArgReducePlan
Sourcepub 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>
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
UnsupportedDTypefor aboolsource or an index dtype other thani32/i64;UnsupportedOpforMax/Minon complex input, and for a zero-lengthaxis, which has no index to return;InvalidAxis,StrideLengthMismatch,ShapeMismatchandNonInjectiveOutputLayoutfor inconsistent layouts;OffsetOverflowwhen a layout’s offsets are not representable, or when the axis is too long for ani32index.
Sourcepub fn dtype(&self) -> KernelDType
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);Sourcepub fn index_dtype(&self) -> KernelDType
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);Sourcepub fn op(&self) -> ArgReduceOp
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);Sourcepub fn execute(
&self,
ctx: &ExecContext,
dest: &mut ErasedRawStridedMut<'_>,
src: &ErasedRawStridedRef<'_>,
) -> Result<(), StridedError>
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.
Sourcepub fn execute_uninit(
&self,
ctx: &ExecContext,
dest: &mut ErasedRawStridedUninitMut<'_>,
src: &ErasedRawStridedPtr<'_>,
) -> Result<(), StridedError>
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
impl Clone for ErasedArgReducePlan
Source§fn clone(&self) -> ErasedArgReducePlan
fn clone(&self) -> ErasedArgReducePlan
1.0.0 (const: unstable) · Source§fn clone_from(&mut self, source: &Self)
fn clone_from(&mut self, source: &Self)
source. Read more