use crate::bindings::{
ONNXTensorElementDataType, OrtCustomOpInputOutputCharacteristic,
OrtCustomOpInputOutputCharacteristic_INPUT_OUTPUT_OPTIONAL,
OrtCustomOpInputOutputCharacteristic_INPUT_OUTPUT_REQUIRED,
OrtCustomOpInputOutputCharacteristic_INPUT_OUTPUT_VARIADIC,
};
use crate::value::Value;
use anyhow::{Result, bail};
use ndarray::ArrayViewD;
pub trait Inputs<'a>: Sized {
const VARIADIC_IS_HOMOGENEOUS: Option<bool>;
const NUM_POSITIONAL: usize;
fn try_from_values(values: Vec<Option<Value<'a>>>) -> Result<Self>;
fn tensor_data_type(index: usize) -> Option<ONNXTensorElementDataType>;
fn characteristic(index: usize) -> OrtCustomOpInputOutputCharacteristic;
}
pub trait Input<'s>: Sized {
fn try_from_value(value: Option<Value<'s>>) -> Result<Self>;
fn characteristic() -> OrtCustomOpInputOutputCharacteristic;
}
trait OnnxTensorDtype {
fn dtype_id() -> Option<ONNXTensorElementDataType>;
}
impl<'s> Input<'s> for ArrayViewD<'s, &'s str> {
fn try_from_value(value: Option<Value<'s>>) -> Result<Self> {
if let Some(Value::TensorStr(arr)) = value {
Ok(arr)
} else {
bail!("Expected 'String' tensor, found {:?}", value)
}
}
fn characteristic() -> OrtCustomOpInputOutputCharacteristic {
OrtCustomOpInputOutputCharacteristic_INPUT_OUTPUT_REQUIRED
}
}
impl<'s, T> Input<'s> for Option<T>
where
T: Input<'s>,
{
fn try_from_value(value: Option<Value<'s>>) -> Result<Self> {
if value.is_none() {
Ok(None)
} else {
T::try_from_value(value).map(Some)
}
}
fn characteristic() -> OrtCustomOpInputOutputCharacteristic {
OrtCustomOpInputOutputCharacteristic_INPUT_OUTPUT_OPTIONAL
}
}
macro_rules! impl_try_from {
($ty:ty, $variant:path) => {
impl<'a> Input<'a> for ArrayViewD<'a, $ty> {
fn try_from_value(value: Option<Value<'a>>) -> Result<Self> {
if let Some($variant(arr)) = value {
Ok(arr)
} else {
bail!("Expected '{}' tensor, found {:?}", stringify!($ty), value)
}
}
fn characteristic() -> OrtCustomOpInputOutputCharacteristic {
OrtCustomOpInputOutputCharacteristic_INPUT_OUTPUT_REQUIRED
}
}
};
}
impl_try_from!(bool, Value::TensorBool);
impl_try_from!(u8, Value::TensorU8);
impl_try_from!(u64, Value::TensorU64);
impl_try_from!(u32, Value::TensorU32);
impl_try_from!(u16, Value::TensorU16);
impl_try_from!(i8, Value::TensorI8);
impl_try_from!(i64, Value::TensorI64);
impl_try_from!(i32, Value::TensorI32);
impl_try_from!(i16, Value::TensorI16);
impl_try_from!(f64, Value::TensorF64);
impl_try_from!(f32, Value::TensorF32);
impl<'s, A> Inputs<'s> for (Vec<A>,)
where
A: Input<'s> + OnnxTensorDtype,
{
const VARIADIC_IS_HOMOGENEOUS: Option<bool> = Some(true);
const NUM_POSITIONAL: usize = 0;
fn try_from_values(values: Vec<Option<Value<'s>>>) -> Result<Self> {
let rest = values
.into_iter()
.map(|el| Input::try_from_value(el))
.collect::<Result<_, _>>()?;
Ok((rest,))
}
fn tensor_data_type(idx: usize) -> Option<ONNXTensorElementDataType> {
[A::dtype_id()][idx]
}
fn characteristic(_index: usize) -> OrtCustomOpInputOutputCharacteristic {
OrtCustomOpInputOutputCharacteristic_INPUT_OUTPUT_VARIADIC
}
}
macro_rules! impl_inputs {
($n_min:literal, $is_variadic:literal, $($var_ty:ident)? | $($positional_ty:ident),*) => {
impl<'s, $($positional_ty,)* $($var_ty)*> Inputs<'s> for ($($positional_ty,)* $(Vec<$var_ty>,)*)
where
$($positional_ty: Input<'s> + OnnxTensorDtype,)*
$($var_ty: Input<'s> + OnnxTensorDtype,)*
{
const VARIADIC_IS_HOMOGENEOUS: Option<bool> = if $is_variadic {Some(true)} else { None };
const NUM_POSITIONAL: usize = $n_min;
fn try_from_values(values: Vec<Option<Value<'s>>>) -> Result<Self>
{
if $is_variadic {
if values.len() < $n_min {
bail!("expected at least {} inputs; found {}", $n_min, values.len())
}
} else if values.len() != $n_min {
bail!("expected {} inputs; found {}", $n_min, values.len())
}
let mut iter = values.into_iter();
Ok((
$(<$positional_ty as Input>::try_from_value(iter.next().unwrap())?,)*
$(iter.map(|el| Input::try_from_value(el)).collect::<Result<Vec<$var_ty>, _>>()?,)*
))
}
fn tensor_data_type(idx: usize) -> Option<ONNXTensorElementDataType> {
[
$($positional_ty::dtype_id(),)*
$($var_ty::dtype_id())*
][idx]
}
fn characteristic(index: usize) -> OrtCustomOpInputOutputCharacteristic {
if index < Self::NUM_POSITIONAL {
[
$($positional_ty::characteristic(),)*
$($var_ty::characteristic())*
][index]
} else if $is_variadic {
OrtCustomOpInputOutputCharacteristic_INPUT_OUTPUT_VARIADIC
} else {
panic!("Provided index '{}' is out of range", index)
}
}
}
};
}
impl_inputs!(1, false, | A);
impl_inputs!(2, false, | A, B);
impl_inputs!(3, false, | A, B, C);
impl_inputs!(4, false, | A, B, C, D);
impl_inputs!(5, false, | A, B, C, D, E);
impl_inputs!(6, false, | A, B, C, D, E, F);
impl_inputs!(7, false, | A, B, C, D, E, F, G);
impl_inputs!(8, false, | A, B, C, D, E, F, G, H);
impl_inputs!(9, false, | A, B, C, D, E, F, G, H, I);
impl_inputs!(10, false, | A, B, C, D, E, F, G, H, I, J);
impl_inputs!(1, true, Z | A);
impl_inputs!(2, true, Z | A, B);
impl_inputs!(3, true, Z | A, B, C);
impl_inputs!(4, true, Z | A, B, C, D);
impl_inputs!(5, true, Z | A, B, C, D, E);
impl_inputs!(6, true, Z | A, B, C, D, E, F);
impl_inputs!(7, true, Z | A, B, C, D, E, F, G);
impl_inputs!(8, true, Z | A, B, C, D, E, F, G, H);
impl_inputs!(9, true, Z | A, B, C, D, E, F, G, H, I);
impl_inputs!(10, true, Z | A, B, C, D, E, F, G, H, I, J);
impl<'s> OnnxTensorDtype for ArrayViewD<'s, &'s str> {
fn dtype_id() -> Option<ONNXTensorElementDataType> {
Some(crate::bindings::ONNXTensorElementDataType_ONNX_TENSOR_ELEMENT_DATA_TYPE_STRING)
}
}
impl<T> OnnxTensorDtype for Option<T>
where
T: OnnxTensorDtype,
{
fn dtype_id() -> Option<ONNXTensorElementDataType> {
T::dtype_id()
}
}
macro_rules! impl_onnx_tensor_dtype {
($ty:ty, $ident:ident) => {
impl<'s> OnnxTensorDtype for ArrayViewD<'s, $ty> {
fn dtype_id() -> Option<ONNXTensorElementDataType> {
Some(crate::bindings::$ident)
}
}
};
}
#[rustfmt::skip] impl_onnx_tensor_dtype!(f32, ONNXTensorElementDataType_ONNX_TENSOR_ELEMENT_DATA_TYPE_FLOAT);
#[rustfmt::skip] impl_onnx_tensor_dtype!(f64, ONNXTensorElementDataType_ONNX_TENSOR_ELEMENT_DATA_TYPE_DOUBLE);
#[rustfmt::skip] impl_onnx_tensor_dtype!(bool, ONNXTensorElementDataType_ONNX_TENSOR_ELEMENT_DATA_TYPE_BOOL);
#[rustfmt::skip] impl_onnx_tensor_dtype!(u8, ONNXTensorElementDataType_ONNX_TENSOR_ELEMENT_DATA_TYPE_UINT8);
#[rustfmt::skip] impl_onnx_tensor_dtype!(u16, ONNXTensorElementDataType_ONNX_TENSOR_ELEMENT_DATA_TYPE_UINT16);
#[rustfmt::skip] impl_onnx_tensor_dtype!(u32, ONNXTensorElementDataType_ONNX_TENSOR_ELEMENT_DATA_TYPE_UINT32);
#[rustfmt::skip] impl_onnx_tensor_dtype!(u64, ONNXTensorElementDataType_ONNX_TENSOR_ELEMENT_DATA_TYPE_UINT64);
#[rustfmt::skip] impl_onnx_tensor_dtype!(i8, ONNXTensorElementDataType_ONNX_TENSOR_ELEMENT_DATA_TYPE_INT8);
#[rustfmt::skip] impl_onnx_tensor_dtype!(i16, ONNXTensorElementDataType_ONNX_TENSOR_ELEMENT_DATA_TYPE_INT16);
#[rustfmt::skip] impl_onnx_tensor_dtype!(i32, ONNXTensorElementDataType_ONNX_TENSOR_ELEMENT_DATA_TYPE_INT32);
#[rustfmt::skip] impl_onnx_tensor_dtype!(i64, ONNXTensorElementDataType_ONNX_TENSOR_ELEMENT_DATA_TYPE_INT64);