mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-27 20:02:15 +00:00
Consolidate enabled/default kernel def type constraints (#13034)
Consolidate enabled/default kernel def type constraint types into enabled.
This commit is contained in:
parent
440f31668f
commit
5c89c37f7f
26 changed files with 89 additions and 282 deletions
|
|
@ -3,17 +3,17 @@
|
|||
|
||||
#pragma once
|
||||
|
||||
#include <limits.h>
|
||||
#include <memory>
|
||||
#include <optional>
|
||||
#include <string>
|
||||
#include <unordered_map>
|
||||
#include <vector>
|
||||
#include <limits.h>
|
||||
|
||||
#include "core/common/common.h"
|
||||
#include "core/common/optional.h"
|
||||
#include "core/graph/basic_types.h"
|
||||
#include "core/framework/data_types.h"
|
||||
#include "core/framework/allocator.h"
|
||||
#include "core/framework/data_types.h"
|
||||
#include "core/graph/basic_types.h"
|
||||
|
||||
namespace onnxruntime {
|
||||
class KernelDefBuilder;
|
||||
|
|
@ -53,16 +53,9 @@ class KernelDef {
|
|||
return provider_type_;
|
||||
}
|
||||
|
||||
// TODO(edgchen1) do we need both TypeConstraints() and EnabledTypeConstraints()?
|
||||
|
||||
// type constraints with types supported by default
|
||||
const std::unordered_map<std::string, std::vector<MLDataType>>& TypeConstraints() const {
|
||||
return default_type_constraints_;
|
||||
}
|
||||
|
||||
// type constraints with types supported in this build
|
||||
const std::unordered_map<std::string, std::vector<MLDataType>>& EnabledTypeConstraints() const {
|
||||
return enabled_type_constraints_;
|
||||
const std::unordered_map<std::string, std::vector<MLDataType>>& TypeConstraints() const {
|
||||
return type_constraints_;
|
||||
}
|
||||
|
||||
const std::vector<std::pair<int, int>>& MayInplace() const {
|
||||
|
|
@ -73,7 +66,7 @@ class KernelDef {
|
|||
return alias_map_;
|
||||
}
|
||||
|
||||
const optional<std::pair<int, int>>& VariadicAlias() const {
|
||||
const std::optional<std::pair<int, int>>& VariadicAlias() const {
|
||||
return variadic_alias_offsets_;
|
||||
}
|
||||
|
||||
|
|
@ -130,12 +123,9 @@ class KernelDef {
|
|||
// The type of the execution provider.
|
||||
std::string provider_type_;
|
||||
|
||||
// The data types that are supported by default for inputs/outputs.
|
||||
// Key is input/output/type constraint name defined in op schema, Value are supported types.
|
||||
std::unordered_map<std::string, std::vector<MLDataType>> default_type_constraints_;
|
||||
|
||||
// the type constraints that are supported in this build (enabled) for the kernel
|
||||
std::unordered_map<std::string, std::vector<MLDataType>> enabled_type_constraints_;
|
||||
// The data types that are supported in this build (enabled) for inputs/outputs.
|
||||
// Key is input/output/type constraint name defined in op schema, Value is supported types.
|
||||
std::unordered_map<std::string, std::vector<MLDataType>> type_constraints_;
|
||||
|
||||
// An element <i, j> means that output j reuses the memory of input i.
|
||||
std::vector<std::pair<int, int>> inplace_map_;
|
||||
|
|
@ -145,7 +135,7 @@ class KernelDef {
|
|||
|
||||
// This variable stores <input_offset, output_offset> for the variadic alias mapping
|
||||
// output 'i + output_offset' is an alias of input 'i + input_offset' for all i >= 0
|
||||
optional<std::pair<int, int>> variadic_alias_offsets_;
|
||||
std::optional<std::pair<int, int>> variadic_alias_offsets_;
|
||||
|
||||
// Require input tensors to be allocated contiguously.
|
||||
bool allocate_inputs_contiguously_ = false;
|
||||
|
|
@ -210,7 +200,7 @@ class KernelDefBuilder {
|
|||
/**
|
||||
The execution provider type of the kernel.
|
||||
*/
|
||||
KernelDefBuilder& Provider(onnxruntime::ProviderType provider_type);
|
||||
KernelDefBuilder& Provider(ProviderType provider_type);
|
||||
KernelDefBuilder& Provider(const char* provider_type);
|
||||
|
||||
/**
|
||||
|
|
@ -219,27 +209,16 @@ class KernelDefBuilder {
|
|||
|
||||
@param arg_name The arg name can be either op formal parameter name, say "X", or type
|
||||
argument name specified in op schema, say "T".
|
||||
@param default_types The types that are supported by default.
|
||||
@param enabled_types The types that are supported in this build.
|
||||
Possibly different from default_types when type reduction is enabled.
|
||||
@param types The types that are supported in this build.
|
||||
*/
|
||||
KernelDefBuilder& TypeConstraint(const std::string& arg_name,
|
||||
const std::vector<MLDataType>& default_types);
|
||||
KernelDefBuilder& TypeConstraint(const char* arg_name,
|
||||
const std::vector<MLDataType>& default_types);
|
||||
|
||||
KernelDefBuilder& TypeConstraint(const std::string& arg_name,
|
||||
const std::vector<MLDataType>& default_types,
|
||||
const std::vector<MLDataType>& enabled_types);
|
||||
KernelDefBuilder& TypeConstraint(const char* arg_name,
|
||||
const std::vector<MLDataType>& default_types,
|
||||
const std::vector<MLDataType>& enabled_types);
|
||||
KernelDefBuilder& TypeConstraint(const std::string& arg_name, std::vector<MLDataType> types);
|
||||
KernelDefBuilder& TypeConstraint(const char* arg_name, std::vector<MLDataType> types);
|
||||
|
||||
/**
|
||||
Like TypeConstraint but supports just a single type.
|
||||
*/
|
||||
KernelDefBuilder& TypeConstraint(const std::string& arg_name, MLDataType default_type);
|
||||
KernelDefBuilder& TypeConstraint(const char* arg_name, MLDataType default_type);
|
||||
KernelDefBuilder& TypeConstraint(const std::string& arg_name, MLDataType type);
|
||||
KernelDefBuilder& TypeConstraint(const char* arg_name, MLDataType type);
|
||||
|
||||
/**
|
||||
Inplace mapping from inputs to outputs allowed.
|
||||
|
|
@ -367,10 +346,6 @@ class KernelDefBuilder {
|
|||
}
|
||||
|
||||
private:
|
||||
KernelDefBuilder& TypeConstraintImpl(const std::string& arg_name,
|
||||
const std::vector<MLDataType>& default_types,
|
||||
const std::vector<MLDataType>* enabled_types = nullptr);
|
||||
|
||||
// we own the KernelDef until Build() is called.
|
||||
std::unique_ptr<KernelDef> kernel_def_;
|
||||
};
|
||||
|
|
|
|||
|
|
@ -52,9 +52,9 @@ bool KernelDef::IsConflict(const KernelDef& other) const {
|
|||
//only one case they don't conflict:
|
||||
//There is a type_constraint, it exists in both hands, but they don't overlap
|
||||
//check types
|
||||
const auto& other_types = other.default_type_constraints_;
|
||||
const auto& other_types = other.type_constraints_;
|
||||
bool type_has_conflict = true;
|
||||
for (const auto& it : default_type_constraints_) {
|
||||
for (const auto& it : type_constraints_) {
|
||||
auto iter = other_types.find(it.first);
|
||||
if (iter != other_types.end()) {
|
||||
if (!AreVectorsOverlap(it.second, iter->second)) {
|
||||
|
|
@ -106,7 +106,7 @@ KernelDefBuilder& KernelDefBuilder::SetName(const std::string& op_name) {
|
|||
}
|
||||
|
||||
KernelDefBuilder& KernelDefBuilder::SetName(const char* op_name) {
|
||||
kernel_def_->op_name_ = std::string(op_name);
|
||||
kernel_def_->op_name_ = std::string{op_name};
|
||||
return *this;
|
||||
}
|
||||
|
||||
|
|
@ -116,61 +116,42 @@ KernelDefBuilder& KernelDefBuilder::SetDomain(const std::string& domain) {
|
|||
}
|
||||
|
||||
KernelDefBuilder& KernelDefBuilder::SetDomain(const char* domain) {
|
||||
kernel_def_->op_domain_ = std::string(domain);
|
||||
kernel_def_->op_domain_ = std::string{domain};
|
||||
return *this;
|
||||
}
|
||||
|
||||
KernelDefBuilder& KernelDefBuilder::Provider(onnxruntime::ProviderType provider_type) {
|
||||
KernelDefBuilder& KernelDefBuilder::Provider(ProviderType provider_type) {
|
||||
kernel_def_->provider_type_ = provider_type;
|
||||
return *this;
|
||||
}
|
||||
|
||||
KernelDefBuilder& KernelDefBuilder::Provider(const char* provider_type) {
|
||||
kernel_def_->provider_type_ = std::string(provider_type);
|
||||
return *this;
|
||||
}
|
||||
|
||||
KernelDefBuilder& KernelDefBuilder::TypeConstraintImpl(const std::string& arg_name,
|
||||
const std::vector<MLDataType>& default_types,
|
||||
const std::vector<MLDataType>* enabled_types) {
|
||||
// use the enabled types list if provided
|
||||
kernel_def_->enabled_type_constraints_[arg_name] = enabled_types ? *enabled_types : default_types;
|
||||
kernel_def_->default_type_constraints_[arg_name] = default_types;
|
||||
kernel_def_->provider_type_ = std::string{provider_type};
|
||||
return *this;
|
||||
}
|
||||
|
||||
KernelDefBuilder& KernelDefBuilder::TypeConstraint(const std::string& arg_name,
|
||||
const std::vector<MLDataType>& default_types) {
|
||||
return TypeConstraintImpl(arg_name, default_types, nullptr);
|
||||
}
|
||||
|
||||
KernelDefBuilder& KernelDefBuilder::TypeConstraint(const char* arg_name,
|
||||
const std::vector<MLDataType>& default_types) {
|
||||
return TypeConstraintImpl(arg_name, default_types, nullptr);
|
||||
}
|
||||
|
||||
KernelDefBuilder& KernelDefBuilder::TypeConstraint(const std::string& arg_name,
|
||||
const std::vector<MLDataType>& default_types,
|
||||
const std::vector<MLDataType>& enabled_types) {
|
||||
return TypeConstraintImpl(arg_name, default_types, &enabled_types);
|
||||
}
|
||||
|
||||
KernelDefBuilder& KernelDefBuilder::TypeConstraint(const char* arg_name,
|
||||
const std::vector<MLDataType>& default_types,
|
||||
const std::vector<MLDataType>& enabled_types) {
|
||||
return TypeConstraintImpl(arg_name, default_types, &enabled_types);
|
||||
}
|
||||
|
||||
KernelDefBuilder& KernelDefBuilder::TypeConstraint(const std::string& arg_name,
|
||||
MLDataType default_type) {
|
||||
kernel_def_->enabled_type_constraints_[arg_name] = std::vector<MLDataType>{default_type};
|
||||
kernel_def_->default_type_constraints_[arg_name] = std::vector<MLDataType>{default_type};
|
||||
std::vector<MLDataType> types) {
|
||||
kernel_def_->type_constraints_.insert_or_assign(arg_name, std::move(types));
|
||||
return *this;
|
||||
}
|
||||
|
||||
KernelDefBuilder& KernelDefBuilder::TypeConstraint(const char* arg_name,
|
||||
MLDataType default_type) {
|
||||
return TypeConstraint(std::string(arg_name), default_type);
|
||||
std::vector<MLDataType> types) {
|
||||
kernel_def_->type_constraints_.insert_or_assign(std::string{arg_name}, std::move(types));
|
||||
return *this;
|
||||
}
|
||||
|
||||
KernelDefBuilder& KernelDefBuilder::TypeConstraint(const std::string& arg_name,
|
||||
MLDataType type) {
|
||||
std::vector<MLDataType> types{type};
|
||||
return TypeConstraint(arg_name, std::move(types));
|
||||
}
|
||||
|
||||
KernelDefBuilder& KernelDefBuilder::TypeConstraint(const char* arg_name,
|
||||
MLDataType type) {
|
||||
std::vector<MLDataType> types{type};
|
||||
return TypeConstraint(arg_name, std::move(types));
|
||||
}
|
||||
|
||||
KernelDefBuilder& KernelDefBuilder::MayInplace(const std::vector<std::pair<int, int>>& inplaces) {
|
||||
|
|
|
|||
|
|
@ -63,7 +63,7 @@ bool MatchKernelDefTypes(const Node& node,
|
|||
// for each type constraint
|
||||
// map type constraint to arg
|
||||
// check arg type against type constraint enabled types
|
||||
const auto& kernel_type_constraints = kernel_def.EnabledTypeConstraints();
|
||||
const auto& kernel_type_constraints = kernel_def.TypeConstraints();
|
||||
for (const auto& [kernel_type_str, enabled_types] : kernel_type_constraints) {
|
||||
gsl::span<const ArgTypeAndIndex> constraint_args{};
|
||||
ORT_THROW_IF_ERROR(kernel_type_str_resolver.ResolveKernelTypeStr(node, kernel_type_str,
|
||||
|
|
|
|||
|
|
@ -20,10 +20,6 @@ ORT_SPECIFY_OP_KERNEL_ARG_REQUIRED_TYPES_ALL_OPSETS(
|
|||
|
||||
namespace {
|
||||
|
||||
using OutputTypes =
|
||||
ORT_OP_KERNEL_ARG_DEFAULT_TYPE_LIST_ALL_OPSETS(
|
||||
kCpuExecutionProvider, kOnnxDomain, ConstantOfShape, Output, 0);
|
||||
|
||||
using EnabledOutputTypes =
|
||||
ORT_OP_KERNEL_ARG_ENABLED_TYPE_LIST_ALL_OPSETS(
|
||||
kCpuExecutionProvider, kOnnxDomain, ConstantOfShape, Output, 0);
|
||||
|
|
@ -76,7 +72,6 @@ ONNX_CPU_OPERATOR_KERNEL(
|
|||
KernelDefBuilder()
|
||||
.TypeConstraint("T1", DataTypeImpl::GetTensorType<int64_t>())
|
||||
.TypeConstraint("T2",
|
||||
BuildKernelDefConstraintsFromTypeList<OutputTypes>(),
|
||||
BuildKernelDefConstraintsFromTypeList<EnabledOutputTypes>()),
|
||||
ConstantOfShape);
|
||||
|
||||
|
|
|
|||
|
|
@ -59,28 +59,18 @@ ORT_SPECIFY_OP_KERNEL_ARG_DEFAULT_TYPES_ALL_OPSETS(
|
|||
int32_t, int64_t);
|
||||
}
|
||||
|
||||
using RandomNormalOutputTypes = ORT_OP_KERNEL_ARG_DEFAULT_TYPE_LIST_ALL_OPSETS(
|
||||
kCpuExecutionProvider, kOnnxDomain, RandomNormal, Output, 0);
|
||||
using EnabledRandomNormalOutputTypes = ORT_OP_KERNEL_ARG_ENABLED_TYPE_LIST_ALL_OPSETS(
|
||||
kCpuExecutionProvider, kOnnxDomain, RandomNormal, Output, 0);
|
||||
|
||||
using RandomUniformOutputTypes = ORT_OP_KERNEL_ARG_DEFAULT_TYPE_LIST_ALL_OPSETS(
|
||||
kCpuExecutionProvider, kOnnxDomain, RandomUniform, Output, 0);
|
||||
using EnabledRandomUniformOutputTypes = ORT_OP_KERNEL_ARG_ENABLED_TYPE_LIST_ALL_OPSETS(
|
||||
kCpuExecutionProvider, kOnnxDomain, RandomUniform, Output, 0);
|
||||
|
||||
using RandomNormalLikeOutputTypes = ORT_OP_KERNEL_ARG_DEFAULT_TYPE_LIST_ALL_OPSETS(
|
||||
kCpuExecutionProvider, kOnnxDomain, RandomNormalLike, Output, 0);
|
||||
using EnabledRandomNormalLikeOutputTypes = ORT_OP_KERNEL_ARG_ENABLED_TYPE_LIST_ALL_OPSETS(
|
||||
kCpuExecutionProvider, kOnnxDomain, RandomNormalLike, Output, 0);
|
||||
|
||||
using RandomUniformLikeOutputTypes = ORT_OP_KERNEL_ARG_DEFAULT_TYPE_LIST_ALL_OPSETS(
|
||||
kCpuExecutionProvider, kOnnxDomain, RandomUniformLike, Output, 0);
|
||||
using EnabledRandomUniformLikeOutputTypes = ORT_OP_KERNEL_ARG_ENABLED_TYPE_LIST_ALL_OPSETS(
|
||||
kCpuExecutionProvider, kOnnxDomain, RandomUniformLike, Output, 0);
|
||||
|
||||
using MultinomialOutputTypes = ORT_OP_KERNEL_ARG_DEFAULT_TYPE_LIST_ALL_OPSETS(
|
||||
kCpuExecutionProvider, kOnnxDomain, Multinomial, Output, 0);
|
||||
using EnabledMultinomialOutputTypes = ORT_OP_KERNEL_ARG_ENABLED_TYPE_LIST_ALL_OPSETS(
|
||||
kCpuExecutionProvider, kOnnxDomain, Multinomial, Output, 0);
|
||||
|
||||
|
|
@ -99,7 +89,6 @@ ONNX_CPU_OPERATOR_KERNEL(
|
|||
1,
|
||||
KernelDefBuilder()
|
||||
.TypeConstraint("T",
|
||||
BuildKernelDefConstraintsFromTypeList<RandomNormalOutputTypes>(),
|
||||
BuildKernelDefConstraintsFromTypeList<EnabledRandomNormalOutputTypes>()),
|
||||
RandomNormal);
|
||||
|
||||
|
|
@ -108,7 +97,6 @@ ONNX_CPU_OPERATOR_KERNEL(
|
|||
1,
|
||||
KernelDefBuilder()
|
||||
.TypeConstraint("T",
|
||||
BuildKernelDefConstraintsFromTypeList<RandomUniformOutputTypes>(),
|
||||
BuildKernelDefConstraintsFromTypeList<EnabledRandomUniformOutputTypes>()),
|
||||
RandomUniform);
|
||||
|
||||
|
|
@ -118,7 +106,6 @@ ONNX_CPU_OPERATOR_KERNEL(
|
|||
KernelDefBuilder()
|
||||
.TypeConstraint("T1", DataTypeImpl::AllTensorTypes())
|
||||
.TypeConstraint("T2",
|
||||
BuildKernelDefConstraintsFromTypeList<RandomNormalLikeOutputTypes>(),
|
||||
BuildKernelDefConstraintsFromTypeList<EnabledRandomNormalLikeOutputTypes>()),
|
||||
RandomNormalLike);
|
||||
|
||||
|
|
@ -128,7 +115,6 @@ ONNX_CPU_OPERATOR_KERNEL(
|
|||
KernelDefBuilder()
|
||||
.TypeConstraint("T1", DataTypeImpl::AllTensorTypes())
|
||||
.TypeConstraint("T2",
|
||||
BuildKernelDefConstraintsFromTypeList<RandomUniformLikeOutputTypes>(),
|
||||
BuildKernelDefConstraintsFromTypeList<EnabledRandomUniformLikeOutputTypes>()),
|
||||
RandomUniformLike);
|
||||
|
||||
|
|
@ -139,7 +125,6 @@ ONNX_CPU_OPERATOR_KERNEL(
|
|||
KernelDefBuilder()
|
||||
.TypeConstraint("T1", DataTypeImpl::GetTensorType<float>())
|
||||
.TypeConstraint("T2",
|
||||
BuildKernelDefConstraintsFromTypeList<MultinomialOutputTypes>(),
|
||||
BuildKernelDefConstraintsFromTypeList<EnabledMultinomialOutputTypes>()),
|
||||
Multinomial);
|
||||
|
||||
|
|
|
|||
|
|
@ -21,8 +21,6 @@ ORT_SPECIFY_OP_KERNEL_ARG_REQUIRED_TYPES_ALL_OPSETS(
|
|||
int32_t, int64_t);
|
||||
} // namespace op_kernel_type_control
|
||||
|
||||
using RangeDataTypes = ORT_OP_KERNEL_ARG_DEFAULT_TYPE_LIST_ALL_OPSETS(
|
||||
kCpuExecutionProvider, kOnnxDomain, Range, Input, 0);
|
||||
using EnabledRangeDataTypes = ORT_OP_KERNEL_ARG_ENABLED_TYPE_LIST_ALL_OPSETS(
|
||||
kCpuExecutionProvider, kOnnxDomain, Range, Input, 0);
|
||||
|
||||
|
|
@ -41,7 +39,6 @@ ONNX_OPERATOR_KERNEL_EX(
|
|||
kCpuExecutionProvider,
|
||||
KernelDefBuilder()
|
||||
.TypeConstraint("T",
|
||||
BuildKernelDefConstraintsFromTypeList<RangeDataTypes>(),
|
||||
BuildKernelDefConstraintsFromTypeList<EnabledRangeDataTypes>()),
|
||||
Range);
|
||||
|
||||
|
|
@ -54,7 +51,6 @@ ONNX_CPU_OPERATOR_KERNEL(
|
|||
11,
|
||||
KernelDefBuilder()
|
||||
.TypeConstraint("T",
|
||||
BuildKernelDefConstraintsFromTypeList<RangeDataTypes>(),
|
||||
BuildKernelDefConstraintsFromTypeList<EnabledRangeDataTypes>()),
|
||||
Range);
|
||||
|
||||
|
|
|
|||
|
|
@ -26,12 +26,8 @@ ORT_SPECIFY_OP_KERNEL_ARG_DEFAULT_TYPES(
|
|||
float, double, int8_t, uint8_t, int64_t, uint64_t);
|
||||
} // namespace op_kernel_type_control
|
||||
|
||||
using Clip11Types = ORT_OP_KERNEL_ARG_DEFAULT_TYPE_LIST(
|
||||
kCpuExecutionProvider, kOnnxDomain, Clip, 11, Input, 0);
|
||||
using EnabledClip11Types = ORT_OP_KERNEL_ARG_ENABLED_TYPE_LIST(
|
||||
kCpuExecutionProvider, kOnnxDomain, Clip, 11, Input, 0);
|
||||
using Clip12Types = ORT_OP_KERNEL_ARG_DEFAULT_TYPE_LIST(
|
||||
kCpuExecutionProvider, kOnnxDomain, Clip, 12, Input, 0);
|
||||
using EnabledClip12Types = ORT_OP_KERNEL_ARG_ENABLED_TYPE_LIST(
|
||||
kCpuExecutionProvider, kOnnxDomain, Clip, 12, Input, 0);
|
||||
|
||||
|
|
@ -41,7 +37,7 @@ using AllEnabledClipTypes =
|
|||
EnabledClip12Types>;
|
||||
|
||||
#define REG_KERNEL_VERSIONED_NONTEMPL( \
|
||||
OP_TYPE, START_VER, END_VER, KERNEL_CLASS, DEFAULT_TYPE_LIST, ENABLED_TYPE_LIST) \
|
||||
OP_TYPE, START_VER, END_VER, KERNEL_CLASS, ENABLED_TYPE_LIST) \
|
||||
ONNX_CPU_OPERATOR_VERSIONED_KERNEL( \
|
||||
OP_TYPE, \
|
||||
START_VER, \
|
||||
|
|
@ -49,25 +45,23 @@ using AllEnabledClipTypes =
|
|||
KernelDefBuilder() \
|
||||
.MayInplace(0, 0) \
|
||||
.TypeConstraint("T", \
|
||||
BuildKernelDefConstraintsFromTypeList<DEFAULT_TYPE_LIST>(), \
|
||||
BuildKernelDefConstraintsFromTypeList<ENABLED_TYPE_LIST>()), \
|
||||
KERNEL_CLASS);
|
||||
|
||||
#define REG_KERNEL_NONTEMPL( \
|
||||
OP_TYPE, VERSION, KERNEL_CLASS, DEFAULT_TYPE_LIST, ENABLED_TYPE_LIST) \
|
||||
OP_TYPE, VERSION, KERNEL_CLASS, ENABLED_TYPE_LIST) \
|
||||
ONNX_CPU_OPERATOR_KERNEL( \
|
||||
OP_TYPE, \
|
||||
VERSION, \
|
||||
KernelDefBuilder() \
|
||||
.MayInplace(0, 0) \
|
||||
.TypeConstraint("T", \
|
||||
BuildKernelDefConstraintsFromTypeList<DEFAULT_TYPE_LIST>(), \
|
||||
BuildKernelDefConstraintsFromTypeList<ENABLED_TYPE_LIST>()), \
|
||||
KERNEL_CLASS);
|
||||
|
||||
REG_KERNEL_VERSIONED_NONTEMPL(Clip, 11, 11, Clip, Clip11Types, EnabledClip11Types);
|
||||
REG_KERNEL_VERSIONED_NONTEMPL(Clip, 12, 12, Clip, Clip12Types, EnabledClip12Types);
|
||||
REG_KERNEL_NONTEMPL(Clip, 13, Clip, Clip12Types, EnabledClip12Types);
|
||||
REG_KERNEL_VERSIONED_NONTEMPL(Clip, 11, 11, Clip, EnabledClip11Types);
|
||||
REG_KERNEL_VERSIONED_NONTEMPL(Clip, 12, 12, Clip, EnabledClip12Types);
|
||||
REG_KERNEL_NONTEMPL(Clip, 13, Clip, EnabledClip12Types);
|
||||
|
||||
#undef REG_KERNEL_VERSIONED_NONTEMPL
|
||||
#undef REG_KERNEL_NONTEMPL
|
||||
|
|
|
|||
|
|
@ -51,23 +51,15 @@ ORT_SPECIFY_OP_KERNEL_ARG_DEFAULT_TYPES(kCpuExecutionProvider, kOnnxDomain, Pow,
|
|||
//
|
||||
// reduce the supported type lists to what's allowed in this build
|
||||
//
|
||||
using Max8Types = ORT_OP_KERNEL_ARG_DEFAULT_TYPE_LIST(kCpuExecutionProvider, kOnnxDomain, Max, 8, Input, 0);
|
||||
using Max12Types = ORT_OP_KERNEL_ARG_DEFAULT_TYPE_LIST(kCpuExecutionProvider, kOnnxDomain, Max, 12, Input, 0);
|
||||
using EnabledMax8Types = ORT_OP_KERNEL_ARG_ENABLED_TYPE_LIST(kCpuExecutionProvider, kOnnxDomain, Max, 8, Input, 0);
|
||||
using EnabledMax12Types = ORT_OP_KERNEL_ARG_ENABLED_TYPE_LIST(kCpuExecutionProvider, kOnnxDomain, Max, 12, Input, 0);
|
||||
|
||||
using Min8Types = ORT_OP_KERNEL_ARG_DEFAULT_TYPE_LIST(kCpuExecutionProvider, kOnnxDomain, Min, 8, Input, 0);
|
||||
using Min12Types = ORT_OP_KERNEL_ARG_DEFAULT_TYPE_LIST(kCpuExecutionProvider, kOnnxDomain, Min, 12, Input, 0);
|
||||
using EnabledMin8Types = ORT_OP_KERNEL_ARG_ENABLED_TYPE_LIST(kCpuExecutionProvider, kOnnxDomain, Min, 8, Input, 0);
|
||||
using EnabledMin12Types = ORT_OP_KERNEL_ARG_ENABLED_TYPE_LIST(kCpuExecutionProvider, kOnnxDomain, Min, 12, Input, 0);
|
||||
|
||||
using ModTypes = ORT_OP_KERNEL_ARG_DEFAULT_TYPE_LIST_ALL_OPSETS(kCpuExecutionProvider, kOnnxDomain, Mod, Input, 0);
|
||||
using EnabledModTypes = ORT_OP_KERNEL_ARG_ENABLED_TYPE_LIST_ALL_OPSETS(
|
||||
kCpuExecutionProvider, kOnnxDomain, Mod, Input, 0);
|
||||
|
||||
using Pow7Types = ORT_OP_KERNEL_ARG_DEFAULT_TYPE_LIST(kCpuExecutionProvider, kOnnxDomain, Pow, 7, Input, 0);
|
||||
using Pow12BaseTypes = ORT_OP_KERNEL_ARG_DEFAULT_TYPE_LIST(kCpuExecutionProvider, kOnnxDomain, Pow, 12, Input, 0);
|
||||
using Pow12ExpTypes = ORT_OP_KERNEL_ARG_DEFAULT_TYPE_LIST(kCpuExecutionProvider, kOnnxDomain, Pow, 12, Input, 1);
|
||||
using EnabledPow7Types = ORT_OP_KERNEL_ARG_ENABLED_TYPE_LIST(kCpuExecutionProvider, kOnnxDomain, Pow, 7, Input, 0);
|
||||
using EnabledPow12BaseTypes = ORT_OP_KERNEL_ARG_ENABLED_TYPE_LIST(kCpuExecutionProvider, kOnnxDomain,
|
||||
Pow, 12, Input, 0);
|
||||
|
|
@ -119,45 +111,45 @@ void Exp<float>::operator()(std::ptrdiff_t first, std::ptrdiff_t last) const {
|
|||
.TypeConstraint("T1", DataTypeImpl::GetTensorType<bool>()), \
|
||||
KERNEL_CLASS<TYPE>);
|
||||
|
||||
#define REG_ELEMENTWISE_KERNEL_NONT(OP_TYPE, VERSION, KERNEL_CLASS, CONSTRAINTS, ENABLED_TYPES_CONSTRAINTS) \
|
||||
ONNX_CPU_OPERATOR_KERNEL( \
|
||||
OP_TYPE, \
|
||||
VERSION, \
|
||||
KernelDefBuilder() \
|
||||
.TypeConstraint("T", CONSTRAINTS, ENABLED_TYPES_CONSTRAINTS), \
|
||||
#define REG_ELEMENTWISE_KERNEL_NONT(OP_TYPE, VERSION, KERNEL_CLASS, CONSTRAINTS) \
|
||||
ONNX_CPU_OPERATOR_KERNEL( \
|
||||
OP_TYPE, \
|
||||
VERSION, \
|
||||
KernelDefBuilder() \
|
||||
.TypeConstraint("T", CONSTRAINTS), \
|
||||
KERNEL_CLASS);
|
||||
|
||||
#define REG_ELEMENTWISE_VERSIONED_KERNEL_NONT(OP_TYPE, VERSION_FROM, VERSION_TO, KERNEL_CLASS, \
|
||||
CONSTRAINTS, ENABLED_TYPES_CONSTRAINTS) \
|
||||
CONSTRAINTS) \
|
||||
ONNX_CPU_OPERATOR_VERSIONED_KERNEL( \
|
||||
OP_TYPE, \
|
||||
VERSION_FROM, \
|
||||
VERSION_TO, \
|
||||
KernelDefBuilder() \
|
||||
.TypeConstraint("T", CONSTRAINTS, ENABLED_TYPES_CONSTRAINTS), \
|
||||
.TypeConstraint("T", CONSTRAINTS), \
|
||||
KERNEL_CLASS);
|
||||
|
||||
#define REG_ELEMENTWISE_KERNEL_NONT_2(OP_TYPE, VERSION, KERNEL_CLASS, \
|
||||
T1_CONSTRAINTS, T1_ENABLED_TYPES_CONSTRAINTS, \
|
||||
T2_CONSTRAINTS, T2_ENABLED_TYPES_CONSTRAINTS) \
|
||||
ONNX_CPU_OPERATOR_KERNEL( \
|
||||
OP_TYPE, \
|
||||
VERSION, \
|
||||
KernelDefBuilder() \
|
||||
.TypeConstraint("T", T1_CONSTRAINTS, T1_ENABLED_TYPES_CONSTRAINTS) \
|
||||
.TypeConstraint("T1", T2_CONSTRAINTS, T2_ENABLED_TYPES_CONSTRAINTS), \
|
||||
#define REG_ELEMENTWISE_KERNEL_NONT_2(OP_TYPE, VERSION, KERNEL_CLASS, \
|
||||
T1_CONSTRAINTS, \
|
||||
T2_CONSTRAINTS) \
|
||||
ONNX_CPU_OPERATOR_KERNEL( \
|
||||
OP_TYPE, \
|
||||
VERSION, \
|
||||
KernelDefBuilder() \
|
||||
.TypeConstraint("T", T1_CONSTRAINTS) \
|
||||
.TypeConstraint("T1", T2_CONSTRAINTS), \
|
||||
KERNEL_CLASS);
|
||||
|
||||
#define REG_ELEMENTWISE_VERSIONED_KERNEL_NONT_2(OP_TYPE, VERSION_FROM, VERSION_TO, KERNEL_CLASS, \
|
||||
T1_CONSTRAINTS, T1_ENABLED_TYPES_CONSTRAINTS, \
|
||||
T2_CONSTRAINTS, T2_ENABLED_TYPES_CONSTRAINTS) \
|
||||
T1_CONSTRAINTS, \
|
||||
T2_CONSTRAINTS) \
|
||||
ONNX_CPU_OPERATOR_VERSIONED_KERNEL( \
|
||||
OP_TYPE, \
|
||||
VERSION_FROM, \
|
||||
VERSION_TO, \
|
||||
KernelDefBuilder() \
|
||||
.TypeConstraint("T", T1_CONSTRAINTS, T1_ENABLED_TYPES_CONSTRAINTS) \
|
||||
.TypeConstraint("T1", T2_CONSTRAINTS, T2_ENABLED_TYPES_CONSTRAINTS), \
|
||||
.TypeConstraint("T", T1_CONSTRAINTS) \
|
||||
.TypeConstraint("T1", T2_CONSTRAINTS), \
|
||||
KERNEL_CLASS);
|
||||
|
||||
REG_ELEMENTWISE_VERSIONED_TYPED_KERNEL(Add, 7, 12, float, Add);
|
||||
|
|
@ -262,25 +254,18 @@ REG_ELEMENTWISE_TYPED_KERNEL(Sqrt, 13, float, Sqrt);
|
|||
REG_ELEMENTWISE_TYPED_KERNEL(Sqrt, 13, double, Sqrt);
|
||||
|
||||
REG_ELEMENTWISE_VERSIONED_KERNEL_NONT(Pow, 7, 11, Pow,
|
||||
BuildKernelDefConstraintsFromTypeList<Pow7Types>(),
|
||||
BuildKernelDefConstraintsFromTypeList<EnabledPow7Types>());
|
||||
|
||||
REG_ELEMENTWISE_VERSIONED_KERNEL_NONT_2(Pow, 12, 12, Pow,
|
||||
BuildKernelDefConstraintsFromTypeList<Pow12BaseTypes>(),
|
||||
BuildKernelDefConstraintsFromTypeList<EnabledPow12BaseTypes>(),
|
||||
BuildKernelDefConstraintsFromTypeList<Pow12ExpTypes>(),
|
||||
BuildKernelDefConstraintsFromTypeList<EnabledPow12ExpTypes>());
|
||||
|
||||
REG_ELEMENTWISE_VERSIONED_KERNEL_NONT_2(Pow, 13, 14, Pow,
|
||||
BuildKernelDefConstraintsFromTypeList<Pow12BaseTypes>(),
|
||||
BuildKernelDefConstraintsFromTypeList<EnabledPow12BaseTypes>(),
|
||||
BuildKernelDefConstraintsFromTypeList<Pow12ExpTypes>(),
|
||||
BuildKernelDefConstraintsFromTypeList<EnabledPow12ExpTypes>());
|
||||
|
||||
REG_ELEMENTWISE_KERNEL_NONT_2(Pow, 15, Pow,
|
||||
BuildKernelDefConstraintsFromTypeList<Pow12BaseTypes>(),
|
||||
BuildKernelDefConstraintsFromTypeList<EnabledPow12BaseTypes>(),
|
||||
BuildKernelDefConstraintsFromTypeList<Pow12ExpTypes>(),
|
||||
BuildKernelDefConstraintsFromTypeList<EnabledPow12ExpTypes>());
|
||||
|
||||
REG_ELEMENTWISE_VERSIONED_TYPED_KERNEL(Exp, 6, 12, float, Exp);
|
||||
|
|
@ -303,16 +288,16 @@ REG_ELEMENTWISE_TYPED_KERNEL(Sum, 13, double, Sum_8);
|
|||
|
||||
REG_ELEMENTWISE_VERSIONED_TYPED_KERNEL(Max, 6, 7, float, Max_6);
|
||||
|
||||
REG_ELEMENTWISE_VERSIONED_KERNEL_NONT(Max, 8, 11, Max_8, BuildKernelDefConstraintsFromTypeList<Max8Types>(), BuildKernelDefConstraintsFromTypeList<EnabledMax8Types>());
|
||||
REG_ELEMENTWISE_VERSIONED_KERNEL_NONT(Max, 12, 12, Max_8, BuildKernelDefConstraintsFromTypeList<Max12Types>(), BuildKernelDefConstraintsFromTypeList<EnabledMax12Types>());
|
||||
REG_ELEMENTWISE_VERSIONED_KERNEL_NONT(Max, 8, 11, Max_8, BuildKernelDefConstraintsFromTypeList<EnabledMax8Types>());
|
||||
REG_ELEMENTWISE_VERSIONED_KERNEL_NONT(Max, 12, 12, Max_8, BuildKernelDefConstraintsFromTypeList<EnabledMax12Types>());
|
||||
// Supposed to add BFloat16 but we are not supporting now, however, separate registration
|
||||
REG_ELEMENTWISE_KERNEL_NONT(Max, 13, Max_8, BuildKernelDefConstraintsFromTypeList<Max12Types>(), BuildKernelDefConstraintsFromTypeList<EnabledMax12Types>());
|
||||
REG_ELEMENTWISE_KERNEL_NONT(Max, 13, Max_8, BuildKernelDefConstraintsFromTypeList<EnabledMax12Types>());
|
||||
|
||||
REG_ELEMENTWISE_VERSIONED_TYPED_KERNEL(Min, 6, 7, float, Min_6);
|
||||
REG_ELEMENTWISE_VERSIONED_KERNEL_NONT(Min, 8, 11, Min_8, BuildKernelDefConstraintsFromTypeList<Min8Types>(), BuildKernelDefConstraintsFromTypeList<EnabledMin8Types>());
|
||||
REG_ELEMENTWISE_VERSIONED_KERNEL_NONT(Min, 12, 12, Min_8, BuildKernelDefConstraintsFromTypeList<Min12Types>(), BuildKernelDefConstraintsFromTypeList<EnabledMin12Types>());
|
||||
REG_ELEMENTWISE_VERSIONED_KERNEL_NONT(Min, 8, 11, Min_8, BuildKernelDefConstraintsFromTypeList<EnabledMin8Types>());
|
||||
REG_ELEMENTWISE_VERSIONED_KERNEL_NONT(Min, 12, 12, Min_8, BuildKernelDefConstraintsFromTypeList<EnabledMin12Types>());
|
||||
// Supposed to add BFloat16 but we are not supporting now, however, separate registration
|
||||
REG_ELEMENTWISE_KERNEL_NONT(Min, 13, Min_8, BuildKernelDefConstraintsFromTypeList<Min12Types>(), BuildKernelDefConstraintsFromTypeList<EnabledMin12Types>());
|
||||
REG_ELEMENTWISE_KERNEL_NONT(Min, 13, Min_8, BuildKernelDefConstraintsFromTypeList<EnabledMin12Types>());
|
||||
|
||||
REG_ELEMENTWISE_LOGICALOP_VERSIONED_TYPED_KERNEL(Less, 7, 8, float, Less);
|
||||
REG_ELEMENTWISE_LOGICALOP_VERSIONED_TYPED_KERNEL(Less, 7, 8, double, Less);
|
||||
|
|
@ -1585,7 +1570,6 @@ ONNX_CPU_OPERATOR_VERSIONED_KERNEL(
|
|||
KernelDefBuilder()
|
||||
.TypeConstraint(
|
||||
"T",
|
||||
BuildKernelDefConstraintsFromTypeList<ModTypes>(),
|
||||
BuildKernelDefConstraintsFromTypeList<EnabledModTypes>()),
|
||||
Mod);
|
||||
|
||||
|
|
@ -1595,7 +1579,6 @@ ONNX_CPU_OPERATOR_KERNEL(
|
|||
KernelDefBuilder()
|
||||
.TypeConstraint(
|
||||
"T",
|
||||
BuildKernelDefConstraintsFromTypeList<ModTypes>(),
|
||||
BuildKernelDefConstraintsFromTypeList<EnabledModTypes>()),
|
||||
Mod);
|
||||
|
||||
|
|
|
|||
|
|
@ -22,8 +22,6 @@ ORT_SPECIFY_OP_KERNEL_ARG_DEFAULT_TYPE_LIST_ALL_OPSETS(
|
|||
kCpuExecutionProvider, kOnnxDomain, Sign, Input, 0, element_type_lists::AllNumeric);
|
||||
}
|
||||
|
||||
using SignDataTypes = ORT_OP_KERNEL_ARG_DEFAULT_TYPE_LIST_ALL_OPSETS(
|
||||
kCpuExecutionProvider, kOnnxDomain, Sign, Input, 0);
|
||||
using EnabledSignDataTypes = ORT_OP_KERNEL_ARG_ENABLED_TYPE_LIST_ALL_OPSETS(
|
||||
kCpuExecutionProvider, kOnnxDomain, Sign, Input, 0);
|
||||
|
||||
|
|
@ -39,7 +37,6 @@ ONNX_CPU_OPERATOR_VERSIONED_KERNEL(
|
|||
9,
|
||||
12,
|
||||
KernelDefBuilder().TypeConstraint("T",
|
||||
BuildKernelDefConstraintsFromTypeList<SignDataTypes>(),
|
||||
BuildKernelDefConstraintsFromTypeList<EnabledSignDataTypes>()),
|
||||
Sign);
|
||||
|
||||
|
|
@ -47,7 +44,6 @@ ONNX_CPU_OPERATOR_KERNEL(
|
|||
Sign,
|
||||
13,
|
||||
KernelDefBuilder().TypeConstraint("T",
|
||||
BuildKernelDefConstraintsFromTypeList<SignDataTypes>(),
|
||||
BuildKernelDefConstraintsFromTypeList<EnabledSignDataTypes>()),
|
||||
Sign);
|
||||
|
||||
|
|
|
|||
|
|
@ -26,12 +26,8 @@ ORT_SPECIFY_OP_KERNEL_ARG_DEFAULT_TYPES(
|
|||
uint8_t);
|
||||
} // namespace op_kernel_type_control
|
||||
|
||||
using MaxPool8DataTypes = ORT_OP_KERNEL_ARG_DEFAULT_TYPE_LIST(
|
||||
kCpuExecutionProvider, kOnnxDomain, MaxPool, 8, Input, 0);
|
||||
using EnabledMaxPool8DataTypes = ORT_OP_KERNEL_ARG_ENABLED_TYPE_LIST(
|
||||
kCpuExecutionProvider, kOnnxDomain, MaxPool, 8, Input, 0);
|
||||
using MaxPool12DataTypes = ORT_OP_KERNEL_ARG_DEFAULT_TYPE_LIST(
|
||||
kCpuExecutionProvider, kOnnxDomain, MaxPool, 12, Input, 0);
|
||||
using EnabledMaxPool12DataTypes = ORT_OP_KERNEL_ARG_ENABLED_TYPE_LIST(
|
||||
kCpuExecutionProvider, kOnnxDomain, MaxPool, 12, Input, 0);
|
||||
|
||||
|
|
@ -272,7 +268,6 @@ ONNX_CPU_OPERATOR_VERSIONED_KERNEL(MaxPool, 8, 11,
|
|||
KernelDefBuilder()
|
||||
.TypeConstraint(
|
||||
"T",
|
||||
BuildKernelDefConstraintsFromTypeList<MaxPool8DataTypes>(),
|
||||
BuildKernelDefConstraintsFromTypeList<EnabledMaxPool8DataTypes>())
|
||||
.TypeConstraint("I", DataTypeImpl::GetTensorType<int64_t>()),
|
||||
MaxPoolV8);
|
||||
|
|
@ -281,7 +276,6 @@ ONNX_CPU_OPERATOR_KERNEL(MaxPool, 12,
|
|||
KernelDefBuilder()
|
||||
.TypeConstraint(
|
||||
"T",
|
||||
BuildKernelDefConstraintsFromTypeList<MaxPool12DataTypes>(),
|
||||
BuildKernelDefConstraintsFromTypeList<EnabledMaxPool12DataTypes>())
|
||||
.TypeConstraint("I", DataTypeImpl::GetTensorType<int64_t>()),
|
||||
MaxPoolV8);
|
||||
|
|
|
|||
|
|
@ -17,8 +17,6 @@ ORT_SPECIFY_OP_KERNEL_ARG_DEFAULT_TYPE_LIST_ALL_OPSETS(
|
|||
element_type_lists::AllNumeric);
|
||||
}
|
||||
|
||||
using ShrinkDataTypes = ORT_OP_KERNEL_ARG_DEFAULT_TYPE_LIST_ALL_OPSETS(
|
||||
kCpuExecutionProvider, kOnnxDomain, Shrink, Input, 0);
|
||||
using EnabledShrinkDataTypes = ORT_OP_KERNEL_ARG_ENABLED_TYPE_LIST_ALL_OPSETS(
|
||||
kCpuExecutionProvider, kOnnxDomain, Shrink, Input, 0);
|
||||
|
||||
|
|
@ -28,7 +26,6 @@ ONNX_CPU_OPERATOR_KERNEL(
|
|||
KernelDefBuilder()
|
||||
.MayInplace(0, 0)
|
||||
.TypeConstraint("T",
|
||||
BuildKernelDefConstraintsFromTypeList<ShrinkDataTypes>(),
|
||||
BuildKernelDefConstraintsFromTypeList<EnabledShrinkDataTypes>()),
|
||||
Shrink);
|
||||
//TODO: fix the warnings
|
||||
|
|
|
|||
|
|
@ -328,8 +328,6 @@ ORT_SPECIFY_OP_KERNEL_ARG_DEFAULT_TYPES_ALL_OPSETS(
|
|||
} // namespace op_kernel_type_control
|
||||
|
||||
namespace {
|
||||
using SplitToSequenceDataTypes = ORT_OP_KERNEL_ARG_DEFAULT_TYPE_LIST_ALL_OPSETS(
|
||||
kCpuExecutionProvider, kOnnxDomain, SplitToSequence, Input, 0);
|
||||
using EnabledSplitToSequenceDataTypes = ORT_OP_KERNEL_ARG_ENABLED_TYPE_LIST_ALL_OPSETS(
|
||||
kCpuExecutionProvider, kOnnxDomain, SplitToSequence, Input, 0);
|
||||
} // namespace
|
||||
|
|
@ -339,7 +337,6 @@ ONNX_CPU_OPERATOR_KERNEL(
|
|||
11,
|
||||
KernelDefBuilder()
|
||||
.TypeConstraint("T",
|
||||
BuildKernelDefConstraintsFromTypeList<SplitToSequenceDataTypes>(),
|
||||
BuildKernelDefConstraintsFromTypeList<EnabledSplitToSequenceDataTypes>())
|
||||
.TypeConstraint("S", DataTypeImpl::AllSequenceTensorTypes())
|
||||
.TypeConstraint("I", std::vector<MLDataType>{
|
||||
|
|
|
|||
|
|
@ -49,10 +49,6 @@ ORT_SPECIFY_OP_KERNEL_ARG_REQUIRED_TYPES_ALL_OPSETS(
|
|||
} // namespace op_kernel_type_control
|
||||
|
||||
namespace {
|
||||
using SrcTypes = ORT_OP_KERNEL_ARG_DEFAULT_TYPE_LIST_ALL_OPSETS(kCpuExecutionProvider, kOnnxDomain,
|
||||
Cast, Input, 0);
|
||||
using DstTypes = ORT_OP_KERNEL_ARG_DEFAULT_TYPE_LIST_ALL_OPSETS(kCpuExecutionProvider, kOnnxDomain,
|
||||
Cast, Output, 0);
|
||||
using EnabledSrcTypes = ORT_OP_KERNEL_ARG_ENABLED_TYPE_LIST_ALL_OPSETS(kCpuExecutionProvider, kOnnxDomain,
|
||||
Cast, Input, 0);
|
||||
using EnabledDstTypes = ORT_OP_KERNEL_ARG_ENABLED_TYPE_LIST_ALL_OPSETS(kCpuExecutionProvider, kOnnxDomain,
|
||||
|
|
@ -324,8 +320,8 @@ ONNX_CPU_OPERATOR_VERSIONED_KERNEL(
|
|||
6,
|
||||
12,
|
||||
KernelDefBuilder()
|
||||
.TypeConstraint("T1", BuildKernelDefConstraintsFromTypeList<SrcTypes>(), BuildKernelDefConstraintsFromTypeList<EnabledSrcTypes>())
|
||||
.TypeConstraint("T2", BuildKernelDefConstraintsFromTypeList<DstTypes>(), BuildKernelDefConstraintsFromTypeList<EnabledDstTypes>())
|
||||
.TypeConstraint("T1", BuildKernelDefConstraintsFromTypeList<EnabledSrcTypes>())
|
||||
.TypeConstraint("T2", BuildKernelDefConstraintsFromTypeList<EnabledDstTypes>())
|
||||
.MayInplace(0, 0), // allocation planner will check input and output sizes match before inplacing
|
||||
Cast);
|
||||
|
||||
|
|
@ -333,8 +329,8 @@ ONNX_CPU_OPERATOR_KERNEL(
|
|||
Cast,
|
||||
13,
|
||||
KernelDefBuilder()
|
||||
.TypeConstraint("T1", BuildKernelDefConstraintsFromTypeList<SrcTypes>(), BuildKernelDefConstraintsFromTypeList<EnabledSrcTypes>())
|
||||
.TypeConstraint("T2", BuildKernelDefConstraintsFromTypeList<DstTypes>(), BuildKernelDefConstraintsFromTypeList<EnabledDstTypes>())
|
||||
.TypeConstraint("T1", BuildKernelDefConstraintsFromTypeList<EnabledSrcTypes>())
|
||||
.TypeConstraint("T2", BuildKernelDefConstraintsFromTypeList<EnabledDstTypes>())
|
||||
.MayInplace(0, 0), // allocation planner will check input and output sizes match before inplacing
|
||||
Cast);
|
||||
|
||||
|
|
|
|||
|
|
@ -45,8 +45,6 @@ ORT_SPECIFY_OP_KERNEL_ARG_REQUIRED_TYPES_ALL_OPSETS(
|
|||
} // namespace op_kernel_type_control
|
||||
|
||||
namespace {
|
||||
using DataTypes = ORT_OP_KERNEL_ARG_DEFAULT_TYPE_LIST_ALL_OPSETS(kCpuExecutionProvider, kOnnxDomain,
|
||||
Concat, Input, 0);
|
||||
using EnabledDataTypes = ORT_OP_KERNEL_ARG_ENABLED_TYPE_LIST_ALL_OPSETS(kCpuExecutionProvider, kOnnxDomain,
|
||||
Concat, Input, 0);
|
||||
} // namespace
|
||||
|
|
|
|||
|
|
@ -15,8 +15,6 @@ ORT_SPECIFY_OP_KERNEL_ARG_DEFAULT_TYPES_ALL_OPSETS(
|
|||
float, double, uint64_t, int64_t, int32_t);
|
||||
}
|
||||
|
||||
using EyeLikeDataTypes = ORT_OP_KERNEL_ARG_DEFAULT_TYPE_LIST_ALL_OPSETS(
|
||||
kCpuExecutionProvider, kOnnxDomain, EyeLike, Output, 0);
|
||||
using EnabledEyeLikeDataTypes = ORT_OP_KERNEL_ARG_ENABLED_TYPE_LIST_ALL_OPSETS(
|
||||
kCpuExecutionProvider, kOnnxDomain, EyeLike, Output, 0);
|
||||
|
||||
|
|
@ -26,11 +24,9 @@ ONNX_CPU_OPERATOR_KERNEL(
|
|||
KernelDefBuilder()
|
||||
.TypeConstraint(
|
||||
"T1",
|
||||
BuildKernelDefConstraintsFromTypeList<EyeLikeDataTypes>(),
|
||||
BuildKernelDefConstraintsFromTypeList<EnabledEyeLikeDataTypes>())
|
||||
.TypeConstraint(
|
||||
"T2",
|
||||
BuildKernelDefConstraintsFromTypeList<EyeLikeDataTypes>(),
|
||||
BuildKernelDefConstraintsFromTypeList<EnabledEyeLikeDataTypes>()),
|
||||
EyeLike);
|
||||
|
||||
|
|
|
|||
|
|
@ -26,8 +26,6 @@ ORT_SPECIFY_OP_KERNEL_ARG_REQUIRED_TYPES_ALL_OPSETS(
|
|||
#endif
|
||||
} // namespace op_kernel_type_control
|
||||
|
||||
using IndexTypes = ORT_OP_KERNEL_ARG_DEFAULT_TYPE_LIST_ALL_OPSETS(kCpuExecutionProvider, kOnnxDomain,
|
||||
Gather, Input, 1);
|
||||
using EnabledIndexTypes = ORT_OP_KERNEL_ARG_ENABLED_TYPE_LIST_ALL_OPSETS(kCpuExecutionProvider, kOnnxDomain,
|
||||
Gather, Input, 1);
|
||||
ONNX_CPU_OPERATOR_VERSIONED_KERNEL(
|
||||
|
|
@ -36,7 +34,7 @@ ONNX_CPU_OPERATOR_VERSIONED_KERNEL(
|
|||
10,
|
||||
KernelDefBuilder()
|
||||
.TypeConstraint("T", DataTypeImpl::AllTensorTypes())
|
||||
.TypeConstraint("Tind", BuildKernelDefConstraintsFromTypeList<IndexTypes>(), BuildKernelDefConstraintsFromTypeList<EnabledIndexTypes>()),
|
||||
.TypeConstraint("Tind", BuildKernelDefConstraintsFromTypeList<EnabledIndexTypes>()),
|
||||
Gather);
|
||||
|
||||
ONNX_CPU_OPERATOR_VERSIONED_KERNEL(
|
||||
|
|
@ -45,7 +43,7 @@ ONNX_CPU_OPERATOR_VERSIONED_KERNEL(
|
|||
12,
|
||||
KernelDefBuilder()
|
||||
.TypeConstraint("T", DataTypeImpl::AllTensorTypes())
|
||||
.TypeConstraint("Tind", BuildKernelDefConstraintsFromTypeList<IndexTypes>(), BuildKernelDefConstraintsFromTypeList<EnabledIndexTypes>()),
|
||||
.TypeConstraint("Tind", BuildKernelDefConstraintsFromTypeList<EnabledIndexTypes>()),
|
||||
Gather);
|
||||
|
||||
ONNX_CPU_OPERATOR_KERNEL(
|
||||
|
|
@ -53,7 +51,7 @@ ONNX_CPU_OPERATOR_KERNEL(
|
|||
13,
|
||||
KernelDefBuilder()
|
||||
.TypeConstraint("T", DataTypeImpl::AllTensorTypes())
|
||||
.TypeConstraint("Tind", BuildKernelDefConstraintsFromTypeList<IndexTypes>(), BuildKernelDefConstraintsFromTypeList<EnabledIndexTypes>()),
|
||||
.TypeConstraint("Tind", BuildKernelDefConstraintsFromTypeList<EnabledIndexTypes>()),
|
||||
Gather);
|
||||
|
||||
Status GatherBase::PrepareForCompute(OpKernelContext* context, Prepare& p) const {
|
||||
|
|
|
|||
|
|
@ -21,9 +21,6 @@ ORT_SPECIFY_OP_KERNEL_ARG_DEFAULT_TYPES_ALL_OPSETS(
|
|||
|
||||
class IsInf final : public OpKernel {
|
||||
public:
|
||||
using DataTypes = ORT_OP_KERNEL_ARG_DEFAULT_TYPE_LIST_ALL_OPSETS(kCpuExecutionProvider, kOnnxDomain,
|
||||
IsInf, Input, 0);
|
||||
|
||||
using EnabledDataTypes = ORT_OP_KERNEL_ARG_ENABLED_TYPE_LIST_ALL_OPSETS(kCpuExecutionProvider, kOnnxDomain,
|
||||
IsInf, Input, 0);
|
||||
|
||||
|
|
@ -40,7 +37,6 @@ ONNX_CPU_OPERATOR_KERNEL(
|
|||
10,
|
||||
KernelDefBuilder()
|
||||
.TypeConstraint("T1",
|
||||
BuildKernelDefConstraintsFromTypeList<IsInf::DataTypes>(),
|
||||
BuildKernelDefConstraintsFromTypeList<IsInf::EnabledDataTypes>())
|
||||
.TypeConstraint("T2", DataTypeImpl::GetTensorType<bool>()),
|
||||
IsInf);
|
||||
|
|
|
|||
|
|
@ -72,16 +72,10 @@ ORT_SPECIFY_OP_KERNEL_ARG_REQUIRED_TYPES(
|
|||
kCpuExecutionProvider, kOnnxDomain, Pad, 13, Input, 0, int32_t, int64_t);
|
||||
} // namespace op_kernel_type_control
|
||||
|
||||
using Pad2Types = ORT_OP_KERNEL_ARG_DEFAULT_TYPE_LIST(
|
||||
kCpuExecutionProvider, kOnnxDomain, Pad, 2, Input, 0);
|
||||
using EnabledPad2Types = ORT_OP_KERNEL_ARG_ENABLED_TYPE_LIST(
|
||||
kCpuExecutionProvider, kOnnxDomain, Pad, 2, Input, 0);
|
||||
using Pad11Types = ORT_OP_KERNEL_ARG_DEFAULT_TYPE_LIST(
|
||||
kCpuExecutionProvider, kOnnxDomain, Pad, 11, Input, 0);
|
||||
using EnabledPad11Types = ORT_OP_KERNEL_ARG_ENABLED_TYPE_LIST(
|
||||
kCpuExecutionProvider, kOnnxDomain, Pad, 11, Input, 0);
|
||||
using Pad13Types = ORT_OP_KERNEL_ARG_DEFAULT_TYPE_LIST(
|
||||
kCpuExecutionProvider, kOnnxDomain, Pad, 13, Input, 0);
|
||||
using EnabledPad13Types = ORT_OP_KERNEL_ARG_ENABLED_TYPE_LIST(
|
||||
kCpuExecutionProvider, kOnnxDomain, Pad, 13, Input, 0);
|
||||
|
||||
|
|
@ -97,7 +91,6 @@ ONNX_CPU_OPERATOR_VERSIONED_KERNEL(
|
|||
2, 10,
|
||||
KernelDefBuilder().TypeConstraint(
|
||||
"T",
|
||||
BuildKernelDefConstraintsFromTypeList<Pad2Types>(),
|
||||
BuildKernelDefConstraintsFromTypeList<EnabledPad2Types>()),
|
||||
Pad);
|
||||
|
||||
|
|
@ -110,7 +103,6 @@ ONNX_CPU_OPERATOR_VERSIONED_KERNEL(
|
|||
11, 12,
|
||||
KernelDefBuilder().TypeConstraint(
|
||||
"T",
|
||||
BuildKernelDefConstraintsFromTypeList<Pad11Types>(),
|
||||
BuildKernelDefConstraintsFromTypeList<EnabledPad11Types>()),
|
||||
Pad);
|
||||
|
||||
|
|
@ -120,7 +112,6 @@ ONNX_CPU_OPERATOR_KERNEL(
|
|||
KernelDefBuilder()
|
||||
.TypeConstraint(
|
||||
"T",
|
||||
BuildKernelDefConstraintsFromTypeList<Pad13Types>(),
|
||||
BuildKernelDefConstraintsFromTypeList<EnabledPad13Types>()),
|
||||
Pad);
|
||||
|
||||
|
|
|
|||
|
|
@ -33,8 +33,6 @@ ORT_SPECIFY_OP_KERNEL_ARG_DEFAULT_TYPE_LIST_ALL_OPSETS(
|
|||
element_type_lists::All);
|
||||
}
|
||||
|
||||
using ReverseSequenceDataTypes = ORT_OP_KERNEL_ARG_DEFAULT_TYPE_LIST_ALL_OPSETS(
|
||||
kCpuExecutionProvider, kOnnxDomain, ReverseSequence, Input, 0);
|
||||
using EnabledReverseSequenceDataTypes = ORT_OP_KERNEL_ARG_ENABLED_TYPE_LIST_ALL_OPSETS(
|
||||
kCpuExecutionProvider, kOnnxDomain, ReverseSequence, Input, 0);
|
||||
|
||||
|
|
@ -44,7 +42,6 @@ ONNX_OPERATOR_KERNEL_EX(ReverseSequence,
|
|||
kCpuExecutionProvider,
|
||||
KernelDefBuilder()
|
||||
.TypeConstraint("T",
|
||||
BuildKernelDefConstraintsFromTypeList<ReverseSequenceDataTypes>(),
|
||||
BuildKernelDefConstraintsFromTypeList<EnabledReverseSequenceDataTypes>()),
|
||||
ReverseSequenceOp);
|
||||
|
||||
|
|
|
|||
|
|
@ -26,13 +26,9 @@ ORT_SPECIFY_OP_KERNEL_ARG_DEFAULT_TYPE_LIST_ALL_OPSETS(
|
|||
kCpuExecutionProvider, kOnnxDomain, ScatterElements, Input, 0, element_type_lists::All);
|
||||
} // namespace op_kernel_type_control
|
||||
|
||||
using ScatterDataTypes = ORT_OP_KERNEL_ARG_DEFAULT_TYPE_LIST_ALL_OPSETS(
|
||||
kCpuExecutionProvider, kOnnxDomain, Scatter, Input, 0);
|
||||
using EnabledScatterDataTypes = ORT_OP_KERNEL_ARG_ENABLED_TYPE_LIST_ALL_OPSETS(
|
||||
kCpuExecutionProvider, kOnnxDomain, Scatter, Input, 0);
|
||||
|
||||
using ScatterElementsDataTypes = ORT_OP_KERNEL_ARG_DEFAULT_TYPE_LIST_ALL_OPSETS(
|
||||
kCpuExecutionProvider, kOnnxDomain, ScatterElements, Input, 0);
|
||||
using EnabledScatterElementsDataTypes = ORT_OP_KERNEL_ARG_ENABLED_TYPE_LIST_ALL_OPSETS(
|
||||
kCpuExecutionProvider, kOnnxDomain, ScatterElements, Input, 0);
|
||||
|
||||
|
|
@ -64,7 +60,6 @@ ONNX_CPU_OPERATOR_VERSIONED_KERNEL(
|
|||
KernelDefBuilder()
|
||||
.MayInplace(0, 0)
|
||||
.TypeConstraint("T",
|
||||
BuildKernelDefConstraintsFromTypeList<ScatterDataTypes>(),
|
||||
BuildKernelDefConstraintsFromTypeList<EnabledScatterDataTypes>())
|
||||
.TypeConstraint("Tind", BuildKernelDefConstraints<int32_t, int64_t>()),
|
||||
Scatter<EnabledScatterDataTypes>);
|
||||
|
|
@ -76,7 +71,6 @@ ONNX_CPU_OPERATOR_VERSIONED_KERNEL(
|
|||
KernelDefBuilder()
|
||||
.MayInplace(0, 0)
|
||||
.TypeConstraint("T",
|
||||
BuildKernelDefConstraintsFromTypeList<ScatterElementsDataTypes>(),
|
||||
BuildKernelDefConstraintsFromTypeList<EnabledScatterElementsDataTypes>())
|
||||
.TypeConstraint("Tind", BuildKernelDefConstraints<int32_t, int64_t>()),
|
||||
Scatter<EnabledScatterElementsDataTypes>);
|
||||
|
|
@ -88,7 +82,6 @@ ONNX_CPU_OPERATOR_VERSIONED_KERNEL(
|
|||
KernelDefBuilder()
|
||||
.MayInplace(0, 0)
|
||||
.TypeConstraint("T",
|
||||
BuildKernelDefConstraintsFromTypeList<ScatterElementsDataTypes>(),
|
||||
BuildKernelDefConstraintsFromTypeList<EnabledScatterElementsDataTypes>())
|
||||
.TypeConstraint("Tind", BuildKernelDefConstraints<int32_t, int64_t>()),
|
||||
Scatter<EnabledScatterElementsDataTypes>);
|
||||
|
|
@ -99,7 +92,6 @@ ONNX_CPU_OPERATOR_KERNEL(
|
|||
KernelDefBuilder()
|
||||
.MayInplace(0, 0)
|
||||
.TypeConstraint("T",
|
||||
BuildKernelDefConstraintsFromTypeList<ScatterElementsDataTypes>(),
|
||||
BuildKernelDefConstraintsFromTypeList<EnabledScatterElementsDataTypes>())
|
||||
.TypeConstraint("Tind", BuildKernelDefConstraints<int32_t, int64_t>()),
|
||||
Scatter<EnabledScatterElementsDataTypes>);
|
||||
|
|
|
|||
|
|
@ -17,8 +17,6 @@ ORT_SPECIFY_OP_KERNEL_ARG_DEFAULT_TYPE_LIST_ALL_OPSETS(
|
|||
element_type_lists::All);
|
||||
}
|
||||
|
||||
using ScatterNDDataTypes = ORT_OP_KERNEL_ARG_DEFAULT_TYPE_LIST_ALL_OPSETS(
|
||||
kCpuExecutionProvider, kOnnxDomain, ScatterND, Input, 0);
|
||||
using EnabledScatterNDDataTypes = ORT_OP_KERNEL_ARG_ENABLED_TYPE_LIST_ALL_OPSETS(
|
||||
kCpuExecutionProvider, kOnnxDomain, ScatterND, Input, 0);
|
||||
|
||||
|
|
@ -28,7 +26,6 @@ ONNX_CPU_OPERATOR_VERSIONED_KERNEL(
|
|||
12,
|
||||
KernelDefBuilder()
|
||||
.TypeConstraint("T",
|
||||
BuildKernelDefConstraintsFromTypeList<ScatterNDDataTypes>(),
|
||||
BuildKernelDefConstraintsFromTypeList<EnabledScatterNDDataTypes>()),
|
||||
ScatterND);
|
||||
|
||||
|
|
@ -38,7 +35,6 @@ ONNX_CPU_OPERATOR_VERSIONED_KERNEL(
|
|||
15,
|
||||
KernelDefBuilder()
|
||||
.TypeConstraint("T",
|
||||
BuildKernelDefConstraintsFromTypeList<ScatterNDDataTypes>(),
|
||||
BuildKernelDefConstraintsFromTypeList<EnabledScatterNDDataTypes>()),
|
||||
ScatterND);
|
||||
|
||||
|
|
@ -47,7 +43,6 @@ ONNX_CPU_OPERATOR_KERNEL(
|
|||
16,
|
||||
KernelDefBuilder()
|
||||
.TypeConstraint("T",
|
||||
BuildKernelDefConstraintsFromTypeList<ScatterNDDataTypes>(),
|
||||
BuildKernelDefConstraintsFromTypeList<EnabledScatterNDDataTypes>()),
|
||||
ScatterND);
|
||||
|
||||
|
|
|
|||
|
|
@ -32,10 +32,6 @@ ORT_SPECIFY_OP_KERNEL_ARG_REQUIRED_TYPES_ALL_OPSETS(
|
|||
} // namespace op_kernel_type_control
|
||||
|
||||
namespace {
|
||||
using DataTypes = ORT_OP_KERNEL_ARG_DEFAULT_TYPE_LIST_ALL_OPSETS(kCpuExecutionProvider, kOnnxDomain,
|
||||
Slice, Input, 0);
|
||||
using IndicesTypes = ORT_OP_KERNEL_ARG_DEFAULT_TYPE_LIST_ALL_OPSETS(kCpuExecutionProvider, kOnnxDomain,
|
||||
Slice, Input, 1);
|
||||
using EnabledDataTypes = ORT_OP_KERNEL_ARG_ENABLED_TYPE_LIST_ALL_OPSETS(kCpuExecutionProvider, kOnnxDomain,
|
||||
Slice, Input, 0);
|
||||
using EnabledIndicesTypes = ORT_OP_KERNEL_ARG_ENABLED_TYPE_LIST_ALL_OPSETS(kCpuExecutionProvider, kOnnxDomain,
|
||||
|
|
@ -45,15 +41,15 @@ using EnabledIndicesTypes = ORT_OP_KERNEL_ARG_ENABLED_TYPE_LIST_ALL_OPSETS(kCpuE
|
|||
ONNX_CPU_OPERATOR_VERSIONED_KERNEL(
|
||||
Slice,
|
||||
1, 9,
|
||||
KernelDefBuilder().TypeConstraint("T", BuildKernelDefConstraintsFromTypeList<DataTypes>(), BuildKernelDefConstraintsFromTypeList<EnabledDataTypes>()),
|
||||
KernelDefBuilder().TypeConstraint("T", BuildKernelDefConstraintsFromTypeList<EnabledDataTypes>()),
|
||||
Slice1);
|
||||
|
||||
ONNX_CPU_OPERATOR_VERSIONED_KERNEL(
|
||||
Slice,
|
||||
10, 10,
|
||||
KernelDefBuilder()
|
||||
.TypeConstraint("T", BuildKernelDefConstraintsFromTypeList<DataTypes>(), BuildKernelDefConstraintsFromTypeList<EnabledDataTypes>())
|
||||
.TypeConstraint("Tind", BuildKernelDefConstraintsFromTypeList<IndicesTypes>(), BuildKernelDefConstraintsFromTypeList<EnabledIndicesTypes>()),
|
||||
.TypeConstraint("T", BuildKernelDefConstraintsFromTypeList<EnabledDataTypes>())
|
||||
.TypeConstraint("Tind", BuildKernelDefConstraintsFromTypeList<EnabledIndicesTypes>()),
|
||||
Slice10);
|
||||
|
||||
ONNX_CPU_OPERATOR_VERSIONED_KERNEL(
|
||||
|
|
@ -61,16 +57,16 @@ ONNX_CPU_OPERATOR_VERSIONED_KERNEL(
|
|||
11,
|
||||
12,
|
||||
KernelDefBuilder()
|
||||
.TypeConstraint("T", BuildKernelDefConstraintsFromTypeList<DataTypes>(), BuildKernelDefConstraintsFromTypeList<EnabledDataTypes>())
|
||||
.TypeConstraint("Tind", BuildKernelDefConstraintsFromTypeList<IndicesTypes>(), BuildKernelDefConstraintsFromTypeList<EnabledIndicesTypes>()),
|
||||
.TypeConstraint("T", BuildKernelDefConstraintsFromTypeList<EnabledDataTypes>())
|
||||
.TypeConstraint("Tind", BuildKernelDefConstraintsFromTypeList<EnabledIndicesTypes>()),
|
||||
Slice10);
|
||||
|
||||
ONNX_CPU_OPERATOR_KERNEL(
|
||||
Slice,
|
||||
13,
|
||||
KernelDefBuilder()
|
||||
.TypeConstraint("T", BuildKernelDefConstraintsFromTypeList<DataTypes>(), BuildKernelDefConstraintsFromTypeList<EnabledDataTypes>())
|
||||
.TypeConstraint("Tind", BuildKernelDefConstraintsFromTypeList<IndicesTypes>(), BuildKernelDefConstraintsFromTypeList<EnabledIndicesTypes>()),
|
||||
.TypeConstraint("T", BuildKernelDefConstraintsFromTypeList<EnabledDataTypes>())
|
||||
.TypeConstraint("Tind", BuildKernelDefConstraintsFromTypeList<EnabledIndicesTypes>()),
|
||||
Slice10);
|
||||
|
||||
// Check if it's possible to combine innermost dimensions so we copy larger blocks.
|
||||
|
|
|
|||
|
|
@ -22,8 +22,6 @@ ORT_SPECIFY_OP_KERNEL_ARG_REQUIRED_TYPES_ALL_OPSETS(
|
|||
int32_t, int64_t);
|
||||
} // namespace op_kernel_type_control
|
||||
|
||||
using SplitDataTypes = ORT_OP_KERNEL_ARG_DEFAULT_TYPE_LIST_ALL_OPSETS(
|
||||
kCpuExecutionProvider, kOnnxDomain, Split, Input, 0);
|
||||
using EnabledSplitDataTypes = ORT_OP_KERNEL_ARG_ENABLED_TYPE_LIST_ALL_OPSETS(
|
||||
kCpuExecutionProvider, kOnnxDomain, Split, Input, 0);
|
||||
|
||||
|
|
@ -32,7 +30,6 @@ ONNX_CPU_OPERATOR_VERSIONED_KERNEL(
|
|||
2,
|
||||
10,
|
||||
KernelDefBuilder().TypeConstraint("T",
|
||||
BuildKernelDefConstraintsFromTypeList<SplitDataTypes>(),
|
||||
BuildKernelDefConstraintsFromTypeList<EnabledSplitDataTypes>()),
|
||||
Split);
|
||||
|
||||
|
|
@ -42,7 +39,6 @@ ONNX_CPU_OPERATOR_VERSIONED_KERNEL(
|
|||
11,
|
||||
12,
|
||||
KernelDefBuilder().TypeConstraint("T",
|
||||
BuildKernelDefConstraintsFromTypeList<SplitDataTypes>(),
|
||||
BuildKernelDefConstraintsFromTypeList<EnabledSplitDataTypes>()),
|
||||
Split);
|
||||
|
||||
|
|
@ -51,7 +47,6 @@ ONNX_CPU_OPERATOR_KERNEL(
|
|||
Split,
|
||||
13,
|
||||
KernelDefBuilder().TypeConstraint("T",
|
||||
BuildKernelDefConstraintsFromTypeList<SplitDataTypes>(),
|
||||
BuildKernelDefConstraintsFromTypeList<EnabledSplitDataTypes>()),
|
||||
Split);
|
||||
|
||||
|
|
|
|||
|
|
@ -33,8 +33,6 @@ ORT_SPECIFY_OP_KERNEL_ARG_REQUIRED_TYPE_LIST_ALL_OPSETS(
|
|||
|
||||
namespace {
|
||||
// reduce the supported types with any global or op specific lists
|
||||
using DataTypes = ORT_OP_KERNEL_ARG_DEFAULT_TYPE_LIST_ALL_OPSETS(kCpuExecutionProvider, kOnnxDomain,
|
||||
Transpose, Input, 0);
|
||||
using EnabledDataTypes = ORT_OP_KERNEL_ARG_ENABLED_TYPE_LIST_ALL_OPSETS(kCpuExecutionProvider, kOnnxDomain,
|
||||
Transpose, Input, 0);
|
||||
} // namespace
|
||||
|
|
@ -418,13 +416,13 @@ ONNX_CPU_OPERATOR_VERSIONED_KERNEL(
|
|||
Transpose,
|
||||
1,
|
||||
12,
|
||||
KernelDefBuilder().TypeConstraint("T", BuildKernelDefConstraintsFromTypeList<DataTypes>(), BuildKernelDefConstraintsFromTypeList<EnabledDataTypes>()),
|
||||
KernelDefBuilder().TypeConstraint("T", BuildKernelDefConstraintsFromTypeList<EnabledDataTypes>()),
|
||||
Transpose);
|
||||
|
||||
ONNX_CPU_OPERATOR_KERNEL(
|
||||
Transpose,
|
||||
13,
|
||||
KernelDefBuilder().TypeConstraint("T", BuildKernelDefConstraintsFromTypeList<DataTypes>(), BuildKernelDefConstraintsFromTypeList<EnabledDataTypes>()),
|
||||
KernelDefBuilder().TypeConstraint("T", BuildKernelDefConstraintsFromTypeList<EnabledDataTypes>()),
|
||||
Transpose);
|
||||
|
||||
} // namespace onnxruntime
|
||||
|
|
|
|||
|
|
@ -17,8 +17,6 @@ ORT_SPECIFY_OP_KERNEL_ARG_DEFAULT_TYPES_ALL_OPSETS(
|
|||
float, int64_t, int8_t, std::string);
|
||||
}
|
||||
|
||||
using UniqueDataTypes = ORT_OP_KERNEL_ARG_DEFAULT_TYPE_LIST_ALL_OPSETS(
|
||||
kCpuExecutionProvider, kOnnxDomain, Unique, Input, 0);
|
||||
using EnabledUniqueDataTypes = ORT_OP_KERNEL_ARG_ENABLED_TYPE_LIST_ALL_OPSETS(
|
||||
kCpuExecutionProvider, kOnnxDomain, Unique, Input, 0);
|
||||
|
||||
|
|
@ -84,7 +82,6 @@ ONNX_CPU_OPERATOR_KERNEL(
|
|||
Unique,
|
||||
11,
|
||||
KernelDefBuilder().TypeConstraint("T",
|
||||
BuildKernelDefConstraintsFromTypeList<UniqueDataTypes>(),
|
||||
BuildKernelDefConstraintsFromTypeList<EnabledUniqueDataTypes>()),
|
||||
Unique);
|
||||
|
||||
|
|
|
|||
|
|
@ -362,37 +362,6 @@ struct EnabledTypes {
|
|||
::onnxruntime::op_kernel_type_control::kAllOpSets, \
|
||||
ArgDirection, ArgIndex, __VA_ARGS__)
|
||||
|
||||
/**
|
||||
* TypeList type with the default types for a given Op kernel argument.
|
||||
*
|
||||
* @param OpProvider The Op provider.
|
||||
* @param OpDomain The Op domain.
|
||||
* @param OpName The Op name.
|
||||
* @param OpSet The opset to use for the default types list.
|
||||
* @param ArgDirection Direction of the given Op kernel argument - Input or Output.
|
||||
* @param ArgIndex Index of the given Op kernel argument.
|
||||
*/
|
||||
#define ORT_OP_KERNEL_ARG_DEFAULT_TYPE_LIST( \
|
||||
OpProvider, OpDomain, OpName, OpSet, ArgDirection, ArgIndex) \
|
||||
::onnxruntime::op_kernel_type_control:: \
|
||||
ORT_OP_KERNEL_TYPE_CTRL_INTERNAL_DEFAULT_TYPES_HOLDER( \
|
||||
OpProvider, OpDomain, OpName, OpSet, ArgDirection, ArgIndex)::types
|
||||
|
||||
/**
|
||||
* TypeList type with the default types for a given Op kernel argument that are valid for all opsets.
|
||||
*
|
||||
* @param OpProvider The Op provider.
|
||||
* @param OpDomain The Op domain.
|
||||
* @param OpName The Op name.
|
||||
* @param ArgDirection Direction of the given Op kernel argument - Input or Output.
|
||||
* @param ArgIndex Index of the given Op kernel argument.
|
||||
*/
|
||||
#define ORT_OP_KERNEL_ARG_DEFAULT_TYPE_LIST_ALL_OPSETS( \
|
||||
OpProvider, OpDomain, OpName, ArgDirection, ArgIndex) \
|
||||
ORT_OP_KERNEL_ARG_DEFAULT_TYPE_LIST(OpProvider, OpDomain, OpName, \
|
||||
::onnxruntime::op_kernel_type_control::kAllOpSets, \
|
||||
ArgDirection, ArgIndex)
|
||||
|
||||
/**
|
||||
* TypeList type with the enabled types for a given Op kernel argument.
|
||||
*
|
||||
|
|
|
|||
Loading…
Reference in a new issue