Consolidate enabled/default kernel def type constraints (#13034)

Consolidate enabled/default kernel def type constraint types into enabled.
This commit is contained in:
Edward Chen 2022-09-27 14:04:15 -07:00 committed by GitHub
parent 440f31668f
commit 5c89c37f7f
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23
26 changed files with 89 additions and 282 deletions

View file

@ -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_;
};

View file

@ -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) {

View file

@ -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,

View file

@ -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);

View file

@ -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);

View file

@ -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);

View file

@ -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

View file

@ -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);

View file

@ -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);

View file

@ -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);

View file

@ -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

View file

@ -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>{

View file

@ -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);

View file

@ -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

View file

@ -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);

View file

@ -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 {

View file

@ -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);

View file

@ -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);

View file

@ -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);

View file

@ -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>);

View file

@ -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);

View file

@ -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.

View file

@ -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);

View file

@ -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

View file

@ -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);

View file

@ -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.
*