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/// The typed tensor behind `tensor` when its dtype is the requested scalar.
25///
26/// The exported dispatch macros match on the dtype first, so a mismatch here means the table and the
27/// runtime dtype disagree; the refusal names the operation and backend the caller supplied.
28///
29/// # Errors
30///
31/// Returns [`Error::UnsupportedDType`](crate::Error::UnsupportedDType) when the tensor's dtype is not `T`,
32/// which for an externally defined scalar reports the caller-owned payload.
33///
34/// # Examples
35///
36/// ```
37/// use tenferro_tensor::{DType, Tensor};
38///
39/// let tensor = Tensor::from_vec_col_major(vec![1], vec![1.0_f64])?;
40/// assert_eq!(tensor.dtype(), DType::F64);
41/// # Ok::<(), tenferro_tensor::Error>(())
42/// ```
43pub fn scalar_operand<'a, T: crate::TensorScalar>(
44    tensor: &'a crate::Tensor,
45    op: &'static str,
46    backend: impl core::fmt::Display,
47) -> crate::Result<&'a crate::TypedTensor<T>> {
48    tensor.as_typed::<T>().ok_or_else(|| {
49        crate::Error::unsupported_dtype(
50            op,
51            tensor.dtype(),
52            format!("backend {backend} does not support an externally defined scalar"),
53        )
54    })
55}
56
57#[macro_export]
58macro_rules! with_scalar {
59    ($tensor:expr, all, backend = $backend:expr, op = $op:expr, |$typed:ident| $(-> $ret:ty)? $body:block) => {{
60        let _ = &$backend;
61        let _ = &$op;
62        match $tensor.dtype() {
63            $crate::DType::F32 => $crate::dispatch::scalar_operand::<f32>($tensor, $op, $backend)
64                .and_then(|$typed| $body),
65            $crate::DType::F64 => $crate::dispatch::scalar_operand::<f64>($tensor, $op, $backend)
66                .and_then(|$typed| $body),
67            $crate::DType::I32 => $crate::dispatch::scalar_operand::<i32>($tensor, $op, $backend)
68                .and_then(|$typed| $body),
69            $crate::DType::I64 => $crate::dispatch::scalar_operand::<i64>($tensor, $op, $backend)
70                .and_then(|$typed| $body),
71            $crate::DType::Bool => $crate::dispatch::scalar_operand::<bool>($tensor, $op, $backend)
72                .and_then(|$typed| $body),
73            $crate::DType::C32 => {
74                $crate::dispatch::scalar_operand::<$crate::Complex32>($tensor, $op, $backend)
75                    .and_then(|$typed| $body)
76            }
77            $crate::DType::C64 => {
78                $crate::dispatch::scalar_operand::<$crate::Complex64>($tensor, $op, $backend)
79                    .and_then(|$typed| $body)
80            }
81            // A caller-owned payload has no runtime value in this crate.
82            $crate::DType::External(type_id) => Err($crate::Error::unsupported_dtype(
83                $op,
84                $crate::DType::External(type_id),
85                format!(
86                    "backend {} does not support an externally defined scalar",
87                    $backend
88                ),
89            )),
90        }
91    }};
92    ($tensor:expr, numeric, backend = $backend:expr, op = $op:expr, |$typed:ident| $(-> $ret:ty)? $body:block) => {{
93        match $tensor.dtype() {
94            $crate::DType::F32 => $crate::dispatch::scalar_operand::<f32>($tensor, $op, $backend)
95                .and_then(|$typed| $body),
96            $crate::DType::F64 => $crate::dispatch::scalar_operand::<f64>($tensor, $op, $backend)
97                .and_then(|$typed| $body),
98            $crate::DType::I32 => $crate::dispatch::scalar_operand::<i32>($tensor, $op, $backend)
99                .and_then(|$typed| $body),
100            $crate::DType::I64 => $crate::dispatch::scalar_operand::<i64>($tensor, $op, $backend)
101                .and_then(|$typed| $body),
102            $crate::DType::C32 => {
103                $crate::dispatch::scalar_operand::<$crate::Complex32>($tensor, $op, $backend)
104                    .and_then(|$typed| $body)
105            }
106            $crate::DType::C64 => {
107                $crate::dispatch::scalar_operand::<$crate::Complex64>($tensor, $op, $backend)
108                    .and_then(|$typed| $body)
109            }
110            // A caller-owned payload has no runtime value in this crate.
111            $crate::DType::External(type_id) => Err($crate::Error::unsupported_dtype(
112                $op,
113                $crate::DType::External(type_id),
114                format!(
115                    "backend {} does not support an externally defined scalar",
116                    $backend
117                ),
118            )),
119            other => Err($crate::Error::unsupported_dtype(
120                $op,
121                other,
122                format!("backend {} does not support this operation/dtype", $backend),
123            )),
124        }
125    }};
126    ($tensor:expr, float_complex, backend = $backend:expr, op = $op:expr, |$typed:ident| $(-> $ret:ty)? $body:block) => {{
127        match $tensor.dtype() {
128            $crate::DType::F32 => $crate::dispatch::scalar_operand::<f32>($tensor, $op, $backend)
129                .and_then(|$typed| $body),
130            $crate::DType::F64 => $crate::dispatch::scalar_operand::<f64>($tensor, $op, $backend)
131                .and_then(|$typed| $body),
132            $crate::DType::C32 => {
133                $crate::dispatch::scalar_operand::<$crate::Complex32>($tensor, $op, $backend)
134                    .and_then(|$typed| $body)
135            }
136            $crate::DType::C64 => {
137                $crate::dispatch::scalar_operand::<$crate::Complex64>($tensor, $op, $backend)
138                    .and_then(|$typed| $body)
139            }
140            // A caller-owned payload has no runtime value in this crate.
141            $crate::DType::External(type_id) => Err($crate::Error::unsupported_dtype(
142                $op,
143                $crate::DType::External(type_id),
144                format!(
145                    "backend {} does not support an externally defined scalar",
146                    $backend
147                ),
148            )),
149            other => Err($crate::Error::unsupported_dtype(
150                $op,
151                other,
152                format!("backend {} does not support this operation/dtype", $backend),
153            )),
154        }
155    }};
156    ($tensor:expr, float_only, backend = $backend:expr, op = $op:expr, |$typed:ident| $(-> $ret:ty)? $body:block) => {{
157        match $tensor.dtype() {
158            $crate::DType::F32 => $crate::dispatch::scalar_operand::<f32>($tensor, $op, $backend)
159                .and_then(|$typed| $body),
160            $crate::DType::F64 => $crate::dispatch::scalar_operand::<f64>($tensor, $op, $backend)
161                .and_then(|$typed| $body),
162            other => Err($crate::Error::unsupported_dtype(
163                $op,
164                other,
165                format!("backend {} does not support this operation/dtype", $backend),
166            )),
167        }
168    }};
169}
170
171/// Dispatch a [`TensorRead`](crate::TensorRead) to a typed tensor view body.
172///
173/// # Examples
174///
175/// ```rust
176/// use tenferro_tensor::{BackendId, Tensor, TensorRead};
177///
178/// let tensor = Tensor::from_vec_col_major(vec![2], vec![1.0_f32, 2.0])?;
179/// let read = TensorRead::from_tensor(&tensor);
180/// let shape = tenferro_tensor::with_scalar_read!(
181///     read,
182///     float_only,
183///     backend = BackendId::Cpu,
184///     op = "shape_probe",
185///     |view| -> tenferro_tensor::Result<Vec<usize>> { Ok(view.shape().to_vec()) }
186/// )?;
187/// assert_eq!(shape, vec![2]);
188/// # Ok::<(), tenferro_tensor::Error>(())
189/// ```
190#[macro_export]
191macro_rules! with_scalar_read {
192    ($read:expr, all, backend = $backend:expr, op = $op:expr, |$view:ident| $(-> $ret:ty)? $body:block) => {{
193        let _ = &$backend;
194        let _ = &$op;
195        match $read {
196            $crate::TensorRead::Tensor(tensor) => match tensor.dtype() {
197                $crate::DType::F32 => $crate::dispatch::scalar_operand::<f32>(
198                    tensor, $op, $backend,
199                )
200                .and_then(|tensor| {
201                    let $view = tensor.as_view();
202                    $body
203                }),
204                $crate::DType::F64 => $crate::dispatch::scalar_operand::<f64>(
205                    tensor, $op, $backend,
206                )
207                .and_then(|tensor| {
208                    let $view = tensor.as_view();
209                    $body
210                }),
211                $crate::DType::I32 => $crate::dispatch::scalar_operand::<i32>(
212                    tensor, $op, $backend,
213                )
214                .and_then(|tensor| {
215                    let $view = tensor.as_view();
216                    $body
217                }),
218                $crate::DType::I64 => $crate::dispatch::scalar_operand::<i64>(
219                    tensor, $op, $backend,
220                )
221                .and_then(|tensor| {
222                    let $view = tensor.as_view();
223                    $body
224                }),
225                $crate::DType::Bool => $crate::dispatch::scalar_operand::<bool>(
226                    tensor, $op, $backend,
227                )
228                .and_then(|tensor| {
229                    let $view = tensor.as_view();
230                    $body
231                }),
232                $crate::DType::C32 => {
233                    $crate::dispatch::scalar_operand::<$crate::Complex32>(tensor, $op, $backend)
234                        .and_then(|tensor| {
235                            let $view = tensor.as_view();
236                            $body
237                        })
238                }
239                $crate::DType::C64 => {
240                    $crate::dispatch::scalar_operand::<$crate::Complex64>(tensor, $op, $backend)
241                        .and_then(|tensor| {
242                            let $view = tensor.as_view();
243                            $body
244                        })
245                }
246                // A caller-owned payload has no runtime read view in this crate.
247                $crate::DType::External(type_id) => Err($crate::Error::unsupported_dtype(
248                    $op,
249                    $crate::DType::External(type_id),
250                    format!(
251                        "backend {} does not support an externally defined scalar",
252                        $backend
253                    ),
254                )),
255            },
256            $crate::TensorRead::View(view) => match view {
257                $crate::TensorView::F32($view) => $body,
258                $crate::TensorView::F64($view) => $body,
259                $crate::TensorView::I32($view) => $body,
260                $crate::TensorView::I64($view) => $body,
261                $crate::TensorView::Bool($view) => $body,
262                $crate::TensorView::C32($view) => $body,
263                $crate::TensorView::C64($view) => $body,
264            },
265        }
266    }};
267    ($read:expr, numeric, backend = $backend:expr, op = $op:expr, |$view:ident| $(-> $ret:ty)? $body:block) => {{
268        let read = $read;
269        let dtype = read.dtype();
270        match read {
271            $crate::TensorRead::Tensor(tensor) => match tensor.dtype() {
272                $crate::DType::F32 => $crate::dispatch::scalar_operand::<f32>(
273                    tensor, $op, $backend,
274                )
275                .and_then(|tensor| {
276                    let $view = tensor.as_view();
277                    $body
278                }),
279                $crate::DType::F64 => $crate::dispatch::scalar_operand::<f64>(
280                    tensor, $op, $backend,
281                )
282                .and_then(|tensor| {
283                    let $view = tensor.as_view();
284                    $body
285                }),
286                $crate::DType::I32 => $crate::dispatch::scalar_operand::<i32>(
287                    tensor, $op, $backend,
288                )
289                .and_then(|tensor| {
290                    let $view = tensor.as_view();
291                    $body
292                }),
293                $crate::DType::I64 => $crate::dispatch::scalar_operand::<i64>(
294                    tensor, $op, $backend,
295                )
296                .and_then(|tensor| {
297                    let $view = tensor.as_view();
298                    $body
299                }),
300                $crate::DType::C32 => {
301                    $crate::dispatch::scalar_operand::<$crate::Complex32>(tensor, $op, $backend)
302                        .and_then(|tensor| {
303                            let $view = tensor.as_view();
304                            $body
305                        })
306                }
307                $crate::DType::C64 => {
308                    $crate::dispatch::scalar_operand::<$crate::Complex64>(tensor, $op, $backend)
309                        .and_then(|tensor| {
310                            let $view = tensor.as_view();
311                            $body
312                        })
313                }
314                $crate::DType::Bool => Err($crate::Error::unsupported_dtype(
315                    $op,
316                    dtype,
317                    format!("backend {} does not support this operation/dtype", $backend),
318                )),
319                $crate::DType::External(type_id) => Err($crate::Error::unsupported_dtype(
320                    $op,
321                    $crate::DType::External(type_id),
322                    format!(
323                        "backend {} does not support an externally defined scalar",
324                        $backend
325                    ),
326                )),
327            },
328            $crate::TensorRead::View(view) => match view {
329                $crate::TensorView::F32($view) => $body,
330                $crate::TensorView::F64($view) => $body,
331                $crate::TensorView::I32($view) => $body,
332                $crate::TensorView::I64($view) => $body,
333                $crate::TensorView::C32($view) => $body,
334                $crate::TensorView::C64($view) => $body,
335                $crate::TensorView::Bool(_) => Err($crate::Error::unsupported_dtype(
336                    $op,
337                    dtype,
338                    format!("backend {} does not support this operation/dtype", $backend),
339                )),
340            },
341        }
342    }};
343    ($read:expr, float_complex, backend = $backend:expr, op = $op:expr, |$view:ident| $(-> $ret:ty)? $body:block) => {{
344        let read = $read;
345        let dtype = read.dtype();
346        match read {
347            $crate::TensorRead::Tensor(tensor) => match tensor.dtype() {
348                $crate::DType::F32 => $crate::dispatch::scalar_operand::<f32>(
349                    tensor, $op, $backend,
350                )
351                .and_then(|tensor| {
352                    let $view = tensor.as_view();
353                    $body
354                }),
355                $crate::DType::F64 => $crate::dispatch::scalar_operand::<f64>(
356                    tensor, $op, $backend,
357                )
358                .and_then(|tensor| {
359                    let $view = tensor.as_view();
360                    $body
361                }),
362                $crate::DType::C32 => {
363                    $crate::dispatch::scalar_operand::<$crate::Complex32>(tensor, $op, $backend)
364                        .and_then(|tensor| {
365                            let $view = tensor.as_view();
366                            $body
367                        })
368                }
369                $crate::DType::C64 => {
370                    $crate::dispatch::scalar_operand::<$crate::Complex64>(tensor, $op, $backend)
371                        .and_then(|tensor| {
372                            let $view = tensor.as_view();
373                            $body
374                        })
375                }
376                $crate::DType::I32 | $crate::DType::I64 | $crate::DType::Bool => {
377                    Err($crate::Error::unsupported_dtype(
378                        $op,
379                        dtype,
380                        format!("backend {} does not support this operation/dtype", $backend),
381                    ))
382                }
383                $crate::DType::External(type_id) => Err($crate::Error::unsupported_dtype(
384                    $op,
385                    $crate::DType::External(type_id),
386                    format!(
387                        "backend {} does not support an externally defined scalar",
388                        $backend
389                    ),
390                )),
391            },
392            $crate::TensorRead::View(view) => match view {
393                $crate::TensorView::F32($view) => $body,
394                $crate::TensorView::F64($view) => $body,
395                $crate::TensorView::C32($view) => $body,
396                $crate::TensorView::C64($view) => $body,
397                $crate::TensorView::I32(_)
398                | $crate::TensorView::I64(_)
399                | $crate::TensorView::Bool(_) => Err($crate::Error::unsupported_dtype(
400                    $op,
401                    dtype,
402                    format!("backend {} does not support this operation/dtype", $backend),
403                )),
404            },
405        }
406    }};
407    ($read:expr, float_only, backend = $backend:expr, op = $op:expr, |$view:ident| $(-> $ret:ty)? $body:block) => {{
408        let read = $read;
409        let dtype = read.dtype();
410        match read {
411            $crate::TensorRead::Tensor(tensor) => match tensor.dtype() {
412                $crate::DType::F32 => $crate::dispatch::scalar_operand::<f32>(
413                    tensor, $op, $backend,
414                )
415                .and_then(|tensor| {
416                    let $view = tensor.as_view();
417                    $body
418                }),
419                $crate::DType::F64 => $crate::dispatch::scalar_operand::<f64>(
420                    tensor, $op, $backend,
421                )
422                .and_then(|tensor| {
423                    let $view = tensor.as_view();
424                    $body
425                }),
426                $crate::DType::I32
427                | $crate::DType::I64
428                | $crate::DType::Bool
429                | $crate::DType::C32
430                | $crate::DType::C64
431                | $crate::DType::External(_) => Err($crate::Error::unsupported_dtype(
432                    $op,
433                    dtype,
434                    format!("backend {} does not support this operation/dtype", $backend),
435                )),
436            },
437            $crate::TensorRead::View(view) => match view {
438                $crate::TensorView::F32($view) => $body,
439                $crate::TensorView::F64($view) => $body,
440                $crate::TensorView::I32(_)
441                | $crate::TensorView::I64(_)
442                | $crate::TensorView::Bool(_)
443                | $crate::TensorView::C32(_)
444                | $crate::TensorView::C64(_) => Err($crate::Error::unsupported_dtype(
445                    $op,
446                    dtype,
447                    format!("backend {} does not support this operation/dtype", $backend),
448                )),
449            },
450        }
451    }};
452}