Support opset-11 Loop CPU kernel (#1816)

* Add support for opset 11 Loop

* Change test name to be more verbose

* Add a new kernel for Loop - 11
This commit is contained in:
Hariharan Seshadri 2019-09-17 14:17:25 -07:00 committed by GitHub
parent dc03ce0278
commit c0d953a268
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23
3 changed files with 98 additions and 4 deletions

View file

@ -79,14 +79,22 @@ ONNX_OPERATOR_SET_SCHEMA(
.TypeAndShapeInferenceFunction(LoopInferenceFunction));
*/
ONNX_CPU_OPERATOR_KERNEL(Loop,
1,
ONNX_CPU_OPERATOR_VERSIONED_KERNEL(Loop,
1, 10,
KernelDefBuilder()
.TypeConstraint("I", DataTypeImpl::GetTensorType<int64_t>())
.TypeConstraint("B", DataTypeImpl::GetTensorType<bool>())
.TypeConstraint("V", DataTypeImpl::AllTensorTypes()),
Loop);
ONNX_CPU_OPERATOR_KERNEL(Loop,
11,
KernelDefBuilder()
.TypeConstraint("I", DataTypeImpl::GetTensorType<int64_t>())
.TypeConstraint("B", DataTypeImpl::GetTensorType<bool>())
.TypeConstraint("V", DataTypeImpl::AllTensorTypes()),
Loop);
struct Loop::Info {
Info(const onnxruntime::Node& node, const GraphViewer& subgraph_in)
: subgraph{subgraph_in} {

View file

@ -218,7 +218,7 @@ class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain,
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 8, MLFloat16, Expand);
class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 8, 8, Scan);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, If);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, Loop);
class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, 10, Loop);
// Opset 9
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, Compress);
@ -314,6 +314,7 @@ class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain,
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, Hardmax);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, LogSoftmax);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, Softmax);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, Loop);
void RegisterOnnxOperatorKernels(KernelRegistry& kernel_registry) {
static const BuildKernelCreateInfoFn function_table[] = {
@ -517,7 +518,7 @@ void RegisterOnnxOperatorKernels(KernelRegistry& kernel_registry) {
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 8, MLFloat16, Expand)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 8, 8, Scan)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, If)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, Loop)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, 10, Loop)>,
// Opset 9
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, Compress)>,
@ -613,6 +614,8 @@ void RegisterOnnxOperatorKernels(KernelRegistry& kernel_registry) {
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, Hardmax)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, LogSoftmax)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, Softmax)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, Loop)>,
};
for (auto& function_table_entry : function_table) {

View file

@ -630,6 +630,89 @@ TEST(Loop, SubgraphInputShadowsOuterScopeValue) {
}
}
TEST(Loop, Opset11WithNoVariadicInputsAndOutputs) {
auto create_subgraph = []() {
Model model("Loop opset 11 op body graph");
auto& graph = model.MainGraph();
std::vector<NodeArg*> inputs;
std::vector<NodeArg*> outputs;
// graph inputs types.
// iteration number
TypeProto int64_scalar;
int64_scalar.mutable_tensor_type()->set_elem_type(TensorProto_DataType_INT64);
int64_scalar.mutable_tensor_type()->mutable_shape()->add_dim()->set_dim_value(1);
// loop condition
TypeProto bool_scalar;
bool_scalar.mutable_tensor_type()->set_elem_type(TensorProto_DataType_BOOL);
bool_scalar.mutable_tensor_type()->mutable_shape()->add_dim()->set_dim_value(1);
// graph output types
// constant_out
TypeProto float_scalar;
float_scalar.mutable_tensor_type()->set_elem_type(TensorProto_DataType_FLOAT);
// graph inputs
auto& iter_num_in = graph.GetOrCreateNodeArg("iter_num_in", &int64_scalar);
auto& cond_in = graph.GetOrCreateNodeArg("cond_in", &bool_scalar);
// graph outputs
auto& cond_out = graph.GetOrCreateNodeArg("cond_out", &bool_scalar);
auto& constant_out = graph.GetOrCreateNodeArg("constant_out", &float_scalar);
// cond_in -> cond_out
{
inputs = {&cond_in};
outputs = {&cond_out};
graph.AddNode("cond_in_identity", "Identity", "Forward cond_in to cond_out", inputs, outputs);
}
// produce constant_out
{
outputs = {&constant_out};
TensorProto constant_tensor_proto;
auto& constant_node = graph.AddNode("constant_out", "Constant", "Produce constant_out", {}, outputs);
AttributeProto attr_proto;
attr_proto.set_name("value");
attr_proto.set_type(AttributeProto_AttributeType_TENSOR);
auto* constant_attribute_tensor_proto = attr_proto.mutable_t();
constant_attribute_tensor_proto->mutable_dims()->Clear(); // scalar
constant_attribute_tensor_proto->set_data_type(TensorProto_DataType_FLOAT); //float scalar
*constant_attribute_tensor_proto->mutable_float_data()->Add() = 1.0f; //float scalar with value 1.0f
constant_node.AddAttribute("value", attr_proto);
}
graph.SetInputs({&iter_num_in, &cond_in});
graph.SetOutputs({&cond_out, &constant_out});
auto status = graph.Resolve();
EXPECT_EQ(status, Status::OK());
return graph.ToGraphProto();
};
OpTester test("Loop", 11);
auto body = create_subgraph();
test.AddAttribute<GraphProto>("body", body);
test.AddInput<int64_t>("M", {1}, {1});
test.AddInput<bool>("cond", {1}, {true});
// This 'Loop' has no variadic inputs to test the spec of 'Loop' opset 11 which allows
// 'Loop' to be used without variadic inputs
test.AddOutput<float>("loop_scan_out", {1}, {1.0f});
// Disable TensorRT on unsupported data type BOOL
test.Run(OpTester::ExpectResult::kExpectSuccess, "", {kTensorrtExecutionProvider});
}
#ifdef USE_CUDA
// test that when part of the subgraph run on CUDA it executes successfully
TEST(Loop, MixedExecutionProviders) {