mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-30 20:18:08 +00:00
Trim DataTypeImpl binary size (#9813)
* De-virtualize DataTypeImpl::AsXType() functions. * Refactor helpers.
This commit is contained in:
parent
567749b2dc
commit
bcc6ab29f6
2 changed files with 118 additions and 174 deletions
|
|
@ -85,11 +85,12 @@ class DataTypeImpl {
|
|||
kTensor = 2,
|
||||
kTensorSequence = 3,
|
||||
kSparseTensor = 4,
|
||||
kOptional = 5
|
||||
kOptional = 5,
|
||||
kPrimitive = 6,
|
||||
};
|
||||
|
||||
GeneralType type_;
|
||||
size_t size_;
|
||||
const GeneralType type_;
|
||||
const size_t size_;
|
||||
|
||||
protected:
|
||||
DataTypeImpl(GeneralType type, size_t size) : type_{type}, size_{size} {}
|
||||
|
|
@ -134,37 +135,33 @@ class DataTypeImpl {
|
|||
return type_ == GeneralType::kOptional;
|
||||
}
|
||||
|
||||
// Returns this if this is of tensor-type and null otherwise
|
||||
virtual const TensorTypeBase* AsTensorType() const {
|
||||
return nullptr;
|
||||
bool IsNonTensorType() const {
|
||||
return type_ == GeneralType::kNonTensor;
|
||||
}
|
||||
|
||||
virtual const SequenceTensorTypeBase* AsSequenceTensorType() const {
|
||||
return nullptr;
|
||||
bool IsPrimitiveDataType() const {
|
||||
return type_ == GeneralType::kPrimitive;
|
||||
}
|
||||
|
||||
// Returns this if this is of tensor-type and null otherwise
|
||||
const TensorTypeBase* AsTensorType() const;
|
||||
|
||||
const SequenceTensorTypeBase* AsSequenceTensorType() const;
|
||||
|
||||
#if !defined(DISABLE_SPARSE_TENSORS)
|
||||
// Returns this if this is of sparse-tensor-type and null otherwise
|
||||
virtual const SparseTensorTypeBase* AsSparseTensorType() const {
|
||||
return nullptr;
|
||||
}
|
||||
const SparseTensorTypeBase* AsSparseTensorType() const;
|
||||
#endif
|
||||
|
||||
#if !defined(DISABLE_OPTIONAL_TYPE)
|
||||
virtual const OptionalTypeBase* AsOptionalType() const {
|
||||
return nullptr;
|
||||
}
|
||||
const OptionalTypeBase* AsOptionalType() const;
|
||||
#endif
|
||||
|
||||
virtual const NonTensorTypeBase* AsNonTensorType() const {
|
||||
return nullptr;
|
||||
}
|
||||
const NonTensorTypeBase* AsNonTensorType() const;
|
||||
|
||||
// Returns this if this is one of the primitive data types (specialization of PrimitiveDataTypeBase)
|
||||
// and null otherwise
|
||||
virtual const PrimitiveDataTypeBase* AsPrimitiveDataType() const {
|
||||
return nullptr;
|
||||
}
|
||||
const PrimitiveDataTypeBase* AsPrimitiveDataType() const;
|
||||
|
||||
// Return the type meta that we are using in the runtime.
|
||||
template <typename T>
|
||||
|
|
@ -233,8 +230,9 @@ namespace data_types_internal {
|
|||
///
|
||||
|
||||
template <typename T>
|
||||
constexpr ONNX_NAMESPACE::TensorProto_DataType ToTensorDataType();
|
||||
|
||||
constexpr ONNX_NAMESPACE::TensorProto_DataType ToTensorDataType() {
|
||||
return ONNX_NAMESPACE::TensorProto_DataType_UNDEFINED;
|
||||
}
|
||||
template <>
|
||||
constexpr ONNX_NAMESPACE::TensorProto_DataType ToTensorDataType<float>() {
|
||||
return ONNX_NAMESPACE::TensorProto_DataType_FLOAT;
|
||||
|
|
@ -242,65 +240,56 @@ constexpr ONNX_NAMESPACE::TensorProto_DataType ToTensorDataType<float>() {
|
|||
template <>
|
||||
constexpr ONNX_NAMESPACE::TensorProto_DataType ToTensorDataType<uint8_t>() {
|
||||
return ONNX_NAMESPACE::TensorProto_DataType_UINT8;
|
||||
};
|
||||
}
|
||||
template <>
|
||||
constexpr ONNX_NAMESPACE::TensorProto_DataType ToTensorDataType<int8_t>() {
|
||||
return ONNX_NAMESPACE::TensorProto_DataType_INT8;
|
||||
};
|
||||
}
|
||||
template <>
|
||||
constexpr ONNX_NAMESPACE::TensorProto_DataType ToTensorDataType<uint16_t>() {
|
||||
return ONNX_NAMESPACE::TensorProto_DataType_UINT16;
|
||||
};
|
||||
}
|
||||
template <>
|
||||
constexpr ONNX_NAMESPACE::TensorProto_DataType ToTensorDataType<int16_t>() {
|
||||
return ONNX_NAMESPACE::TensorProto_DataType_INT16;
|
||||
};
|
||||
}
|
||||
template <>
|
||||
constexpr ONNX_NAMESPACE::TensorProto_DataType ToTensorDataType<int32_t>() {
|
||||
return ONNX_NAMESPACE::TensorProto_DataType_INT32;
|
||||
};
|
||||
}
|
||||
template <>
|
||||
constexpr ONNX_NAMESPACE::TensorProto_DataType ToTensorDataType<int64_t>() {
|
||||
return ONNX_NAMESPACE::TensorProto_DataType_INT64;
|
||||
};
|
||||
}
|
||||
template <>
|
||||
constexpr ONNX_NAMESPACE::TensorProto_DataType ToTensorDataType<std::string>() {
|
||||
return ONNX_NAMESPACE::TensorProto_DataType_STRING;
|
||||
};
|
||||
}
|
||||
template <>
|
||||
constexpr ONNX_NAMESPACE::TensorProto_DataType ToTensorDataType<bool>() {
|
||||
return ONNX_NAMESPACE::TensorProto_DataType_BOOL;
|
||||
};
|
||||
}
|
||||
template <>
|
||||
constexpr ONNX_NAMESPACE::TensorProto_DataType ToTensorDataType<MLFloat16>() {
|
||||
return ONNX_NAMESPACE::TensorProto_DataType_FLOAT16;
|
||||
};
|
||||
}
|
||||
template <>
|
||||
constexpr ONNX_NAMESPACE::TensorProto_DataType ToTensorDataType<double>() {
|
||||
return ONNX_NAMESPACE::TensorProto_DataType_DOUBLE;
|
||||
};
|
||||
}
|
||||
template <>
|
||||
constexpr ONNX_NAMESPACE::TensorProto_DataType ToTensorDataType<uint32_t>() {
|
||||
return ONNX_NAMESPACE::TensorProto_DataType_UINT32;
|
||||
};
|
||||
}
|
||||
template <>
|
||||
constexpr ONNX_NAMESPACE::TensorProto_DataType ToTensorDataType<uint64_t>() {
|
||||
return ONNX_NAMESPACE::TensorProto_DataType_UINT64;
|
||||
};
|
||||
}
|
||||
template <>
|
||||
constexpr ONNX_NAMESPACE::TensorProto_DataType ToTensorDataType<BFloat16>() {
|
||||
return ONNX_NAMESPACE::TensorProto_DataType_BFLOAT16;
|
||||
}
|
||||
|
||||
// There is a specialization only for one
|
||||
// type argument.
|
||||
template <typename... Types>
|
||||
struct TensorElementTypeSetter {
|
||||
static void SetTensorElementType(ONNX_NAMESPACE::TypeProto&);
|
||||
static void SetMapKeyType(ONNX_NAMESPACE::TypeProto&);
|
||||
static int32_t GetElementType();
|
||||
};
|
||||
|
||||
/// Is a given type on the list of types?
|
||||
/// Accepts a list of types and the first argument is the type
|
||||
/// We are checking if it is listed among those that follow
|
||||
|
|
@ -367,34 +356,47 @@ struct GetMLDataType<T, false> {
|
|||
}
|
||||
};
|
||||
|
||||
struct TensorTypeHelper {
|
||||
static void Set(ONNX_NAMESPACE::TensorProto_DataType element_type,
|
||||
ONNX_NAMESPACE::TypeProto& proto) {
|
||||
proto.mutable_tensor_type()->set_elem_type(element_type);
|
||||
}
|
||||
};
|
||||
|
||||
#if !defined(DISABLE_SPARSE_TENSORS)
|
||||
struct SparseTensorTypeHelper {
|
||||
static void Set(ONNX_NAMESPACE::TensorProto_DataType element_type,
|
||||
ONNX_NAMESPACE::TypeProto& proto) {
|
||||
proto.mutable_sparse_tensor_type()->set_elem_type(element_type);
|
||||
}
|
||||
};
|
||||
#endif // !defined(DISABLE_SPARSE_TENSORS)
|
||||
|
||||
#if !defined(DISABLE_ML_OPS)
|
||||
/// MapTypes helper API
|
||||
/// K should always be one of the primitive data types
|
||||
/// V can be either a primitive type (in which case it is a tensor)
|
||||
/// or other preregistered types
|
||||
/// Map helpers
|
||||
|
||||
void CopyMutableMapValue(const ONNX_NAMESPACE::TypeProto&,
|
||||
ONNX_NAMESPACE::TypeProto&);
|
||||
|
||||
template <typename K, typename V>
|
||||
struct SetMapTypes {
|
||||
static void Set(ONNX_NAMESPACE::TypeProto& proto) {
|
||||
TensorElementTypeSetter<K>::SetMapKeyType(proto);
|
||||
MLDataType dt = GetMLDataType<V, IsTensorContainedType<V>::value>::Get();
|
||||
const auto* value_proto = dt->GetTypeProto();
|
||||
#ifdef ORT_NO_RTTI
|
||||
struct MapTypeHelper {
|
||||
// V can be either a primitive type (in which case it is a tensor)
|
||||
// or other preregistered types
|
||||
template <typename V>
|
||||
static MLDataType GetValueType() {
|
||||
return GetMLDataType<V, IsTensorContainedType<V>::value>::Get();
|
||||
}
|
||||
|
||||
static void Set(ONNX_NAMESPACE::TensorProto_DataType key_type, const ONNX_NAMESPACE::TypeProto* value_proto,
|
||||
ONNX_NAMESPACE::TypeProto& proto) {
|
||||
ORT_ENFORCE(value_proto != nullptr, "expected a registered ONNX type");
|
||||
#else
|
||||
ORT_ENFORCE(value_proto != nullptr, typeid(V).name(),
|
||||
" expected to be a registered ONNX type");
|
||||
#endif
|
||||
proto.mutable_map_type()->set_key_type(key_type);
|
||||
CopyMutableMapValue(*value_proto, proto);
|
||||
}
|
||||
};
|
||||
#endif
|
||||
|
||||
/// Sequence helpers
|
||||
///
|
||||
|
||||
// Element type is a primitive type so we set it to a tensor<elemT>
|
||||
void CopyMutableSeqElement(const ONNX_NAMESPACE::TypeProto&,
|
||||
ONNX_NAMESPACE::TypeProto&);
|
||||
|
|
@ -402,13 +404,12 @@ void CopyMutableSeqElement(const ONNX_NAMESPACE::TypeProto&,
|
|||
// helper to create TypeProto with minimal binary size impact
|
||||
struct SequenceTypeHelper {
|
||||
template <typename T>
|
||||
static const onnx::TypeProto* GetElemType() {
|
||||
MLDataType dt = GetMLDataType<T, IsTensorContainedType<T>::value>::Get();
|
||||
const auto* elem_proto = dt->GetTypeProto();
|
||||
return elem_proto; // check for nullptr is in Set
|
||||
static MLDataType GetElemType() {
|
||||
return GetMLDataType<T, IsTensorContainedType<T>::value>::Get();
|
||||
}
|
||||
|
||||
static void Set(const onnx::TypeProto* elem_proto, ONNX_NAMESPACE::TypeProto& proto) {
|
||||
static void Set(const ONNX_NAMESPACE::TypeProto* elem_proto,
|
||||
ONNX_NAMESPACE::TypeProto& proto) {
|
||||
ORT_ENFORCE(elem_proto != nullptr, "expected a registered ONNX type");
|
||||
CopyMutableSeqElement(*elem_proto, proto);
|
||||
}
|
||||
|
|
@ -422,27 +423,23 @@ void CopyMutableOptionalElement(const ONNX_NAMESPACE::TypeProto&,
|
|||
// helper to create TypeProto with minimal binary size impact
|
||||
struct OptionalTypeHelper {
|
||||
template <typename T, typename elemT>
|
||||
static const onnx::TypeProto* GetElemType() {
|
||||
const onnx::TypeProto* elem_proto = nullptr;
|
||||
if (std::is_same<T, Tensor>::value) {
|
||||
MLDataType dt = DataTypeImpl::GetTensorType<elemT>();
|
||||
elem_proto = dt->GetTypeProto();
|
||||
} else if (std::is_same<T, TensorSeq>::value) {
|
||||
MLDataType dt = DataTypeImpl::GetSequenceTensorType<elemT>();
|
||||
elem_proto = dt->GetTypeProto();
|
||||
static MLDataType GetElemType() {
|
||||
if constexpr (std::is_same<T, Tensor>::value) {
|
||||
return DataTypeImpl::GetTensorType<elemT>();
|
||||
} else {
|
||||
static_assert(std::is_same<T, TensorSeq>::value, "Unsupported element type for optional type");
|
||||
return DataTypeImpl::GetSequenceTensorType<elemT>();
|
||||
}
|
||||
|
||||
return elem_proto; // check for nullptr is in Set
|
||||
}
|
||||
|
||||
static void Set(const onnx::TypeProto* elem_proto, ONNX_NAMESPACE::TypeProto& proto) {
|
||||
ORT_ENFORCE(elem_proto != nullptr, " unregistered or unsupported ORT type for Optional");
|
||||
ORT_ENFORCE(elem_proto != nullptr, "expected a registered ONNX type");
|
||||
CopyMutableOptionalElement(*elem_proto, proto);
|
||||
}
|
||||
};
|
||||
|
||||
/// OpaqueTypes helpers
|
||||
///
|
||||
|
||||
void AssignOpaqueDomainName(const char* domain, const char* name,
|
||||
ONNX_NAMESPACE::TypeProto& proto);
|
||||
|
||||
|
|
@ -458,10 +455,6 @@ class TensorTypeBase : public DataTypeImpl {
|
|||
/// where TypeProto was created ad-hoc and not queried from MLDataType
|
||||
bool IsCompatible(const ONNX_NAMESPACE::TypeProto& type_proto) const override;
|
||||
|
||||
const TensorTypeBase* AsTensorType() const override {
|
||||
return this;
|
||||
}
|
||||
|
||||
DeleteFunc GetDeleteFunc() const override;
|
||||
|
||||
const ONNX_NAMESPACE::TypeProto* GetTypeProto() const override;
|
||||
|
|
@ -513,7 +506,7 @@ class TensorType : public TensorTypeBase {
|
|||
private:
|
||||
TensorType() {
|
||||
using namespace data_types_internal;
|
||||
TensorElementTypeSetter<elemT>::SetTensorElementType(this->MutableTypeProto());
|
||||
TensorTypeHelper::Set(ToTensorDataType<elemT>(), MutableTypeProto());
|
||||
}
|
||||
};
|
||||
|
||||
|
|
@ -562,10 +555,6 @@ class SparseTensorTypeBase : public DataTypeImpl {
|
|||
public:
|
||||
static MLDataType Type();
|
||||
|
||||
const SparseTensorTypeBase* AsSparseTensorType() const override {
|
||||
return this;
|
||||
}
|
||||
|
||||
bool IsCompatible(const ONNX_NAMESPACE::TypeProto& type_proto) const override;
|
||||
|
||||
DeleteFunc GetDeleteFunc() const override;
|
||||
|
|
@ -606,7 +595,7 @@ class SparseTensorType : public SparseTensorTypeBase {
|
|||
private:
|
||||
SparseTensorType() {
|
||||
using namespace data_types_internal;
|
||||
TensorElementTypeSetter<elemT>::SetSparseTensorElementType(MutableTypeProto());
|
||||
SparseTensorTypeHelper::Set(ToTensorDataType<elemT>(), MutableTypeProto());
|
||||
}
|
||||
};
|
||||
|
||||
|
|
@ -619,10 +608,6 @@ class OptionalTypeBase : public DataTypeImpl {
|
|||
public:
|
||||
static MLDataType Type();
|
||||
|
||||
const OptionalTypeBase* AsOptionalType() const override {
|
||||
return this;
|
||||
}
|
||||
|
||||
bool IsCompatible(const ONNX_NAMESPACE::TypeProto& type_proto) const override;
|
||||
|
||||
DeleteFunc GetDeleteFunc() const override {
|
||||
|
|
@ -673,14 +658,7 @@ class OptionalType :
|
|||
"Requires one of the tensor fundamental types");
|
||||
|
||||
MLDataType GetElementType() const override {
|
||||
if (std::is_same<T, Tensor>::value) {
|
||||
return DataTypeImpl::GetTensorType<elemT>();
|
||||
} else if (std::is_same<T, TensorSeq>::value) {
|
||||
return DataTypeImpl::GetSequenceTensorType<elemT>();
|
||||
} else {
|
||||
// Will not reach here
|
||||
ORT_ENFORCE(false, "Unsupported optional type");
|
||||
}
|
||||
return data_types_internal::OptionalTypeHelper::GetElemType<T, elemT>();
|
||||
}
|
||||
#endif
|
||||
|
||||
|
|
@ -692,7 +670,7 @@ class OptionalType :
|
|||
#endif
|
||||
{
|
||||
using namespace data_types_internal;
|
||||
OptionalTypeHelper::Set(OptionalTypeHelper::GetElemType<T, elemT>(), MutableTypeProto());
|
||||
OptionalTypeHelper::Set(OptionalTypeHelper::GetElemType<T, elemT>()->GetTypeProto(), MutableTypeProto());
|
||||
}
|
||||
}; // namespace onnxruntime
|
||||
|
||||
|
|
@ -726,10 +704,6 @@ class NonTensorTypeBase : public DataTypeImpl {
|
|||
|
||||
const ONNX_NAMESPACE::TypeProto* GetTypeProto() const override;
|
||||
|
||||
const NonTensorTypeBase* AsNonTensorType() const override {
|
||||
return this;
|
||||
}
|
||||
|
||||
// \brief Override for Non-tensor types to initialize non-tensor CPP
|
||||
// data representation from data. The caller of the interface
|
||||
// should have a shared definition of the data which is used to initialize
|
||||
|
|
@ -819,7 +793,9 @@ class MapType : public NonTensorType<CPPType> {
|
|||
private:
|
||||
MapType() {
|
||||
using namespace data_types_internal;
|
||||
SetMapTypes<typename CPPType::key_type, typename CPPType::mapped_type>::Set(this->MutableTypeProto());
|
||||
MapTypeHelper::Set(ToTensorDataType<typename CPPType::key_type>(),
|
||||
MapTypeHelper::GetValueType<typename CPPType::mapped_type>()->GetTypeProto(),
|
||||
this->MutableTypeProto());
|
||||
}
|
||||
};
|
||||
#endif
|
||||
|
|
@ -845,14 +821,14 @@ class SequenceType : public NonTensorType<CPPType> {
|
|||
private:
|
||||
SequenceType() {
|
||||
using namespace data_types_internal;
|
||||
SequenceTypeHelper::Set(SequenceTypeHelper::GetElemType<typename CPPType::value_type>(),
|
||||
SequenceTypeHelper::Set(SequenceTypeHelper::GetElemType<typename CPPType::value_type>()->GetTypeProto(),
|
||||
this->MutableTypeProto());
|
||||
}
|
||||
};
|
||||
|
||||
/**
|
||||
* \brief SequenceTensorTypeBase serves as a base type class for
|
||||
* Tensor sequences. Akin TensorTypeBase.
|
||||
* Tensor sequences. Akin to TensorTypeBase.
|
||||
* Runtime representation is always TensorSeq.
|
||||
*/
|
||||
class SequenceTensorTypeBase : public DataTypeImpl {
|
||||
|
|
@ -861,10 +837,6 @@ class SequenceTensorTypeBase : public DataTypeImpl {
|
|||
|
||||
bool IsCompatible(const ONNX_NAMESPACE::TypeProto& type_proto) const override;
|
||||
|
||||
const SequenceTensorTypeBase* AsSequenceTensorType() const override {
|
||||
return this;
|
||||
}
|
||||
|
||||
virtual MLDataType GetElementType() const {
|
||||
// should never reach here.
|
||||
ORT_NOT_IMPLEMENTED(__FUNCTION__, " is not implemented");
|
||||
|
|
@ -914,18 +886,19 @@ class SequenceTensorType : public SequenceTensorTypeBase {
|
|||
private:
|
||||
SequenceTensorType() {
|
||||
using namespace data_types_internal;
|
||||
SequenceTypeHelper::Set(SequenceTypeHelper::GetElemType<TensorElemType>(), MutableTypeProto());
|
||||
SequenceTypeHelper::Set(SequenceTypeHelper::GetElemType<TensorElemType>()->GetTypeProto(),
|
||||
MutableTypeProto());
|
||||
}
|
||||
};
|
||||
|
||||
/**
|
||||
* \brief OpaqueType
|
||||
*
|
||||
* \param T - cpp runtume that implements the Opaque type
|
||||
* \tparam T - cpp runtume that implements the Opaque type
|
||||
*
|
||||
* \param const char D[] - domain must be extern to be unique
|
||||
* \tparam const char D[] - domain must be extern to be unique
|
||||
*
|
||||
* \param const char N[] - name must be extern to be unique
|
||||
* \tparam const char N[] - name must be extern to be unique
|
||||
*
|
||||
* \details Only one CPP type can be associated with a particular
|
||||
* OpaqueType registration
|
||||
|
|
@ -968,10 +941,6 @@ class PrimitiveDataTypeBase : public DataTypeImpl {
|
|||
return false;
|
||||
}
|
||||
|
||||
const PrimitiveDataTypeBase* AsPrimitiveDataType() const override final {
|
||||
return this;
|
||||
}
|
||||
|
||||
const ONNX_NAMESPACE::TypeProto* GetTypeProto() const final {
|
||||
return nullptr;
|
||||
}
|
||||
|
|
@ -982,10 +951,10 @@ class PrimitiveDataTypeBase : public DataTypeImpl {
|
|||
|
||||
protected:
|
||||
PrimitiveDataTypeBase(size_t size, int32_t data_type)
|
||||
: DataTypeImpl{GeneralType::kNonTensor, size}, data_type_{data_type} {}
|
||||
: DataTypeImpl{GeneralType::kPrimitive, size}, data_type_{data_type} {}
|
||||
|
||||
private:
|
||||
int32_t data_type_;
|
||||
const int32_t data_type_;
|
||||
};
|
||||
|
||||
/**
|
||||
|
|
@ -1013,10 +982,38 @@ class PrimitiveDataType : public PrimitiveDataTypeBase {
|
|||
private:
|
||||
PrimitiveDataType()
|
||||
: PrimitiveDataTypeBase{sizeof(T),
|
||||
data_types_internal::TensorElementTypeSetter<T>::GetElementType()} {
|
||||
data_types_internal::ToTensorDataType<T>()} {
|
||||
}
|
||||
};
|
||||
|
||||
inline const TensorTypeBase* DataTypeImpl::AsTensorType() const {
|
||||
return IsTensorType() ? static_cast<const TensorTypeBase*>(this) : nullptr;
|
||||
}
|
||||
|
||||
inline const SequenceTensorTypeBase* DataTypeImpl::AsSequenceTensorType() const {
|
||||
return IsTensorSequenceType() ? static_cast<const SequenceTensorTypeBase*>(this) : nullptr;
|
||||
}
|
||||
|
||||
#if !defined(DISABLE_SPARSE_TENSORS)
|
||||
inline const SparseTensorTypeBase* DataTypeImpl::AsSparseTensorType() const {
|
||||
return IsSparseTensorType() ? static_cast<const SparseTensorTypeBase*>(this) : nullptr;
|
||||
}
|
||||
#endif
|
||||
|
||||
#if !defined(DISABLE_OPTIONAL_TYPE)
|
||||
inline const OptionalTypeBase* DataTypeImpl::AsOptionalType() const {
|
||||
return IsOptionalType() ? static_cast<const OptionalTypeBase*>(this) : nullptr;
|
||||
}
|
||||
#endif
|
||||
|
||||
inline const NonTensorTypeBase* DataTypeImpl::AsNonTensorType() const {
|
||||
return IsNonTensorType() ? static_cast<const NonTensorTypeBase*>(this) : nullptr;
|
||||
}
|
||||
|
||||
inline const PrimitiveDataTypeBase* DataTypeImpl::AsPrimitiveDataType() const {
|
||||
return IsPrimitiveDataType() ? static_cast<const PrimitiveDataTypeBase*>(this) : nullptr;
|
||||
}
|
||||
|
||||
// Explicit specialization of base class template function
|
||||
// is only possible within the enclosing namespace scope,
|
||||
// thus a simple way to pre-instantiate a given template
|
||||
|
|
|
|||
|
|
@ -66,59 +66,6 @@ MLDataType DataTypeImpl::GetType<TensorSeq>() {
|
|||
|
||||
namespace data_types_internal {
|
||||
|
||||
template <typename T>
|
||||
struct TensorElementTypeSetter<T> {
|
||||
static void SetTensorElementType(ONNX_NAMESPACE::TypeProto& proto) {
|
||||
proto.mutable_tensor_type()->set_elem_type(utils::ToTensorProtoElementType<T>());
|
||||
}
|
||||
|
||||
#if !defined(DISABLE_SPARSE_TENSORS)
|
||||
static void SetSparseTensorElementType(ONNX_NAMESPACE::TypeProto& proto) {
|
||||
proto.mutable_sparse_tensor_type()->set_elem_type(utils::ToTensorProtoElementType<T>());
|
||||
}
|
||||
#endif
|
||||
|
||||
#if !defined(DISABLE_ML_OPS)
|
||||
static void SetMapKeyType(ONNX_NAMESPACE::TypeProto& proto) {
|
||||
proto.mutable_map_type()->set_key_type(utils::ToTensorProtoElementType<T>());
|
||||
}
|
||||
#endif
|
||||
|
||||
constexpr static int32_t GetElementType() {
|
||||
return utils::ToTensorProtoElementType<T>();
|
||||
}
|
||||
};
|
||||
|
||||
// Pre-instantiate
|
||||
template struct
|
||||
TensorElementTypeSetter<float>;
|
||||
template struct
|
||||
TensorElementTypeSetter<uint8_t>;
|
||||
template struct
|
||||
TensorElementTypeSetter<int8_t>;
|
||||
template struct
|
||||
TensorElementTypeSetter<uint16_t>;
|
||||
template struct
|
||||
TensorElementTypeSetter<int16_t>;
|
||||
template struct
|
||||
TensorElementTypeSetter<int32_t>;
|
||||
template struct
|
||||
TensorElementTypeSetter<int64_t>;
|
||||
template struct
|
||||
TensorElementTypeSetter<std::string>;
|
||||
template struct
|
||||
TensorElementTypeSetter<bool>;
|
||||
template struct
|
||||
TensorElementTypeSetter<MLFloat16>;
|
||||
template struct
|
||||
TensorElementTypeSetter<double>;
|
||||
template struct
|
||||
TensorElementTypeSetter<uint32_t>;
|
||||
template struct
|
||||
TensorElementTypeSetter<uint64_t>;
|
||||
template struct
|
||||
TensorElementTypeSetter<BFloat16>;
|
||||
|
||||
#if !defined(DISABLE_ML_OPS)
|
||||
void CopyMutableMapValue(const ONNX_NAMESPACE::TypeProto& value_proto,
|
||||
ONNX_NAMESPACE::TypeProto& map_proto) {
|
||||
|
|
|
|||
Loading…
Reference in a new issue