diff --git a/include/custom_op/tensor_api.h b/include/custom_op/tensor_api.h index 0e6f221d9..7ff128609 100644 --- a/include/custom_op/tensor_api.h +++ b/include/custom_op/tensor_api.h @@ -136,13 +136,19 @@ class OrtEagerTensorStorage : public ITensorStorage { }; template -ONNXTensorElementDataType GetOrtDType(){ +constexpr ONNXTensorElementDataType GetOrtDType(){ if constexpr (std::is_same::value) return ONNX_TENSOR_ELEMENT_DATA_TYPE_BOOL; else if constexpr (std::is_same::value) return ONNX_TENSOR_ELEMENT_DATA_TYPE_FLOAT; else if constexpr (std::is_same::value) return ONNX_TENSOR_ELEMENT_DATA_TYPE_DOUBLE; +#if ORT_API_VERSION >= 16 + else if constexpr (std::is_same::value) + return ONNX_TENSOR_ELEMENT_DATA_TYPE_FLOAT16; + else if constexpr (std::is_same::value) + return ONNX_TENSOR_ELEMENT_DATA_TYPE_BFLOAT16; +#endif // ORT_API_VERSION >= 16 else if constexpr (std::is_same::value) return ONNX_TENSOR_ELEMENT_DATA_TYPE_UINT8; else if constexpr (std::is_same::value) @@ -159,10 +165,13 @@ ONNXTensorElementDataType GetOrtDType(){ return ONNX_TENSOR_ELEMENT_DATA_TYPE_UINT64; else if constexpr (std::is_same::value) return ONNX_TENSOR_ELEMENT_DATA_TYPE_INT64; - else if constexpr (std::is_same::value) + else if constexpr (std::is_same::value || std::is_same::value) return ONNX_TENSOR_ELEMENT_DATA_TYPE_STRING; - ORTX_CXX_API_THROW("Unexpected type", ORT_RUNTIME_EXCEPTION); - return ONNX_TENSOR_ELEMENT_DATA_TYPE_FLOAT; + else { + // Type-dependent false value. Needed for compilers which don't allow static_assert(false) here (see CWG 2518). + constexpr bool always_false = !std::is_same_v; + static_assert(always_false, "Invalid type"); + } } class TensorBase : public Arg {