Skip to main content

same_variant_pair

Macro same_variant_pair 

Source
macro_rules! same_variant_pair {
    ($lhs:expr, $rhs:expr, |$a:ident, $b:ident| $call:expr, $fallback:expr) => { ... };
}
Expand description

Dispatch a same-variant pair of tensors to a typed kernel.

The four real and complex arms are declared once. The caller supplies the operands, the typed kernel, and the expression to return for an unsupported pair.

ยงExamples

use tenferro_internal_cpu_kernels::same_variant_pair;
use tenferro_tensor::{Error, Tensor};

fn first_matching_pair(lhs: &Tensor, rhs: &Tensor) -> tenferro_tensor::Result<Tensor> {
    same_variant_pair!(
        lhs,
        rhs,
        |a, b| {
            let _ = b.shape();
            a.duplicate()
        },
        Err(Error::dtype_mismatch(
            "first_matching_pair",
            lhs.dtype(),
            rhs.dtype()
        ))
    )
}

let lhs = Tensor::from_vec_col_major(vec![1], vec![1.0_f64])?;
let rhs = Tensor::from_vec_col_major(vec![1], vec![2.0_f64])?;
let out = first_matching_pair(&lhs, &rhs)?;
assert_eq!(out.as_slice::<f64>()?, &[1.0]);