Skip to main content

strided_basic/
erased_common.rs

1use crate::*;
2/// Reject an input overlapping any byte of the output backing allocation.
3///
4/// # Examples
5///
6/// ```
7/// use strided_basic::{ErasedRawStridedRef, ErasedRawStridedPtr, ErasedRawStridedUninitMut, execution::validate_uninit_no_overlap};
8/// use core::mem::MaybeUninit;
9/// let input = ErasedRawStridedRef::from_slice(&[2_i32], &[1], &[1], 0).unwrap();
10/// let input = ErasedRawStridedPtr::from_ref(&input);
11/// let mut values = [MaybeUninit::<i32>::uninit()];
12/// let output = ErasedRawStridedUninitMut::from_uninit_slice(&mut values, &[1], &[1], 0).unwrap();
13/// validate_uninit_no_overlap(&output, &input, 0).unwrap();
14/// ```
15///
16/// # Errors
17/// Returns an overlap or byte-range overflow error without reading element values.
18pub fn validate_uninit_no_overlap(
19    dest: &ErasedRawStridedUninitMut<'_>,
20    input: &ErasedRawStridedPtr<'_>,
21    input_index: usize,
22) -> Result<()> {
23    if input.overlaps_uninit_mut(dest)? {
24        Err(StridedError::OverlappingInputOutput { input: input_index })
25    } else {
26        Ok(())
27    }
28}
29/// Check a descriptor dtype against a prepared operation.
30///
31/// # Examples
32///
33/// ```
34/// use strided_basic::{KernelDType, execution::check_dtype};
35/// check_dtype(KernelDType::F64, KernelDType::F64).unwrap();
36/// assert!(check_dtype(KernelDType::F64, KernelDType::F32).is_err());
37/// ```
38///
39/// # Errors
40/// Returns `DTypeMismatch` when the tags differ.
41pub fn check_dtype(expected: KernelDType, actual: KernelDType) -> Result<()> {
42    if actual != expected {
43        return Err(StridedError::DTypeMismatch {
44            expected: expected.label(),
45            actual: actual.label(),
46        });
47    }
48    Ok(())
49}
50/// Check whether a dtype has a static-indexing implementation.
51///
52/// # Examples
53///
54/// ```
55/// use strided_basic::{KernelDType, execution::check_static_indexing_dtype};
56/// check_static_indexing_dtype(KernelDType::Bool).unwrap();
57/// ```
58///
59/// # Errors
60/// Returns `UnsupportedDType` for unimplemented tags.
61pub fn check_static_indexing_dtype(dtype: KernelDType) -> Result<()> {
62    match dtype {
63        KernelDType::F32
64        | KernelDType::F64
65        | KernelDType::I32
66        | KernelDType::I64
67        | KernelDType::Bool
68        | KernelDType::C32
69        | KernelDType::C64 => Ok(()),
70        _ => Err(StridedError::UnsupportedDType {
71            dtype: dtype.label(),
72        }),
73    }
74}
75/// Borrow initialized element storage with owning view metadata.
76///
77/// # Examples
78///
79/// ```
80/// use strided_basic::{ErasedRawStridedRef, execution::erased_view};
81/// let values = [2.0_f64, 3.0];
82/// let raw = ErasedRawStridedRef::from_slice(&values, &[2], &[1], 0).unwrap();
83/// let view = erased_view::<f64>(&raw).unwrap();
84/// assert_eq!(view.get(&[1]), 3.0);
85/// ```
86///
87/// # Errors
88/// Returns a dtype mismatch if `T` differs from the descriptor tag.
89pub fn erased_view<'a, T: KernelStorageElement>(
90    src: &'a ErasedRawStridedRef<'a>,
91) -> Result<StridedView<'a, T>> {
92    let data = src.data_as::<T>()?;
93    Ok(unsafe { StridedView::new_unchecked(data, src.dims(), src.strides(), src.offset()) })
94}