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}