From aae18a3fe36d0b8e85272382822cdf76a0e0ac6a Mon Sep 17 00:00:00 2001 From: Nathan <7902510+ybrnathan@users.noreply.github.com> Date: Sun, 20 Oct 2019 10:44:20 -0700 Subject: [PATCH] Upgrade onehot to OpSet 11 (#2185) * Upgrade onehot to OpSet 11 * Move Onehot test out of blacklist * Add negative indices support besides negative axis. * PR comments - 1 * PR comments-2 --- .../providers/cpu/cpu_execution_provider.cc | 69 ++++++++++++++----- .../core/providers/cpu/tensor/onehot.cc | 41 +++++++++-- onnxruntime/test/onnx/main.cc | 4 -- .../providers/cpu/tensor/onehot_op_test.cc | 20 ++++++ .../test/python/onnx_backend_test_series.py | 1 - 5 files changed, 106 insertions(+), 29 deletions(-) diff --git a/onnxruntime/core/providers/cpu/cpu_execution_provider.cc b/onnxruntime/core/providers/cpu/cpu_execution_provider.cc index 7b062cc218..7b93a1538d 100644 --- a/onnxruntime/core/providers/cpu/cpu_execution_provider.cc +++ b/onnxruntime/core/providers/cpu/cpu_execution_provider.cc @@ -239,15 +239,16 @@ class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, Sign); class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, Shrink); class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, float, Erf); -class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, int64_t_int64_t_int64_t, OneHot); -class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, float_int64_t_int64_t, OneHot); -class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, int64_t_string_int64_t, OneHot); -class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, float_string_int64_t, OneHot); -class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, float_float_float, OneHot); -class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, int64_t_int32_t_float, OneHot); -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_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, 10, int64_t_int64_t_int64_t, OneHot); +class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, 10, float_int64_t_int64_t, OneHot); +class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, 10, int64_t_string_int64_t, OneHot); +class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, 10, float_string_int64_t, OneHot); +class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, 10, float_float_float, OneHot); +class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, 10, int64_t_int32_t_float, OneHot); +class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, 10, int64_t_float_int64_t, OneHot); +class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, 10, int32_t_float_int32_t, OneHot); +class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, 10, int32_t_float_float, OneHot); +class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, 10, int64_t_float_float, OneHot); 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); @@ -410,6 +411,16 @@ class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, Ga class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, Range); class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, Unique); class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, TopK); +class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, int64_t_int64_t_int64_t, OneHot); +class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, float_int64_t_int64_t, OneHot); +class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, int64_t_string_int64_t, OneHot); +class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, float_string_int64_t, OneHot); +class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, float_float_float, OneHot); +class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, int64_t_int32_t_float, OneHot); +class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, int64_t_float_int64_t, OneHot); +class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, int32_t_float_int32_t, OneHot); +class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, int32_t_float_float, OneHot); +class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, int64_t_float_float, OneHot); class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, float, Resize); Status RegisterOnnxOperatorKernels(KernelRegistry& kernel_registry) { @@ -814,24 +825,26 @@ Status RegisterOnnxOperatorKernels(KernelRegistry& kernel_registry) { BuildKernelCreateInfo, BuildKernelCreateInfo, BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, + BuildKernelCreateInfo, BuildKernelCreateInfo, BuildKernelCreateInfo, @@ -1047,6 +1060,26 @@ Status RegisterOnnxOperatorKernels(KernelRegistry& kernel_registry) { BuildKernelCreateInfo, BuildKernelCreateInfo, BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, BuildKernelCreateInfo, }; diff --git a/onnxruntime/core/providers/cpu/tensor/onehot.cc b/onnxruntime/core/providers/cpu/tensor/onehot.cc index c4f0c2479a..e033e7aed8 100644 --- a/onnxruntime/core/providers/cpu/tensor/onehot.cc +++ b/onnxruntime/core/providers/cpu/tensor/onehot.cc @@ -28,10 +28,22 @@ namespace onnxruntime { // spec: https://github.com/onnx/onnx/blob/master/docs/Operators.md#OneHot // T1: indices, T2: depth, T3: values -#define REG_TYPED_ONE_HOT_OP(types_str, in_type, out_type, depth_type) \ +#define REG_TYPED_ONE_HOT_OP_V9_10(types_str, in_type, out_type, depth_type) \ + ONNX_CPU_OPERATOR_VERSIONED_TYPED_KERNEL( \ + OneHot, \ + 9, 10, \ + types_str, \ + KernelDefBuilder() \ + .TypeConstraint("T1", DataTypeImpl::GetTensorType()) \ + .TypeConstraint("T2", DataTypeImpl::GetTensorType()) \ + .TypeConstraint("T3", DataTypeImpl::GetTensorType()), \ + OneHotOp); + +// T1: indices, T2: depth, T3: values +#define REG_TYPED_ONE_HOT_OP_V11(types_str, in_type, out_type, depth_type) \ ONNX_CPU_OPERATOR_TYPED_KERNEL( \ OneHot, \ - 9, \ + 11, \ types_str, \ KernelDefBuilder() \ .TypeConstraint("T1", DataTypeImpl::GetTensorType()) \ @@ -39,8 +51,9 @@ namespace onnxruntime { .TypeConstraint("T3", DataTypeImpl::GetTensorType()), \ OneHotOp); -#define REG_ONE_HOT_OP(in_type, out_type, depth_type) \ - REG_TYPED_ONE_HOT_OP(in_type##_##out_type##_##depth_type, in_type, out_type, depth_type) +#define REG_ONE_HOT_OP(in_type, out_type, depth_type) \ + REG_TYPED_ONE_HOT_OP_V9_10(in_type##_##out_type##_##depth_type, in_type, out_type, depth_type); \ + REG_TYPED_ONE_HOT_OP_V11(in_type##_##out_type##_##depth_type, in_type, out_type, depth_type) REG_ONE_HOT_OP(int64_t, int64_t, int64_t); REG_ONE_HOT_OP(float, int64_t, int64_t); @@ -51,6 +64,7 @@ REG_ONE_HOT_OP(int32_t, float, int32_t); REG_ONE_HOT_OP(int32_t, float, float); REG_ONE_HOT_OP(float, float, float); // added this to satisfy onnx model tests REG_ONE_HOT_OP(int64_t, int32_t, float); // added this to satisfy onnx model tests +REG_ONE_HOT_OP(int64_t, float, float); // added this to satisfy onnx model tests Status ValidateInputs(const Tensor* depth, const Tensor* values) { @@ -129,7 +143,7 @@ Status OneHotOp::Compute(OpKernelContext* p_op_ke const auto output_rank = static_cast(indices_num_dims + 1); if (axis_ >= output_rank || axis_ < -output_rank) { std::ostringstream oss; - oss << "'axis' attribute must have a value in the range [" << -output_rank + oss << "'axis' attribute must have a value in the range [" << -output_rank << "," << indices_num_dims << "]"; return Status(ONNXRUNTIME, INVALID_ARGUMENT, oss.str()); } @@ -152,8 +166,23 @@ Status OneHotOp::Compute(OpKernelContext* p_op_ke // Split indices into matrix of size prefix_dim_size x suffix_dim_size Eigen::array indices_dims_e = { - {static_cast(prefix_dim_size), static_cast(suffix_dim_size)}}; + {static_cast(prefix_dim_size), static_cast(suffix_dim_size)}}; + + // Handle negative indices. It's faster to create a new indices instead of comparing in generator + // since generator has much larger loops. const auto* indices_data = indices->Data(); + const auto indices_size = indices_shape.Size(); + std::vector adjusted_indices; + adjusted_indices.reserve(indices_size); + for (int64_t i = 0; i < indices_size; ++i) + { + if (indices_data[i] < 0) + adjusted_indices.push_back(indices_data[i] + static_cast(depth_val)); + else + adjusted_indices.push_back(indices_data[i]); + } + indices_data = adjusted_indices.data(); + typename EigenTensorTypes::ConstEigenTensorMap indices_tensor_e(indices_data, indices_dims_e); // Split output into 3-Tensor of size: diff --git a/onnxruntime/test/onnx/main.cc b/onnxruntime/test/onnx/main.cc index 48733ecd69..f3115be11c 100644 --- a/onnxruntime/test/onnx/main.cc +++ b/onnxruntime/test/onnx/main.cc @@ -446,10 +446,6 @@ int real_main(int argc, char* argv[], Ort::Env& env) { {"sequence_model2", "SequenceConstruct not implemented yet"}, {"sequence_model1", "Sequence* not implemented yet"}, {"scatter_elements_with_negative_indices", "ScatterElements(11) not implemented yet"}, - {"onehot_without_axis", "OneHot(11) not implemented yet"}, - {"onehot_with_negative_axis", "OneHot(11) not implemented yet"}, - {"onehot_with_axis", "OneHot(11) not implemented yet"}, - {"onehot_negative_indices", "OneHot(11) not implemented yet"}, {"bitshift_right_uint8", "BitShift(11) uint8 support not enabled currently"}, {"bitshift_right_uint16", "BitShift(11) uint16 support not enabled currently"}, {"bitshift_left_uint8", "BitShift(11) uint8 support not enabled currently"}, diff --git a/onnxruntime/test/providers/cpu/tensor/onehot_op_test.cc b/onnxruntime/test/providers/cpu/tensor/onehot_op_test.cc index 63cebfa1e7..a5eef11481 100644 --- a/onnxruntime/test/providers/cpu/tensor/onehot_op_test.cc +++ b/onnxruntime/test/providers/cpu/tensor/onehot_op_test.cc @@ -206,5 +206,25 @@ TEST(OneHotOpTest, FloatString) { "off", "off", "off", "off", "off", "off", "on", "off", "off", "off",}); test.Run(); } + +TEST(OneHotOpTest, Axis_Negative_NegIndex_NonDefault) { + OpTester test("OneHot", 11); + int64_t axis = -3; + test.AddAttribute("axis", axis); + test.AddInput("indices", {2, 3}, {1, -1, 8, 2, 4, 6}); + test.AddInput("depth", {1}, {10}); + test.AddInput("values", {2}, {0, 1}); + test.AddOutput("output", {10, 2, 3}, { 0, 0, 0, 0, 0, 0, + 1, 0, 0, 0, 0, 0, + 0, 0, 0, 1, 0, 0, + 0, 0, 0, 0, 0, 0, + 0, 0, 0, 0, 1, 0, + 0, 0, 0, 0, 0, 0, + 0, 0, 0, 0, 0, 1, + 0, 0, 0, 0, 0, 0, + 0, 0, 1, 0, 0, 0, + 0, 1, 0, 0, 0, 0,}); + test.Run(); +} } } // namespace onnxruntime diff --git a/onnxruntime/test/python/onnx_backend_test_series.py b/onnxruntime/test/python/onnx_backend_test_series.py index 7a532f1c09..d1801e6b3f 100644 --- a/onnxruntime/test/python/onnx_backend_test_series.py +++ b/onnxruntime/test/python/onnx_backend_test_series.py @@ -151,7 +151,6 @@ def create_backend_test(testname=None): '^test_resize_upsample_sizes_nearest_round_prefer_ceil_asymmetric_cpu', '^test_sequence_*', '^test_scatter_.*', - '^test_onehot_.*', '^test_edge_pad_cpu.*', # test data type `int32_t` not supported yet, the `float` equivalent is covered via unit tests '^test_reflect_pad_cpu.*' # test data type `int32_t` not supported yet, the `float` equivalent is covered via unit tests ]