Skip to main content

tenferro_tensor/
dispatch.rs

1//! Dtype-stripping dispatch macros for erased tensor values.
2
3/// Dispatch a dtype-erased [`Tensor`](crate::Tensor) to a typed tensor body.
4///
5/// The dtype-set guard keeps unsupported dtype rejection at the boundary where
6/// the backend and operation name are still visible.
7///
8/// # Examples
9///
10/// ```rust
11/// use tenferro_tensor::{BackendId, Tensor};
12///
13/// let tensor = Tensor::from_vec_col_major(vec![2], vec![1.0_f64, 2.0])?;
14/// let shape = tenferro_tensor::with_scalar!(
15///     &tensor,
16///     float_only,
17///     backend = BackendId::Cpu,
18///     op = "shape_probe",
19///     |typed| -> tenferro_tensor::Result<Vec<usize>> { Ok(typed.shape().to_vec()) }
20/// )?;
21/// assert_eq!(shape, vec![2]);
22/// # Ok::<(), tenferro_tensor::Error>(())
23/// ```
24#[macro_export]
25macro_rules! with_scalar {
26    ($tensor:expr, all, backend = $backend:expr, op = $op:expr, |$typed:ident| $(-> $ret:ty)? $body:block) => {{
27        let _ = &$backend;
28        let _ = &$op;
29        match $tensor {
30            $crate::Tensor::F32($typed) => $body,
31            $crate::Tensor::F64($typed) => $body,
32            $crate::Tensor::I32($typed) => $body,
33            $crate::Tensor::I64($typed) => $body,
34            $crate::Tensor::Bool($typed) => $body,
35            $crate::Tensor::C32($typed) => $body,
36            $crate::Tensor::C64($typed) => $body,
37        }
38    }};
39    ($tensor:expr, numeric, backend = $backend:expr, op = $op:expr, |$typed:ident| $(-> $ret:ty)? $body:block) => {{
40        match $tensor {
41            $crate::Tensor::F32($typed) => $body,
42            $crate::Tensor::F64($typed) => $body,
43            $crate::Tensor::I32($typed) => $body,
44            $crate::Tensor::I64($typed) => $body,
45            $crate::Tensor::C32($typed) => $body,
46            $crate::Tensor::C64($typed) => $body,
47            other => Err($crate::Error::unsupported_dtype(
48                $op,
49                other.dtype(),
50                format!("backend {} does not support this operation/dtype", $backend),
51            )),
52        }
53    }};
54    ($tensor:expr, float_complex, backend = $backend:expr, op = $op:expr, |$typed:ident| $(-> $ret:ty)? $body:block) => {{
55        match $tensor {
56            $crate::Tensor::F32($typed) => $body,
57            $crate::Tensor::F64($typed) => $body,
58            $crate::Tensor::C32($typed) => $body,
59            $crate::Tensor::C64($typed) => $body,
60            other => Err($crate::Error::unsupported_dtype(
61                $op,
62                other.dtype(),
63                format!("backend {} does not support this operation/dtype", $backend),
64            )),
65        }
66    }};
67    ($tensor:expr, float_only, backend = $backend:expr, op = $op:expr, |$typed:ident| $(-> $ret:ty)? $body:block) => {{
68        match $tensor {
69            $crate::Tensor::F32($typed) => $body,
70            $crate::Tensor::F64($typed) => $body,
71            other => Err($crate::Error::unsupported_dtype(
72                $op,
73                other.dtype(),
74                format!("backend {} does not support this operation/dtype", $backend),
75            )),
76        }
77    }};
78}
79
80/// Dispatch a [`TensorRead`](crate::TensorRead) to a typed tensor view body.
81///
82/// # Examples
83///
84/// ```rust
85/// use tenferro_tensor::{BackendId, Tensor, TensorRead};
86///
87/// let tensor = Tensor::from_vec_col_major(vec![2], vec![1.0_f32, 2.0])?;
88/// let read = TensorRead::from_tensor(&tensor);
89/// let shape = tenferro_tensor::with_scalar_read!(
90///     read,
91///     float_only,
92///     backend = BackendId::Cpu,
93///     op = "shape_probe",
94///     |view| -> tenferro_tensor::Result<Vec<usize>> { Ok(view.shape().to_vec()) }
95/// )?;
96/// assert_eq!(shape, vec![2]);
97/// # Ok::<(), tenferro_tensor::Error>(())
98/// ```
99#[macro_export]
100macro_rules! with_scalar_read {
101    ($read:expr, all, backend = $backend:expr, op = $op:expr, |$view:ident| $(-> $ret:ty)? $body:block) => {{
102        let _ = &$backend;
103        let _ = &$op;
104        match $read {
105            $crate::TensorRead::Tensor(tensor) => match tensor {
106                $crate::Tensor::F32(tensor) => {
107                    let $view = tensor.as_view();
108                    $body
109                }
110                $crate::Tensor::F64(tensor) => {
111                    let $view = tensor.as_view();
112                    $body
113                }
114                $crate::Tensor::I32(tensor) => {
115                    let $view = tensor.as_view();
116                    $body
117                }
118                $crate::Tensor::I64(tensor) => {
119                    let $view = tensor.as_view();
120                    $body
121                }
122                $crate::Tensor::Bool(tensor) => {
123                    let $view = tensor.as_view();
124                    $body
125                }
126                $crate::Tensor::C32(tensor) => {
127                    let $view = tensor.as_view();
128                    $body
129                }
130                $crate::Tensor::C64(tensor) => {
131                    let $view = tensor.as_view();
132                    $body
133                }
134            },
135            $crate::TensorRead::View(view) => match view {
136                $crate::TensorView::F32($view) => $body,
137                $crate::TensorView::F64($view) => $body,
138                $crate::TensorView::I32($view) => $body,
139                $crate::TensorView::I64($view) => $body,
140                $crate::TensorView::Bool($view) => $body,
141                $crate::TensorView::C32($view) => $body,
142                $crate::TensorView::C64($view) => $body,
143            },
144        }
145    }};
146    ($read:expr, numeric, backend = $backend:expr, op = $op:expr, |$view:ident| $(-> $ret:ty)? $body:block) => {{
147        let read = $read;
148        let dtype = read.dtype();
149        match read {
150            $crate::TensorRead::Tensor(tensor) => match tensor {
151                $crate::Tensor::F32(tensor) => {
152                    let $view = tensor.as_view();
153                    $body
154                }
155                $crate::Tensor::F64(tensor) => {
156                    let $view = tensor.as_view();
157                    $body
158                }
159                $crate::Tensor::I32(tensor) => {
160                    let $view = tensor.as_view();
161                    $body
162                }
163                $crate::Tensor::I64(tensor) => {
164                    let $view = tensor.as_view();
165                    $body
166                }
167                $crate::Tensor::C32(tensor) => {
168                    let $view = tensor.as_view();
169                    $body
170                }
171                $crate::Tensor::C64(tensor) => {
172                    let $view = tensor.as_view();
173                    $body
174                }
175                $crate::Tensor::Bool(_) => Err($crate::Error::unsupported_dtype(
176                    $op,
177                    dtype,
178                    format!("backend {} does not support this operation/dtype", $backend),
179                )),
180            },
181            $crate::TensorRead::View(view) => match view {
182                $crate::TensorView::F32($view) => $body,
183                $crate::TensorView::F64($view) => $body,
184                $crate::TensorView::I32($view) => $body,
185                $crate::TensorView::I64($view) => $body,
186                $crate::TensorView::C32($view) => $body,
187                $crate::TensorView::C64($view) => $body,
188                $crate::TensorView::Bool(_) => Err($crate::Error::unsupported_dtype(
189                    $op,
190                    dtype,
191                    format!("backend {} does not support this operation/dtype", $backend),
192                )),
193            },
194        }
195    }};
196    ($read:expr, float_complex, backend = $backend:expr, op = $op:expr, |$view:ident| $(-> $ret:ty)? $body:block) => {{
197        let read = $read;
198        let dtype = read.dtype();
199        match read {
200            $crate::TensorRead::Tensor(tensor) => match tensor {
201                $crate::Tensor::F32(tensor) => {
202                    let $view = tensor.as_view();
203                    $body
204                }
205                $crate::Tensor::F64(tensor) => {
206                    let $view = tensor.as_view();
207                    $body
208                }
209                $crate::Tensor::C32(tensor) => {
210                    let $view = tensor.as_view();
211                    $body
212                }
213                $crate::Tensor::C64(tensor) => {
214                    let $view = tensor.as_view();
215                    $body
216                }
217                $crate::Tensor::I32(_) | $crate::Tensor::I64(_) | $crate::Tensor::Bool(_) => {
218                    Err($crate::Error::unsupported_dtype(
219                        $op,
220                        dtype,
221                        format!("backend {} does not support this operation/dtype", $backend),
222                    ))
223                }
224            },
225            $crate::TensorRead::View(view) => match view {
226                $crate::TensorView::F32($view) => $body,
227                $crate::TensorView::F64($view) => $body,
228                $crate::TensorView::C32($view) => $body,
229                $crate::TensorView::C64($view) => $body,
230                $crate::TensorView::I32(_)
231                | $crate::TensorView::I64(_)
232                | $crate::TensorView::Bool(_) => Err($crate::Error::unsupported_dtype(
233                    $op,
234                    dtype,
235                    format!("backend {} does not support this operation/dtype", $backend),
236                )),
237            },
238        }
239    }};
240    ($read:expr, float_only, backend = $backend:expr, op = $op:expr, |$view:ident| $(-> $ret:ty)? $body:block) => {{
241        let read = $read;
242        let dtype = read.dtype();
243        match read {
244            $crate::TensorRead::Tensor(tensor) => match tensor {
245                $crate::Tensor::F32(tensor) => {
246                    let $view = tensor.as_view();
247                    $body
248                }
249                $crate::Tensor::F64(tensor) => {
250                    let $view = tensor.as_view();
251                    $body
252                }
253                $crate::Tensor::I32(_)
254                | $crate::Tensor::I64(_)
255                | $crate::Tensor::Bool(_)
256                | $crate::Tensor::C32(_)
257                | $crate::Tensor::C64(_) => Err($crate::Error::unsupported_dtype(
258                    $op,
259                    dtype,
260                    format!("backend {} does not support this operation/dtype", $backend),
261                )),
262            },
263            $crate::TensorRead::View(view) => match view {
264                $crate::TensorView::F32($view) => $body,
265                $crate::TensorView::F64($view) => $body,
266                $crate::TensorView::I32(_)
267                | $crate::TensorView::I64(_)
268                | $crate::TensorView::Bool(_)
269                | $crate::TensorView::C32(_)
270                | $crate::TensorView::C64(_) => Err($crate::Error::unsupported_dtype(
271                    $op,
272                    dtype,
273                    format!("backend {} does not support this operation/dtype", $backend),
274                )),
275            },
276        }
277    }};
278}