mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-30 20:18:08 +00:00
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:
parent
dc03ce0278
commit
c0d953a268
3 changed files with 98 additions and 4 deletions
|
|
@ -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} {
|
||||
|
|
|
|||
|
|
@ -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) {
|
||||
|
|
|
|||
|
|
@ -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) {
|
||||
|
|
|
|||
Loading…
Reference in a new issue