Update ops that had strides/dilations documentation updates to default to 1 (#1913)

* Update ops that had strides/dilations documentation updates to default to 1.
Code was already doing this.
Add tests to explicilty test.

* Update optimizers to add opset 11 support where possible
This commit is contained in:
Scott McKay 2019-10-01 08:11:15 +10:00 committed by GitHub
parent df472cbfbd
commit aef49d1f22
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23
17 changed files with 308 additions and 134 deletions

View file

@ -38,7 +38,7 @@ Status ConvActivationFusion::ApplyImpl(Graph& graph, bool& modified, int graph_l
auto* node = graph.GetNode(index);
ORT_RETURN_IF_ERROR(Recurse(*node, modified, graph_level));
if (!graph_utils::IsSupportedOptypeVersionAndDomain(*node, "Conv", {1}) ||
if (!graph_utils::IsSupportedOptypeVersionAndDomain(*node, "Conv", {1, 11}) ||
!graph_utils::IsSupportedProvider(*node, GetCompatibleExecutionProviders()) ||
node->GetOutputEdgesCount() != 1) {
continue;
@ -84,7 +84,6 @@ Status ConvActivationFusion::ApplyImpl(Graph& graph, bool& modified, int graph_l
}
if (!graph.IsNodeOutputsInGraphOutputs(next_node)) {
HandleActivationNodeEdges(graph, next_node, fused_conv);
// Replace the input of the node following activation node

View file

@ -105,7 +105,7 @@ Status ConvAddFusion::Apply(Graph& graph, Node& node, RewriteRuleEffect& modifie
}
bool ConvAddFusion::SatisfyCondition(const Graph& graph, const Node& node) const {
if (!graph_utils::IsSupportedOptypeVersionAndDomain(node, "Conv", {1}) ||
if (!graph_utils::IsSupportedOptypeVersionAndDomain(node, "Conv", {1, 11}) ||
node.GetOutputEdgesCount() != 1) {
return false;
}

View file

@ -141,7 +141,7 @@ Status ConvBNFusion::Apply(Graph& graph, Node& node, RewriteRuleEffect& rule_eff
}
bool ConvBNFusion::SatisfyCondition(const Graph& graph, const Node& node) const {
if (!graph_utils::IsSupportedOptypeVersionAndDomain(node, "Conv", {1}) ||
if (!graph_utils::IsSupportedOptypeVersionAndDomain(node, "Conv", {1, 11}) ||
node.GetOutputEdgesCount() != 1) {
return false;
}

View file

@ -103,7 +103,7 @@ Status ConvMulFusion::Apply(Graph& graph, Node& node, RewriteRuleEffect& rule_ef
}
bool ConvMulFusion::SatisfyCondition(const Graph& graph, const Node& node) const {
if (!graph_utils::IsSupportedOptypeVersionAndDomain(node, "Conv", {1}) ||
if (!graph_utils::IsSupportedOptypeVersionAndDomain(node, "Conv", {1, 11}) ||
node.GetOutputEdgesCount() != 1) {
return false;
}

View file

@ -661,11 +661,11 @@ void NchwcTransformerImpl::TransformActivation(Node& node) {
}
void NchwcTransformerImpl::Transform(Node& node) {
if (graph_utils::IsSupportedOptypeVersionAndDomain(node, "Conv", {1}) ||
if (graph_utils::IsSupportedOptypeVersionAndDomain(node, "Conv", {1, 11}) ||
graph_utils::IsSupportedOptypeVersionAndDomain(node, "FusedConv", {1}, kMSDomain)) {
TransformConv(node);
} else if (graph_utils::IsSupportedOptypeVersionAndDomain(node, "MaxPool", {1, 8, 10}) ||
graph_utils::IsSupportedOptypeVersionAndDomain(node, "AveragePool", {1, 7, 10}) ||
} else if (graph_utils::IsSupportedOptypeVersionAndDomain(node, "MaxPool", {1, 8, 10, 11}) ||
graph_utils::IsSupportedOptypeVersionAndDomain(node, "AveragePool", {1, 7, 10, 11}) ||
graph_utils::IsSupportedOptypeVersionAndDomain(node, "GlobalMaxPool", {1}) ||
graph_utils::IsSupportedOptypeVersionAndDomain(node, "GlobalAveragePool", {1})) {
TransformPool(node);

View file

@ -19,10 +19,10 @@ Status EliminateSlice::Apply(Graph& graph, Node& node, RewriteRuleEffect& rule_e
bool EliminateSlice::SatisfyCondition(const Graph& graph, const Node& node) const {
// We currently support elimination for Slice operator v1.
// TODO Extend to support Slice operator v10, which includes "steps" and all attributes are now given as inputs.
if (!graph_utils::IsSupportedOptypeVersionAndDomain(node, "Slice", {1})) {
if (!graph_utils::IsSupportedOptypeVersionAndDomain(node, "Slice", {1, 11})) {
return false;
}
if (!graph_utils::IsSingleInSingleOutNode(node) ||
graph.IsNodeOutputsInGraphOutputs(node)) {
return false;

View file

@ -113,17 +113,16 @@ class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOn
class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, 10, Softmax);
class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, 9, TopK);
class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 7, 9, BatchNormalization);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, Conv);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, ConvTranspose);
class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, 10, Conv);
class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, 10, ConvTranspose);
class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, 8, Flatten);
class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, 10, Flatten);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 6, InstanceNormalization);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, LpNormalization);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, LRN);
class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 7, 9, AveragePool);
class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, 7, MaxPool);
class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 8, 9, MaxPool);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 2, LpPool);
class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 2, 10, LpPool);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 2, GlobalLpPool);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, GlobalAveragePool);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, GlobalMaxPool);
@ -244,7 +243,7 @@ class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain,
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, int64_t_float_int64_t, OneHot);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, int32_t_float_int32_t, OneHot);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, int32_t_float_float, OneHot);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, MaxUnpool);
class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, 10, MaxUnpool);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, Sinh);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, Cosh);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, Asinh);
@ -264,12 +263,13 @@ class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain,
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, float, Where);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, int32_t, Where);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, int64_t, Where);
class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, 10, Flatten);
// Opset 10
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, StringNormalizer);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, TopK);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, MaxPool);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, AveragePool);
class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, 10, MaxPool);
class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, 10, AveragePool);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, Mod);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, float, Resize);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, int32_t, Resize);
@ -347,7 +347,12 @@ class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, Un
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, Det);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, ScatterElements);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, NonMaxSuppression);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, AveragePool);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, MaxPool);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, MaxUnpool);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, LpPool);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, Conv);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, ConvTranspose);
void RegisterOnnxOperatorKernels(KernelRegistry& kernel_registry) {
static const BuildKernelCreateInfoFn function_table[] = {
@ -446,17 +451,16 @@ void RegisterOnnxOperatorKernels(KernelRegistry& kernel_registry) {
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, 10, Softmax)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, 9, TopK)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 7, 9, BatchNormalization)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, Conv)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, ConvTranspose)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, 10, Conv)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, 10, ConvTranspose)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, 8, Flatten)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, 10, Flatten)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 6, InstanceNormalization)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, LpNormalization)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, LRN)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 7, 9, AveragePool)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, 7, MaxPool)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 8, 9, MaxPool)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 2, LpPool)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 2, 10, LpPool)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 2, GlobalLpPool)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, GlobalAveragePool)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, GlobalMaxPool)>,
@ -577,7 +581,7 @@ void RegisterOnnxOperatorKernels(KernelRegistry& kernel_registry) {
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, int64_t_float_int64_t, OneHot)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, int32_t_float_int32_t, OneHot)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, int32_t_float_float, OneHot)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, MaxUnpool)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, 10, MaxUnpool)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, Sinh)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, Cosh)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, Asinh)>,
@ -597,12 +601,13 @@ void RegisterOnnxOperatorKernels(KernelRegistry& kernel_registry) {
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, float, Where)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, int32_t, Where)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, int64_t, Where)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, 10, Flatten)>,
// Opset 10
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, StringNormalizer)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, TopK)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, MaxPool)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, AveragePool)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, 10, MaxPool)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, 10, AveragePool)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, Mod)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, float, Resize)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, int32_t, Resize)>,
@ -680,6 +685,12 @@ void RegisterOnnxOperatorKernels(KernelRegistry& kernel_registry) {
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, Det)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, ScatterElements)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, NonMaxSuppression)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, AveragePool)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, MaxPool)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, MaxUnpool)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, LpPool)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, Conv)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, ConvTranspose)>,
};
for (auto& function_table_entry : function_table) {

View file

@ -14,13 +14,20 @@ using namespace ::onnxruntime::common;
namespace onnxruntime {
ONNX_CPU_OPERATOR_KERNEL(
ONNX_CPU_OPERATOR_VERSIONED_KERNEL(
MaxUnpool,
9,
9, 10,
KernelDefBuilder()
.TypeConstraint("T1", DataTypeImpl::GetTensorType<float>())
.TypeConstraint("T2", DataTypeImpl::GetTensorType<int64_t>()),
MaxUnpool);
ONNX_CPU_OPERATOR_KERNEL(
MaxUnpool,
11,
KernelDefBuilder()
.TypeConstraint("T1", DataTypeImpl::GetTensorType<float>())
.TypeConstraint("T2", DataTypeImpl::GetTensorType<int64_t>()),
// .TypeConstraint("Y", DataTypeImpl::GetTensorType<float>()),
MaxUnpool);
Status MaxUnpool::Compute(OpKernelContext* context) const {

View file

@ -278,9 +278,15 @@ Status Conv<float>::Compute(OpKernelContext* context) const {
return Status::OK();
}
ONNX_CPU_OPERATOR_VERSIONED_KERNEL(
Conv,
1, 10,
KernelDefBuilder().TypeConstraint("T", DataTypeImpl::GetTensorType<float>()),
Conv<float>);
ONNX_CPU_OPERATOR_KERNEL(
Conv,
1,
11,
KernelDefBuilder().TypeConstraint("T", DataTypeImpl::GetTensorType<float>()),
Conv<float>);

View file

@ -68,7 +68,7 @@ struct ConvAttributes {
kernel_shape_specified = info.GetAttrs<int64_t>("kernel_shape", kernel_shape_).IsOK();
status = info.GetAttrs<int64_t>("strides", strides);
if (!status.IsOK()) {
if (!status.IsOK() || strides.empty()) {
strides.resize(kernel_shape_.size(), 1);
}
@ -78,7 +78,7 @@ struct ConvAttributes {
}
status = info.GetAttrs<int64_t>("dilations", dilations);
if (!status.IsOK()) {
if (!status.IsOK() || dilations.empty()) {
dilations.resize(kernel_shape_.size(), 1);
}
@ -89,6 +89,7 @@ struct ConvAttributes {
#if false
// TODO: Re-enable when attributes values are guaranteed to be filled.
// Make sure empty strides or dilations are defaulted to 1 if necessary
std::string auto_pad_str;
ORT_ENFORCE(info.GetAttr<std::string>("auto_pad", &auto_pad_str).IsOK());
auto_pad = StringToAutoPadType(auto_pad_str);

View file

@ -23,9 +23,15 @@
namespace onnxruntime {
ONNX_CPU_OPERATOR_VERSIONED_KERNEL(
ConvTranspose,
1, 10,
KernelDefBuilder().TypeConstraint("T", DataTypeImpl::GetTensorType<float>()),
ConvTranspose<float>);
ONNX_CPU_OPERATOR_KERNEL(
ConvTranspose,
1,
11,
KernelDefBuilder().TypeConstraint("T", DataTypeImpl::GetTensorType<float>()),
ConvTranspose<float>);

View file

@ -415,9 +415,15 @@ ONNX_CPU_OPERATOR_VERSIONED_KERNEL(
KernelDefBuilder().TypeConstraint("T", DataTypeImpl::GetTensorType<float>()),
Pool<float, AveragePool>);
ONNX_CPU_OPERATOR_VERSIONED_KERNEL(
AveragePool,
10, 10,
KernelDefBuilder().TypeConstraint("T", DataTypeImpl::GetTensorType<float>()),
Pool<float, AveragePool>);
ONNX_CPU_OPERATOR_KERNEL(
AveragePool,
10,
11,
KernelDefBuilder().TypeConstraint("T", DataTypeImpl::GetTensorType<float>()),
Pool<float, AveragePool>);
@ -433,15 +439,27 @@ ONNX_CPU_OPERATOR_VERSIONED_KERNEL(
KernelDefBuilder().TypeConstraint("T", DataTypeImpl::GetTensorType<float>()).TypeConstraint("I", DataTypeImpl::GetTensorType<int64_t>()),
Pool<float, MaxPool<8 /*VERSION*/>>);
ONNX_CPU_OPERATOR_KERNEL(
ONNX_CPU_OPERATOR_VERSIONED_KERNEL(
MaxPool,
10,
10, 10,
KernelDefBuilder().TypeConstraint("T", DataTypeImpl::GetTensorType<float>()).TypeConstraint("I", DataTypeImpl::GetTensorType<int64_t>()),
Pool<float, MaxPool<8 /*VERSION*/>>);
ONNX_CPU_OPERATOR_KERNEL(
MaxPool,
11,
KernelDefBuilder().TypeConstraint("T", DataTypeImpl::GetTensorType<float>()).TypeConstraint("I", DataTypeImpl::GetTensorType<int64_t>()),
Pool<float, MaxPool<8 /*VERSION*/>>);
ONNX_CPU_OPERATOR_VERSIONED_KERNEL(
LpPool,
2,
2, 10,
KernelDefBuilder().TypeConstraint("T", DataTypeImpl::GetTensorType<float>()),
Pool<float, LpPool>);
ONNX_CPU_OPERATOR_KERNEL(
LpPool,
11,
KernelDefBuilder().TypeConstraint("T", DataTypeImpl::GetTensorType<float>()),
Pool<float, LpPool>);

View file

@ -54,7 +54,6 @@ class MaxUnpool : public OpKernel {
}
~MaxUnpool() override = default;
;
Status Compute(OpKernelContext* context) const override;

View file

@ -29,13 +29,20 @@ void TestConvOp(const ConvOpAttributes& attributes,
const std::string& err_str = "") {
OpTester test("Conv");
test.AddAttribute("auto_pad", attributes.auto_pad);
test.AddAttribute("dilations", attributes.dilations);
test.AddAttribute("group", attributes.group);
test.AddAttribute("kernel_shape", attributes.kernel_shape);
if (!attributes.dilations.empty()) {
test.AddAttribute("dilations", attributes.dilations);
}
if (!attributes.pads.empty()) {
test.AddAttribute("pads", attributes.pads);
}
test.AddAttribute("strides", attributes.strides);
if (!attributes.strides.empty()) {
test.AddAttribute("strides", attributes.strides);
}
ORT_ENFORCE(inputs.size() <= 3, "Our name array is only setup to handle 3 inputs");
const char* szNames[] = {"X", "W", "B"};
@ -50,7 +57,7 @@ void TestConvOp(const ConvOpAttributes& attributes,
if (!is_mkldnn_supported) {
excluded_providers.insert(kMklDnnExecutionProvider);
}
excluded_providers.insert(kTensorrtExecutionProvider);// Disable TensorRT because weight as input is not supported
excluded_providers.insert(kTensorrtExecutionProvider); // Disable TensorRT because weight as input is not supported
test.Run(expect_result, err_str, excluded_providers);
}
@ -78,6 +85,27 @@ TEST(ConvTest, Conv1D_1) {
TestConvOp(attrs, {X, W}, {X_shape, W_shape}, expected_vals, Y_shape);
}
TEST(ConvTest, Conv1D_1_DefaultStridesAndDilations) {
ConvOpAttributes attrs = {
"", // auto_pad
vector<int64_t>{}, // dilations
1, // group
vector<int64_t>{1}, // kernel_shape
vector<int64_t>{0, 0}, // pads
vector<int64_t>{} // strides
};
vector<float> X = {-0.21559301018714905f, 0.4691687822341919f, 0.4426700472831726f, -0.4517466723918915f,
-0.05216419696807861f, 0.29067182540893555f, 0.251010000705719f};
vector<int64_t> X_shape = {1, 1, 7};
vector<float> W = {0.24472862482070923f};
vector<int64_t> W_shape = {1, 1, 1};
vector<int64_t> Y_shape = {1, 1, 7};
auto expected_vals = {-0.052761781960725784f, 0.11481902748346329f, 0.10833403468132019f, -0.11055534332990646f,
-0.012766072526574135f, 0.07113571465015411f, 0.061429332941770554f};
TestConvOp(attrs, {X, W}, {X_shape, W_shape}, expected_vals, Y_shape);
}
// Conv3
TEST(ConvTest, Conv1D_2) {
ConvOpAttributes attrs = {
@ -179,7 +207,8 @@ TEST(ConvTest, Conv1D_Invalid_Input_Shape) {
vector<int64_t> dummy_shape = {1, 1, 2};
auto dummy_vals = {0.0f, 0.0f};
TestConvOp(attrs, {X, dummy_vals}, {X_shape, dummy_shape}, dummy_vals, dummy_shape, true, true,
OpTester::ExpectResult::kExpectFailure, "Node:node1 Output:Y [ShapeInferenceError] Can't merge shape info. "
OpTester::ExpectResult::kExpectFailure,
"Node:node1 Output:Y [ShapeInferenceError] Can't merge shape info. "
"Both source and target dimension have values but they differ. Source=0 Target=2 Dimension=2");
}
@ -198,7 +227,8 @@ TEST(ConvTest, Conv2D_Invalid_Input_Shape) {
auto dummy_vals = {-0.0f, 0.0f, -0.0f, -0.0f,
-0.0f, 0.0f, -0.0f, -0.0f};
TestConvOp(attrs, {X, dummy_vals}, {X_shape, dummy_shape}, dummy_vals, dummy_shape, true, true,
OpTester::ExpectResult::kExpectFailure, "Node:node1 Output:Y [ShapeInferenceError] Can't merge shape info. "
OpTester::ExpectResult::kExpectFailure,
"Node:node1 Output:Y [ShapeInferenceError] Can't merge shape info. "
"Both source and target dimension have values but they differ. Source=1 Target=2 Dimension=0");
}

View file

@ -28,16 +28,23 @@ void TestConvTransposeOp(const ConvTransposeOpAttributes& attributes,
const std::string& err_str = "") {
OpTester test("ConvTranspose");
test.AddAttribute("kernel_shape", attributes.kernel_shape);
test.AddAttribute("pads", attributes.pads);
test.AddAttribute("group", attributes.group);
if (!attributes.output_padding.empty()) {
test.AddAttribute("output_padding", attributes.output_padding);
}
if (!attributes.output_shape.empty()) {
test.AddAttribute("output_shape", attributes.output_shape);
}
test.AddAttribute("pads", attributes.pads);
test.AddAttribute("strides", attributes.strides);
test.AddAttribute("dilations", attributes.dilations);
test.AddAttribute("group", attributes.group);
if (!attributes.strides.empty()) {
test.AddAttribute("strides", attributes.strides);
}
if (!attributes.dilations.empty()) {
test.AddAttribute("dilations", attributes.dilations);
}
ORT_ENFORCE(inputs.size() <= 3, "Our name array is only setup to handle 3 inputs");
const char* szNames[] = {"X", "W", "B"};
@ -45,7 +52,7 @@ void TestConvTransposeOp(const ConvTransposeOpAttributes& attributes,
test.AddInput<float>(szNames[i], input_shapes[i], inputs[i]);
}
test.AddOutput<float>("Y", expected_output_shape, expected_output);
test.Run(expect_result, err_str, {kTensorrtExecutionProvider});// Disable TensorRT because weight as input is not supported
test.Run(expect_result, err_str, {kTensorrtExecutionProvider}); // Disable TensorRT because weight as input is not supported
}
} // namespace
@ -350,127 +357,158 @@ TEST(ConvTransposeTest, ConvTranspose_onnx_group) {
TEST(ConvTransposeTest, ConvTranspose_2D_Dilation_1) {
ConvTransposeOpAttributes attrs = {
vector<int64_t>{2, 2},
{}, {},
vector<int64_t>{0,0,0,0},
vector<int64_t>{1,1},
{2,2},
1
};
vector<int64_t>{2, 2},
{},
{},
vector<int64_t>{0, 0, 0, 0},
vector<int64_t>{1, 1},
{2, 2},
1};
vector<float> X = {11.0f,12.0f,21.0f,22.0f};
vector<int64_t> X_shape = {1,1,2,2};
vector<float> W = {1.0f,1.0f,1.0f,1.0f};
vector<int64_t> W_shape = {1,1,2,2};
vector<int64_t> Y_shape = {1,1,4,4};
auto expected_vals = {11.0f,12.0f,11.0f,12.0f,
21.0f,22.0f,21.0f,22.0f,
11.0f,12.0f,11.0f,12.0f,
21.0f,22.0f,21.0f,22.0f};
vector<float> X = {11.0f, 12.0f, 21.0f, 22.0f};
vector<int64_t> X_shape = {1, 1, 2, 2};
vector<float> W = {1.0f, 1.0f, 1.0f, 1.0f};
vector<int64_t> W_shape = {1, 1, 2, 2};
vector<int64_t> Y_shape = {1, 1, 4, 4};
auto expected_vals = {11.0f, 12.0f, 11.0f, 12.0f,
21.0f, 22.0f, 21.0f, 22.0f,
11.0f, 12.0f, 11.0f, 12.0f,
21.0f, 22.0f, 21.0f, 22.0f};
TestConvTransposeOp(attrs, {X, W}, {X_shape, W_shape}, expected_vals, Y_shape);
}
TEST(ConvTransposeTest, ConvTranspose_2D_Dilation_2) {
ConvTransposeOpAttributes attrs = {
vector<int64_t>{2, 2},
{}, {},
vector<int64_t>{0,0,0,0},
vector<int64_t>{1,1},
{3,3},
1
};
vector<int64_t>{2, 2},
{},
{},
vector<int64_t>{0, 0, 0, 0},
vector<int64_t>{1, 1},
{3, 3},
1};
vector<float> X = {11.0f,12.0f,21.0f,22.0f};
vector<int64_t> X_shape = {1,1,2,2};
vector<float> W = {1.0f,1.0f,1.0f,1.0f};
vector<int64_t> W_shape = {1,1,2,2};
vector<int64_t> Y_shape = {1,1,5,5};
auto expected_vals = {11.0f,12.0f,0.0f,11.0f,12.0f,
21.0f,22.0f,0.0f,21.0f,22.0f,
0.0f, 0.0f, 0.0f,0.0f, 0.0f,
11.0f,12.0f,0.0f,11.0f,12.0f,
21.0f,22.0f,0.0f,21.0f,22.0f};
vector<float> X = {11.0f, 12.0f, 21.0f, 22.0f};
vector<int64_t> X_shape = {1, 1, 2, 2};
vector<float> W = {1.0f, 1.0f, 1.0f, 1.0f};
vector<int64_t> W_shape = {1, 1, 2, 2};
vector<int64_t> Y_shape = {1, 1, 5, 5};
auto expected_vals = {11.0f, 12.0f, 0.0f, 11.0f, 12.0f,
21.0f, 22.0f, 0.0f, 21.0f, 22.0f,
0.0f, 0.0f, 0.0f, 0.0f, 0.0f,
11.0f, 12.0f, 0.0f, 11.0f, 12.0f,
21.0f, 22.0f, 0.0f, 21.0f, 22.0f};
TestConvTransposeOp(attrs, {X, W}, {X_shape, W_shape}, expected_vals, Y_shape);
}
TEST(ConvTransposeTest, ConvTranspose_2D_Dilation_3) {
ConvTransposeOpAttributes attrs = {
vector<int64_t>{2, 2},
{}, {},
vector<int64_t>{0,0,0,0},
vector<int64_t>{1,1},
{2,2},
1
};
vector<int64_t>{2, 2},
{},
{},
vector<int64_t>{0, 0, 0, 0},
vector<int64_t>{1, 1},
{2, 2},
1};
vector<float> X = {3.0f,8.0f,1.0f,9.0f,5.0f,7.0f,3.0f,2.0f,6.0f};
vector<int64_t> X_shape = {1,1,3,3};
vector<float> W = {7.0f,2.0f,1.0f,9.0f};
vector<int64_t> W_shape = {1,1,2,2};
vector<int64_t> Y_shape = {1,1,5,5};
auto expected_vals = {21.0f, 56.0f, 13.0f, 16.0f, 2.0f,
vector<float> X = {3.0f, 8.0f, 1.0f, 9.0f, 5.0f, 7.0f, 3.0f, 2.0f, 6.0f};
vector<int64_t> X_shape = {1, 1, 3, 3};
vector<float> W = {7.0f, 2.0f, 1.0f, 9.0f};
vector<int64_t> W_shape = {1, 1, 2, 2};
vector<int64_t> Y_shape = {1, 1, 5, 5};
auto expected_vals = {21.0f, 56.0f, 13.0f, 16.0f, 2.0f,
63.0f, 35.0f, 67.0f, 10.0f, 14.0f,
24.0f, 22.0f, 76.0f, 76.0f, 21.0f,
9.0f, 5.0f, 88.0f, 45.0f, 63.0f,
3.0f, 2.0f, 33.0f, 18.0f, 54.0f};
9.0f, 5.0f, 88.0f, 45.0f, 63.0f,
3.0f, 2.0f, 33.0f, 18.0f, 54.0f};
TestConvTransposeOp(attrs, {X, W}, {X_shape, W_shape}, expected_vals, Y_shape);
}
TEST(ConvTransposeTest, ConvTranspose_2D_Dilation_4) {
ConvTransposeOpAttributes attrs = {
vector<int64_t>{2, 2},
{}, {},
vector<int64_t>{0,0,0,0},
vector<int64_t>{1,1},
{3,3},
1
};
vector<int64_t>{2, 2},
{},
{},
vector<int64_t>{0, 0, 0, 0},
vector<int64_t>{1, 1},
{3, 3},
1};
vector<float> X = {3.0f,8.0f,1.0f,9.0f,5.0f,7.0f,3.0f,2.0f,6.0f};
vector<int64_t> X_shape = {1,1,3,3};
vector<float> W = {7.0f,2.0f,1.0f,9.0f};
vector<int64_t> W_shape = {1,1,2,2};
vector<int64_t> Y_shape = {1,1,6,6};
auto expected_vals = {21.0f, 56.0f, 7.0f, 6.0f, 16.0f, 2.0f,
vector<float> X = {3.0f, 8.0f, 1.0f, 9.0f, 5.0f, 7.0f, 3.0f, 2.0f, 6.0f};
vector<int64_t> X_shape = {1, 1, 3, 3};
vector<float> W = {7.0f, 2.0f, 1.0f, 9.0f};
vector<int64_t> W_shape = {1, 1, 2, 2};
vector<int64_t> Y_shape = {1, 1, 6, 6};
auto expected_vals = {21.0f, 56.0f, 7.0f, 6.0f, 16.0f, 2.0f,
63.0f, 35.0f, 49.0f, 18.0f, 10.0f, 14.0f,
21.0f, 14.0f, 42.0f, 6.0f, 4.0f, 12.0f,
3.0f, 8.0f, 1.0f, 27.0f, 72.0f, 9.0f,
9.0f, 5.0f, 7.0f, 81.0f, 45.0f, 63.0f,
3.0f, 2.0f, 6.0f, 27.0f, 18.0f, 54.0f};
21.0f, 14.0f, 42.0f, 6.0f, 4.0f, 12.0f,
3.0f, 8.0f, 1.0f, 27.0f, 72.0f, 9.0f,
9.0f, 5.0f, 7.0f, 81.0f, 45.0f, 63.0f,
3.0f, 2.0f, 6.0f, 27.0f, 18.0f, 54.0f};
TestConvTransposeOp(attrs, {X, W}, {X_shape, W_shape}, expected_vals, Y_shape);
}
TEST(ConvTransposeTest, ConvTranspose_2D_Dilation_Group_1) {
ConvTransposeOpAttributes attrs = {
vector<int64_t>{2, 2},
{}, {},
vector<int64_t>{0,0,0,0},
vector<int64_t>{1,1},
{2,2},
2
};
vector<int64_t>{2, 2},
{},
{},
vector<int64_t>{0, 0, 0, 0},
vector<int64_t>{1, 1},
{2, 2},
2};
vector<float> X = {3.0f,8.0f,1.0f,9.0f,5.0f,7.0f,3.0f,2.0f,3.0f,7.0f,9.0f,1.0f,5.0f,2.0f,3.0f,9.0f,0.0f,2.0f};
vector<int64_t> X_shape = {1,2,3,3};
vector<float> W = {9.0f,3.0f,1.0f,2.0f,3.0f,7.0f,0.0f,8.0f};
vector<int64_t> W_shape = {2,1,2,2};
vector<int64_t> Y_shape = {1,2,5,5};
auto expected_vals = {27.0f, 72.0f, 18.0f, 24.0f, 3.0f,
81.0f, 45.0f, 90.0f, 15.0f, 21.0f,
30.0f, 26.0f, 43.0f, 22.0f, 11.0f,
9.0f, 5.0f, 25.0f, 10.0f, 14.0f,
3.0f, 2.0f, 9.0f, 4.0f, 6.0f,
21.0f, 27.0f, 52.0f, 63.0f, 7.0f,
15.0f, 6.0f, 44.0f, 14.0f, 21.0f,
27.0f, 0.0f, 125.0f, 72.0f, 22.0f,
0.0f, 0.0f, 40.0f, 16.0f, 24.0f,
0.0f, 0.0f, 72.0f, 0.0f, 16.0f};
vector<float> X = {3.0f, 8.0f, 1.0f, 9.0f, 5.0f, 7.0f, 3.0f, 2.0f, 3.0f, 7.0f, 9.0f, 1.0f, 5.0f, 2.0f, 3.0f, 9.0f, 0.0f, 2.0f};
vector<int64_t> X_shape = {1, 2, 3, 3};
vector<float> W = {9.0f, 3.0f, 1.0f, 2.0f, 3.0f, 7.0f, 0.0f, 8.0f};
vector<int64_t> W_shape = {2, 1, 2, 2};
vector<int64_t> Y_shape = {1, 2, 5, 5};
auto expected_vals = {27.0f, 72.0f, 18.0f, 24.0f, 3.0f,
81.0f, 45.0f, 90.0f, 15.0f, 21.0f,
30.0f, 26.0f, 43.0f, 22.0f, 11.0f,
9.0f, 5.0f, 25.0f, 10.0f, 14.0f,
3.0f, 2.0f, 9.0f, 4.0f, 6.0f,
21.0f, 27.0f, 52.0f, 63.0f, 7.0f,
15.0f, 6.0f, 44.0f, 14.0f, 21.0f,
27.0f, 0.0f, 125.0f, 72.0f, 22.0f,
0.0f, 0.0f, 40.0f, 16.0f, 24.0f,
0.0f, 0.0f, 72.0f, 0.0f, 16.0f};
TestConvTransposeOp(attrs, {X, W}, {X_shape, W_shape}, expected_vals, Y_shape);
}
TEST(ConvTransposeTest, ConvTranspose_DefaultStridesAndDilations) {
ConvTransposeOpAttributes attrs = {
vector<int64_t>{2, 2}, // kernel_shape
{}, // output_padding
{}, // output_shape
vector<int64_t>{0, 0, 0, 0}, // pads
vector<int64_t>{}, // strides
vector<int64_t>{}, // dilations
1 // group
};
vector<float> X = {0., 1., 2., 3., 4., 5., 6., 7., 8., 9., 10., 11., 12., 13., 14., 15., 16., 17.};
vector<int64_t> X_shape = {1, 2, 3, 3};
vector<float> W = {0., 1., 2., 3., 4., 5., 6., 7., 8., 9., 10., 11., 12., 13., 14., 15., 16., 17., 18., 19., 20., 21., 22., 23.};
vector<int64_t> W_shape = {2, 3, 2, 2}; // this requires weight transpose
vector<int64_t> Y_shape = {1, 3, 4, 4};
auto expected_vals = {
108.f, 237.f, 263.f, 145.f,
270.f, 592.f, 652.f, 358.f,
354.f, 772.f, 832.f, 454.f,
222.f, 481.f, 515.f, 279.f,
144.f, 317.f, 359.f, 197.f,
366.f, 800.f, 892.f, 486.f,
498.f, 1076.f, 1168.f, 630.f,
306.f, 657.f, 707.f, 379.f,
180.f, 397.f, 455.f, 249.f,
462.f, 1008.f, 1132.f, 614.f,
642.f, 1380.f, 1504.f, 806.f,
390.f, 833.f, 899.f, 479.f};
TestConvTransposeOp(attrs, {X, W}, {X_shape, W_shape}, expected_vals, Y_shape);
}
} // namespace test
} // namespace onnxruntime

View file

@ -229,6 +229,26 @@ TEST(PoolTest, MaxPool_10_Dilation_1d) {
test.Run(OpTester::ExpectResult::kExpectSuccess, "", {kTensorrtExecutionProvider});
}
TEST(PoolTest, MaxPool_DefaultDilations) {
OpTester test("MaxPool");
test.AddAttribute("kernel_shape", vector<int64_t>{2});
std::vector<int64_t> x_dims = {1, 3, 3};
std::vector<float> x_vals = {0.f, 1.f, 2.f,
3.f, 4.f, 5.f,
6.f, 7.f, 8.f};
std::vector<int64_t> expected_dims = {1, 3, 2};
std::vector<float> expected_vals = {1.f, 2.f,
4.f, 5.f,
7.f, 8.f};
test.AddInput<float>("X", x_dims, x_vals);
test.AddOutput<float>("Y", expected_dims, expected_vals);
test.Run(OpTester::ExpectResult::kExpectSuccess, "", {kTensorrtExecutionProvider});
}
TEST(PoolTest, MaxPool_10_DilationPadding_1d) {
OpTester test("MaxPool", 10);
@ -638,6 +658,25 @@ TEST(PoolTest, AveragePool_IncludePadPixel) {
test.Run(OpTester::ExpectResult::kExpectSuccess, "", {kTensorrtExecutionProvider});
}
// test 'strides' attribute not specified
TEST(PoolTest, AveragePool_DefaultStrides) {
OpTester test("AveragePool");
test.AddAttribute("kernel_shape", vector<int64_t>{2});
std::vector<float> x_vals = {0.f, 1.f, 2.f,
3.f, 4.f, 5.f,
6.f, 7.f, 8.f};
std::vector<int64_t> x_dims = {1, 3, 3};
std::vector<int64_t> expected_dims = {1, 3, 2};
std::vector<float> expected_vals = {0.5f, 1.5f,
3.5f, 4.5f,
6.5f, 7.5f};
test.AddInput<float>("X", x_dims, x_vals);
test.AddOutput<float>("Y", expected_dims, expected_vals);
test.Run(OpTester::ExpectResult::kExpectSuccess, "", {kTensorrtExecutionProvider});
}
TEST(PoolTest, AveragePool_10_ceil1_2d) {
OpTester test("AveragePool", 10);

View file

@ -378,5 +378,25 @@ TEST(UnpoolTest, MaxUnPool3D_WithPaddedOutput) {
test.Run();
}
TEST(UnpoolTest, MaxUnPool_DefaultStrides) {
OpTester test("MaxUnpool", 11);
test.AddAttribute("kernel_shape", vector<int64_t>{2});
std::vector<float> t_vals = {1, 2, 4, 8};
std::vector<int64_t> t_dims = {1, 1, 4};
std::vector<int64_t> i_vals = {1, 2, 3, 4};
std::vector<int64_t> i_dims = {1, 1, 4};
std::vector<int64_t> expected_dims = {1, 1, 5};
std::vector<float> expected_vals = {0, 1, 2, 4, 8};
test.AddInput<float>("xT", t_dims, t_vals);
test.AddInput<int64_t>("xI", i_dims, i_vals);
test.AddOutput<float>("Y", expected_dims, expected_vals);
test.Run();
}
} // namespace test
} // namespace onnxruntime