Skip to main content

tenferro_internal_cpu_kernels/
dispatch.rs

1//! Shared erased-dispatch macros for concrete tensor values.
2//!
3//! Every crate that carries the concrete value type needs to reach a typed kernel
4//! from it. Declaring the matching once here keeps each call site to the kernel it
5//! calls, so a call site stops naming the variants and does not change again when
6//! the value type stops being a closed enum.
7//!
8//! The macros expand to paths that resolve at the call site, so a caller must have
9//! `Tensor` in scope.
10
11/// Dispatch a same-variant pair of tensors to a typed kernel.
12///
13/// The four real and complex arms are declared once. The caller supplies the
14/// operands, the typed kernel, and the expression to return for an unsupported
15/// pair.
16///
17/// # Examples
18///
19/// ```rust
20/// use tenferro_internal_cpu_kernels::same_variant_pair;
21/// use tenferro_tensor::{Error, Tensor};
22///
23/// fn first_matching_pair(lhs: &Tensor, rhs: &Tensor) -> tenferro_tensor::Result<Tensor> {
24///     same_variant_pair!(
25///         lhs,
26///         rhs,
27///         |a, b| {
28///             let _ = b.shape();
29///             a.duplicate()
30///         },
31///         Err(Error::dtype_mismatch(
32///             "first_matching_pair",
33///             lhs.dtype(),
34///             rhs.dtype()
35///         ))
36///     )
37/// }
38///
39/// let lhs = Tensor::from_vec_col_major(vec![1], vec![1.0_f64])?;
40/// let rhs = Tensor::from_vec_col_major(vec![1], vec![2.0_f64])?;
41/// let out = first_matching_pair(&lhs, &rhs)?;
42/// assert_eq!(out.as_slice::<f64>()?, &[1.0]);
43/// # Ok::<(), tenferro_tensor::Error>(())
44/// ```
45#[macro_export]
46macro_rules! same_variant_pair {
47    ($lhs:expr, $rhs:expr, |$a:ident, $b:ident| $call:expr, $fallback:expr) => {
48        match ($lhs.dtype(), $rhs.dtype()) {
49            ($crate::DType::F32, $crate::DType::F32) => {
50                match ($lhs.as_typed::<f32>(), $rhs.as_typed::<f32>()) {
51                    (Some($a), Some($b)) => $call.map(Tensor::from_typed::<f32>),
52                    _ => $fallback,
53                }
54            }
55            ($crate::DType::F64, $crate::DType::F64) => {
56                match ($lhs.as_typed::<f64>(), $rhs.as_typed::<f64>()) {
57                    (Some($a), Some($b)) => $call.map(Tensor::from_typed::<f64>),
58                    _ => $fallback,
59                }
60            }
61            ($crate::DType::C32, $crate::DType::C32) => {
62                match (
63                    $lhs.as_typed::<$crate::Complex32>(),
64                    $rhs.as_typed::<$crate::Complex32>(),
65                ) {
66                    (Some($a), Some($b)) => $call.map(Tensor::from_typed::<$crate::Complex32>),
67                    _ => $fallback,
68                }
69            }
70            ($crate::DType::C64, $crate::DType::C64) => {
71                match (
72                    $lhs.as_typed::<$crate::Complex64>(),
73                    $rhs.as_typed::<$crate::Complex64>(),
74                ) {
75                    (Some($a), Some($b)) => $call.map(Tensor::from_typed::<$crate::Complex64>),
76                    _ => $fallback,
77                }
78            }
79            _ => $fallback,
80        }
81    };
82}
83
84/// Dispatch a single tensor to a typed kernel and erase its result.
85///
86/// The typed result is wrapped by the scalar's own constructor, so a real-valued
87/// result of a complex input still lands in the matching real variant without the
88/// call site naming a variant.
89///
90/// # Examples
91///
92/// ```rust
93/// use tenferro_internal_cpu_kernels::same_variant_unary;
94/// use tenferro_tensor::{Error, Tensor, TensorScalar};
95///
96/// fn negate(input: &Tensor) -> tenferro_tensor::Result<Tensor> {
97///     same_variant_unary!(
98///         input,
99///         |t| t.duplicate(),
100///         |value| TensorScalar::typed_tensor_into_tensor(value),
101///         Err(Error::dtype_mismatch("negate", input.dtype(), input.dtype()))
102///     )
103/// }
104///
105/// let value = Tensor::from_vec_col_major(vec![1], vec![3.0_f64])?;
106/// assert_eq!(negate(&value)?.as_slice::<f64>()?, &[3.0]);
107/// # Ok::<(), tenferro_tensor::Error>(())
108/// ```
109#[macro_export]
110macro_rules! same_variant_unary {
111    ($input:expr, |$t:ident| $call:expr, |$value:ident| $wrap:expr, $fallback:expr) => {
112        match $input.dtype() {
113            $crate::DType::F32 => match $input.as_typed::<f32>() {
114                Some($t) => $call.map(|$value| $wrap),
115                None => $fallback,
116            },
117            $crate::DType::F64 => match $input.as_typed::<f64>() {
118                Some($t) => $call.map(|$value| $wrap),
119                None => $fallback,
120            },
121            $crate::DType::C32 => match $input.as_typed::<$crate::Complex32>() {
122                Some($t) => $call.map(|$value| $wrap),
123                None => $fallback,
124            },
125            $crate::DType::C64 => match $input.as_typed::<$crate::Complex64>() {
126                Some($t) => $call.map(|$value| $wrap),
127                None => $fallback,
128            },
129            _ => $fallback,
130        }
131    };
132}